Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -0,0 +1,72 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user