Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -2,6 +2,8 @@ import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
MODULE_DIR = Path(__file__).resolve().parents[1]
|
||||
sys.path.insert(0, str(MODULE_DIR))
|
||||
@@ -42,5 +44,45 @@ class MetricsTests(unittest.TestCase):
|
||||
)
|
||||
|
||||
|
||||
class FixedShapeTokenizerTests(unittest.TestCase):
|
||||
class Tokenizer:
|
||||
pad_token_id = 0
|
||||
cls_token_id = 1
|
||||
sep_token_id = 2
|
||||
padding_side = "right"
|
||||
model_input_names = ["input_ids", "attention_mask", "token_type_ids"]
|
||||
|
||||
def __call__(self, texts, **_):
|
||||
return {
|
||||
"input_ids": [
|
||||
list(range(10, 10 + int(text.split()[-1]))) for text in texts
|
||||
]
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def num_special_tokens_to_add(pair=False):
|
||||
return 3 if pair else 2
|
||||
|
||||
def test_long_input_keeps_head_and_tail_in_fixed_pair_shape(self):
|
||||
encoded = train.encode_fixed_shape(
|
||||
self.Tokenizer(), ["tokens 200"], torch
|
||||
)
|
||||
self.assertEqual((1, train.MAX_LENGTH), tuple(encoded["input_ids"].shape))
|
||||
row = encoded["input_ids"][0].tolist()
|
||||
self.assertEqual(list(range(10, 10 + train.HEAD_TOKENS)), row[1:64])
|
||||
self.assertEqual(2, row[64])
|
||||
self.assertEqual(
|
||||
list(range(10 + 200 - train.TAIL_TOKENS, 10 + 200)),
|
||||
row[65:127],
|
||||
)
|
||||
self.assertEqual(2, row[127])
|
||||
self.assertEqual(1, encoded["token_type_ids"][0, 65].item())
|
||||
|
||||
def test_short_input_is_padded_as_one_sequence(self):
|
||||
encoded = train.encode_fixed_shape(self.Tokenizer(), ["tokens 3"], torch)
|
||||
self.assertEqual([1, 10, 11, 12, 2], encoded["input_ids"][0, :5].tolist())
|
||||
self.assertEqual(0, encoded["attention_mask"][0, 5].item())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user