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()