Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -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;
|
||||
`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
|
||||
|
||||
Core ML Tools no longer maintains the legacy ONNX converter, so the Apple artifact is
|
||||
|
||||
@@ -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]
|
||||
@@ -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"},
|
||||
)
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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())
|
||||
@@ -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())
|
||||
Reference in New Issue
Block a user