"""Shared contracts for the ``purpose-deep`` multi-task classifier. The deep tier deliberately has a separate contract from purpose-lite: a pinned ModernBERT backbone, a fixed 512-token head/tail input, and four jointly-trained outputs. This module stays NumPy-only so tokenization, metrics, and selection can be tested without loading either MLX or a 150M-parameter checkpoint. """ from __future__ import annotations import math from dataclasses import dataclass from typing import Any, Sequence import numpy as np from purpose_data import LABELS, DataError, normalize_prompt, validate_source_record from train import classification_metrics MAX_LENGTH = 512 HEAD_TAIL_SPECIAL_TOKENS = 3 HEAD_TOKENS = (MAX_LENGTH - HEAD_TAIL_SPECIAL_TOKENS + 1) // 2 TAIL_TOKENS = MAX_LENGTH - HEAD_TAIL_SPECIAL_TOKENS - HEAD_TOKENS SCORABLE_HARD_SLICES = frozenset({"boundary", "mixed", "pasted-context"}) @dataclass(frozen=True) class DeepVariant: name: str model_id: str revision: str hidden_size: int intermediate_size: int layers: int attention_heads: int parameter_class: str # Immutable upstream revisions. Refreshing either is an explicit experiment, never # an accidental consequence of a mutable Hub ``main`` branch moving. DEEP_VARIANTS = { "base": DeepVariant( name="base", model_id="answerdotai/ModernBERT-base", revision="8949b909ec900327062f0ebf497f51aef5e6f0c8", hidden_size=768, intermediate_size=1152, layers=22, attention_heads=12, parameter_class="149M", ), "large": DeepVariant( name="large", model_id="answerdotai/ModernBERT-large", revision="45bb4654a4d5aaff24dd11d4781fa46d39bf8c13", hidden_size=1024, intermediate_size=2624, layers=28, attention_heads=16, parameter_class="395M", ), } @dataclass(frozen=True) class DeepTargets: primary: np.ndarray secondary: np.ndarray secondary_mask: np.ndarray mixed: np.ndarray difficulty: np.ndarray def validate_variant_config(config: dict[str, Any], variant: DeepVariant) -> None: """Fail closed if an upstream checkpoint no longer matches the pinned rung.""" expected = { "model_type": "modernbert", "hidden_size": variant.hidden_size, "intermediate_size": variant.intermediate_size, "num_hidden_layers": variant.layers, "num_attention_heads": variant.attention_heads, "vocab_size": 50368, } mismatches = [ f"{key}={config.get(key)!r} (expected {value!r})" for key, value in expected.items() if config.get(key) != value ] if mismatches: raise DataError( f"purpose-deep-{variant.name} checkpoint contract changed: " + "; ".join(mismatches) ) if int(config.get("max_position_embeddings", 0)) < MAX_LENGTH: raise DataError("purpose-deep backbone cannot represent the 512-token contract") def validate_deep_records( records: Sequence[dict[str, Any]], location: str, *, training: bool, ) -> None: if not records: raise DataError(f"{location}: split is empty") for index, record in enumerate(records, 1): validate_source_record(record, f"{location}:{index}") if training and record["slice"] == "vague-eval": raise DataError(f"{location}:{index}: vague-eval must never enter training") def encode_fixed_shape_numpy( tokenizer: Any, texts: Sequence[str], ) -> dict[str, np.ndarray]: """Encode the fixed 512-token ModernBERT head/tail input contract. ModernBERT has no token-type input. Long prompts use BERT pair framing so both the leading context and the often-tail-buried request survive: ``[CLS] + 255 head + [SEP] + 254 tail + [SEP]``. """ normalized = [normalize_prompt(text) for text in texts] raw = tokenizer( normalized, add_special_tokens=False, padding=False, truncation=False, return_attention_mask=False, return_token_type_ids=False, verbose=False, ) contents = raw.get("input_ids") if not isinstance(contents, list): raise DataError("ModernBERT tokenizer did not return input_ids") if tokenizer.pad_token_id is None: raise DataError("ModernBERT tokenizer must define a padding token") if tokenizer.cls_token_id is None or tokenizer.sep_token_id is None: raise DataError("ModernBERT tokenizer must define CLS and SEP tokens") if tokenizer.padding_side != "right": raise DataError("purpose-deep tokenizer must use right padding") if "token_type_ids" in tokenizer.model_input_names: raise DataError("purpose-deep ModernBERT must not expose token_type_ids") if tokenizer.num_special_tokens_to_add(pair=False) != 2: raise DataError("ModernBERT single-input special-token layout changed") if tokenizer.num_special_tokens_to_add(pair=True) != 3: raise DataError("ModernBERT pair special-token layout changed") input_rows: list[list[int]] = [] mask_rows: list[list[int]] = [] for content in contents: if len(content) <= MAX_LENGTH - 2: input_ids = [tokenizer.cls_token_id] + content + [tokenizer.sep_token_id] else: input_ids = ( [tokenizer.cls_token_id] + content[:HEAD_TOKENS] + [tokenizer.sep_token_id] + content[-TAIL_TOKENS:] + [tokenizer.sep_token_id] ) if len(input_ids) > MAX_LENGTH: raise DataError("fixed-shape tokenizer exceeded its 512-token contract") padding = MAX_LENGTH - len(input_ids) input_rows.append(input_ids + [tokenizer.pad_token_id] * padding) mask_rows.append([1] * len(input_ids) + [0] * padding) return { "input_ids": np.asarray(input_rows, dtype=np.int32), "attention_mask": np.asarray(mask_rows, dtype=np.int32), } def encode_targets(records: Sequence[dict[str, Any]]) -> DeepTargets: label_to_id = {label: index for index, label in enumerate(LABELS)} secondary_mask = np.asarray( [record["secondary"] is not None for record in records], dtype=np.bool_ ) # MLX cross entropy gathers every index before the mask is applied, so non-mixed # rows use a safe placeholder class rather than an ignore index such as -100. secondary = np.asarray( [ label_to_id[record["secondary"]] if record["secondary"] is not None else 0 for record in records ], dtype=np.int32, ) return DeepTargets( primary=np.asarray( [label_to_id[record["purpose"]] for record in records], dtype=np.int32, ), secondary=secondary, secondary_mask=secondary_mask, mixed=np.asarray([record["mixed"] for record in records], dtype=np.float32), difficulty=np.asarray( [record["difficulty"] for record in records], dtype=np.float32 ), ) def _binary_metrics(actual: np.ndarray, predicted: np.ndarray) -> dict[str, float | int]: true_positive = int(np.sum((actual == 1) & (predicted == 1))) false_positive = int(np.sum((actual == 0) & (predicted == 1))) false_negative = int(np.sum((actual == 1) & (predicted == 0))) true_negative = int(np.sum((actual == 0) & (predicted == 0))) precision = true_positive / max(true_positive + false_positive, 1) recall = true_positive / max(true_positive + false_negative, 1) specificity = true_negative / max(true_negative + false_positive, 1) return { "records": int(len(actual)), "accuracy": float(np.mean(actual == predicted)), "precision": precision, "recall": recall, "f1": 2 * precision * recall / max(precision + recall, 1e-12), "balancedAccuracy": (recall + specificity) / 2, "truePositive": true_positive, "falsePositive": false_positive, "falseNegative": false_negative, "trueNegative": true_negative, } def multitask_metrics( outputs: dict[str, np.ndarray], records: Sequence[dict[str, Any]], *, mixed_threshold: float = 0.5, ) -> dict[str, Any]: """Score every deep head without letting auxiliary heads hide primary quality.""" targets = encode_targets(records) size = len(records) required = { "purpose_logits": (size, len(LABELS)), "secondary_logits": (size, len(LABELS)), "mixed_logits": (size,), "difficulty": (size,), } for key, shape in required.items(): if key not in outputs or outputs[key].shape != shape: raise ValueError(f"{key} must have shape {shape}") scorable = np.asarray( [record["slice"] != "vague-eval" for record in records], dtype=np.bool_ ) if not np.any(scorable): raise ValueError("deep metrics require label-scorable records") hard = np.asarray( [record["slice"] in SCORABLE_HARD_SLICES for record in records], dtype=np.bool_, ) primary_predictions = outputs["purpose_logits"].argmax(axis=-1) primary = classification_metrics( targets.primary[scorable].tolist(), primary_predictions[scorable].tolist() ) hard_mask = scorable & hard primary_hard = ( classification_metrics( targets.primary[hard_mask].tolist(), primary_predictions[hard_mask].tolist(), ) if np.any(hard_mask) else primary ) secondary_predictions = outputs["secondary_logits"].argmax(axis=-1) if np.any(targets.secondary_mask): secondary = classification_metrics( targets.secondary[targets.secondary_mask].tolist(), secondary_predictions[targets.secondary_mask].tolist(), ) supported_labels = [ label for label, row in zip( secondary["confusionMatrix"]["labels"], secondary["confusionMatrix"]["rows"], ) if sum(row) ] secondary["supportedLabels"] = supported_labels secondary["supportedMacroRecall"] = float( np.mean( [secondary["perPurposeRecall"][label] for label in supported_labels] ) ) else: secondary = None mixed_probabilities = 1.0 / (1.0 + np.exp(-outputs["mixed_logits"])) mixed_predictions = (mixed_probabilities >= mixed_threshold).astype(np.int32) mixed = _binary_metrics(targets.mixed.astype(np.int32), mixed_predictions) difficulty_error = outputs["difficulty"] - targets.difficulty difficulty = { "records": size, "mae": float(np.mean(np.abs(difficulty_error))), "rmse": float(math.sqrt(float(np.mean(np.square(difficulty_error))))), } # Checkpoint selection is exclusively primary-task quality: half overall and half # hard-slice accuracy. Auxiliary-head health remains explicit in the report/gates. selection_score = (primary["accuracy"] + primary_hard["accuracy"]) / 2 return { "primary": primary, "primaryHardSlice": primary_hard, "secondary": secondary, "mixed": mixed, "difficulty": difficulty, "selectionScore": selection_score, } def best_mixed_threshold(logits: np.ndarray, actual: np.ndarray) -> float: """Choose the validation F1 threshold, preferring the conservative higher tie.""" probabilities = 1.0 / (1.0 + np.exp(-logits)) candidates = sorted({0.5, *probabilities.tolist()}, reverse=True) best = (float("-inf"), 0.5) for threshold in candidates: metrics = _binary_metrics( actual.astype(np.int32), (probabilities >= threshold).astype(np.int32), ) candidate = (float(metrics["f1"]), float(threshold)) if candidate > best: best = candidate return best[1]