129 lines
4.5 KiB
Python
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()
|