73 lines
2.0 KiB
Python
73 lines
2.0 KiB
Python
import json
|
|
import sys
|
|
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
import numpy as np
|
|
import torch
|
|
|
|
|
|
MODULE_DIR = Path(__file__).resolve().parents[1]
|
|
sys.path.insert(0, str(MODULE_DIR))
|
|
|
|
import train
|
|
import train_mlx
|
|
|
|
|
|
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_mlx_encoding_matches_pytorch_encoding(self):
|
|
texts = ["tokens 3", "tokens 200"]
|
|
pytorch = train.encode_fixed_shape(self.Tokenizer(), texts, torch)
|
|
mlx = train_mlx.encode_fixed_shape_numpy(self.Tokenizer(), texts)
|
|
self.assertEqual(set(pytorch), set(mlx))
|
|
for key in pytorch:
|
|
with self.subTest(key=key):
|
|
np.testing.assert_array_equal(pytorch[key].numpy(), mlx[key])
|
|
|
|
|
|
class CheckpointConfigTests(unittest.TestCase):
|
|
def test_rejects_changed_label_order(self):
|
|
config = {
|
|
"model_type": "bert",
|
|
"hidden_size": 384,
|
|
"num_hidden_layers": 6,
|
|
"id2label": {
|
|
str(index): label
|
|
for index, label in enumerate(reversed(train.LABELS))
|
|
},
|
|
}
|
|
with tempfile.TemporaryDirectory() as temp:
|
|
model_dir = Path(temp)
|
|
(model_dir / "config.json").write_text(
|
|
json.dumps(config),
|
|
encoding="utf-8",
|
|
)
|
|
with self.assertRaisesRegex(
|
|
train.DataError,
|
|
"label order",
|
|
):
|
|
train_mlx._checkpoint_config(model_dir)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|