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 MLXDeviceTests(unittest.TestCase): class Metal: def __init__(self, available): self.available = available def is_available(self): return self.available class MLX: cpu = "cpu" gpu = "gpu" def __init__(self, metal_available): self.metal = MLXDeviceTests.Metal(metal_available) self.selected = None def set_default_device(self, device): self.selected = device def test_cpu_is_an_explicit_fallback(self): mlx = self.MLX(metal_available=False) train_mlx._configure_mlx_device(mlx, "cpu") self.assertEqual("cpu", mlx.selected) def test_metal_fails_closed_when_unavailable(self): with self.assertRaisesRegex(train.DataError, "requires Apple Silicon"): train_mlx._configure_mlx_device( self.MLX(metal_available=False), "metal", ) 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()