import tempfile import sys import unittest from pathlib import Path import mlx.core as mx import mlx.nn as nn import mlx.optimizers as optim from mlx.utils import tree_flatten MODULE_DIR = Path(__file__).resolve().parents[1] sys.path.insert(0, str(MODULE_DIR)) from deep_model_mlx import ( ModernBertForPurposeClassification, ModernBertPurposeConfig, load_checkpoint_weights, load_pretrained_weights, save_weights, ) from purpose_data import DataError, LABELS from train_deep_mlx import ( _distillation_loss, _load_optimizer_state, _save_training_resume, ) def tiny_config(*, checkpointing=False): return ModernBertPurposeConfig( vocab_size=64, hidden_size=16, intermediate_size=24, num_hidden_layers=3, num_attention_heads=4, max_position_embeddings=32, pad_token_id=0, norm_eps=1e-5, norm_bias=False, attention_bias=False, attention_dropout=0.0, layer_types=("full_attention", "sliding_attention", "sliding_attention"), local_attention=4, embedding_dropout=0.0, mlp_bias=False, mlp_dropout=0.0, classifier_bias=False, classifier_dropout=0.0, full_rope_theta=160_000.0, local_rope_theta=10_000.0, gradient_checkpointing=checkpointing, ) class DeepModelTests(unittest.TestCase): def setUp(self): mx.random.seed(7) def test_all_four_heads_have_the_expected_shapes_and_ranges(self): model = ModernBertForPurposeClassification(tiny_config()) output = model( mx.array([[1, 3, 4, 2, 0, 0], [1, 5, 6, 7, 8, 2]]), mx.array([[1, 1, 1, 1, 0, 0], [1, 1, 1, 1, 1, 1]]), ) mx.eval(*output.values()) self.assertEqual((2, len(LABELS)), output["purpose_logits"].shape) self.assertEqual((2, len(LABELS)), output["secondary_logits"].shape) self.assertEqual((2,), output["mixed_logits"].shape) self.assertEqual((2,), output["difficulty"].shape) self.assertTrue(bool(mx.all(output["difficulty"] >= 0).item())) self.assertTrue(bool(mx.all(output["difficulty"] <= 1).item())) def test_masked_padding_tokens_do_not_change_cls_outputs(self): model = ModernBertForPurposeClassification(tiny_config()) model.eval() mask = mx.array([[1, 1, 1, 1, 0, 0]]) first = model(mx.array([[1, 3, 4, 2, 0, 0]]), mask) second = model(mx.array([[1, 3, 4, 2, 9, 10]]), mask) mx.eval(*first.values(), *second.values()) for key in first: with self.subTest(head=key): self.assertLess(float(mx.max(mx.abs(first[key] - second[key])).item()), 1e-5) def test_gradient_checkpointed_multitask_smoke(self): model = ModernBertForPurposeClassification( tiny_config(checkpointing=True) ) model.train() def loss(ids, mask): output = model(ids, mask) return ( mx.mean(output["purpose_logits"] ** 2) + mx.mean(output["secondary_logits"] ** 2) + mx.mean(output["mixed_logits"] ** 2) + mx.mean(output["difficulty"] ** 2) ) value_and_grad = nn.value_and_grad(model, loss) value, gradients = value_and_grad( mx.array([[1, 3, 4, 2]]), mx.ones((1, 4), dtype=mx.int32) ) mx.eval(value, gradients) self.assertTrue(float(value.item()) > 0) def test_distillation_loss_matches_teacher_and_backpropagates(self): teacher = mx.array([[2.0, 0.0, -1.0]]) def loss(student): return _distillation_loss( mx, student, teacher, temperature=2.0, weights=mx.ones((1,)), ) value, gradient = mx.value_and_grad(loss)(teacher) mx.eval(value, gradient) self.assertAlmostEqual(0.0, float(value.item()), places=6) self.assertEqual(teacher.shape, gradient.shape) def test_training_resume_round_trips_model_and_optimizer(self): class Tokenizer: @staticmethod def save_pretrained(destination): (destination / "tokenizer_config.json").write_text("{}\n") model = ModernBertForPurposeClassification(tiny_config()) optimizer = optim.AdamW(learning_rate=1e-3) optimizer.init(model.trainable_parameters()) def loss(ids, mask): return mx.mean(model(ids, mask)["purpose_logits"] ** 2) value_and_grad = nn.value_and_grad(model, loss) value, gradients = value_and_grad( mx.array([[1, 3, 4, 2]]), mx.ones((1, 4), dtype=mx.int32) ) optimizer.update(model, gradients) mx.eval(value, model.parameters(), optimizer.state) with tempfile.TemporaryDirectory() as temp: checkpoint = _save_training_resume( mx, model, optimizer, Tokenizer(), Path(temp), {}, {"schemaVersion": 1, "status": "paused"}, ) restored_model = ModernBertForPurposeClassification(tiny_config()) restored_model.load_weights( str(checkpoint / "model.safetensors"), strict=True ) restored_optimizer = optim.AdamW(learning_rate=1e-3) restored_optimizer.init(restored_model.trainable_parameters()) _load_optimizer_state(mx, restored_optimizer, checkpoint) original = tree_flatten(optimizer.state, destination={}) restored = tree_flatten(restored_optimizer.state, destination={}) self.assertEqual(set(original), set(restored)) for key in original: with self.subTest(optimizer_tensor=key): self.assertEqual( 0, float(mx.max(mx.abs(original[key] - restored[key])).item()), ) def test_checkpoint_round_trip(self): model = ModernBertForPurposeClassification(tiny_config()) with tempfile.TemporaryDirectory() as temp: path = Path(temp) / "model.safetensors" save_weights(model, path) restored = ModernBertForPurposeClassification(tiny_config()) restored.load_weights(str(path), strict=True) ids = mx.array([[1, 3, 4, 2]]) mask = mx.ones((1, 4), dtype=mx.int32) first = model(ids, mask) second = restored(ids, mask) mx.eval(*first.values(), *second.values()) for key in first: with self.subTest(head=key): self.assertEqual( 0, float(mx.max(mx.abs(first[key] - second[key])).item()) ) def test_trained_checkpoint_loader_restores_every_task_head(self): model = ModernBertForPurposeClassification(tiny_config()) with tempfile.TemporaryDirectory() as temp: path = Path(temp) / "model.safetensors" save_weights(model, path) restored = ModernBertForPurposeClassification(tiny_config()) report = load_checkpoint_weights(restored, path) self.assertEqual(0, report["freshTaskHeads"]) self.assertEqual(0, report["ignored"]) ids = mx.array([[1, 3, 4, 2]]) mask = mx.ones((1, 4), dtype=mx.int32) first = model(ids, mask) second = restored(ids, mask) mx.eval(*first.values(), *second.values()) for key in first: with self.subTest(head=key): self.assertEqual( 0, float(mx.max(mx.abs(first[key] - second[key])).item()) ) def test_trained_checkpoint_loader_rejects_missing_task_heads(self): model = ModernBertForPurposeClassification(tiny_config()) with tempfile.TemporaryDirectory() as temp: complete = Path(temp) / "complete.safetensors" partial = Path(temp) / "partial.safetensors" save_weights(model, complete) weights = mx.load(str(complete)) weights.pop("purpose_classifier.weight") mx.save_safetensors(str(partial), weights) with self.assertRaisesRegex(DataError, "missing 1 tensors"): load_checkpoint_weights( ModernBertForPurposeClassification(tiny_config()), partial ) def test_pretrained_loader_rejects_partial_backbone(self): with tempfile.TemporaryDirectory() as temp: path = Path(temp) / "partial.safetensors" mx.save_safetensors(str(path), {"model.final_norm.weight": mx.ones((16,))}) with self.assertRaisesRegex(DataError, "missing"): load_pretrained_weights( ModernBertForPurposeClassification(tiny_config()), path ) if __name__ == "__main__": unittest.main()