From 9a1228efbb261e4e229dedefabc7b84e5e30c4b0 Mon Sep 17 00:00:00 2001 From: Nucleic Date: Fri, 31 Jul 2026 01:24:01 -0700 Subject: [PATCH] Merge nucleic/sleek-ember-seal-uady into dev --- README.md | 47 +++ deep_contract.py | 317 +++++++++++++++ deep_model_mlx.py | 360 +++++++++++++++++ tests/test_deep_contract.py | 152 +++++++ tests/test_deep_model_mlx.py | 128 ++++++ train_deep_mlx.py | 752 +++++++++++++++++++++++++++++++++++ verify_deep_mlx.py | 120 ++++++ 7 files changed, 1876 insertions(+) create mode 100644 deep_contract.py create mode 100644 deep_model_mlx.py create mode 100644 tests/test_deep_contract.py create mode 100644 tests/test_deep_model_mlx.py create mode 100644 train_deep_mlx.py create mode 100644 verify_deep_mlx.py diff --git a/README.md b/README.md index 703b2d3..8920876 100644 --- a/README.md +++ b/README.md @@ -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 diff --git a/deep_contract.py b/deep_contract.py new file mode 100644 index 0000000..d5bb4f7 --- /dev/null +++ b/deep_contract.py @@ -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] diff --git a/deep_model_mlx.py b/deep_model_mlx.py new file mode 100644 index 0000000..cd3d6b3 --- /dev/null +++ b/deep_model_mlx.py @@ -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"}, + ) diff --git a/tests/test_deep_contract.py b/tests/test_deep_contract.py new file mode 100644 index 0000000..7cabcc0 --- /dev/null +++ b/tests/test_deep_contract.py @@ -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() diff --git a/tests/test_deep_model_mlx.py b/tests/test_deep_model_mlx.py new file mode 100644 index 0000000..fd9278d --- /dev/null +++ b/tests/test_deep_model_mlx.py @@ -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() diff --git a/train_deep_mlx.py b/train_deep_mlx.py new file mode 100644 index 0000000..602a22d --- /dev/null +++ b/train_deep_mlx.py @@ -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()) diff --git a/verify_deep_mlx.py b/verify_deep_mlx.py new file mode 100644 index 0000000..6fe5398 --- /dev/null +++ b/verify_deep_mlx.py @@ -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())