Fix truncation when no max length is specified
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user