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

This commit is contained in:
2026-07-31 01:24:01 -07:00
parent 0bee416a88
commit 9a1228efbb
7 changed files with 1876 additions and 0 deletions
+752
View File
@@ -0,0 +1,752 @@
#!/usr/bin/env python3
"""Fine-tune the multi-task purpose-deep ModernBERT classifier with MLX."""
from __future__ import annotations
import argparse
import json
import math
import random
import shutil
import sys
import time
from collections import Counter
from pathlib import Path
from typing import Any, Iterator, Sequence
import numpy as np
from deep_contract import (
DEEP_VARIANTS,
HEAD_TOKENS,
MAX_LENGTH,
SCORABLE_HARD_SLICES,
TAIL_TOKENS,
DeepTargets,
DeepVariant,
best_mixed_threshold,
encode_fixed_shape_numpy,
encode_targets,
multitask_metrics,
validate_deep_records,
validate_variant_config,
)
from purpose_data import LABELS, DataError, load_jsonl, write_json
from train import (
_fit_temperature,
choose_confidence_thresholds,
expected_calibration_error,
)
from train_mlx import _configure_mlx_device, _linear_schedule
SCRIPT_DIR = Path(__file__).resolve().parent
DEFAULT_DATASET_DIR = SCRIPT_DIR / ".artifacts" / "dataset-v1"
DEFAULT_OUTPUT_ROOT = SCRIPT_DIR / "outputs"
def _load_mlx(device: str) -> tuple[Any, Any, Any]:
try:
import mlx.core as mx
import mlx.nn as nn
import mlx.optimizers as optim
except ImportError as exc:
raise DataError(
"purpose-deep MLX training requires requirements-mlx.txt"
) from exc
_configure_mlx_device(mx, device)
return mx, nn, optim
def _resolve_source(variant: DeepVariant, local_model: Path | None) -> Path:
if local_model is not None:
source = local_model.expanduser().resolve()
if not source.is_dir():
raise DataError(f"{source}: --model must be a local checkpoint directory")
return source
try:
from huggingface_hub import snapshot_download
except ImportError as exc:
raise DataError("downloading ModernBERT requires huggingface_hub") from exc
print(
f"resolving {variant.model_id}@{variant.revision} ({variant.parameter_class})",
flush=True,
)
return Path(
snapshot_download(
repo_id=variant.model_id,
revision=variant.revision,
allow_patterns=[
"config.json",
"model.safetensors",
"tokenizer.json",
"tokenizer_config.json",
"special_tokens_map.json",
],
)
)
def _load_config(source: Path, variant: DeepVariant) -> dict[str, Any]:
path = source / "config.json"
try:
config = json.loads(path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
raise DataError(f"{path}: cannot load ModernBERT config: {exc}") from exc
validate_variant_config(config, variant)
return config
def _prepare_output(path: Path, source: Path, overwrite: bool) -> None:
try:
source.resolve().relative_to(path.resolve())
except ValueError:
pass
else:
raise DataError("--model must not be inside --output-dir")
if path.exists() and any(path.iterdir()):
if not overwrite:
raise DataError(
f"{path}: output is not empty; pass --overwrite-output intentionally"
)
shutil.rmtree(path)
path.mkdir(parents=True, exist_ok=True)
def _encode_records(
tokenizer: Any,
records: Sequence[dict[str, Any]],
*,
chunk_size: int = 256,
) -> dict[str, np.ndarray]:
chunks: dict[str, list[np.ndarray]] = {}
for start in range(0, len(records), chunk_size):
encoded = encode_fixed_shape_numpy(
tokenizer,
[record["prompt"] for record in records[start : start + chunk_size]],
)
for key, value in encoded.items():
chunks.setdefault(key, []).append(value)
return {key: np.concatenate(values) for key, values in chunks.items()}
def _batch_indexes(
size: int,
batch_size: int,
*,
permutation: np.ndarray | None = None,
) -> Iterator[np.ndarray]:
indexes = permutation if permutation is not None else np.arange(size)
for start in range(0, size, batch_size):
yield indexes[start : start + batch_size]
def _mlx_batch(
mx: Any,
encoded: dict[str, np.ndarray],
targets: DeepTargets,
weights: np.ndarray,
indexes: np.ndarray,
) -> dict[str, Any]:
return {
"input_ids": mx.array(encoded["input_ids"][indexes]),
"attention_mask": mx.array(encoded["attention_mask"][indexes]),
"primary": mx.array(targets.primary[indexes]),
"secondary": mx.array(targets.secondary[indexes]),
"secondary_mask": mx.array(targets.secondary_mask[indexes]),
"mixed": mx.array(targets.mixed[indexes]),
"difficulty": mx.array(targets.difficulty[indexes]),
"sample_weights": mx.array(weights[indexes]),
}
def _evaluate(
mx: Any,
model: Any,
encoded: dict[str, np.ndarray],
batch_size: int,
) -> dict[str, np.ndarray]:
model.eval()
collected: dict[str, list[np.ndarray]] = {}
for indexes in _batch_indexes(len(encoded["input_ids"]), batch_size):
output = model(
input_ids=mx.array(encoded["input_ids"][indexes]),
attention_mask=mx.array(encoded["attention_mask"][indexes]),
)
mx.eval(*output.values())
for key, value in output.items():
collected.setdefault(key, []).append(np.asarray(value))
return {key: np.concatenate(values) for key, values in collected.items()}
def _secondary_class_weights(records: Sequence[dict[str, Any]]) -> np.ndarray:
counts = Counter(
record["secondary"] for record in records if record["secondary"] is not None
)
present = [counts[label] for label in LABELS if counts[label]]
if not present:
raise DataError("purpose-deep needs mixed records with secondary labels")
reference = sum(present) / len(present)
# Square-root balancing corrects the known skew without letting a five-example
# secondary class dominate the shared encoder's primary-purpose gradients.
raw = np.asarray(
[math.sqrt(reference / max(counts[label], 1)) for label in LABELS],
dtype=np.float32,
)
return raw / raw.mean()
def _sample_weights(
records: Sequence[dict[str, Any]], hard_weight: float
) -> np.ndarray:
return np.asarray(
[hard_weight if record["slice"] in SCORABLE_HARD_SLICES else 1.0 for record in records],
dtype=np.float32,
)
def _checkpoint_config(
source_config: dict[str, Any],
variant: DeepVariant,
) -> dict[str, Any]:
config = dict(source_config)
config.update(
{
"architectures": ["ModernBertForPurposeClassification"],
"id2label": {str(index): label for index, label in enumerate(LABELS)},
"label2id": {label: index for index, label in enumerate(LABELS)},
"num_labels": len(LABELS),
"purpose_classifier": {
"schemaVersion": 1,
"modelVersion": f"purpose-deep-v1-{variant.name}",
"trainingBackend": "mlx",
"fixedInputShape": [1, MAX_LENGTH],
"heads": ["purpose", "secondary", "mixed", "difficulty"],
},
}
)
return config
def _save_checkpoint(
mx: Any,
model: Any,
tokenizer: Any,
destination: Path,
config: dict[str, Any],
) -> None:
from deep_model_mlx import save_weights
if destination.exists():
shutil.rmtree(destination)
destination.mkdir(parents=True)
tokenizer.save_pretrained(destination)
(destination / "config.json").write_text(
json.dumps(config, indent=2, sort_keys=True) + "\n", encoding="utf-8"
)
save_weights(model, destination / "model.safetensors")
mx.eval(model.parameters())
def _softmax(values: np.ndarray) -> np.ndarray:
shifted = values - values.max(axis=-1, keepdims=True)
exponentials = np.exp(shifted)
return exponentials / exponentials.sum(axis=-1, keepdims=True)
def _calibration(
outputs: dict[str, np.ndarray],
records: Sequence[dict[str, Any]],
args: argparse.Namespace,
model_version: str,
) -> tuple[dict[str, Any], dict[str, Any]]:
try:
import torch
except ImportError as exc:
raise DataError("final purpose-deep calibration requires PyTorch") from exc
targets = encode_targets(records)
scorable = np.asarray(
[record["slice"] != "vague-eval" for record in records], dtype=np.bool_
)
temperature = _fit_temperature(
torch,
torch.from_numpy(outputs["purpose_logits"][scorable]),
torch.from_numpy(targets.primary[scorable].astype(np.int64)),
)
probabilities = _softmax(outputs["purpose_logits"] / temperature)
ranked = np.argsort(probabilities, axis=-1)
row_indexes = np.arange(len(records))
top = ranked[:, -1]
top_probabilities = probabilities[row_indexes, top]
margins = top_probabilities - probabilities[row_indexes, ranked[:, -2]]
correct = ((top == targets.primary) & scorable).tolist()
confidence = choose_confidence_thresholds(
top_probabilities.tolist(),
margins.tolist(),
correct,
high_precision=args.high_precision,
accepted_precision=args.accepted_precision,
)
mixed_mask = targets.secondary_mask
secondary_temperature = _fit_temperature(
torch,
torch.from_numpy(outputs["secondary_logits"][mixed_mask]),
torch.from_numpy(targets.secondary[mixed_mask].astype(np.int64)),
)
mixed_threshold = best_mixed_threshold(
outputs["mixed_logits"][scorable], targets.mixed[scorable]
)
calibrated_metrics = multitask_metrics(
outputs, records, mixed_threshold=mixed_threshold
)
vague = ~scorable
score = top_probabilities * (0.5 + 0.5 * margins)
vague_low_rate = (
float(np.mean(score[vague] < confidence["medium"]["minimumScore"]))
if np.any(vague)
else None
)
calibration = {
"schemaVersion": 1,
"modelVersion": model_version,
"labels": list(LABELS),
"temperature": temperature,
"confidence": confidence,
"validationECE": expected_calibration_error(
top_probabilities.tolist(), correct
),
"secondary": {
"temperature": secondary_temperature,
"labels": list(LABELS),
},
"mixed": {
"threshold": mixed_threshold,
"validationF1": calibrated_metrics["mixed"]["f1"],
},
"difficulty": {
"activation": "sigmoid",
"advisoryOnly": True,
},
}
return calibration, {
"multitask": calibrated_metrics,
"vagueLowRate": vague_low_rate,
}
def train(args: argparse.Namespace) -> dict[str, Any]:
mx, nn, optim = _load_mlx(args.device)
try:
from transformers import AutoTokenizer
from deep_model_mlx import (
ModernBertForPurposeClassification,
ModernBertPurposeConfig,
load_pretrained_weights,
)
except ImportError as exc:
raise DataError(
"purpose-deep dependencies are missing; install requirements-base.txt "
"and requirements-mlx.txt"
) from exc
variant = DEEP_VARIANTS[args.variant]
source = _resolve_source(variant, args.model)
source_config = _load_config(source, variant)
output_dir = args.output_dir or (
DEFAULT_OUTPUT_ROOT / f"purpose-deep-v1-{variant.name}-mlx"
)
_prepare_output(output_dir, source, args.overwrite_output)
train_path = args.dataset_dir / "train.jsonl"
validation_path = args.dataset_dir / "validation.jsonl"
train_records = load_jsonl(train_path)
validation_records = load_jsonl(validation_path)
validate_deep_records(train_records, str(train_path), training=True)
validate_deep_records(validation_records, str(validation_path), training=False)
if args.max_train_records:
train_records = train_records[: args.max_train_records]
if args.max_validation_records:
validation_records = validation_records[: args.max_validation_records]
tokenizer = AutoTokenizer.from_pretrained(source, local_files_only=True)
print("tokenizing fixed 1x512 train and validation splits", flush=True)
encoded_train = _encode_records(tokenizer, train_records)
encoded_validation = _encode_records(tokenizer, validation_records)
train_targets = encode_targets(train_records)
validation_targets = encode_targets(validation_records)
sample_weights = _sample_weights(train_records, args.hard_weight)
secondary_class_weights = _secondary_class_weights(train_records)
non_mixed = len(train_records) - int(train_targets.mixed.sum())
mixed_positive_weight = math.sqrt(
non_mixed / max(float(train_targets.mixed.sum()), 1.0)
)
random.seed(args.seed)
np.random.seed(args.seed)
mx.random.seed(args.seed)
model_config = ModernBertPurposeConfig.from_hugging_face(
source_config, gradient_checkpointing=not args.no_gradient_checkpointing
)
model = ModernBertForPurposeClassification(model_config)
load_report = load_pretrained_weights(model, source / "model.safetensors")
print(
f"loaded ModernBERT tensors={load_report['loaded']} "
f"ignored_mlm_tensors={load_report['ignored']} "
f"fresh_task_tensors={load_report['freshTaskHeads']}",
flush=True,
)
batch_size = args.batch_size or (4 if variant.name == "base" else 2)
eval_batch_size = args.eval_batch_size or (8 if variant.name == "base" else 4)
learning_rate = args.learning_rate or (
2e-5 if variant.name == "base" else 1e-5
)
steps_per_epoch = math.ceil(len(train_records) / batch_size)
total_steps = steps_per_epoch * args.epochs
optimizer = optim.AdamW(
learning_rate=_linear_schedule(
mx,
learning_rate,
total_steps,
round(total_steps * args.warmup_ratio),
),
weight_decay=args.weight_decay,
bias_correction=True,
)
class_weights_mx = mx.array(secondary_class_weights)
def loss_function(
input_ids: Any,
attention_mask: Any,
primary: Any,
secondary: Any,
secondary_mask: Any,
mixed: Any,
difficulty: Any,
weights: Any,
) -> tuple[Any, Any, Any, Any, Any]:
output = model(input_ids=input_ids, attention_mask=attention_mask)
primary_per_record = nn.losses.cross_entropy(
output["purpose_logits"],
primary,
label_smoothing=args.label_smoothing,
reduction="none",
)
primary_loss = mx.sum(primary_per_record * weights) / mx.sum(weights)
secondary_per_record = nn.losses.cross_entropy(
output["secondary_logits"],
secondary,
label_smoothing=args.label_smoothing,
reduction="none",
)
secondary_weights = (
weights
* secondary_mask.astype(weights.dtype)
* class_weights_mx[secondary]
)
secondary_loss = mx.sum(secondary_per_record * secondary_weights) / mx.maximum(
mx.sum(secondary_weights), 1.0
)
mixed_per_record = nn.losses.binary_cross_entropy(
output["mixed_logits"], mixed, reduction="none"
)
mixed_balance = mx.where(mixed > 0.5, mixed_positive_weight, 1.0)
mixed_loss = mx.sum(mixed_per_record * mixed_balance * weights) / mx.sum(
mixed_balance * weights
)
difficulty_per_record = nn.losses.smooth_l1_loss(
output["difficulty"], difficulty, beta=0.1, reduction="none"
)
difficulty_loss = mx.sum(difficulty_per_record * weights) / mx.sum(weights)
total = (
primary_loss
+ args.secondary_loss_weight * secondary_loss
+ args.mixed_loss_weight * mixed_loss
+ args.difficulty_loss_weight * difficulty_loss
)
return total, primary_loss, secondary_loss, mixed_loss, difficulty_loss
loss_and_grad = nn.value_and_grad(model, loss_function)
rng = np.random.default_rng(args.seed)
checkpoint_config = _checkpoint_config(source_config, variant)
best_dir = output_dir / "model"
best_score = float("-inf")
best_metrics: dict[str, Any] | None = None
epochs_without_improvement = 0
stopped_early = False
history: list[dict[str, Any]] = []
started = time.perf_counter()
for epoch in range(1, args.epochs + 1):
epoch_started = time.perf_counter()
model.train()
running = np.zeros(5, dtype=np.float64)
permutation = rng.permutation(len(train_records))
for step, indexes in enumerate(
_batch_indexes(
len(train_records), batch_size, permutation=permutation
),
1,
):
batch = _mlx_batch(
mx, encoded_train, train_targets, sample_weights, indexes
)
losses, gradients = loss_and_grad(
batch["input_ids"],
batch["attention_mask"],
batch["primary"],
batch["secondary"],
batch["secondary_mask"],
batch["mixed"],
batch["difficulty"],
batch["sample_weights"],
)
gradients, _ = optim.clip_grad_norm(gradients, args.max_grad_norm)
optimizer.update(model, gradients)
mx.eval(model.parameters(), optimizer.state, *losses)
running += np.asarray([float(value.item()) for value in losses])
if args.progress_steps and (
step % args.progress_steps == 0 or step == steps_per_epoch
):
mean = running / step
print(
f"epoch {epoch} step {step}/{steps_per_epoch} "
f"loss={mean[0]:.4f} primary={mean[1]:.4f} "
f"secondary={mean[2]:.4f} mixed={mean[3]:.4f} "
f"difficulty={mean[4]:.4f} "
f"elapsed={time.perf_counter() - epoch_started:.1f}s",
flush=True,
)
outputs = _evaluate(
mx, model, encoded_validation, eval_batch_size
)
metrics = multitask_metrics(outputs, validation_records)
metrics["epoch"] = epoch
metrics["meanTrainingLoss"] = (running / steps_per_epoch).tolist()
history.append(metrics)
score = float(metrics["selectionScore"])
secondary_macro = (
metrics["secondary"]["macroRecall"]
if metrics["secondary"] is not None
else 0.0
)
print(
f"epoch {epoch}: primary_accuracy={metrics['primary']['accuracy']:.4%} "
f"hard_accuracy={metrics['primaryHardSlice']['accuracy']:.4%} "
f"secondary_macro_recall={secondary_macro:.4%} "
f"mixed_f1={metrics['mixed']['f1']:.4%} "
f"difficulty_mae={metrics['difficulty']['mae']:.4f} "
f"selection_score={score:.4%}",
flush=True,
)
improvement = score - best_score
if improvement > args.minimum_improvement:
best_score = score
best_metrics = metrics
epochs_without_improvement = 0
_save_checkpoint(
mx, model, tokenizer, best_dir, checkpoint_config
)
write_json(
output_dir / "training-state.json",
{
"bestEpoch": epoch,
"bestSelectionScore": best_score,
"elapsedSeconds": time.perf_counter() - started,
"complete": False,
},
)
else:
epochs_without_improvement += 1
if epochs_without_improvement >= args.early_stopping_patience:
stopped_early = True
print(
f"early stopping after epoch {epoch}: no hard-aware "
f"selection improvement greater than "
f"{args.minimum_improvement:.4%} for "
f"{args.early_stopping_patience} epoch(s)",
flush=True,
)
break
if best_metrics is None:
raise DataError("purpose-deep training did not produce a checkpoint")
# Release the optimizer graph before opening the selected checkpoint; base and
# especially large should never hold two full optimizer states at calibration time.
del optimizer, loss_and_grad, model
mx.clear_cache()
selected_model = ModernBertForPurposeClassification(model_config)
selected_model.load_weights(str(best_dir / "model.safetensors"), strict=True)
selected_outputs = _evaluate(
mx, selected_model, encoded_validation, eval_batch_size
)
model_version = f"purpose-deep-v1-{variant.name}"
calibration, calibrated = _calibration(
selected_outputs, validation_records, args, model_version
)
metrics = {
"modelVersion": model_version,
"variant": variant.name,
"baseModel": variant.model_id,
"baseModelRevision": variant.revision,
"parameterClass": variant.parameter_class,
"trainingBackend": "mlx",
"device": args.device,
"fixedInputShape": [1, MAX_LENGTH],
"truncation": {
"strategy": "head-tail-pair",
"headTokens": HEAD_TOKENS,
"tailTokens": TAIL_TOKENS,
},
"trainingSeconds": time.perf_counter() - started,
"trainRecords": len(train_records),
"validationRecords": len(validation_records),
"mixedTrainRecords": int(train_targets.mixed.sum()),
"mixedValidationRecords": int(validation_targets.mixed.sum()),
"hardTrainingWeight": args.hard_weight,
"lossWeights": {
"purpose": 1.0,
"secondary": args.secondary_loss_weight,
"mixed": args.mixed_loss_weight,
"difficulty": args.difficulty_loss_weight,
},
"secondaryClassWeights": {
label: float(secondary_class_weights[index])
for index, label in enumerate(LABELS)
},
"mixedPositiveWeight": mixed_positive_weight,
"gradientCheckpointing": not args.no_gradient_checkpointing,
"batchSize": batch_size,
"learningRate": learning_rate,
"bestValidationSelectionScore": best_score,
"bestValidation": best_metrics,
"selectedValidation": calibrated,
"epochsCompleted": len(history),
"stoppedEarly": stopped_early,
"history": history,
"calibration": calibration,
}
write_json(output_dir / "calibration.json", calibration)
write_json(output_dir / "metrics.json", metrics)
write_json(
output_dir / "training-config.json",
{
key: str(value) if isinstance(value, Path) else value
for key, value in vars(args).items()
},
)
write_json(
output_dir / "training-state.json",
{
"bestEpoch": int(best_metrics["epoch"]),
"bestSelectionScore": best_score,
"elapsedSeconds": metrics["trainingSeconds"],
"complete": True,
},
)
return metrics
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--variant", choices=tuple(DEEP_VARIANTS), default="base")
parser.add_argument(
"--model",
type=Path,
help="local pinned ModernBERT checkpoint (default: download the pinned revision)",
)
parser.add_argument("--dataset-dir", type=Path, default=DEFAULT_DATASET_DIR)
parser.add_argument("--output-dir", type=Path)
parser.add_argument(
"--device",
choices=("metal", "cpu"),
default="metal",
help="MLX execution device (Metal by default; CPU is diagnostic only)",
)
parser.add_argument("--seed", type=int, default=20260731)
parser.add_argument("--epochs", type=int, default=3)
parser.add_argument("--batch-size", type=int)
parser.add_argument("--eval-batch-size", type=int)
parser.add_argument("--learning-rate", type=float)
parser.add_argument("--weight-decay", type=float, default=0.01)
parser.add_argument("--warmup-ratio", type=float, default=0.1)
parser.add_argument("--max-grad-norm", type=float, default=1.0)
parser.add_argument("--label-smoothing", type=float, default=0.05)
parser.add_argument("--hard-weight", type=float, default=2.0)
parser.add_argument("--secondary-loss-weight", type=float, default=0.25)
parser.add_argument("--mixed-loss-weight", type=float, default=0.25)
parser.add_argument("--difficulty-loss-weight", type=float, default=0.10)
parser.add_argument("--progress-steps", type=int, default=25)
parser.add_argument("--early-stopping-patience", type=int, default=1)
parser.add_argument("--minimum-improvement", type=float, default=0.0005)
parser.add_argument("--high-precision", type=float, default=0.98)
parser.add_argument("--accepted-precision", type=float, default=0.95)
parser.add_argument("--no-gradient-checkpointing", action="store_true")
parser.add_argument("--max-train-records", type=int)
parser.add_argument("--max-validation-records", type=int)
parser.add_argument("--overwrite-output", action="store_true")
return parser
def _positive(parser: argparse.ArgumentParser, name: str, value: Any) -> None:
if value is not None and value <= 0:
parser.error(f"--{name.replace('_', '-')} must be positive")
def main(argv: Sequence[str] | None = None) -> int:
parser = build_parser()
args = parser.parse_args(argv)
for name in (
"epochs",
"batch_size",
"eval_batch_size",
"learning_rate",
"max_grad_norm",
"hard_weight",
"early_stopping_patience",
):
_positive(parser, name, getattr(args, name))
if args.progress_steps < 0:
parser.error("--progress-steps must be non-negative")
if not 0 <= args.warmup_ratio < 1:
parser.error("--warmup-ratio must be in [0, 1)")
if not 0 <= args.label_smoothing < 1:
parser.error("--label-smoothing must be in [0, 1)")
for name in (
"secondary_loss_weight",
"mixed_loss_weight",
"difficulty_loss_weight",
):
if getattr(args, name) < 0:
parser.error(f"--{name.replace('_', '-')} must be non-negative")
if not 0 < args.accepted_precision <= args.high_precision <= 1:
parser.error(
"confidence precision targets must satisfy 0 < accepted <= high <= 1"
)
try:
metrics = train(args)
except (DataError, OSError, RuntimeError, ValueError) as exc:
print(f"error: {exc}", file=sys.stderr)
return 1
selected = metrics["selectedValidation"]["multitask"]
print(
f"selected validation: primary={selected['primary']['accuracy']:.4%} "
f"hard={selected['primaryHardSlice']['accuracy']:.4%} "
f"mixed_f1={selected['mixed']['f1']:.4%}",
flush=True,
)
return 0
if __name__ == "__main__":
raise SystemExit(main())