2026-07-30 20:14:37 -07:00
|
|
|
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
|
|
|
|
|
|
|
|
|
|
|
2026-07-30 23:50:55 -07:00
|
|
|
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",
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
2026-07-30 20:14:37 -07:00
|
|
|
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()
|