import tempfile import sys import unittest from pathlib import Path import mlx.core as mx import mlx.nn as nn MODULE_DIR = Path(__file__).resolve().parents[1] sys.path.insert(0, str(MODULE_DIR)) from deep_model_mlx import ( ModernBertForPurposeClassification, ModernBertPurposeConfig, load_pretrained_weights, save_weights, ) from purpose_data import DataError, LABELS 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_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_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()