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

129 lines
4.5 KiB
Python

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()