Files
nucleic-purpose-classifier/tests/test_train_mlx.py
T

105 lines
3.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 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()