Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -0,0 +1,128 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user