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()
|
||||
Reference in New Issue
Block a user