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

236 lines
8.9 KiB
Python

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