Compare commits
4
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
201db8051d | ||
|
|
b8458ada48 | ||
|
|
806d03332f | ||
|
|
c0fe313b61 |
@@ -2853,14 +2853,7 @@ class PreTrainedTokenizerBase(SpecialTokensMixin):
|
||||
encoded_inputs["special_tokens_mask"] = [0] * len(sequence)
|
||||
|
||||
# Check lengths
|
||||
if max_length is None and len(encoded_inputs["input_ids"]) > self.model_max_length and verbose:
|
||||
if not self.deprecation_warnings.get("sequence-length-is-longer-than-the-specified-maximum", False):
|
||||
logger.warning(
|
||||
"Token indices sequence length is longer than the specified maximum sequence length "
|
||||
"for this model ({} > {}). Running this sequence through the model will result in "
|
||||
"indexing errors".format(len(encoded_inputs["input_ids"]), self.model_max_length)
|
||||
)
|
||||
self.deprecation_warnings["sequence-length-is-longer-than-the-specified-maximum"] = True
|
||||
self._eventual_warn_about_too_long_sequence(encoded_inputs["input_ids"], max_length, verbose)
|
||||
|
||||
# Padding
|
||||
if padding_strategy != PaddingStrategy.DO_NOT_PAD or return_attention_mask:
|
||||
@@ -3173,10 +3166,10 @@ class PreTrainedTokenizerBase(SpecialTokensMixin):
|
||||
Clean up a list of simple English tokenization artifacts like spaces before punctuations and abbreviated forms.
|
||||
|
||||
Args:
|
||||
out_string (:obj:`str`): The text to clean up.
|
||||
out_string (:obj:`str`): the text to clean up.
|
||||
|
||||
Returns:
|
||||
:obj:`str`: The cleaned-up string.
|
||||
:obj:`str`: the cleaned-up string.
|
||||
"""
|
||||
out_string = (
|
||||
out_string.replace(" .", ".")
|
||||
@@ -3191,3 +3184,23 @@ class PreTrainedTokenizerBase(SpecialTokensMixin):
|
||||
.replace(" 're", "'re")
|
||||
)
|
||||
return out_string
|
||||
|
||||
def _eventual_warn_about_too_long_sequence(self, ids: List[int], max_length: Optional[int], verbose: bool):
|
||||
"""
|
||||
Depending on the input and internal state we might trigger a warning about a sequence that is too long for it's
|
||||
corresponding model
|
||||
|
||||
Args:
|
||||
ids (:obj:`List[str]`): The ids produced by the tokenization
|
||||
max_length (:obj:`int`, `optional`): The max_length desired (does not trigger a warning if it is set)
|
||||
verbose (:obj:`bool`): Whether or not to print more information and warnings.
|
||||
|
||||
"""
|
||||
if max_length is None and len(ids) > self.model_max_length and verbose:
|
||||
if not self.deprecation_warnings.get("sequence-length-is-longer-than-the-specified-maximum", False):
|
||||
logger.warning(
|
||||
"Token indices sequence length is longer than the specified maximum sequence length "
|
||||
"for this model ({} > {}). Running this sequence through the model will result in "
|
||||
"indexing errors".format(len(ids), self.model_max_length)
|
||||
)
|
||||
self.deprecation_warnings["sequence-length-is-longer-than-the-specified-maximum"] = True
|
||||
|
||||
@@ -418,6 +418,8 @@ class PreTrainedTokenizerFast(PreTrainedTokenizerBase):
|
||||
overflow_to_sample_mapping += [i] * len(toks["input_ids"])
|
||||
sanitized_tokens["overflow_to_sample_mapping"] = overflow_to_sample_mapping
|
||||
|
||||
for input_ids in sanitized_tokens["input_ids"]:
|
||||
self._eventual_warn_about_too_long_sequence(input_ids, max_length, verbose)
|
||||
return BatchEncoding(sanitized_tokens, sanitized_encodings, tensor_type=return_tensors)
|
||||
|
||||
def _encode_plus(
|
||||
@@ -474,6 +476,8 @@ class PreTrainedTokenizerFast(PreTrainedTokenizerBase):
|
||||
batched_output.encodings,
|
||||
)
|
||||
|
||||
self._eventual_warn_about_too_long_sequence(batched_output["input_ids"], max_length, verbose)
|
||||
|
||||
return batched_output
|
||||
|
||||
def convert_tokens_to_string(self, tokens: List[str]) -> str:
|
||||
|
||||
@@ -666,11 +666,28 @@ class TokenizerTesterMixin:
|
||||
self.assertEqual(len(output["input_ids"][0]), model_max_length)
|
||||
|
||||
# Simple with no truncation
|
||||
output = tokenizer(seq_1, padding=padding_state, truncation=False)
|
||||
self.assertNotEqual(len(output["input_ids"]), model_max_length)
|
||||
# Reset warnings
|
||||
tokenizer.deprecation_warnings = {}
|
||||
with self.assertLogs("transformers", level="WARNING") as cm:
|
||||
output = tokenizer(seq_1, padding=padding_state, truncation=False)
|
||||
self.assertNotEqual(len(output["input_ids"]), model_max_length)
|
||||
self.assertEqual(len(cm.records), 1)
|
||||
self.assertTrue(
|
||||
cm.records[0].message.startswith(
|
||||
"Token indices sequence length is longer than the specified maximum sequence length for this model"
|
||||
)
|
||||
)
|
||||
|
||||
output = tokenizer([seq_1], padding=padding_state, truncation=False)
|
||||
self.assertNotEqual(len(output["input_ids"][0]), model_max_length)
|
||||
tokenizer.deprecation_warnings = {}
|
||||
with self.assertLogs("transformers", level="WARNING") as cm:
|
||||
output = tokenizer([seq_1], padding=padding_state, truncation=False)
|
||||
self.assertNotEqual(len(output["input_ids"][0]), model_max_length)
|
||||
self.assertEqual(len(cm.records), 1)
|
||||
self.assertTrue(
|
||||
cm.records[0].message.startswith(
|
||||
"Token indices sequence length is longer than the specified maximum sequence length for this model"
|
||||
)
|
||||
)
|
||||
|
||||
# Overflowing tokens
|
||||
stride = 2
|
||||
@@ -770,11 +787,28 @@ class TokenizerTesterMixin:
|
||||
self.assertEqual(len(output["input_ids"][0]), model_max_length)
|
||||
|
||||
# Simple with no truncation
|
||||
output = tokenizer(seq_1, seq_2, padding=padding_state, truncation=False)
|
||||
self.assertNotEqual(len(output["input_ids"]), model_max_length)
|
||||
# Reset warnings
|
||||
tokenizer.deprecation_warnings = {}
|
||||
with self.assertLogs("transformers", level="WARNING") as cm:
|
||||
output = tokenizer(seq_1, seq_2, padding=padding_state, truncation=False)
|
||||
self.assertNotEqual(len(output["input_ids"]), model_max_length)
|
||||
self.assertEqual(len(cm.records), 1)
|
||||
self.assertTrue(
|
||||
cm.records[0].message.startswith(
|
||||
"Token indices sequence length is longer than the specified maximum sequence length for this model"
|
||||
)
|
||||
)
|
||||
|
||||
output = tokenizer([seq_1], [seq_2], padding=padding_state, truncation=False)
|
||||
self.assertNotEqual(len(output["input_ids"][0]), model_max_length)
|
||||
tokenizer.deprecation_warnings = {}
|
||||
with self.assertLogs("transformers", level="WARNING") as cm:
|
||||
output = tokenizer([seq_1], [seq_2], padding=padding_state, truncation=False)
|
||||
self.assertNotEqual(len(output["input_ids"][0]), model_max_length)
|
||||
self.assertEqual(len(cm.records), 1)
|
||||
self.assertTrue(
|
||||
cm.records[0].message.startswith(
|
||||
"Token indices sequence length is longer than the specified maximum sequence length for this model"
|
||||
)
|
||||
)
|
||||
|
||||
truncated_first_sequence = tokenizer.encode(seq_0, add_special_tokens=False)[:-2] + tokenizer.encode(
|
||||
seq_1, add_special_tokens=False
|
||||
|
||||
Reference in New Issue
Block a user