Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4c2263b94e |
@@ -1852,8 +1852,8 @@ class PreTrainedTokenizerFast(PreTrainedTokenizer):
|
||||
stack = tf.stack(stack, axis=0)
|
||||
elif return_tensors == "pt":
|
||||
stack = torch.stack(stack, dim=0)
|
||||
elif not return_tensors and len(stack) == 1:
|
||||
stack = stack[0]
|
||||
# elif not return_tensors and len(stack) == 1:
|
||||
# stack = stack[0]
|
||||
|
||||
sanitized[key] = stack
|
||||
|
||||
@@ -1902,7 +1902,10 @@ class PreTrainedTokenizerFast(PreTrainedTokenizer):
|
||||
|
||||
# Return tensor is None, then we can remove the leading batch axis
|
||||
if not return_tensors:
|
||||
return {key: value[0] if isinstance(value[0], list) else value for key, value in batched_output.items()}
|
||||
return {
|
||||
key: value[0] if len(value) > 0 and isinstance(value[0], list) else value
|
||||
for key, value in batched_output.items()
|
||||
}
|
||||
else:
|
||||
return batched_output
|
||||
|
||||
|
||||
@@ -272,6 +272,30 @@ class FastTokenizerMatchingTest(unittest.TestCase):
|
||||
# self.assertEqual(getattr(tokenizer_rp, key), getattr(tokenizer_pp, key))
|
||||
# self.assertEqual(getattr(tokenizer_rp, key + "_id"), getattr(tokenizer_pp, key + "_id"))
|
||||
|
||||
def assert_empty_output_no_special_tokens(self, ru_class, py_class, model):
|
||||
tokenizer_r = ru_class.from_pretrained(model, add_special_tokens=False)
|
||||
tokenizer_p = py_class.from_pretrained(model)
|
||||
|
||||
# add_special_tokens=False makes nothing for now.
|
||||
self.assertEqual(
|
||||
tokenizer_p.tokenize(" ", add_special_tokens=False), tokenizer_r.tokenize(" ", add_special_tokens=False)
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
tokenizer_p.encode_plus(" ", add_special_tokens=False),
|
||||
tokenizer_r.encode_plus(" ", add_special_tokens=False),
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
tokenizer_p.encode_plus(" ", add_special_tokens=False),
|
||||
tokenizer_r.encode_plus(" ", add_special_tokens=False),
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
tokenizer_p.batch_encode_plus([" "], add_special_tokens=False),
|
||||
tokenizer_r.batch_encode_plus([" "], add_special_tokens=False),
|
||||
)
|
||||
|
||||
def test_bert(self):
|
||||
for tokenizer_name in BertTokenizer.pretrained_vocab_files_map["vocab_file"].keys():
|
||||
tokenizer_p = BertTokenizer.from_pretrained(tokenizer_name)
|
||||
@@ -313,6 +337,9 @@ class FastTokenizerMatchingTest(unittest.TestCase):
|
||||
# Check for padding
|
||||
self.assert_padding(tokenizer_r, tokenizer_p)
|
||||
|
||||
# Check for space-only input
|
||||
self.assert_empty_output_no_special_tokens(tokenizer_r.__class__, tokenizer_p.__class__, tokenizer_name)
|
||||
|
||||
@require_torch
|
||||
def test_transfoxl(self):
|
||||
for tokenizer_name in TransfoXLTokenizer.pretrained_vocab_files_map["pretrained_vocab_file"].keys():
|
||||
@@ -369,6 +396,9 @@ class FastTokenizerMatchingTest(unittest.TestCase):
|
||||
# self.assertIsNotNone(tokenizer_p.__class__.from_pretrained('./'))
|
||||
self.assertIsNotNone(tokenizer_r.__class__.from_pretrained("./"))
|
||||
|
||||
# Check for space-only input
|
||||
self.assert_empty_output_no_special_tokens(tokenizer_r.__class__, tokenizer_p.__class__, tokenizer_name)
|
||||
|
||||
def test_distilbert(self):
|
||||
for tokenizer_name in DistilBertTokenizer.pretrained_vocab_files_map["vocab_file"].keys():
|
||||
tokenizer_p = DistilBertTokenizer.from_pretrained(tokenizer_name)
|
||||
@@ -411,6 +441,9 @@ class FastTokenizerMatchingTest(unittest.TestCase):
|
||||
# Check for padding
|
||||
self.assert_padding(tokenizer_r, tokenizer_p)
|
||||
|
||||
# Check for space-only input
|
||||
self.assert_empty_output_no_special_tokens(tokenizer_r.__class__, tokenizer_p.__class__, tokenizer_name)
|
||||
|
||||
def test_gpt2(self):
|
||||
for tokenizer_name in GPT2Tokenizer.pretrained_vocab_files_map["vocab_file"].keys():
|
||||
tokenizer_p = GPT2Tokenizer.from_pretrained(tokenizer_name)
|
||||
@@ -452,6 +485,9 @@ class FastTokenizerMatchingTest(unittest.TestCase):
|
||||
# Check for padding
|
||||
self.assertRaises(ValueError, self.assert_padding, tokenizer_r, tokenizer_p)
|
||||
|
||||
# Check for space-only input
|
||||
self.assert_empty_output_no_special_tokens(tokenizer_r.__class__, tokenizer_p.__class__, tokenizer_name)
|
||||
|
||||
def test_roberta(self):
|
||||
for tokenizer_name in RobertaTokenizer.pretrained_vocab_files_map["vocab_file"].keys():
|
||||
tokenizer_p = RobertaTokenizer.from_pretrained(tokenizer_name)
|
||||
@@ -494,6 +530,9 @@ class FastTokenizerMatchingTest(unittest.TestCase):
|
||||
# TODO: Re-enable this test as soon as Roberta align with the python tokenizer.
|
||||
# self.assert_padding(tokenizer_r, tokenizer_p)
|
||||
|
||||
# Check for space-only input
|
||||
self.assert_empty_output_no_special_tokens(tokenizer_r.__class__, tokenizer_p.__class__, tokenizer_name)
|
||||
|
||||
def test_openai(self):
|
||||
for tokenizer_name in OpenAIGPTTokenizer.pretrained_vocab_files_map["vocab_file"].keys():
|
||||
tokenizer_p = OpenAIGPTTokenizer.from_pretrained(tokenizer_name)
|
||||
@@ -536,3 +575,6 @@ class FastTokenizerMatchingTest(unittest.TestCase):
|
||||
|
||||
# Check the number of returned files for save_vocabulary
|
||||
self.assert_save_pretrained(tokenizer_r, tokenizer_p)
|
||||
|
||||
# Check for space-only input
|
||||
self.assert_empty_output_no_special_tokens(tokenizer_r.__class__, tokenizer_p.__class__, tokenizer_name)
|
||||
|
||||
Reference in New Issue
Block a user