Merge nucleic/sleek-ember-seal-uady into dev

This commit is contained in:
2026-07-31 01:24:01 -07:00
parent 0bee416a88
commit 9a1228efbb
7 changed files with 1876 additions and 0 deletions
+152
View File
@@ -0,0 +1,152 @@
import sys
import unittest
from pathlib import Path
import numpy as np
MODULE_DIR = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(MODULE_DIR))
from deep_contract import (
DEEP_VARIANTS,
HEAD_TOKENS,
MAX_LENGTH,
TAIL_TOKENS,
best_mixed_threshold,
encode_fixed_shape_numpy,
encode_targets,
multitask_metrics,
validate_variant_config,
)
from purpose_data import DataError, LABELS
def record(
purpose,
*,
secondary=None,
slice="core",
difficulty=0.5,
):
return {
"prompt": f"a {purpose} prompt",
"purpose": purpose,
"secondary": secondary,
"mixed": secondary is not None,
"difficulty": difficulty,
"slice": "mixed" if secondary is not None else slice,
"lang": "en",
}
class FixedShapeTests(unittest.TestCase):
class Tokenizer:
pad_token_id = 0
cls_token_id = 1
sep_token_id = 2
padding_side = "right"
model_input_names = ["input_ids", "attention_mask"]
def __call__(self, texts, **_):
return {
"input_ids": [
list(range(10, 10 + int(text.split()[-1]))) for text in texts
]
}
@staticmethod
def num_special_tokens_to_add(pair=False):
return 3 if pair else 2
def test_short_and_long_inputs_are_fixed_and_preserve_both_ends(self):
encoded = encode_fixed_shape_numpy(
self.Tokenizer(), ["tokens 3", "tokens 700"]
)
self.assertEqual((2, MAX_LENGTH), encoded["input_ids"].shape)
self.assertEqual([1, 10, 11, 12, 2], encoded["input_ids"][0, :5].tolist())
self.assertEqual(5, int(encoded["attention_mask"][0].sum()))
long = encoded["input_ids"][1]
self.assertEqual(1, long[0])
self.assertEqual(2, long[HEAD_TOKENS + 1])
self.assertEqual(10 + 700 - TAIL_TOKENS, long[HEAD_TOKENS + 2])
self.assertEqual(2, long[-1])
self.assertEqual(HEAD_TOKENS + TAIL_TOKENS + 3, len(long))
def test_token_type_ids_fail_closed(self):
tokenizer = self.Tokenizer()
tokenizer.model_input_names = [
"input_ids",
"attention_mask",
"token_type_ids",
]
with self.assertRaisesRegex(DataError, "token_type_ids"):
encode_fixed_shape_numpy(tokenizer, ["tokens 3"])
class TargetAndMetricTests(unittest.TestCase):
def test_non_mixed_secondary_uses_safe_index_plus_mask(self):
targets = encode_targets(
[
record("planning"),
record("review", secondary="writing"),
]
)
self.assertEqual([0, LABELS.index("writing")], targets.secondary.tolist())
self.assertEqual([False, True], targets.secondary_mask.tolist())
def test_selection_is_half_overall_half_hard_primary_accuracy(self):
records = [
record("planning"),
record("backendImpl", slice="boundary"),
record("review", secondary="writing"),
]
purpose = np.full((3, len(LABELS)), -4.0, dtype=np.float32)
# Core is right; both hard records are wrong.
purpose[0, LABELS.index("planning")] = 4
purpose[1, LABELS.index("planning")] = 4
purpose[2, LABELS.index("planning")] = 4
secondary = np.zeros_like(purpose)
secondary[2, LABELS.index("writing")] = 4
metrics = multitask_metrics(
{
"purpose_logits": purpose,
"secondary_logits": secondary,
"mixed_logits": np.asarray([-4.0, -4.0, 4.0]),
"difficulty": np.asarray([0.5, 0.5, 0.5]),
},
records,
)
self.assertAlmostEqual(1 / 3, metrics["primary"]["accuracy"])
self.assertEqual(0, metrics["primaryHardSlice"]["accuracy"])
self.assertAlmostEqual(1 / 6, metrics["selectionScore"])
self.assertEqual(1, metrics["secondary"]["accuracy"])
self.assertEqual(1, metrics["mixed"]["f1"])
def test_mixed_threshold_prefers_higher_threshold_on_f1_tie(self):
logits = np.asarray([-4.0, 0.2, 2.0], dtype=np.float32)
actual = np.asarray([0.0, 0.0, 1.0], dtype=np.float32)
threshold = best_mixed_threshold(logits, actual)
self.assertGreater(threshold, 0.5)
class VariantTests(unittest.TestCase):
def test_pinned_base_contract(self):
variant = DEEP_VARIANTS["base"]
config = {
"model_type": "modernbert",
"hidden_size": 768,
"intermediate_size": 1152,
"num_hidden_layers": 22,
"num_attention_heads": 12,
"vocab_size": 50368,
"max_position_embeddings": 8192,
}
validate_variant_config(config, variant)
config["num_hidden_layers"] = 23
with self.assertRaisesRegex(DataError, "contract changed"):
validate_variant_config(config, variant)
if __name__ == "__main__":
unittest.main()
+128
View File
@@ -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()