Files

155 lines
5.1 KiB
Python
Raw Permalink Normal View History

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(["writing"], metrics["secondary"]["supportedLabels"])
self.assertEqual(1, metrics["secondary"]["supportedMacroRecall"])
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()