155 lines
5.1 KiB
Python
155 lines
5.1 KiB
Python
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()
|