Fix truncation when no max length is specified

This commit is contained in:
Lysandre
2020-12-09 21:18:21 +01:00
committed by Rogge Niels
parent 71890e5082
commit 212f4b0d59
2 changed files with 12 additions and 1 deletions
@@ -1152,7 +1152,7 @@ class TapasTokenizer(PreTrainedTokenizer):
num_columns = self._get_num_columns(raw_table)
_, _, num_tokens = self._get_table_boundaries(tokenized_table)
if truncation != TapasTruncationStrategy.DO_NOT_TRUNCATE and max_length:
if truncation != TapasTruncationStrategy.DO_NOT_TRUNCATE:
num_rows, num_tokens = self._get_truncated_table_rows(query_tokens, tokenized_table, num_rows, num_columns,
max_length, truncation_strategy=truncation)
table_data = list(self._get_table_values(tokenized_table, num_columns, num_rows, num_tokens))
@@ -1306,6 +1306,9 @@ class TapasTokenizer(PreTrainedTokenizer):
if not isinstance(truncation_strategy, TapasTruncationStrategy):
truncation_strategy = TapasTruncationStrategy(truncation_strategy)
if max_length is None:
max_length = self.model_max_length
if truncation_strategy == TapasTruncationStrategy.DROP_ROWS_TO_FIT:
while True:
num_tokens = self._get_max_num_tokens(
+8
View File
@@ -1144,6 +1144,14 @@ class TapasTokenizationTest(TokenizerTesterMixin, unittest.TestCase):
# Ensure that the input IDs are less than the max length defined.
self.assertLessEqual(len(new_encoded_inputs), i)
tokenizer.model_max_length = 20
new_encoded_inputs = tokenizer.encode(table=table, query=queries[0], truncation=True)
dropped_encoded_inputs = tokenizer.encode(table=table, query=queries[0], truncation="drop_rows_to_fit")
# Ensure that the input IDs are still truncated when no max_length is specified
self.assertListEqual(new_encoded_inputs, dropped_encoded_inputs)
self.assertLessEqual(len(new_encoded_inputs), 20)
@is_pt_tf_cross_test
def test_batch_encode_plus_tensors(self):
tokenizers = self.get_tokenizers(do_lower_case=False)