2026-07-31 01:24:01 -07:00
|
|
|
#!/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,
|
|
|
|
|
)
|
2026-07-31 14:30:10 -07:00
|
|
|
from train_mlx import _configure_mlx_device, _linear_schedule, _teacher_cache
|
2026-07-31 01:24:01 -07:00
|
|
|
|
|
|
|
|
|
|
|
|
|
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:
|
2026-07-31 03:05:29 -07:00
|
|
|
raise DataError("the input checkpoint must not be inside --output-dir")
|
2026-07-31 01:24:01 -07:00
|
|
|
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)
|
|
|
|
|
|
|
|
|
|
|
2026-07-31 14:30:10 -07:00
|
|
|
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)
|
|
|
|
|
|
|
|
|
|
|
2026-07-31 01:24:01 -07:00
|
|
|
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,
|
2026-07-31 03:05:29 -07:00
|
|
|
load_checkpoint_weights,
|
2026-07-31 01:24:01 -07:00
|
|
|
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]
|
2026-07-31 03:05:29 -07:00
|
|
|
source = _resolve_source(variant, args.resume_from or args.model)
|
2026-07-31 01:24:01 -07:00
|
|
|
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)
|
2026-07-31 14:30:10 -07:00
|
|
|
validation_scorable = np.asarray(
|
|
|
|
|
[record["slice"] != "vague-eval" for record in validation_records],
|
|
|
|
|
dtype=np.bool_,
|
|
|
|
|
)
|
2026-07-31 01:24:01 -07:00
|
|
|
sample_weights = _sample_weights(train_records, args.hard_weight)
|
|
|
|
|
secondary_class_weights = _secondary_class_weights(train_records)
|
2026-07-31 14:30:10 -07:00
|
|
|
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,
|
|
|
|
|
)
|
2026-07-31 01:24:01 -07:00
|
|
|
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)
|
2026-07-31 03:05:29 -07:00
|
|
|
if 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; 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,
|
|
|
|
|
)
|
2026-07-31 01:24:01 -07:00
|
|
|
|
|
|
|
|
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,
|
2026-07-31 14:30:10 -07:00
|
|
|
teacher_logits: Any | None,
|
|
|
|
|
) -> tuple[Any, Any, Any, Any, Any, Any]:
|
2026-07-31 01:24:01 -07:00
|
|
|
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",
|
|
|
|
|
)
|
2026-07-31 14:30:10 -07:00
|
|
|
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
|
|
|
|
|
)
|
2026-07-31 01:24:01 -07:00
|
|
|
|
|
|
|
|
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
|
|
|
|
|
)
|
2026-07-31 14:30:10 -07:00
|
|
|
return (
|
|
|
|
|
total,
|
|
|
|
|
primary_label_loss,
|
|
|
|
|
secondary_loss,
|
|
|
|
|
mixed_loss,
|
|
|
|
|
difficulty_loss,
|
|
|
|
|
primary_distillation_loss,
|
|
|
|
|
)
|
2026-07-31 01:24:01 -07:00
|
|
|
|
|
|
|
|
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"
|
|
|
|
|
epochs_without_improvement = 0
|
|
|
|
|
stopped_early = False
|
|
|
|
|
history: list[dict[str, Any]] = []
|
|
|
|
|
started = time.perf_counter()
|
|
|
|
|
|
2026-07-31 14:30:10 -07:00
|
|
|
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,
|
|
|
|
|
)
|
|
|
|
|
|
2026-07-31 01:24:01 -07:00
|
|
|
for epoch in range(1, args.epochs + 1):
|
|
|
|
|
epoch_started = time.perf_counter()
|
|
|
|
|
model.train()
|
2026-07-31 14:30:10 -07:00
|
|
|
running = np.zeros(6, dtype=np.float64)
|
2026-07-31 01:24:01 -07:00
|
|
|
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
|
|
|
|
|
)
|
2026-07-31 14:30:10 -07:00
|
|
|
teacher_logits = (
|
|
|
|
|
mx.array(teacher_train_logits[indexes])
|
|
|
|
|
if teacher_train_logits is not None
|
|
|
|
|
else None
|
|
|
|
|
)
|
2026-07-31 01:24:01 -07:00
|
|
|
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"],
|
2026-07-31 14:30:10 -07:00
|
|
|
teacher_logits,
|
2026-07-31 01:24:01 -07:00
|
|
|
)
|
|
|
|
|
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} "
|
2026-07-31 14:30:10 -07:00
|
|
|
f"difficulty={mean[4]:.4f}"
|
|
|
|
|
+ (
|
|
|
|
|
f" distillation={mean[5]:.4f}"
|
|
|
|
|
if teacher_train_logits is not None
|
|
|
|
|
else ""
|
|
|
|
|
)
|
|
|
|
|
+ " "
|
2026-07-31 01:24:01 -07:00
|
|
|
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)
|
2026-07-31 14:30:10 -07:00
|
|
|
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)
|
|
|
|
|
)
|
|
|
|
|
)
|
2026-07-31 01:24:01 -07:00
|
|
|
metrics["epoch"] = epoch
|
|
|
|
|
metrics["meanTrainingLoss"] = (running / steps_per_epoch).tolist()
|
|
|
|
|
history.append(metrics)
|
|
|
|
|
score = float(metrics["selectionScore"])
|
|
|
|
|
secondary_macro = (
|
2026-07-31 14:30:10 -07:00
|
|
|
metrics["secondary"]["supportedMacroRecall"]
|
2026-07-31 01:24:01 -07:00
|
|
|
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%} "
|
2026-07-31 14:30:10 -07:00
|
|
|
f"secondary_supported_macro_recall={secondary_macro:.4%} "
|
2026-07-31 01:24:01 -07:00
|
|
|
f"mixed_f1={metrics['mixed']['f1']:.4%} "
|
|
|
|
|
f"difficulty_mae={metrics['difficulty']['mae']:.4f} "
|
2026-07-31 14:30:10 -07:00
|
|
|
f"selection_score={score:.4%}"
|
|
|
|
|
+ (
|
|
|
|
|
f" teacher_agreement={metrics['teacherAgreement']:.4%}"
|
|
|
|
|
if "teacherAgreement" in metrics
|
|
|
|
|
else ""
|
|
|
|
|
),
|
2026-07-31 01:24:01 -07:00
|
|
|
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
|
|
|
|
|
|
|
|
|
|
# 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,
|
2026-07-31 03:05:29 -07:00
|
|
|
"resumedFrom": str(source) if args.resume_from is not None else None,
|
2026-07-31 01:24:01 -07:00
|
|
|
"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,
|
|
|
|
|
},
|
2026-07-31 14:30:10 -07:00
|
|
|
"distillation": {
|
|
|
|
|
"cache": (
|
|
|
|
|
str(args.distillation_cache)
|
|
|
|
|
if args.distillation_cache is not None
|
|
|
|
|
else None
|
|
|
|
|
),
|
|
|
|
|
"weight": args.distillation_weight,
|
|
|
|
|
"temperature": args.distillation_temperature,
|
|
|
|
|
},
|
2026-07-31 01:24:01 -07:00
|
|
|
"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,
|
2026-07-31 14:30:10 -07:00
|
|
|
"initialValidation": initial_metrics,
|
2026-07-31 01:24:01 -07:00
|
|
|
"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")
|
2026-07-31 03:05:29 -07:00
|
|
|
source_group = parser.add_mutually_exclusive_group()
|
|
|
|
|
source_group.add_argument(
|
2026-07-31 01:24:01 -07:00
|
|
|
"--model",
|
|
|
|
|
type=Path,
|
|
|
|
|
help="local pinned ModernBERT checkpoint (default: download the pinned revision)",
|
|
|
|
|
)
|
2026-07-31 03:05:29 -07:00
|
|
|
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"
|
|
|
|
|
),
|
|
|
|
|
)
|
2026-07-31 01:24:01 -07:00
|
|
|
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)
|
2026-07-31 14:30:10 -07:00
|
|
|
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)
|
2026-07-31 01:24:01 -07:00
|
|
|
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")
|
2026-07-31 14:30:10 -07:00
|
|
|
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"
|
|
|
|
|
)
|
2026-07-31 01:24:01 -07:00
|
|
|
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())
|