Files
nucleic-purpose-classifier/train_deep_mlx.py
T

1288 lines
46 KiB
Python

#!/usr/bin/env python3
"""Fine-tune the multi-task purpose-deep ModernBERT classifier with MLX."""
from __future__ import annotations
import argparse
import hashlib
import json
import math
import os
import random
import shutil
import shlex
import signal
import sys
import tempfile
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, _teacher_cache
SCRIPT_DIR = Path(__file__).resolve().parent
DEFAULT_DATASET_DIR = SCRIPT_DIR / ".artifacts" / "dataset-v1"
DEFAULT_OUTPUT_ROOT = SCRIPT_DIR / "outputs"
RESUME_SCHEMA_VERSION = 1
PATH_ARGUMENTS = {"dataset_dir", "distillation_cache"}
class TrainingPaused(Exception):
"""Raised after a signal-requested training checkpoint is durable."""
def __init__(self, checkpoint: Path, signum: int) -> None:
super().__init__(str(checkpoint))
self.checkpoint = checkpoint
self.signum = signum
class _ShutdownController:
def __init__(self) -> None:
self.signum: int | None = None
self._previous: dict[int, Any] = {}
def _handle(self, signum: int, _frame: Any) -> None:
if self.signum is not None:
raise KeyboardInterrupt
self.signum = signum
def install(self) -> None:
for signum in (signal.SIGINT, signal.SIGTERM):
self._previous[signum] = signal.getsignal(signum)
signal.signal(signum, self._handle)
def restore(self) -> None:
for signum, handler in self._previous.items():
signal.signal(signum, handler)
self._previous.clear()
def _sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as source:
for chunk in iter(lambda: source.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def _read_resume_state(checkpoint: Path) -> dict[str, Any]:
state_path = checkpoint / "resume-state.json"
try:
state = json.loads(state_path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
raise DataError(f"{state_path}: cannot load training resume state: {exc}") from exc
if state.get("schemaVersion") != RESUME_SCHEMA_VERSION:
raise DataError(f"{state_path}: unsupported training resume schema")
if state.get("status") != "paused":
raise DataError(f"{state_path}: checkpoint is not paused training state")
return state
def _serialized_resume_arguments(args: argparse.Namespace) -> dict[str, Any]:
excluded = {
"model",
"resume_from",
"resume_training",
"output_dir",
"overwrite_output",
}
return {
key: (
str(value.expanduser().resolve())
if isinstance(value, Path)
else value
)
for key, value in vars(args).items()
if key not in excluded
}
def _restore_resume_arguments(
args: argparse.Namespace, checkpoint: Path, state: dict[str, Any]
) -> None:
saved = state.get("arguments")
if not isinstance(saved, dict):
raise DataError("training resume state has no saved arguments")
for key, value in saved.items():
if not hasattr(args, key):
raise DataError(f"training resume state has unknown argument {key!r}")
setattr(args, key, Path(value) if key in PATH_ARGUMENTS and value else value)
args.model = None
args.resume_from = None
args.resume_training = checkpoint
args.output_dir = checkpoint.parent
args.overwrite_output = False
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("the input checkpoint 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 _save_training_resume(
mx: Any,
model: Any,
optimizer: Any,
tokenizer: Any,
output_dir: Path,
checkpoint_config: dict[str, Any],
state: dict[str, Any],
) -> Path:
"""Atomically save the current model, optimizer, and loop cursor."""
from mlx.utils import tree_flatten
output_dir.mkdir(parents=True, exist_ok=True)
temporary = Path(
tempfile.mkdtemp(prefix=".training-resume-", dir=output_dir)
)
destination = output_dir / "resume"
backup = output_dir / ".training-resume-backup"
try:
_save_checkpoint(mx, model, tokenizer, temporary, checkpoint_config)
mx.eval(optimizer.state)
optimizer_state = tree_flatten(optimizer.state, destination={})
if not optimizer_state:
raise DataError("optimizer state is empty; refusing an incomplete resume")
mx.save_safetensors(
str(temporary / "optimizer.safetensors"), optimizer_state
)
write_json(temporary / "resume-state.json", state)
if backup.exists():
shutil.rmtree(backup)
if destination.exists():
os.replace(destination, backup)
os.replace(temporary, destination)
if backup.exists():
shutil.rmtree(backup)
except BaseException:
if temporary.exists():
shutil.rmtree(temporary)
if not destination.exists() and backup.exists():
os.replace(backup, destination)
raise
return destination
def _load_optimizer_state(mx: Any, optimizer: Any, checkpoint: Path) -> None:
from mlx.utils import tree_unflatten
path = checkpoint / "optimizer.safetensors"
if not path.is_file():
raise DataError(f"{path}: optimizer resume state is missing")
optimizer.state = tree_unflatten(mx.load(str(path)))
mx.eval(optimizer.state)
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 _distillation_loss(
mx: Any,
student_logits: Any,
teacher_logits: Any,
*,
temperature: float,
weights: Any,
) -> Any:
"""Return weighted teacher-to-student KL loss for one MLX batch."""
student_log_probabilities = (
student_logits / temperature
- mx.logsumexp(student_logits / temperature, axis=-1, keepdims=True)
)
teacher_probabilities = mx.softmax(
teacher_logits / temperature,
axis=-1,
)
teacher_log_probabilities = mx.log(
mx.maximum(teacher_probabilities, 1e-12)
)
per_record = (
mx.sum(
teacher_probabilities
* (teacher_log_probabilities - student_log_probabilities),
axis=-1,
)
* temperature
* temperature
)
return mx.sum(per_record * weights) / mx.sum(weights)
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_checkpoint_weights,
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]
resume_state = (
_read_resume_state(args.resume_training)
if args.resume_training is not None
else None
)
source = _resolve_source(
variant, args.resume_training or args.resume_from or args.model
)
source_config = _load_config(source, variant)
if args.resume_training is not None:
output_dir = args.resume_training.parent
if args.output_dir is not None and args.output_dir.resolve() != output_dir:
raise DataError("--resume-training must use its original output directory")
if not (output_dir / "model" / "model.safetensors").is_file():
raise DataError(f"{output_dir}: selected model checkpoint is missing")
else:
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]
input_hashes = {
"train": _sha256(train_path),
"validation": _sha256(validation_path),
"distillationCache": (
_sha256(args.distillation_cache.expanduser())
if args.distillation_cache is not None
else None
),
}
if resume_state is not None and resume_state.get("inputHashes") != input_hashes:
raise DataError(
"training inputs changed after the pause; refusing a non-deterministic resume"
)
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)
validation_scorable = np.asarray(
[record["slice"] != "vague-eval" for record in validation_records],
dtype=np.bool_,
)
sample_weights = _sample_weights(train_records, args.hard_weight)
secondary_class_weights = _secondary_class_weights(train_records)
teacher_train_logits = None
teacher_validation_logits = None
if args.distillation_cache is not None:
teacher_train_logits, teacher_validation_logits = _teacher_cache(
args.distillation_cache.expanduser(),
train_records,
validation_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)
if args.resume_training is not None or args.resume_from is not None:
load_report = load_checkpoint_weights(
model, source / "model.safetensors"
)
print(
f"resumed purpose-deep tensors={load_report['loaded']} "
"including all task heads"
+ (
" and exact optimizer/loop state"
if args.resume_training is not None
else "; optimizer state starts fresh"
),
flush=True,
)
else:
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,
)
optimizer.init(model.trainable_parameters())
mx.eval(optimizer.state)
if args.resume_training is not None:
_load_optimizer_state(mx, optimizer, source)
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,
teacher_logits: Any | None,
) -> tuple[Any, 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_label_loss = mx.sum(primary_per_record * weights) / mx.sum(weights)
primary_distillation_loss = mx.zeros_like(primary_label_loss)
if teacher_logits is not None:
primary_distillation_loss = _distillation_loss(
mx,
output["purpose_logits"],
teacher_logits,
temperature=args.distillation_temperature,
weights=weights,
)
primary_loss = (
(1.0 - args.distillation_weight) * primary_label_loss
+ args.distillation_weight * primary_distillation_loss
)
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_label_loss,
secondary_loss,
mixed_loss,
difficulty_loss,
primary_distillation_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"
stopped_early = False
write_json(
output_dir / "training-config.json",
{
key: str(value) if isinstance(value, Path) else value
for key, value in vars(args).items()
},
)
if resume_state is not None:
try:
initial_metrics = resume_state["initialMetrics"]
best_score = float(resume_state["bestScore"])
best_metrics = resume_state["bestMetrics"]
epochs_without_improvement = int(
resume_state["epochsWithoutImprovement"]
)
history = list(resume_state["history"])
start_epoch = int(resume_state["epoch"])
resume_next_step = int(resume_state["nextStep"])
resume_permutation = resume_state.get("permutation")
resume_running = np.asarray(
resume_state["runningLosses"], dtype=np.float64
)
resume_epoch_elapsed = float(resume_state["epochElapsedSeconds"])
rng.bit_generator.state = resume_state["numpyRngState"]
started = time.perf_counter() - float(resume_state["elapsedSeconds"])
except (KeyError, TypeError, ValueError) as exc:
raise DataError(f"invalid training loop resume state: {exc}") from exc
if not 1 <= resume_next_step <= steps_per_epoch + 1:
raise DataError("training resume step is outside the epoch")
if resume_running.shape != (6,):
raise DataError("training resume loss accumulator has the wrong shape")
print(
f"continuing epoch {start_epoch} at step "
f"{resume_next_step}/{steps_per_epoch} after "
f"{resume_state['elapsedSeconds']:.1f}s of saved training",
flush=True,
)
else:
epochs_without_improvement = 0
history: list[dict[str, Any]] = []
started = time.perf_counter()
initial_outputs = _evaluate(
mx, model, encoded_validation, eval_batch_size
)
initial_metrics = multitask_metrics(initial_outputs, validation_records)
if teacher_validation_logits is not None:
initial_metrics["teacherAgreement"] = float(
np.mean(
initial_outputs["purpose_logits"][validation_scorable].argmax(
axis=-1
)
== teacher_validation_logits[validation_scorable].argmax(axis=-1)
)
)
initial_metrics["epoch"] = 0
best_score = float(initial_metrics["selectionScore"])
best_metrics: dict[str, Any] = initial_metrics
_save_checkpoint(mx, model, tokenizer, best_dir, checkpoint_config)
write_json(
output_dir / "training-state.json",
{
"bestEpoch": 0,
"bestSelectionScore": best_score,
"elapsedSeconds": time.perf_counter() - started,
"complete": False,
},
)
print(
f"epoch 0: primary_accuracy={initial_metrics['primary']['accuracy']:.4%} "
f"hard_accuracy={initial_metrics['primaryHardSlice']['accuracy']:.4%} "
f"mixed_f1={initial_metrics['mixed']['f1']:.4%} "
f"selection_score={best_score:.4%}"
+ (
f" teacher_agreement={initial_metrics['teacherAgreement']:.4%}"
if "teacherAgreement" in initial_metrics
else ""
),
flush=True,
)
start_epoch = 1
resume_next_step = 1
resume_permutation = None
resume_running = np.zeros(6, dtype=np.float64)
resume_epoch_elapsed = 0.0
shutdown = _ShutdownController()
def pause_training(
epoch: int,
next_step: int,
permutation: np.ndarray | None,
running: np.ndarray,
epoch_elapsed: float,
) -> None:
signum = shutdown.signum or signal.SIGINT
print(
f"shutdown requested; saving exact training state after epoch {epoch} "
f"step {max(next_step - 1, 0)}",
flush=True,
)
state = {
"schemaVersion": RESUME_SCHEMA_VERSION,
"status": "paused",
"signal": signal.Signals(signum).name,
"arguments": _serialized_resume_arguments(args),
"inputHashes": input_hashes,
"epoch": epoch,
"nextStep": next_step,
"permutation": permutation.tolist() if permutation is not None else None,
"runningLosses": running.tolist(),
"epochElapsedSeconds": epoch_elapsed,
"elapsedSeconds": time.perf_counter() - started,
"numpyRngState": rng.bit_generator.state,
"initialMetrics": initial_metrics,
"bestScore": best_score,
"bestMetrics": best_metrics,
"epochsWithoutImprovement": epochs_without_improvement,
"history": history,
}
checkpoint = _save_training_resume(
mx,
model,
optimizer,
tokenizer,
output_dir,
checkpoint_config,
state,
)
write_json(
output_dir / "training-state.json",
{
"bestEpoch": int(best_metrics["epoch"]),
"bestSelectionScore": best_score,
"elapsedSeconds": state["elapsedSeconds"],
"complete": False,
"paused": True,
"resumeCheckpoint": str(checkpoint),
},
)
print(
"training paused safely; resume with:\n"
f" {shlex.quote(sys.executable)} -u "
f"{shlex.quote(str(Path(__file__).resolve()))} "
f"--resume-training {shlex.quote(str(checkpoint))}",
flush=True,
)
raise TrainingPaused(checkpoint, signum)
shutdown.install()
try:
for epoch in range(start_epoch, args.epochs + 1):
if resume_state is not None and epoch == start_epoch:
permutation = (
np.asarray(resume_permutation, dtype=np.int64)
if resume_permutation is not None
else rng.permutation(len(train_records))
)
running = resume_running.copy()
first_step = resume_next_step
epoch_started = time.perf_counter() - resume_epoch_elapsed
else:
permutation = rng.permutation(len(train_records))
running = np.zeros(6, dtype=np.float64)
first_step = 1
epoch_started = time.perf_counter()
if permutation.shape != (len(train_records),):
raise DataError("training resume permutation has the wrong shape")
if shutdown.signum is not None:
pause_training(
epoch,
first_step,
permutation,
running,
time.perf_counter() - epoch_started,
)
model.train()
for step, indexes in enumerate(
_batch_indexes(
len(train_records), batch_size, permutation=permutation
),
1,
):
if step < first_step:
continue
batch = _mlx_batch(
mx, encoded_train, train_targets, sample_weights, indexes
)
teacher_logits = (
mx.array(teacher_train_logits[indexes])
if teacher_train_logits is not None
else None
)
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"],
teacher_logits,
)
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" distillation={mean[5]:.4f}"
if teacher_train_logits is not None
else ""
)
+ " "
f"elapsed={time.perf_counter() - epoch_started:.1f}s",
flush=True,
)
if shutdown.signum is not None:
pause_training(
epoch,
step + 1,
permutation,
running,
time.perf_counter() - epoch_started,
)
outputs = _evaluate(
mx, model, encoded_validation, eval_batch_size
)
metrics = multitask_metrics(outputs, validation_records)
if teacher_validation_logits is not None:
metrics["teacherAgreement"] = float(
np.mean(
outputs["purpose_logits"][validation_scorable].argmax(
axis=-1
)
== teacher_validation_logits[validation_scorable].argmax(
axis=-1
)
)
)
metrics["epoch"] = epoch
metrics["meanTrainingLoss"] = (running / steps_per_epoch).tolist()
history.append(metrics)
score = float(metrics["selectionScore"])
secondary_macro = (
metrics["secondary"]["supportedMacroRecall"]
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_supported_macro_recall={secondary_macro:.4%} "
f"mixed_f1={metrics['mixed']['f1']:.4%} "
f"difficulty_mae={metrics['difficulty']['mae']:.4f} "
f"selection_score={score:.4%}"
+ (
f" teacher_agreement={metrics['teacherAgreement']:.4%}"
if "teacherAgreement" in metrics
else ""
),
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
resume_state = None
if shutdown.signum is not None:
pause_training(
epoch + 1,
1,
None,
np.zeros(6, dtype=np.float64),
0.0,
)
finally:
shutdown.restore()
# 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,
"resumedFrom": (
str(source)
if args.resume_from is not None or args.resume_training is not None
else None
),
"exactTrainingResume": args.resume_training is not None,
"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,
},
"distillation": {
"cache": (
str(args.distillation_cache)
if args.distillation_cache is not None
else None
),
"weight": args.distillation_weight,
"temperature": args.distillation_temperature,
},
"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,
"initialValidation": initial_metrics,
"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,
},
)
for stale_resume in (
output_dir / "resume",
output_dir / ".training-resume-backup",
):
if stale_resume.exists():
shutil.rmtree(stale_resume)
return metrics
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--variant", choices=tuple(DEEP_VARIANTS), default="base")
source_group = parser.add_mutually_exclusive_group()
source_group.add_argument(
"--model",
type=Path,
help="local pinned ModernBERT checkpoint (default: download the pinned revision)",
)
source_group.add_argument(
"--resume-from",
type=Path,
help=(
"selected purpose-deep model directory to continue from; restores "
"the backbone and all four task heads with a fresh optimizer"
),
)
source_group.add_argument(
"--resume-training",
type=Path,
help=(
"exact signal-created resume checkpoint; restores the saved arguments, "
"model, optimizer, shuffle order, and next batch"
),
)
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("--distillation-cache", type=Path)
parser.add_argument("--distillation-weight", type=float, default=0.0)
parser.add_argument("--distillation-temperature", type=float, default=2.0)
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)
if args.resume_training is not None:
checkpoint = args.resume_training.expanduser().resolve()
try:
resume_state = _read_resume_state(checkpoint)
_restore_resume_arguments(args, checkpoint, resume_state)
except DataError as exc:
parser.error(str(exc))
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.distillation_weight <= 1:
parser.error("--distillation-weight must be in [0, 1]")
if args.distillation_temperature <= 0:
parser.error("--distillation-temperature must be positive")
if (args.distillation_cache is None) != (args.distillation_weight == 0):
parser.error(
"--distillation-cache and a positive --distillation-weight "
"must be supplied together"
)
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 TrainingPaused as paused:
return 128 + paused.signum
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())