Merge nucleic/sleek-ember-seal-uady into dev

This commit is contained in:
2026-07-31 01:24:01 -07:00
parent 0bee416a88
commit 9a1228efbb
7 changed files with 1876 additions and 0 deletions
+47
View File
@@ -283,6 +283,53 @@ the accepted baseline's float checkpoint. That gain does not survive export: the
baseline and fails the 95% shipping gate. Preserve the artifact as a rejected experiment; baseline and fails the 95% shipping gate. Preserve the artifact as a rejected experiment;
`purpose-lite-v1-distilled-qat-mlx-4e` remains the candidate of record. `purpose-lite-v1-distilled-qat-mlx-4e` remains the candidate of record.
## Train purpose-deep
The next classifier tier is an MLX-native ModernBERT multi-task model. It keeps the
primary eight-way purpose output and jointly learns secondary purpose, a mixed-intent
flag, and advisory difficulty. Both upstream rungs are immutable: base is ModernBERT
149M at revision `8949b909ec900327062f0ebf497f51aef5e6f0c8`; large is ModernBERT
395M at revision `45bb4654a4d5aaff24dd11d4781fa46d39bf8c13`.
Before the first run on a new MLX/Transformers version, compare the real pinned backbone
against Hugging Face. The check crosses ModernBERT's local-attention window and fails if
pooled-representation drift exceeds `5e-4`:
```bash
ml/purpose-classifier/venv/bin/python ml/purpose-classifier/verify_deep_mlx.py \
--variant base
```
Start with the base ablation rung and the validated first-prompt history augmentation.
The trainer downloads the pinned checkpoint on first use, fixes every input at 512 tokens
(`255` head + `254` tail + three special tokens for long prompts), and uses gradient
checkpointing by default. Checkpoint selection is half scored overall accuracy and half
scored hard-slice accuracy; auxiliary heads are reported independently and cannot hide a
primary-purpose regression.
```bash
ml/purpose-classifier/venv/bin/python -u \
ml/purpose-classifier/train_deep_mlx.py \
--variant base \
--dataset-dir \
ml/purpose-classifier/.artifacts/dataset-v1-history-first-prompts \
--epochs 3 --early-stopping-patience 1 \
--progress-steps 10 \
--output-dir ml/purpose-classifier/outputs/purpose-deep-v1-base-mlx \
--overwrite-output
```
The defaults use batch size 4 and learning rate `2e-5` for base (2 and `1e-5` for
large). If unified memory is tight, lower `--batch-size` before disabling gradient
checkpointing. `--device cpu` is diagnostic only: a real 512-token backward pass is
expected to be extremely slow there. Each improved epoch atomically rewrites `model/`
and updates `training-state.json`, so progress is visible and an interrupted run retains
the last selected checkpoint.
Do not launch the large rung yet. It is justified only after base is evaluated on the
frozen set; large must beat base by at least two hard-slice points, while deep itself must
reach 97% scored overall and beat the shipping lite artifact by five hard-slice points.
### Convert and validate Core ML ### Convert and validate Core ML
Core ML Tools no longer maintains the legacy ONNX converter, so the Apple artifact is Core ML Tools no longer maintains the legacy ONNX converter, so the Apple artifact is
+317
View File
@@ -0,0 +1,317 @@
"""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]
+360
View File
@@ -0,0 +1,360 @@
"""MLX implementation of the ModernBERT ``purpose-deep`` multi-task model.
Parameter names and tensor layouts match Hugging Face ModernBERT. The masked-LM
backbone and prediction head therefore load directly from the pinned safetensors
checkpoint; only the four small task heads start fresh.
"""
from __future__ import annotations
import math
from dataclasses import dataclass
from pathlib import Path
from typing import Any
import mlx.core as mx
import mlx.nn as nn
from mlx.utils import tree_flatten
from purpose_data import LABELS, DataError
@dataclass(frozen=True)
class ModernBertPurposeConfig:
vocab_size: int
hidden_size: int
intermediate_size: int
num_hidden_layers: int
num_attention_heads: int
max_position_embeddings: int
pad_token_id: int
norm_eps: float
norm_bias: bool
attention_bias: bool
attention_dropout: float
layer_types: tuple[str, ...]
local_attention: int
embedding_dropout: float
mlp_bias: bool
mlp_dropout: float
classifier_bias: bool
classifier_dropout: float
full_rope_theta: float
local_rope_theta: float
gradient_checkpointing: bool = True
@classmethod
def from_hugging_face(
cls,
config: dict[str, Any],
*,
gradient_checkpointing: bool = True,
) -> "ModernBertPurposeConfig":
layer_types = config.get("layer_types")
if layer_types is None:
every = int(config.get("global_attn_every_n_layers", 3))
layer_types = [
"sliding_attention" if index % every else "full_attention"
for index in range(int(config["num_hidden_layers"]))
]
rope = config.get("rope_parameters") or {}
return cls(
vocab_size=int(config["vocab_size"]),
hidden_size=int(config["hidden_size"]),
intermediate_size=int(config["intermediate_size"]),
num_hidden_layers=int(config["num_hidden_layers"]),
num_attention_heads=int(config["num_attention_heads"]),
max_position_embeddings=int(config["max_position_embeddings"]),
pad_token_id=int(config["pad_token_id"]),
norm_eps=float(config.get("norm_eps", 1e-5)),
norm_bias=bool(config.get("norm_bias", False)),
attention_bias=bool(config.get("attention_bias", False)),
attention_dropout=float(config.get("attention_dropout", 0.0)),
layer_types=tuple(layer_types),
local_attention=int(config.get("local_attention", 128)),
embedding_dropout=float(config.get("embedding_dropout", 0.0)),
mlp_bias=bool(config.get("mlp_bias", False)),
mlp_dropout=float(config.get("mlp_dropout", 0.0)),
classifier_bias=bool(config.get("classifier_bias", False)),
classifier_dropout=float(config.get("classifier_dropout", 0.0)),
full_rope_theta=float(
(rope.get("full_attention") or {}).get(
"rope_theta", config.get("global_rope_theta", 160_000.0)
)
),
local_rope_theta=float(
(rope.get("sliding_attention") or {}).get(
"rope_theta", config.get("local_rope_theta", 10_000.0)
)
),
gradient_checkpointing=gradient_checkpointing,
)
def _rotate_half(value: Any) -> Any:
half = value.shape[-1] // 2
return mx.concatenate((-value[..., half:], value[..., :half]), axis=-1)
class ModernBertEmbeddings(nn.Module):
def __init__(self, config: ModernBertPurposeConfig) -> None:
super().__init__()
self.tok_embeddings = nn.Embedding(config.vocab_size, config.hidden_size)
self.norm = nn.LayerNorm(
config.hidden_size, eps=config.norm_eps, bias=config.norm_bias
)
self.drop = nn.Dropout(config.embedding_dropout)
def __call__(self, input_ids: Any) -> Any:
return self.drop(self.norm(self.tok_embeddings(input_ids)))
class ModernBertMLP(nn.Module):
def __init__(self, config: ModernBertPurposeConfig) -> None:
super().__init__()
self.Wi = nn.Linear(
config.hidden_size, config.intermediate_size * 2, bias=config.mlp_bias
)
self.Wo = nn.Linear(
config.intermediate_size, config.hidden_size, bias=config.mlp_bias
)
self.drop = nn.Dropout(config.mlp_dropout)
self.act = nn.GELU(approx="none")
def __call__(self, hidden_states: Any) -> Any:
projected = self.Wi(hidden_states)
content, gate = mx.split(projected, 2, axis=-1)
return self.Wo(self.drop(self.act(content) * gate))
class ModernBertAttention(nn.Module):
def __init__(
self, config: ModernBertPurposeConfig, attention_type: str
) -> None:
super().__init__()
if config.hidden_size % config.num_attention_heads:
raise DataError("ModernBERT hidden size must divide evenly into heads")
self.num_heads = config.num_attention_heads
self.head_dim = config.hidden_size // config.num_attention_heads
self.attention_type = attention_type
self.Wqkv = nn.Linear(
config.hidden_size, config.hidden_size * 3, bias=config.attention_bias
)
self.Wo = nn.Linear(
config.hidden_size, config.hidden_size, bias=config.attention_bias
)
self.out_drop = nn.Dropout(config.attention_dropout)
def __call__(
self,
hidden_states: Any,
mask: Any | None,
cos: Any,
sin: Any,
) -> Any:
batch, length, hidden = hidden_states.shape
qkv = self.Wqkv(hidden_states).reshape(
batch, length, 3, self.num_heads, self.head_dim
)
query, key, value = (
qkv[:, :, index].transpose(0, 2, 1, 3) for index in range(3)
)
cos = cos[None, None, :, :].astype(query.dtype)
sin = sin[None, None, :, :].astype(query.dtype)
query = query * cos + _rotate_half(query) * sin
key = key * cos + _rotate_half(key) * sin
context = mx.fast.scaled_dot_product_attention(
query,
key,
value,
scale=self.head_dim**-0.5,
mask=mask,
)
context = context.transpose(0, 2, 1, 3).reshape(batch, length, hidden)
return self.out_drop(self.Wo(context))
class ModernBertEncoderLayer(nn.Module):
def __init__(
self,
config: ModernBertPurposeConfig,
layer_index: int,
) -> None:
super().__init__()
self.attention_type = config.layer_types[layer_index]
self.attn_norm = (
nn.Identity()
if layer_index == 0
else nn.LayerNorm(
config.hidden_size, eps=config.norm_eps, bias=config.norm_bias
)
)
self.attn = ModernBertAttention(config, self.attention_type)
self.mlp_norm = nn.LayerNorm(
config.hidden_size, eps=config.norm_eps, bias=config.norm_bias
)
self.mlp = ModernBertMLP(config)
def __call__(self, hidden_states: Any, mask: Any, cos: Any, sin: Any) -> Any:
hidden_states = hidden_states + self.attn(
self.attn_norm(hidden_states), mask, cos, sin
)
return hidden_states + self.mlp(self.mlp_norm(hidden_states))
class ModernBertModel(nn.Module):
def __init__(self, config: ModernBertPurposeConfig) -> None:
super().__init__()
self.config = config
self.embeddings = ModernBertEmbeddings(config)
self.layers = [
ModernBertEncoderLayer(config, index)
for index in range(config.num_hidden_layers)
]
self.final_norm = nn.LayerNorm(
config.hidden_size, eps=config.norm_eps, bias=config.norm_bias
)
def _rotary(self, length: int, theta: float) -> tuple[Any, Any]:
head_dim = self.config.hidden_size // self.config.num_attention_heads
inverse_frequency = 1.0 / (
theta ** (mx.arange(0, head_dim, 2).astype(mx.float32) / head_dim)
)
frequency = mx.arange(length).astype(mx.float32)[:, None] * inverse_frequency[None]
embedding = mx.concatenate((frequency, frequency), axis=-1)
return mx.cos(embedding), mx.sin(embedding)
def _masks(self, attention_mask: Any, dtype: Any) -> dict[str, Any]:
length = attention_mask.shape[1]
visible = attention_mask.astype(mx.bool_)[:, None, None, :]
zero = mx.array(0.0, dtype=dtype)
blocked = mx.array(-1e4, dtype=dtype)
full = mx.where(visible, zero, blocked)
positions = mx.arange(length)
half_window = self.config.local_attention // 2
local_visible = mx.abs(positions[:, None] - positions[None, :]) <= half_window
local = mx.where(
visible & local_visible[None, None, :, :], zero, blocked
)
return {"full_attention": full, "sliding_attention": local}
def __call__(self, input_ids: Any, attention_mask: Any | None = None) -> Any:
if input_ids.shape[1] > self.config.max_position_embeddings:
raise DataError("ModernBERT input exceeds its position contract")
if attention_mask is None:
attention_mask = mx.ones_like(input_ids)
hidden_states = self.embeddings(input_ids)
masks = self._masks(attention_mask, hidden_states.dtype)
rotary = {
"full_attention": self._rotary(
input_ids.shape[1], self.config.full_rope_theta
),
"sliding_attention": self._rotary(
input_ids.shape[1], self.config.local_rope_theta
),
}
for layer in self.layers:
cos, sin = rotary[layer.attention_type]
if self.config.gradient_checkpointing and self.training:
hidden_states = mx.checkpoint(layer)(
hidden_states, masks[layer.attention_type], cos, sin
)
else:
hidden_states = layer(
hidden_states, masks[layer.attention_type], cos, sin
)
return self.final_norm(hidden_states)
class ModernBertPredictionHead(nn.Module):
def __init__(self, config: ModernBertPurposeConfig) -> None:
super().__init__()
self.dense = nn.Linear(
config.hidden_size, config.hidden_size, bias=config.classifier_bias
)
self.act = nn.GELU(approx="none")
self.norm = nn.LayerNorm(
config.hidden_size, eps=config.norm_eps, bias=config.norm_bias
)
def __call__(self, hidden_states: Any) -> Any:
return self.norm(self.act(self.dense(hidden_states)))
class ModernBertForPurposeClassification(nn.Module):
def __init__(self, config: ModernBertPurposeConfig) -> None:
super().__init__()
self.config = config
self.model = ModernBertModel(config)
# Reuse ModernBERT's pretrained masked-LM prediction head as the shared
# representation adapter before the four task-specific readouts.
self.head = ModernBertPredictionHead(config)
self.drop = nn.Dropout(config.classifier_dropout)
self.purpose_classifier = nn.Linear(
config.hidden_size, len(LABELS), bias=True
)
self.secondary_classifier = nn.Linear(
config.hidden_size, len(LABELS), bias=True
)
self.mixed_classifier = nn.Linear(config.hidden_size, 1, bias=True)
self.difficulty_regressor = nn.Linear(config.hidden_size, 1, bias=True)
def __call__(self, input_ids: Any, attention_mask: Any | None = None) -> dict[str, Any]:
sequence = self.model(input_ids, attention_mask)
pooled = self.drop(self.head(sequence[:, 0]))
return {
"purpose_logits": self.purpose_classifier(pooled),
"secondary_logits": self.secondary_classifier(pooled),
"mixed_logits": self.mixed_classifier(pooled).squeeze(-1),
"difficulty": mx.sigmoid(
self.difficulty_regressor(pooled).squeeze(-1)
),
}
def load_pretrained_weights(
model: ModernBertForPurposeClassification,
checkpoint: Path,
) -> dict[str, int]:
"""Load every pretrained backbone/head tensor and reject partial checkpoints."""
if not checkpoint.is_file():
raise DataError(f"{checkpoint}: ModernBERT safetensors checkpoint is missing")
weights = mx.load(str(checkpoint))
available = {
key: value
for key, value in weights.items()
if key.startswith("model.") or key.startswith("head.")
}
parameters = dict(tree_flatten(model.parameters()))
expected = {
key for key in parameters if key.startswith("model.") or key.startswith("head.")
}
missing = sorted(expected - set(available))
if missing:
preview = ", ".join(missing[:5])
raise DataError(
f"ModernBERT checkpoint is missing {len(missing)} required tensors: {preview}"
)
for key in expected:
if tuple(available[key].shape) != tuple(parameters[key].shape):
raise DataError(
f"ModernBERT tensor {key} has shape {available[key].shape}; "
f"expected {parameters[key].shape}"
)
model.load_weights(list(available.items()), strict=False)
mx.eval(model.parameters())
return {
"loaded": len(available),
"ignored": len(weights) - len(available),
"freshTaskHeads": len(parameters) - len(expected),
}
def save_weights(model: ModernBertForPurposeClassification, checkpoint: Path) -> None:
mx.eval(model.parameters())
mx.save_safetensors(
str(checkpoint),
dict(tree_flatten(model.parameters())),
metadata={"format": "pt"},
)
+152
View File
@@ -0,0 +1,152 @@
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(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()
+128
View File
@@ -0,0 +1,128 @@
import tempfile
import sys
import unittest
from pathlib import Path
import mlx.core as mx
import mlx.nn as nn
MODULE_DIR = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(MODULE_DIR))
from deep_model_mlx import (
ModernBertForPurposeClassification,
ModernBertPurposeConfig,
load_pretrained_weights,
save_weights,
)
from purpose_data import DataError, LABELS
def tiny_config(*, checkpointing=False):
return ModernBertPurposeConfig(
vocab_size=64,
hidden_size=16,
intermediate_size=24,
num_hidden_layers=3,
num_attention_heads=4,
max_position_embeddings=32,
pad_token_id=0,
norm_eps=1e-5,
norm_bias=False,
attention_bias=False,
attention_dropout=0.0,
layer_types=("full_attention", "sliding_attention", "sliding_attention"),
local_attention=4,
embedding_dropout=0.0,
mlp_bias=False,
mlp_dropout=0.0,
classifier_bias=False,
classifier_dropout=0.0,
full_rope_theta=160_000.0,
local_rope_theta=10_000.0,
gradient_checkpointing=checkpointing,
)
class DeepModelTests(unittest.TestCase):
def setUp(self):
mx.random.seed(7)
def test_all_four_heads_have_the_expected_shapes_and_ranges(self):
model = ModernBertForPurposeClassification(tiny_config())
output = model(
mx.array([[1, 3, 4, 2, 0, 0], [1, 5, 6, 7, 8, 2]]),
mx.array([[1, 1, 1, 1, 0, 0], [1, 1, 1, 1, 1, 1]]),
)
mx.eval(*output.values())
self.assertEqual((2, len(LABELS)), output["purpose_logits"].shape)
self.assertEqual((2, len(LABELS)), output["secondary_logits"].shape)
self.assertEqual((2,), output["mixed_logits"].shape)
self.assertEqual((2,), output["difficulty"].shape)
self.assertTrue(bool(mx.all(output["difficulty"] >= 0).item()))
self.assertTrue(bool(mx.all(output["difficulty"] <= 1).item()))
def test_masked_padding_tokens_do_not_change_cls_outputs(self):
model = ModernBertForPurposeClassification(tiny_config())
model.eval()
mask = mx.array([[1, 1, 1, 1, 0, 0]])
first = model(mx.array([[1, 3, 4, 2, 0, 0]]), mask)
second = model(mx.array([[1, 3, 4, 2, 9, 10]]), mask)
mx.eval(*first.values(), *second.values())
for key in first:
with self.subTest(head=key):
self.assertLess(float(mx.max(mx.abs(first[key] - second[key])).item()), 1e-5)
def test_gradient_checkpointed_multitask_smoke(self):
model = ModernBertForPurposeClassification(
tiny_config(checkpointing=True)
)
model.train()
def loss(ids, mask):
output = model(ids, mask)
return (
mx.mean(output["purpose_logits"] ** 2)
+ mx.mean(output["secondary_logits"] ** 2)
+ mx.mean(output["mixed_logits"] ** 2)
+ mx.mean(output["difficulty"] ** 2)
)
value_and_grad = nn.value_and_grad(model, loss)
value, gradients = value_and_grad(
mx.array([[1, 3, 4, 2]]), mx.ones((1, 4), dtype=mx.int32)
)
mx.eval(value, gradients)
self.assertTrue(float(value.item()) > 0)
def test_checkpoint_round_trip(self):
model = ModernBertForPurposeClassification(tiny_config())
with tempfile.TemporaryDirectory() as temp:
path = Path(temp) / "model.safetensors"
save_weights(model, path)
restored = ModernBertForPurposeClassification(tiny_config())
restored.load_weights(str(path), strict=True)
ids = mx.array([[1, 3, 4, 2]])
mask = mx.ones((1, 4), dtype=mx.int32)
first = model(ids, mask)
second = restored(ids, mask)
mx.eval(*first.values(), *second.values())
for key in first:
with self.subTest(head=key):
self.assertEqual(
0, float(mx.max(mx.abs(first[key] - second[key])).item())
)
def test_pretrained_loader_rejects_partial_backbone(self):
with tempfile.TemporaryDirectory() as temp:
path = Path(temp) / "partial.safetensors"
mx.save_safetensors(str(path), {"model.final_norm.weight": mx.ones((16,))})
with self.assertRaisesRegex(DataError, "missing"):
load_pretrained_weights(
ModernBertForPurposeClassification(tiny_config()), path
)
if __name__ == "__main__":
unittest.main()
+752
View File
@@ -0,0 +1,752 @@
#!/usr/bin/env python3
"""Fine-tune the multi-task purpose-deep ModernBERT classifier with MLX."""
from __future__ import annotations
import argparse
import json
import math
import random
import shutil
import sys
import time
from collections import Counter
from pathlib import Path
from typing import Any, Iterator, Sequence
import numpy as np
from deep_contract import (
DEEP_VARIANTS,
HEAD_TOKENS,
MAX_LENGTH,
SCORABLE_HARD_SLICES,
TAIL_TOKENS,
DeepTargets,
DeepVariant,
best_mixed_threshold,
encode_fixed_shape_numpy,
encode_targets,
multitask_metrics,
validate_deep_records,
validate_variant_config,
)
from purpose_data import LABELS, DataError, load_jsonl, write_json
from train import (
_fit_temperature,
choose_confidence_thresholds,
expected_calibration_error,
)
from train_mlx import _configure_mlx_device, _linear_schedule
SCRIPT_DIR = Path(__file__).resolve().parent
DEFAULT_DATASET_DIR = SCRIPT_DIR / ".artifacts" / "dataset-v1"
DEFAULT_OUTPUT_ROOT = SCRIPT_DIR / "outputs"
def _load_mlx(device: str) -> tuple[Any, Any, Any]:
try:
import mlx.core as mx
import mlx.nn as nn
import mlx.optimizers as optim
except ImportError as exc:
raise DataError(
"purpose-deep MLX training requires requirements-mlx.txt"
) from exc
_configure_mlx_device(mx, device)
return mx, nn, optim
def _resolve_source(variant: DeepVariant, local_model: Path | None) -> Path:
if local_model is not None:
source = local_model.expanduser().resolve()
if not source.is_dir():
raise DataError(f"{source}: --model must be a local checkpoint directory")
return source
try:
from huggingface_hub import snapshot_download
except ImportError as exc:
raise DataError("downloading ModernBERT requires huggingface_hub") from exc
print(
f"resolving {variant.model_id}@{variant.revision} ({variant.parameter_class})",
flush=True,
)
return Path(
snapshot_download(
repo_id=variant.model_id,
revision=variant.revision,
allow_patterns=[
"config.json",
"model.safetensors",
"tokenizer.json",
"tokenizer_config.json",
"special_tokens_map.json",
],
)
)
def _load_config(source: Path, variant: DeepVariant) -> dict[str, Any]:
path = source / "config.json"
try:
config = json.loads(path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
raise DataError(f"{path}: cannot load ModernBERT config: {exc}") from exc
validate_variant_config(config, variant)
return config
def _prepare_output(path: Path, source: Path, overwrite: bool) -> None:
try:
source.resolve().relative_to(path.resolve())
except ValueError:
pass
else:
raise DataError("--model must not be inside --output-dir")
if path.exists() and any(path.iterdir()):
if not overwrite:
raise DataError(
f"{path}: output is not empty; pass --overwrite-output intentionally"
)
shutil.rmtree(path)
path.mkdir(parents=True, exist_ok=True)
def _encode_records(
tokenizer: Any,
records: Sequence[dict[str, Any]],
*,
chunk_size: int = 256,
) -> dict[str, np.ndarray]:
chunks: dict[str, list[np.ndarray]] = {}
for start in range(0, len(records), chunk_size):
encoded = encode_fixed_shape_numpy(
tokenizer,
[record["prompt"] for record in records[start : start + chunk_size]],
)
for key, value in encoded.items():
chunks.setdefault(key, []).append(value)
return {key: np.concatenate(values) for key, values in chunks.items()}
def _batch_indexes(
size: int,
batch_size: int,
*,
permutation: np.ndarray | None = None,
) -> Iterator[np.ndarray]:
indexes = permutation if permutation is not None else np.arange(size)
for start in range(0, size, batch_size):
yield indexes[start : start + batch_size]
def _mlx_batch(
mx: Any,
encoded: dict[str, np.ndarray],
targets: DeepTargets,
weights: np.ndarray,
indexes: np.ndarray,
) -> dict[str, Any]:
return {
"input_ids": mx.array(encoded["input_ids"][indexes]),
"attention_mask": mx.array(encoded["attention_mask"][indexes]),
"primary": mx.array(targets.primary[indexes]),
"secondary": mx.array(targets.secondary[indexes]),
"secondary_mask": mx.array(targets.secondary_mask[indexes]),
"mixed": mx.array(targets.mixed[indexes]),
"difficulty": mx.array(targets.difficulty[indexes]),
"sample_weights": mx.array(weights[indexes]),
}
def _evaluate(
mx: Any,
model: Any,
encoded: dict[str, np.ndarray],
batch_size: int,
) -> dict[str, np.ndarray]:
model.eval()
collected: dict[str, list[np.ndarray]] = {}
for indexes in _batch_indexes(len(encoded["input_ids"]), batch_size):
output = model(
input_ids=mx.array(encoded["input_ids"][indexes]),
attention_mask=mx.array(encoded["attention_mask"][indexes]),
)
mx.eval(*output.values())
for key, value in output.items():
collected.setdefault(key, []).append(np.asarray(value))
return {key: np.concatenate(values) for key, values in collected.items()}
def _secondary_class_weights(records: Sequence[dict[str, Any]]) -> np.ndarray:
counts = Counter(
record["secondary"] for record in records if record["secondary"] is not None
)
present = [counts[label] for label in LABELS if counts[label]]
if not present:
raise DataError("purpose-deep needs mixed records with secondary labels")
reference = sum(present) / len(present)
# Square-root balancing corrects the known skew without letting a five-example
# secondary class dominate the shared encoder's primary-purpose gradients.
raw = np.asarray(
[math.sqrt(reference / max(counts[label], 1)) for label in LABELS],
dtype=np.float32,
)
return raw / raw.mean()
def _sample_weights(
records: Sequence[dict[str, Any]], hard_weight: float
) -> np.ndarray:
return np.asarray(
[hard_weight if record["slice"] in SCORABLE_HARD_SLICES else 1.0 for record in records],
dtype=np.float32,
)
def _checkpoint_config(
source_config: dict[str, Any],
variant: DeepVariant,
) -> dict[str, Any]:
config = dict(source_config)
config.update(
{
"architectures": ["ModernBertForPurposeClassification"],
"id2label": {str(index): label for index, label in enumerate(LABELS)},
"label2id": {label: index for index, label in enumerate(LABELS)},
"num_labels": len(LABELS),
"purpose_classifier": {
"schemaVersion": 1,
"modelVersion": f"purpose-deep-v1-{variant.name}",
"trainingBackend": "mlx",
"fixedInputShape": [1, MAX_LENGTH],
"heads": ["purpose", "secondary", "mixed", "difficulty"],
},
}
)
return config
def _save_checkpoint(
mx: Any,
model: Any,
tokenizer: Any,
destination: Path,
config: dict[str, Any],
) -> None:
from deep_model_mlx import save_weights
if destination.exists():
shutil.rmtree(destination)
destination.mkdir(parents=True)
tokenizer.save_pretrained(destination)
(destination / "config.json").write_text(
json.dumps(config, indent=2, sort_keys=True) + "\n", encoding="utf-8"
)
save_weights(model, destination / "model.safetensors")
mx.eval(model.parameters())
def _softmax(values: np.ndarray) -> np.ndarray:
shifted = values - values.max(axis=-1, keepdims=True)
exponentials = np.exp(shifted)
return exponentials / exponentials.sum(axis=-1, keepdims=True)
def _calibration(
outputs: dict[str, np.ndarray],
records: Sequence[dict[str, Any]],
args: argparse.Namespace,
model_version: str,
) -> tuple[dict[str, Any], dict[str, Any]]:
try:
import torch
except ImportError as exc:
raise DataError("final purpose-deep calibration requires PyTorch") from exc
targets = encode_targets(records)
scorable = np.asarray(
[record["slice"] != "vague-eval" for record in records], dtype=np.bool_
)
temperature = _fit_temperature(
torch,
torch.from_numpy(outputs["purpose_logits"][scorable]),
torch.from_numpy(targets.primary[scorable].astype(np.int64)),
)
probabilities = _softmax(outputs["purpose_logits"] / temperature)
ranked = np.argsort(probabilities, axis=-1)
row_indexes = np.arange(len(records))
top = ranked[:, -1]
top_probabilities = probabilities[row_indexes, top]
margins = top_probabilities - probabilities[row_indexes, ranked[:, -2]]
correct = ((top == targets.primary) & scorable).tolist()
confidence = choose_confidence_thresholds(
top_probabilities.tolist(),
margins.tolist(),
correct,
high_precision=args.high_precision,
accepted_precision=args.accepted_precision,
)
mixed_mask = targets.secondary_mask
secondary_temperature = _fit_temperature(
torch,
torch.from_numpy(outputs["secondary_logits"][mixed_mask]),
torch.from_numpy(targets.secondary[mixed_mask].astype(np.int64)),
)
mixed_threshold = best_mixed_threshold(
outputs["mixed_logits"][scorable], targets.mixed[scorable]
)
calibrated_metrics = multitask_metrics(
outputs, records, mixed_threshold=mixed_threshold
)
vague = ~scorable
score = top_probabilities * (0.5 + 0.5 * margins)
vague_low_rate = (
float(np.mean(score[vague] < confidence["medium"]["minimumScore"]))
if np.any(vague)
else None
)
calibration = {
"schemaVersion": 1,
"modelVersion": model_version,
"labels": list(LABELS),
"temperature": temperature,
"confidence": confidence,
"validationECE": expected_calibration_error(
top_probabilities.tolist(), correct
),
"secondary": {
"temperature": secondary_temperature,
"labels": list(LABELS),
},
"mixed": {
"threshold": mixed_threshold,
"validationF1": calibrated_metrics["mixed"]["f1"],
},
"difficulty": {
"activation": "sigmoid",
"advisoryOnly": True,
},
}
return calibration, {
"multitask": calibrated_metrics,
"vagueLowRate": vague_low_rate,
}
def train(args: argparse.Namespace) -> dict[str, Any]:
mx, nn, optim = _load_mlx(args.device)
try:
from transformers import AutoTokenizer
from deep_model_mlx import (
ModernBertForPurposeClassification,
ModernBertPurposeConfig,
load_pretrained_weights,
)
except ImportError as exc:
raise DataError(
"purpose-deep dependencies are missing; install requirements-base.txt "
"and requirements-mlx.txt"
) from exc
variant = DEEP_VARIANTS[args.variant]
source = _resolve_source(variant, args.model)
source_config = _load_config(source, variant)
output_dir = args.output_dir or (
DEFAULT_OUTPUT_ROOT / f"purpose-deep-v1-{variant.name}-mlx"
)
_prepare_output(output_dir, source, args.overwrite_output)
train_path = args.dataset_dir / "train.jsonl"
validation_path = args.dataset_dir / "validation.jsonl"
train_records = load_jsonl(train_path)
validation_records = load_jsonl(validation_path)
validate_deep_records(train_records, str(train_path), training=True)
validate_deep_records(validation_records, str(validation_path), training=False)
if args.max_train_records:
train_records = train_records[: args.max_train_records]
if args.max_validation_records:
validation_records = validation_records[: args.max_validation_records]
tokenizer = AutoTokenizer.from_pretrained(source, local_files_only=True)
print("tokenizing fixed 1x512 train and validation splits", flush=True)
encoded_train = _encode_records(tokenizer, train_records)
encoded_validation = _encode_records(tokenizer, validation_records)
train_targets = encode_targets(train_records)
validation_targets = encode_targets(validation_records)
sample_weights = _sample_weights(train_records, args.hard_weight)
secondary_class_weights = _secondary_class_weights(train_records)
non_mixed = len(train_records) - int(train_targets.mixed.sum())
mixed_positive_weight = math.sqrt(
non_mixed / max(float(train_targets.mixed.sum()), 1.0)
)
random.seed(args.seed)
np.random.seed(args.seed)
mx.random.seed(args.seed)
model_config = ModernBertPurposeConfig.from_hugging_face(
source_config, gradient_checkpointing=not args.no_gradient_checkpointing
)
model = ModernBertForPurposeClassification(model_config)
load_report = load_pretrained_weights(model, source / "model.safetensors")
print(
f"loaded ModernBERT tensors={load_report['loaded']} "
f"ignored_mlm_tensors={load_report['ignored']} "
f"fresh_task_tensors={load_report['freshTaskHeads']}",
flush=True,
)
batch_size = args.batch_size or (4 if variant.name == "base" else 2)
eval_batch_size = args.eval_batch_size or (8 if variant.name == "base" else 4)
learning_rate = args.learning_rate or (
2e-5 if variant.name == "base" else 1e-5
)
steps_per_epoch = math.ceil(len(train_records) / batch_size)
total_steps = steps_per_epoch * args.epochs
optimizer = optim.AdamW(
learning_rate=_linear_schedule(
mx,
learning_rate,
total_steps,
round(total_steps * args.warmup_ratio),
),
weight_decay=args.weight_decay,
bias_correction=True,
)
class_weights_mx = mx.array(secondary_class_weights)
def loss_function(
input_ids: Any,
attention_mask: Any,
primary: Any,
secondary: Any,
secondary_mask: Any,
mixed: Any,
difficulty: Any,
weights: Any,
) -> tuple[Any, Any, Any, Any, Any]:
output = model(input_ids=input_ids, attention_mask=attention_mask)
primary_per_record = nn.losses.cross_entropy(
output["purpose_logits"],
primary,
label_smoothing=args.label_smoothing,
reduction="none",
)
primary_loss = mx.sum(primary_per_record * weights) / mx.sum(weights)
secondary_per_record = nn.losses.cross_entropy(
output["secondary_logits"],
secondary,
label_smoothing=args.label_smoothing,
reduction="none",
)
secondary_weights = (
weights
* secondary_mask.astype(weights.dtype)
* class_weights_mx[secondary]
)
secondary_loss = mx.sum(secondary_per_record * secondary_weights) / mx.maximum(
mx.sum(secondary_weights), 1.0
)
mixed_per_record = nn.losses.binary_cross_entropy(
output["mixed_logits"], mixed, reduction="none"
)
mixed_balance = mx.where(mixed > 0.5, mixed_positive_weight, 1.0)
mixed_loss = mx.sum(mixed_per_record * mixed_balance * weights) / mx.sum(
mixed_balance * weights
)
difficulty_per_record = nn.losses.smooth_l1_loss(
output["difficulty"], difficulty, beta=0.1, reduction="none"
)
difficulty_loss = mx.sum(difficulty_per_record * weights) / mx.sum(weights)
total = (
primary_loss
+ args.secondary_loss_weight * secondary_loss
+ args.mixed_loss_weight * mixed_loss
+ args.difficulty_loss_weight * difficulty_loss
)
return total, primary_loss, secondary_loss, mixed_loss, difficulty_loss
loss_and_grad = nn.value_and_grad(model, loss_function)
rng = np.random.default_rng(args.seed)
checkpoint_config = _checkpoint_config(source_config, variant)
best_dir = output_dir / "model"
best_score = float("-inf")
best_metrics: dict[str, Any] | None = None
epochs_without_improvement = 0
stopped_early = False
history: list[dict[str, Any]] = []
started = time.perf_counter()
for epoch in range(1, args.epochs + 1):
epoch_started = time.perf_counter()
model.train()
running = np.zeros(5, dtype=np.float64)
permutation = rng.permutation(len(train_records))
for step, indexes in enumerate(
_batch_indexes(
len(train_records), batch_size, permutation=permutation
),
1,
):
batch = _mlx_batch(
mx, encoded_train, train_targets, sample_weights, indexes
)
losses, gradients = loss_and_grad(
batch["input_ids"],
batch["attention_mask"],
batch["primary"],
batch["secondary"],
batch["secondary_mask"],
batch["mixed"],
batch["difficulty"],
batch["sample_weights"],
)
gradients, _ = optim.clip_grad_norm(gradients, args.max_grad_norm)
optimizer.update(model, gradients)
mx.eval(model.parameters(), optimizer.state, *losses)
running += np.asarray([float(value.item()) for value in losses])
if args.progress_steps and (
step % args.progress_steps == 0 or step == steps_per_epoch
):
mean = running / step
print(
f"epoch {epoch} step {step}/{steps_per_epoch} "
f"loss={mean[0]:.4f} primary={mean[1]:.4f} "
f"secondary={mean[2]:.4f} mixed={mean[3]:.4f} "
f"difficulty={mean[4]:.4f} "
f"elapsed={time.perf_counter() - epoch_started:.1f}s",
flush=True,
)
outputs = _evaluate(
mx, model, encoded_validation, eval_batch_size
)
metrics = multitask_metrics(outputs, validation_records)
metrics["epoch"] = epoch
metrics["meanTrainingLoss"] = (running / steps_per_epoch).tolist()
history.append(metrics)
score = float(metrics["selectionScore"])
secondary_macro = (
metrics["secondary"]["macroRecall"]
if metrics["secondary"] is not None
else 0.0
)
print(
f"epoch {epoch}: primary_accuracy={metrics['primary']['accuracy']:.4%} "
f"hard_accuracy={metrics['primaryHardSlice']['accuracy']:.4%} "
f"secondary_macro_recall={secondary_macro:.4%} "
f"mixed_f1={metrics['mixed']['f1']:.4%} "
f"difficulty_mae={metrics['difficulty']['mae']:.4f} "
f"selection_score={score:.4%}",
flush=True,
)
improvement = score - best_score
if improvement > args.minimum_improvement:
best_score = score
best_metrics = metrics
epochs_without_improvement = 0
_save_checkpoint(
mx, model, tokenizer, best_dir, checkpoint_config
)
write_json(
output_dir / "training-state.json",
{
"bestEpoch": epoch,
"bestSelectionScore": best_score,
"elapsedSeconds": time.perf_counter() - started,
"complete": False,
},
)
else:
epochs_without_improvement += 1
if epochs_without_improvement >= args.early_stopping_patience:
stopped_early = True
print(
f"early stopping after epoch {epoch}: no hard-aware "
f"selection improvement greater than "
f"{args.minimum_improvement:.4%} for "
f"{args.early_stopping_patience} epoch(s)",
flush=True,
)
break
if best_metrics is None:
raise DataError("purpose-deep training did not produce a checkpoint")
# Release the optimizer graph before opening the selected checkpoint; base and
# especially large should never hold two full optimizer states at calibration time.
del optimizer, loss_and_grad, model
mx.clear_cache()
selected_model = ModernBertForPurposeClassification(model_config)
selected_model.load_weights(str(best_dir / "model.safetensors"), strict=True)
selected_outputs = _evaluate(
mx, selected_model, encoded_validation, eval_batch_size
)
model_version = f"purpose-deep-v1-{variant.name}"
calibration, calibrated = _calibration(
selected_outputs, validation_records, args, model_version
)
metrics = {
"modelVersion": model_version,
"variant": variant.name,
"baseModel": variant.model_id,
"baseModelRevision": variant.revision,
"parameterClass": variant.parameter_class,
"trainingBackend": "mlx",
"device": args.device,
"fixedInputShape": [1, MAX_LENGTH],
"truncation": {
"strategy": "head-tail-pair",
"headTokens": HEAD_TOKENS,
"tailTokens": TAIL_TOKENS,
},
"trainingSeconds": time.perf_counter() - started,
"trainRecords": len(train_records),
"validationRecords": len(validation_records),
"mixedTrainRecords": int(train_targets.mixed.sum()),
"mixedValidationRecords": int(validation_targets.mixed.sum()),
"hardTrainingWeight": args.hard_weight,
"lossWeights": {
"purpose": 1.0,
"secondary": args.secondary_loss_weight,
"mixed": args.mixed_loss_weight,
"difficulty": args.difficulty_loss_weight,
},
"secondaryClassWeights": {
label: float(secondary_class_weights[index])
for index, label in enumerate(LABELS)
},
"mixedPositiveWeight": mixed_positive_weight,
"gradientCheckpointing": not args.no_gradient_checkpointing,
"batchSize": batch_size,
"learningRate": learning_rate,
"bestValidationSelectionScore": best_score,
"bestValidation": best_metrics,
"selectedValidation": calibrated,
"epochsCompleted": len(history),
"stoppedEarly": stopped_early,
"history": history,
"calibration": calibration,
}
write_json(output_dir / "calibration.json", calibration)
write_json(output_dir / "metrics.json", metrics)
write_json(
output_dir / "training-config.json",
{
key: str(value) if isinstance(value, Path) else value
for key, value in vars(args).items()
},
)
write_json(
output_dir / "training-state.json",
{
"bestEpoch": int(best_metrics["epoch"]),
"bestSelectionScore": best_score,
"elapsedSeconds": metrics["trainingSeconds"],
"complete": True,
},
)
return metrics
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--variant", choices=tuple(DEEP_VARIANTS), default="base")
parser.add_argument(
"--model",
type=Path,
help="local pinned ModernBERT checkpoint (default: download the pinned revision)",
)
parser.add_argument("--dataset-dir", type=Path, default=DEFAULT_DATASET_DIR)
parser.add_argument("--output-dir", type=Path)
parser.add_argument(
"--device",
choices=("metal", "cpu"),
default="metal",
help="MLX execution device (Metal by default; CPU is diagnostic only)",
)
parser.add_argument("--seed", type=int, default=20260731)
parser.add_argument("--epochs", type=int, default=3)
parser.add_argument("--batch-size", type=int)
parser.add_argument("--eval-batch-size", type=int)
parser.add_argument("--learning-rate", type=float)
parser.add_argument("--weight-decay", type=float, default=0.01)
parser.add_argument("--warmup-ratio", type=float, default=0.1)
parser.add_argument("--max-grad-norm", type=float, default=1.0)
parser.add_argument("--label-smoothing", type=float, default=0.05)
parser.add_argument("--hard-weight", type=float, default=2.0)
parser.add_argument("--secondary-loss-weight", type=float, default=0.25)
parser.add_argument("--mixed-loss-weight", type=float, default=0.25)
parser.add_argument("--difficulty-loss-weight", type=float, default=0.10)
parser.add_argument("--progress-steps", type=int, default=25)
parser.add_argument("--early-stopping-patience", type=int, default=1)
parser.add_argument("--minimum-improvement", type=float, default=0.0005)
parser.add_argument("--high-precision", type=float, default=0.98)
parser.add_argument("--accepted-precision", type=float, default=0.95)
parser.add_argument("--no-gradient-checkpointing", action="store_true")
parser.add_argument("--max-train-records", type=int)
parser.add_argument("--max-validation-records", type=int)
parser.add_argument("--overwrite-output", action="store_true")
return parser
def _positive(parser: argparse.ArgumentParser, name: str, value: Any) -> None:
if value is not None and value <= 0:
parser.error(f"--{name.replace('_', '-')} must be positive")
def main(argv: Sequence[str] | None = None) -> int:
parser = build_parser()
args = parser.parse_args(argv)
for name in (
"epochs",
"batch_size",
"eval_batch_size",
"learning_rate",
"max_grad_norm",
"hard_weight",
"early_stopping_patience",
):
_positive(parser, name, getattr(args, name))
if args.progress_steps < 0:
parser.error("--progress-steps must be non-negative")
if not 0 <= args.warmup_ratio < 1:
parser.error("--warmup-ratio must be in [0, 1)")
if not 0 <= args.label_smoothing < 1:
parser.error("--label-smoothing must be in [0, 1)")
for name in (
"secondary_loss_weight",
"mixed_loss_weight",
"difficulty_loss_weight",
):
if getattr(args, name) < 0:
parser.error(f"--{name.replace('_', '-')} must be non-negative")
if not 0 < args.accepted_precision <= args.high_precision <= 1:
parser.error(
"confidence precision targets must satisfy 0 < accepted <= high <= 1"
)
try:
metrics = train(args)
except (DataError, OSError, RuntimeError, ValueError) as exc:
print(f"error: {exc}", file=sys.stderr)
return 1
selected = metrics["selectedValidation"]["multitask"]
print(
f"selected validation: primary={selected['primary']['accuracy']:.4%} "
f"hard={selected['primaryHardSlice']['accuracy']:.4%} "
f"mixed_f1={selected['mixed']['f1']:.4%}",
flush=True,
)
return 0
if __name__ == "__main__":
raise SystemExit(main())
+120
View File
@@ -0,0 +1,120 @@
#!/usr/bin/env python3
"""Verify pinned Hugging Face ModernBERT -> MLX backbone parity."""
from __future__ import annotations
import argparse
import json
import sys
from pathlib import Path
from typing import Sequence
import numpy as np
from deep_contract import DEEP_VARIANTS, validate_variant_config
from purpose_data import DataError
from train_deep_mlx import _configure_mlx_device, _resolve_source
def verify(variant_name: str, local_model: Path | None, device: str) -> float:
try:
import mlx.core as mx
import torch
from transformers import ModernBertForMaskedLM
from deep_model_mlx import (
ModernBertForPurposeClassification,
ModernBertPurposeConfig,
load_pretrained_weights,
)
except ImportError as exc:
raise DataError(
"deep parity requires PyTorch, Transformers, and MLX"
) from exc
_configure_mlx_device(mx, device)
variant = DEEP_VARIANTS[variant_name]
source = _resolve_source(variant, local_model)
try:
config_json = json.loads(
(source / "config.json").read_text(encoding="utf-8")
)
except (OSError, json.JSONDecodeError) as exc:
raise DataError(f"cannot read ModernBERT config: {exc}") from exc
validate_variant_config(config_json, variant)
torch.set_num_threads(1)
reference = ModernBertForMaskedLM.from_pretrained(
source,
local_files_only=True,
attn_implementation="eager",
)
reference.eval()
candidate = ModernBertForPurposeClassification(
ModernBertPurposeConfig.from_hugging_face(
config_json, gradient_checkpointing=False
)
)
report = load_pretrained_weights(candidate, source / "model.safetensors")
candidate.eval()
# 96 tokens crosses the local layer's 64-token half-window, so this catches
# both full-attention and sliding-window-mask parity without a slow 512-token
# CPU reference pass.
rng = np.random.default_rng(20260731)
input_ids = rng.integers(
3, int(config_json["vocab_size"]) - 1, size=(2, 96), dtype=np.int32
)
input_ids[:, 0] = int(config_json["bos_token_id"])
input_ids[:, -1] = int(config_json["eos_token_id"])
attention_mask = np.ones_like(input_ids, dtype=np.int32)
attention_mask[0, -11:] = 0
input_ids[0, -11:] = int(config_json["pad_token_id"])
with torch.inference_mode():
sequence = reference.model(
input_ids=torch.from_numpy(input_ids.astype(np.int64)),
attention_mask=torch.from_numpy(attention_mask.astype(np.int64)),
).last_hidden_state
reference_pooled = reference.head(sequence[:, 0]).cpu().numpy()
candidate_sequence = candidate.model(
mx.array(input_ids), mx.array(attention_mask)
)
candidate_pooled = candidate.head(candidate_sequence[:, 0])
mx.eval(candidate_pooled)
error = float(
np.max(np.abs(reference_pooled - np.asarray(candidate_pooled)))
)
print(
f"{variant.model_id}@{variant.revision}: "
f"loaded={report['loaded']} ignored={report['ignored']} "
f"pooled_max_abs_error={error:.3g}",
flush=True,
)
if not np.isfinite(error) or error > 5e-4:
raise DataError(
f"ModernBERT MLX parity failed: max abs error {error:.6g} > 0.0005"
)
return error
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--variant", choices=tuple(DEEP_VARIANTS), default="base")
parser.add_argument("--model", type=Path)
parser.add_argument("--device", choices=("metal", "cpu"), default="metal")
return parser
def main(argv: Sequence[str] | None = None) -> int:
args = build_parser().parse_args(argv)
try:
verify(args.variant, args.model, args.device)
except (DataError, OSError, RuntimeError, ValueError) as exc:
print(f"error: {exc}", file=sys.stderr)
return 1
return 0
if __name__ == "__main__":
raise SystemExit(main())