Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -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()
|
||||
@@ -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