Files
nucleic-purpose-classifier/deep_contract.py
T

318 lines
11 KiB
Python
Raw Normal View History

"""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(),
)
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]