318 lines
11 KiB
Python
318 lines
11 KiB
Python
"""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]
|