2026-07-30 20:14:37 -07:00
|
|
|
#!/usr/bin/env python3
|
2026-07-30 23:50:55 -07:00
|
|
|
"""Fine-tune purpose-lite with MLX, using Metal by default."""
|
2026-07-30 20:14:37 -07:00
|
|
|
|
|
|
|
|
from __future__ import annotations
|
|
|
|
|
|
|
|
|
|
import argparse
|
|
|
|
|
import json
|
|
|
|
|
import math
|
|
|
|
|
import random
|
|
|
|
|
import shutil
|
|
|
|
|
import sys
|
|
|
|
|
import time
|
|
|
|
|
from pathlib import Path
|
|
|
|
|
from typing import Any, Iterator, Sequence
|
|
|
|
|
|
|
|
|
|
import numpy as np
|
|
|
|
|
|
|
|
|
|
from purpose_data import LABELS, DataError, load_jsonl, write_json
|
|
|
|
|
from train import (
|
|
|
|
|
HEAD_TOKENS,
|
|
|
|
|
MAX_LENGTH,
|
|
|
|
|
TAIL_TOKENS,
|
|
|
|
|
_fit_temperature,
|
|
|
|
|
_validate_split,
|
|
|
|
|
choose_confidence_thresholds,
|
|
|
|
|
classification_metrics,
|
|
|
|
|
distillation_record_keys,
|
|
|
|
|
expected_calibration_error,
|
|
|
|
|
prepare_text,
|
|
|
|
|
training_weight,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
SCRIPT_DIR = Path(__file__).resolve().parent
|
|
|
|
|
DEFAULT_DATASET_DIR = SCRIPT_DIR / ".artifacts" / "dataset-v1"
|
|
|
|
|
DEFAULT_OUTPUT_DIR = SCRIPT_DIR / "outputs" / "purpose-lite-v1-mlx"
|
|
|
|
|
|
|
|
|
|
|
2026-07-30 23:50:55 -07:00
|
|
|
def _configure_mlx_device(mx: Any, device: str) -> None:
|
|
|
|
|
if device == "metal":
|
|
|
|
|
if not mx.metal.is_available():
|
|
|
|
|
raise DataError("MLX Metal training requires Apple Silicon")
|
|
|
|
|
mx.set_default_device(mx.gpu)
|
|
|
|
|
return
|
|
|
|
|
if device == "cpu":
|
|
|
|
|
mx.set_default_device(mx.cpu)
|
|
|
|
|
return
|
|
|
|
|
raise DataError(f"unsupported MLX device {device!r}")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _load_mlx(device: str) -> tuple[Any, Any, Any]:
|
2026-07-30 20:14:37 -07:00
|
|
|
try:
|
|
|
|
|
import mlx.core as mx
|
|
|
|
|
import mlx.nn as nn
|
|
|
|
|
import mlx.optimizers as optim
|
|
|
|
|
except ImportError as exc:
|
|
|
|
|
raise DataError(
|
2026-07-30 23:50:55 -07:00
|
|
|
"MLX training requires requirements-mlx.txt"
|
2026-07-30 20:14:37 -07:00
|
|
|
) from exc
|
2026-07-30 23:50:55 -07:00
|
|
|
_configure_mlx_device(mx, device)
|
2026-07-30 20:14:37 -07:00
|
|
|
return mx, nn, optim
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def encode_fixed_shape_numpy(
|
|
|
|
|
tokenizer: Any,
|
|
|
|
|
texts: Sequence[str],
|
|
|
|
|
) -> dict[str, np.ndarray]:
|
|
|
|
|
"""Apply the same fixed 128-token head-tail contract as train.py."""
|
|
|
|
|
|
|
|
|
|
normalized = [prepare_text(text) for text in texts]
|
|
|
|
|
raw = tokenizer(
|
|
|
|
|
normalized,
|
|
|
|
|
add_special_tokens=False,
|
|
|
|
|
padding=False,
|
|
|
|
|
truncation=False,
|
|
|
|
|
return_attention_mask=False,
|
|
|
|
|
return_token_type_ids=False,
|
|
|
|
|
verbose=False,
|
|
|
|
|
)
|
|
|
|
|
if not isinstance(raw.get("input_ids"), list):
|
|
|
|
|
raise DataError("tokenizer did not return input_ids")
|
|
|
|
|
if tokenizer.pad_token_id is None:
|
|
|
|
|
raise DataError("purpose-lite tokenizer must define a padding token")
|
|
|
|
|
if tokenizer.cls_token_id is None or tokenizer.sep_token_id is None:
|
|
|
|
|
raise DataError("purpose-lite tokenizer must define BERT CLS and SEP tokens")
|
|
|
|
|
if tokenizer.padding_side != "right":
|
|
|
|
|
raise DataError("purpose-lite tokenizer must use right padding")
|
|
|
|
|
|
|
|
|
|
input_rows: list[list[int]] = []
|
|
|
|
|
mask_rows: list[list[int]] = []
|
|
|
|
|
type_rows: list[list[int]] = []
|
|
|
|
|
include_token_types = "token_type_ids" in tokenizer.model_input_names
|
|
|
|
|
single_budget = MAX_LENGTH - tokenizer.num_special_tokens_to_add(pair=False)
|
|
|
|
|
pair_budget = MAX_LENGTH - tokenizer.num_special_tokens_to_add(pair=True)
|
|
|
|
|
if pair_budget != HEAD_TOKENS + TAIL_TOKENS:
|
|
|
|
|
raise DataError(
|
|
|
|
|
"purpose-lite tokenizer special-token layout changed; expected three "
|
|
|
|
|
"tokens for head-tail inputs"
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
for content in raw["input_ids"]:
|
|
|
|
|
if len(content) <= single_budget:
|
|
|
|
|
first = content
|
|
|
|
|
second = None
|
|
|
|
|
else:
|
|
|
|
|
first = content[:HEAD_TOKENS]
|
|
|
|
|
second = content[-TAIL_TOKENS:]
|
|
|
|
|
if second is None:
|
|
|
|
|
input_ids = [tokenizer.cls_token_id] + first + [tokenizer.sep_token_id]
|
|
|
|
|
token_types = [0] * len(input_ids)
|
|
|
|
|
else:
|
|
|
|
|
input_ids = (
|
|
|
|
|
[tokenizer.cls_token_id]
|
|
|
|
|
+ first
|
|
|
|
|
+ [tokenizer.sep_token_id]
|
|
|
|
|
+ second
|
|
|
|
|
+ [tokenizer.sep_token_id]
|
|
|
|
|
)
|
|
|
|
|
token_types = [0] * (len(first) + 2) + [1] * (len(second) + 1)
|
|
|
|
|
if len(input_ids) > MAX_LENGTH:
|
|
|
|
|
raise DataError("fixed-shape tokenizer exceeded its 128-token contract")
|
|
|
|
|
padding = MAX_LENGTH - len(input_ids)
|
|
|
|
|
input_rows.append(input_ids + [tokenizer.pad_token_id] * padding)
|
|
|
|
|
mask_rows.append([1] * len(input_ids) + [0] * padding)
|
|
|
|
|
if include_token_types:
|
|
|
|
|
type_rows.append(token_types + [0] * padding)
|
|
|
|
|
|
|
|
|
|
encoded = {
|
|
|
|
|
"input_ids": np.asarray(input_rows, dtype=np.int32),
|
|
|
|
|
"attention_mask": np.asarray(mask_rows, dtype=np.int32),
|
|
|
|
|
}
|
|
|
|
|
if include_token_types:
|
|
|
|
|
encoded["token_type_ids"] = np.asarray(type_rows, dtype=np.int32)
|
|
|
|
|
return encoded
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
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],
|
|
|
|
|
indexes: np.ndarray,
|
|
|
|
|
) -> dict[str, Any]:
|
|
|
|
|
return {key: mx.array(value[indexes]) for key, value in encoded.items()}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _evaluate(
|
|
|
|
|
mx: Any,
|
|
|
|
|
model: Any,
|
|
|
|
|
encoded: dict[str, np.ndarray],
|
|
|
|
|
labels: np.ndarray,
|
|
|
|
|
batch_size: int,
|
|
|
|
|
) -> tuple[np.ndarray, np.ndarray]:
|
|
|
|
|
model.eval()
|
|
|
|
|
logits: list[np.ndarray] = []
|
|
|
|
|
for indexes in _batch_indexes(len(labels), batch_size):
|
|
|
|
|
batch = _mlx_batch(mx, encoded, indexes)
|
|
|
|
|
output = model(**batch)
|
|
|
|
|
mx.eval(output)
|
|
|
|
|
logits.append(np.asarray(output))
|
|
|
|
|
return np.concatenate(logits), labels.copy()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _teacher_cache(
|
|
|
|
|
path: Path,
|
|
|
|
|
train_records: Sequence[dict[str, Any]],
|
|
|
|
|
validation_records: Sequence[dict[str, Any]],
|
|
|
|
|
) -> tuple[np.ndarray, np.ndarray]:
|
|
|
|
|
try:
|
|
|
|
|
import torch
|
|
|
|
|
except ImportError as exc:
|
|
|
|
|
raise DataError(
|
|
|
|
|
"loading the existing teacher cache requires PyTorch"
|
|
|
|
|
) from exc
|
|
|
|
|
if not path.is_file():
|
|
|
|
|
raise DataError(f"{path}: distillation cache is missing")
|
|
|
|
|
try:
|
|
|
|
|
cache = torch.load(path, map_location="cpu", weights_only=True)
|
|
|
|
|
if cache["schemaVersion"] != 1 or cache["labels"] != list(LABELS):
|
|
|
|
|
raise DataError("distillation cache contract does not match purpose-lite")
|
|
|
|
|
if cache["trainRecordKeys"][: len(train_records)] != distillation_record_keys(
|
|
|
|
|
train_records
|
|
|
|
|
):
|
|
|
|
|
raise DataError("distillation cache does not match the training split")
|
|
|
|
|
if cache["validationRecordKeys"][
|
|
|
|
|
: len(validation_records)
|
|
|
|
|
] != distillation_record_keys(validation_records):
|
|
|
|
|
raise DataError("distillation cache does not match the validation split")
|
|
|
|
|
train_logits = (
|
|
|
|
|
cache["trainLogits"][: len(train_records)].float().numpy().copy()
|
|
|
|
|
)
|
|
|
|
|
validation_logits = (
|
|
|
|
|
cache["validationLogits"][: len(validation_records)]
|
|
|
|
|
.float()
|
|
|
|
|
.numpy()
|
|
|
|
|
.copy()
|
|
|
|
|
)
|
|
|
|
|
except DataError:
|
|
|
|
|
raise
|
|
|
|
|
except (KeyError, TypeError, ValueError, RuntimeError) as exc:
|
|
|
|
|
raise DataError(f"{path}: cannot load distillation cache: {exc}") from exc
|
|
|
|
|
if train_logits.shape != (len(train_records), len(LABELS)):
|
|
|
|
|
raise DataError("distillation training logits have the wrong shape")
|
|
|
|
|
if validation_logits.shape != (len(validation_records), len(LABELS)):
|
|
|
|
|
raise DataError("distillation validation logits have the wrong shape")
|
|
|
|
|
return train_logits, validation_logits
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _checkpoint_config(model_dir: Path) -> dict[str, Any]:
|
|
|
|
|
config_path = model_dir / "config.json"
|
|
|
|
|
try:
|
|
|
|
|
config = json.loads(config_path.read_text(encoding="utf-8"))
|
|
|
|
|
except (OSError, json.JSONDecodeError) as exc:
|
|
|
|
|
raise DataError(f"{config_path}: cannot load model config: {exc}") from exc
|
|
|
|
|
if (
|
|
|
|
|
config.get("model_type") != "bert"
|
|
|
|
|
or config.get("hidden_size") != 384
|
|
|
|
|
or config.get("num_hidden_layers") != 6
|
|
|
|
|
or len(config.get("id2label", {})) != len(LABELS)
|
|
|
|
|
):
|
|
|
|
|
raise DataError("MLX purpose-lite requires the 6-layer 384-wide BERT classifier")
|
|
|
|
|
configured_labels = [
|
|
|
|
|
config["id2label"].get(str(index), config["id2label"].get(index))
|
|
|
|
|
for index in range(len(LABELS))
|
|
|
|
|
]
|
|
|
|
|
if configured_labels != list(LABELS):
|
|
|
|
|
raise DataError("MLX checkpoint label order does not match purpose-lite")
|
|
|
|
|
return config
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _save_checkpoint(
|
|
|
|
|
mx: Any,
|
|
|
|
|
model: Any,
|
|
|
|
|
source_dir: Path,
|
|
|
|
|
destination: Path,
|
|
|
|
|
config: dict[str, Any],
|
|
|
|
|
*,
|
|
|
|
|
quantization_aware: bool,
|
|
|
|
|
) -> None:
|
|
|
|
|
from mlx_model import save_hugging_face_weights
|
|
|
|
|
|
|
|
|
|
if destination.exists():
|
|
|
|
|
shutil.rmtree(destination)
|
|
|
|
|
destination.mkdir(parents=True)
|
|
|
|
|
for source in source_dir.iterdir():
|
|
|
|
|
if source.name.startswith("model") and source.suffix == ".safetensors":
|
|
|
|
|
continue
|
|
|
|
|
target = destination / source.name
|
|
|
|
|
if source.is_dir():
|
|
|
|
|
shutil.copytree(source, target)
|
|
|
|
|
else:
|
|
|
|
|
shutil.copy2(source, target)
|
|
|
|
|
|
|
|
|
|
output_config = dict(config)
|
|
|
|
|
output_config["purpose_classifier_training_backend"] = "mlx"
|
|
|
|
|
output_config["purpose_classifier_quantization_aware_training"] = bool(
|
|
|
|
|
quantization_aware
|
|
|
|
|
)
|
|
|
|
|
(destination / "config.json").write_text(
|
|
|
|
|
json.dumps(output_config, indent=2, sort_keys=True) + "\n",
|
|
|
|
|
encoding="utf-8",
|
|
|
|
|
)
|
|
|
|
|
save_hugging_face_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 _linear_schedule(
|
|
|
|
|
mx: Any,
|
|
|
|
|
learning_rate: float,
|
|
|
|
|
total_steps: int,
|
|
|
|
|
warmup_steps: int,
|
|
|
|
|
) -> Any:
|
|
|
|
|
def schedule(step: Any) -> Any:
|
|
|
|
|
step = step.astype(mx.float32)
|
|
|
|
|
if warmup_steps:
|
|
|
|
|
warmup = learning_rate * step / warmup_steps
|
|
|
|
|
else:
|
|
|
|
|
warmup = mx.array(learning_rate)
|
|
|
|
|
remaining = max(total_steps - warmup_steps, 1)
|
|
|
|
|
decay = learning_rate * mx.maximum(
|
|
|
|
|
0.0,
|
|
|
|
|
(total_steps - step) / remaining,
|
|
|
|
|
)
|
|
|
|
|
if warmup_steps:
|
|
|
|
|
return mx.where(step < warmup_steps, warmup, decay)
|
|
|
|
|
return decay
|
|
|
|
|
|
|
|
|
|
return schedule
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def train(args: argparse.Namespace) -> dict[str, Any]:
|
2026-07-30 23:50:55 -07:00
|
|
|
mx, nn, optim = _load_mlx(args.device)
|
2026-07-30 20:14:37 -07:00
|
|
|
try:
|
|
|
|
|
from transformers import AutoTokenizer
|
|
|
|
|
|
|
|
|
|
from mlx_model import (
|
|
|
|
|
BertClassifierConfig,
|
|
|
|
|
BertForSequenceClassification,
|
|
|
|
|
QATEmbedding,
|
|
|
|
|
QATLinear,
|
|
|
|
|
load_hugging_face_weights,
|
|
|
|
|
)
|
|
|
|
|
except ImportError as exc:
|
|
|
|
|
raise DataError(
|
|
|
|
|
"MLX training dependencies are missing; install requirements-mlx.txt"
|
|
|
|
|
) from exc
|
|
|
|
|
|
|
|
|
|
model_dir = args.model.expanduser()
|
|
|
|
|
checkpoint = model_dir / "model.safetensors"
|
|
|
|
|
if not checkpoint.is_file():
|
|
|
|
|
raise DataError("--model must be a local Hugging Face safetensors checkpoint")
|
|
|
|
|
output_dir: Path = args.output_dir
|
|
|
|
|
try:
|
|
|
|
|
model_dir.resolve().relative_to(output_dir.resolve())
|
|
|
|
|
except ValueError:
|
|
|
|
|
pass
|
|
|
|
|
else:
|
|
|
|
|
raise DataError("--model must not be inside --output-dir")
|
|
|
|
|
if output_dir.exists() and any(output_dir.iterdir()):
|
|
|
|
|
if not args.overwrite_output:
|
|
|
|
|
raise DataError(
|
|
|
|
|
f"{output_dir}: output is not empty; pass --overwrite-output intentionally"
|
|
|
|
|
)
|
|
|
|
|
shutil.rmtree(output_dir)
|
|
|
|
|
output_dir.mkdir(parents=True, exist_ok=True)
|
|
|
|
|
|
|
|
|
|
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_split(train_records, train_path)
|
|
|
|
|
_validate_split(validation_records, validation_path)
|
|
|
|
|
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]
|
|
|
|
|
|
|
|
|
|
config_json = _checkpoint_config(model_dir)
|
|
|
|
|
config = BertClassifierConfig.from_hugging_face(config_json)
|
|
|
|
|
tokenizer = AutoTokenizer.from_pretrained(model_dir, local_files_only=True)
|
|
|
|
|
print("tokenizing train and validation splits", flush=True)
|
|
|
|
|
encoded_train = _encode_records(tokenizer, train_records)
|
|
|
|
|
encoded_validation = _encode_records(tokenizer, validation_records)
|
|
|
|
|
label_to_id = {label: index for index, label in enumerate(LABELS)}
|
|
|
|
|
train_labels = np.asarray(
|
|
|
|
|
[label_to_id[record["purpose"]] for record in train_records],
|
|
|
|
|
dtype=np.int32,
|
|
|
|
|
)
|
|
|
|
|
validation_labels = np.asarray(
|
|
|
|
|
[label_to_id[record["purpose"]] for record in validation_records],
|
|
|
|
|
dtype=np.int32,
|
|
|
|
|
)
|
|
|
|
|
sample_weights = np.asarray(
|
|
|
|
|
[
|
|
|
|
|
training_weight(record, args.boundary_weight)
|
|
|
|
|
for record in train_records
|
|
|
|
|
],
|
|
|
|
|
dtype=np.float32,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
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,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
random.seed(args.seed)
|
|
|
|
|
np.random.seed(args.seed)
|
|
|
|
|
mx.random.seed(args.seed)
|
|
|
|
|
model = BertForSequenceClassification(
|
|
|
|
|
config,
|
|
|
|
|
quantization_aware=args.quantization_aware,
|
|
|
|
|
)
|
|
|
|
|
load_hugging_face_weights(model, checkpoint)
|
|
|
|
|
qat_modules = {
|
|
|
|
|
"linear": sum(isinstance(module, QATLinear) for module in model.modules()),
|
|
|
|
|
"embedding": sum(
|
|
|
|
|
isinstance(module, QATEmbedding) for module in model.modules()
|
|
|
|
|
),
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
validation_scorable = np.asarray(
|
|
|
|
|
[record.get("slice") != "vague-eval" for record in validation_records],
|
|
|
|
|
dtype=np.bool_,
|
|
|
|
|
)
|
|
|
|
|
initial_logits, _ = _evaluate(
|
|
|
|
|
mx,
|
|
|
|
|
model,
|
|
|
|
|
encoded_validation,
|
|
|
|
|
validation_labels,
|
|
|
|
|
args.eval_batch_size,
|
|
|
|
|
)
|
|
|
|
|
initial_predictions = initial_logits.argmax(axis=-1)
|
|
|
|
|
initial_metrics = classification_metrics(
|
|
|
|
|
validation_labels[validation_scorable].tolist(),
|
|
|
|
|
initial_predictions[validation_scorable].tolist(),
|
|
|
|
|
)
|
|
|
|
|
teacher_validation_predictions = (
|
|
|
|
|
teacher_validation_logits.argmax(axis=-1)
|
|
|
|
|
if teacher_validation_logits is not None
|
|
|
|
|
else None
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
def teacher_agreement(predictions: np.ndarray) -> float | None:
|
|
|
|
|
if teacher_validation_predictions is None:
|
|
|
|
|
return None
|
|
|
|
|
return float(
|
|
|
|
|
np.mean(
|
|
|
|
|
predictions[validation_scorable]
|
|
|
|
|
== teacher_validation_predictions[validation_scorable]
|
|
|
|
|
)
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
def selection_score(accuracy: float, agreement: float | None) -> float:
|
|
|
|
|
if agreement is None:
|
|
|
|
|
return accuracy
|
|
|
|
|
weight = args.distillation_selection_weight
|
|
|
|
|
return (accuracy + weight * agreement) / (1.0 + weight)
|
|
|
|
|
|
|
|
|
|
initial_agreement = teacher_agreement(initial_predictions)
|
|
|
|
|
initial_selection_score = selection_score(
|
|
|
|
|
initial_metrics["accuracy"],
|
|
|
|
|
initial_agreement,
|
|
|
|
|
)
|
|
|
|
|
if initial_agreement is not None:
|
|
|
|
|
initial_metrics["teacherAgreement"] = initial_agreement
|
|
|
|
|
initial_metrics["selectionScore"] = initial_selection_score
|
|
|
|
|
|
|
|
|
|
best_accuracy = initial_metrics["accuracy"]
|
|
|
|
|
best_selection_score = initial_selection_score
|
|
|
|
|
best_dir = output_dir / "model"
|
|
|
|
|
_save_checkpoint(
|
|
|
|
|
mx,
|
|
|
|
|
model,
|
|
|
|
|
model_dir,
|
|
|
|
|
best_dir,
|
|
|
|
|
config_json,
|
|
|
|
|
quantization_aware=args.quantization_aware,
|
|
|
|
|
)
|
|
|
|
|
print(
|
|
|
|
|
f"epoch 0: validation_accuracy={initial_metrics['accuracy']:.4%} "
|
|
|
|
|
f"macro_recall={initial_metrics['macroRecall']:.4%}"
|
|
|
|
|
+ (
|
|
|
|
|
f" teacher_agreement={initial_agreement:.4%} "
|
|
|
|
|
f"selection_score={initial_selection_score:.4%}"
|
|
|
|
|
if initial_agreement is not None
|
|
|
|
|
else ""
|
|
|
|
|
),
|
|
|
|
|
flush=True,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
steps_per_epoch = math.ceil(len(train_records) / args.batch_size)
|
|
|
|
|
total_steps = steps_per_epoch * args.epochs
|
|
|
|
|
schedule = _linear_schedule(
|
|
|
|
|
mx,
|
|
|
|
|
args.learning_rate,
|
|
|
|
|
total_steps,
|
|
|
|
|
round(total_steps * args.warmup_ratio),
|
|
|
|
|
)
|
|
|
|
|
optimizer = optim.AdamW(
|
|
|
|
|
learning_rate=schedule,
|
|
|
|
|
weight_decay=args.weight_decay,
|
|
|
|
|
bias_correction=True,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
def loss_function(
|
|
|
|
|
input_ids: Any,
|
|
|
|
|
attention_mask: Any,
|
|
|
|
|
token_type_ids: Any,
|
|
|
|
|
labels: Any,
|
|
|
|
|
weights: Any,
|
|
|
|
|
teacher_logits: Any | None,
|
|
|
|
|
) -> tuple[Any, Any, Any]:
|
|
|
|
|
logits = model(
|
|
|
|
|
input_ids=input_ids,
|
|
|
|
|
attention_mask=attention_mask,
|
|
|
|
|
token_type_ids=token_type_ids,
|
|
|
|
|
)
|
|
|
|
|
label_loss = nn.losses.cross_entropy(
|
|
|
|
|
logits,
|
|
|
|
|
labels,
|
|
|
|
|
reduction="none",
|
|
|
|
|
)
|
|
|
|
|
distillation_loss = mx.zeros_like(label_loss)
|
|
|
|
|
if teacher_logits is not None:
|
|
|
|
|
temperature = args.distillation_temperature
|
|
|
|
|
student_log_probabilities = (
|
|
|
|
|
logits / temperature
|
|
|
|
|
- mx.logsumexp(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)
|
|
|
|
|
)
|
|
|
|
|
distillation_loss = (
|
|
|
|
|
mx.sum(
|
|
|
|
|
teacher_probabilities
|
|
|
|
|
* (teacher_log_probabilities - student_log_probabilities),
|
|
|
|
|
axis=-1,
|
|
|
|
|
)
|
|
|
|
|
* temperature
|
|
|
|
|
* temperature
|
|
|
|
|
)
|
|
|
|
|
per_record_loss = (
|
|
|
|
|
(1.0 - args.distillation_weight) * label_loss
|
|
|
|
|
+ args.distillation_weight * distillation_loss
|
|
|
|
|
)
|
|
|
|
|
denominator = mx.sum(weights)
|
|
|
|
|
loss = mx.sum(per_record_loss * weights) / denominator
|
|
|
|
|
mean_label = mx.sum(label_loss * weights) / denominator
|
|
|
|
|
mean_distillation = mx.sum(distillation_loss * weights) / denominator
|
|
|
|
|
return loss, mean_label, mean_distillation
|
|
|
|
|
|
|
|
|
|
loss_and_grad = nn.value_and_grad(model, loss_function)
|
|
|
|
|
rng = np.random.default_rng(args.seed)
|
|
|
|
|
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_loss = 0.0
|
|
|
|
|
running_label_loss = 0.0
|
|
|
|
|
running_distillation_loss = 0.0
|
|
|
|
|
permutation = rng.permutation(len(train_records))
|
|
|
|
|
for step, indexes in enumerate(
|
|
|
|
|
_batch_indexes(
|
|
|
|
|
len(train_records),
|
|
|
|
|
args.batch_size,
|
|
|
|
|
permutation=permutation,
|
|
|
|
|
),
|
|
|
|
|
1,
|
|
|
|
|
):
|
|
|
|
|
batch = _mlx_batch(mx, encoded_train, indexes)
|
|
|
|
|
labels = mx.array(train_labels[indexes])
|
|
|
|
|
weights = mx.array(sample_weights[indexes])
|
|
|
|
|
teacher_logits = (
|
|
|
|
|
mx.array(teacher_train_logits[indexes])
|
|
|
|
|
if teacher_train_logits is not None
|
|
|
|
|
else None
|
|
|
|
|
)
|
|
|
|
|
(loss, label_loss, distillation_loss), gradients = loss_and_grad(
|
|
|
|
|
batch["input_ids"],
|
|
|
|
|
batch["attention_mask"],
|
|
|
|
|
batch.get("token_type_ids"),
|
|
|
|
|
labels,
|
|
|
|
|
weights,
|
|
|
|
|
teacher_logits,
|
|
|
|
|
)
|
|
|
|
|
gradients, _ = optim.clip_grad_norm(gradients, args.max_grad_norm)
|
|
|
|
|
optimizer.update(model, gradients)
|
|
|
|
|
mx.eval(
|
|
|
|
|
model.parameters(),
|
|
|
|
|
optimizer.state,
|
|
|
|
|
loss,
|
|
|
|
|
label_loss,
|
|
|
|
|
distillation_loss,
|
|
|
|
|
)
|
|
|
|
|
running_loss += float(loss.item())
|
|
|
|
|
running_label_loss += float(label_loss.item())
|
|
|
|
|
running_distillation_loss += float(distillation_loss.item())
|
|
|
|
|
if args.progress_steps and (
|
|
|
|
|
step % args.progress_steps == 0 or step == steps_per_epoch
|
|
|
|
|
):
|
|
|
|
|
print(
|
|
|
|
|
f"epoch {epoch} step {step}/{steps_per_epoch} "
|
|
|
|
|
f"mean_loss={running_loss / step:.4f} "
|
|
|
|
|
f"label_loss={running_label_loss / step:.4f} "
|
|
|
|
|
f"distill_loss={running_distillation_loss / step:.4f} "
|
|
|
|
|
f"elapsed={time.perf_counter() - epoch_started:.1f}s",
|
|
|
|
|
flush=True,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
logits, _ = _evaluate(
|
|
|
|
|
mx,
|
|
|
|
|
model,
|
|
|
|
|
encoded_validation,
|
|
|
|
|
validation_labels,
|
|
|
|
|
args.eval_batch_size,
|
|
|
|
|
)
|
|
|
|
|
predictions = logits.argmax(axis=-1)
|
|
|
|
|
metrics = classification_metrics(
|
|
|
|
|
validation_labels[validation_scorable].tolist(),
|
|
|
|
|
predictions[validation_scorable].tolist(),
|
|
|
|
|
)
|
|
|
|
|
agreement = teacher_agreement(predictions)
|
|
|
|
|
candidate_selection_score = selection_score(metrics["accuracy"], agreement)
|
|
|
|
|
if agreement is not None:
|
|
|
|
|
metrics["teacherAgreement"] = agreement
|
|
|
|
|
metrics["selectionScore"] = candidate_selection_score
|
|
|
|
|
metrics["epoch"] = epoch
|
|
|
|
|
metrics["meanTrainingLoss"] = running_loss / steps_per_epoch
|
|
|
|
|
metrics["meanLabelLoss"] = running_label_loss / steps_per_epoch
|
|
|
|
|
metrics["meanDistillationLoss"] = (
|
|
|
|
|
running_distillation_loss / steps_per_epoch
|
|
|
|
|
)
|
|
|
|
|
history.append(metrics)
|
|
|
|
|
print(
|
|
|
|
|
f"epoch {epoch}: loss={metrics['meanTrainingLoss']:.4f} "
|
|
|
|
|
f"validation_accuracy={metrics['accuracy']:.4%} "
|
|
|
|
|
f"macro_recall={metrics['macroRecall']:.4%}"
|
|
|
|
|
+ (
|
|
|
|
|
f" teacher_agreement={agreement:.4%} "
|
|
|
|
|
f"selection_score={candidate_selection_score:.4%}"
|
|
|
|
|
if agreement is not None
|
|
|
|
|
else ""
|
|
|
|
|
),
|
|
|
|
|
flush=True,
|
|
|
|
|
)
|
|
|
|
|
improvement = candidate_selection_score - best_selection_score
|
|
|
|
|
if improvement > args.minimum_improvement:
|
|
|
|
|
best_accuracy = metrics["accuracy"]
|
|
|
|
|
best_selection_score = candidate_selection_score
|
|
|
|
|
epochs_without_improvement = 0
|
|
|
|
|
_save_checkpoint(
|
|
|
|
|
mx,
|
|
|
|
|
model,
|
|
|
|
|
model_dir,
|
|
|
|
|
best_dir,
|
|
|
|
|
config_json,
|
|
|
|
|
quantization_aware=args.quantization_aware,
|
|
|
|
|
)
|
|
|
|
|
else:
|
|
|
|
|
epochs_without_improvement += 1
|
|
|
|
|
if epochs_without_improvement >= args.early_stopping_patience:
|
|
|
|
|
stopped_early = True
|
|
|
|
|
print(
|
|
|
|
|
f"early stopping after epoch {epoch}: no selection-score "
|
|
|
|
|
f"improvement greater than {args.minimum_improvement:.4%} "
|
|
|
|
|
f"for {args.early_stopping_patience} epoch(s)",
|
|
|
|
|
flush=True,
|
|
|
|
|
)
|
|
|
|
|
break
|
|
|
|
|
|
|
|
|
|
final_model = BertForSequenceClassification(
|
|
|
|
|
config,
|
|
|
|
|
quantization_aware=False,
|
|
|
|
|
)
|
|
|
|
|
load_hugging_face_weights(final_model, best_dir / "model.safetensors")
|
|
|
|
|
logits, labels = _evaluate(
|
|
|
|
|
mx,
|
|
|
|
|
final_model,
|
|
|
|
|
encoded_validation,
|
|
|
|
|
validation_labels,
|
|
|
|
|
args.eval_batch_size,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
try:
|
|
|
|
|
import torch
|
|
|
|
|
except ImportError as exc:
|
|
|
|
|
raise DataError("final calibration requires PyTorch") from exc
|
|
|
|
|
temperature = _fit_temperature(
|
|
|
|
|
torch,
|
|
|
|
|
torch.from_numpy(logits),
|
|
|
|
|
torch.from_numpy(labels.astype(np.int64)),
|
|
|
|
|
)
|
|
|
|
|
calibrated = _softmax(logits / temperature)
|
|
|
|
|
sorted_indexes = np.argsort(calibrated, axis=-1)
|
|
|
|
|
top_indexes = sorted_indexes[:, -1]
|
|
|
|
|
second_indexes = sorted_indexes[:, -2]
|
|
|
|
|
row_indexes = np.arange(len(labels))
|
|
|
|
|
top_probabilities = calibrated[row_indexes, top_indexes]
|
|
|
|
|
margins = (
|
|
|
|
|
top_probabilities - calibrated[row_indexes, second_indexes]
|
|
|
|
|
)
|
|
|
|
|
correct = (
|
|
|
|
|
(top_indexes == labels) & validation_scorable
|
|
|
|
|
).tolist()
|
|
|
|
|
thresholds = choose_confidence_thresholds(
|
|
|
|
|
top_probabilities.tolist(),
|
|
|
|
|
margins.tolist(),
|
|
|
|
|
correct,
|
|
|
|
|
high_precision=args.high_precision,
|
|
|
|
|
accepted_precision=args.accepted_precision,
|
|
|
|
|
)
|
|
|
|
|
calibration = {
|
|
|
|
|
"schemaVersion": 1,
|
|
|
|
|
"modelVersion": "purpose-lite-v1",
|
|
|
|
|
"labels": list(LABELS),
|
|
|
|
|
"temperature": temperature,
|
|
|
|
|
"confidence": thresholds,
|
|
|
|
|
"validationECE": expected_calibration_error(
|
|
|
|
|
top_probabilities.tolist(),
|
|
|
|
|
correct,
|
|
|
|
|
),
|
|
|
|
|
}
|
|
|
|
|
metrics = {
|
|
|
|
|
"modelVersion": "purpose-lite-v1",
|
|
|
|
|
"baseModel": str(model_dir),
|
|
|
|
|
"baseModelRevision": "local-checkpoint",
|
|
|
|
|
"trainingBackend": "mlx",
|
2026-07-30 23:50:55 -07:00
|
|
|
"device": args.device,
|
2026-07-30 20:14:37 -07:00
|
|
|
"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),
|
|
|
|
|
"scoredValidationRecords": int(validation_scorable.sum()),
|
|
|
|
|
"vagueAbstentionValidationRecords": int(
|
|
|
|
|
(~validation_scorable).sum()
|
|
|
|
|
),
|
|
|
|
|
"boundaryTrainingWeight": args.boundary_weight,
|
|
|
|
|
"quantizationAwareTraining": args.quantization_aware,
|
|
|
|
|
"quantizationAwareModules": qat_modules,
|
|
|
|
|
"distillation": {
|
|
|
|
|
"cache": (
|
|
|
|
|
str(args.distillation_cache)
|
|
|
|
|
if args.distillation_cache is not None
|
|
|
|
|
else None
|
|
|
|
|
),
|
|
|
|
|
"weight": args.distillation_weight,
|
|
|
|
|
"temperature": args.distillation_temperature,
|
|
|
|
|
"selectionAgreementWeight": args.distillation_selection_weight,
|
|
|
|
|
},
|
|
|
|
|
"bestValidationAccuracy": best_accuracy,
|
|
|
|
|
"bestValidationSelectionScore": best_selection_score,
|
|
|
|
|
"initialValidation": initial_metrics,
|
|
|
|
|
"epochsCompleted": len(history),
|
|
|
|
|
"stoppedEarly": stopped_early,
|
|
|
|
|
"bestValidation": classification_metrics(
|
|
|
|
|
labels[validation_scorable].tolist(),
|
|
|
|
|
top_indexes[validation_scorable].tolist(),
|
|
|
|
|
),
|
|
|
|
|
"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()
|
|
|
|
|
},
|
|
|
|
|
)
|
|
|
|
|
return metrics
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def build_parser() -> argparse.ArgumentParser:
|
|
|
|
|
parser = argparse.ArgumentParser(description=__doc__)
|
|
|
|
|
parser.add_argument("--dataset-dir", type=Path, default=DEFAULT_DATASET_DIR)
|
|
|
|
|
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR)
|
|
|
|
|
parser.add_argument("--model", type=Path, required=True)
|
2026-07-30 23:50:55 -07:00
|
|
|
parser.add_argument(
|
|
|
|
|
"--device",
|
|
|
|
|
choices=("metal", "cpu"),
|
|
|
|
|
default="metal",
|
|
|
|
|
help="MLX execution device (default: metal; cpu is a diagnostic fallback)",
|
|
|
|
|
)
|
2026-07-30 20:14:37 -07:00
|
|
|
parser.add_argument("--seed", type=int, default=20260730)
|
|
|
|
|
parser.add_argument("--epochs", type=int, default=3)
|
|
|
|
|
parser.add_argument("--batch-size", type=int, default=32)
|
|
|
|
|
parser.add_argument("--eval-batch-size", type=int, default=64)
|
|
|
|
|
parser.add_argument("--learning-rate", type=float, default=2e-5)
|
|
|
|
|
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("--progress-steps", type=int, default=50)
|
|
|
|
|
parser.add_argument("--early-stopping-patience", type=int, default=2)
|
|
|
|
|
parser.add_argument("--minimum-improvement", type=float, default=0.0005)
|
|
|
|
|
parser.add_argument("--boundary-weight", type=float, default=1.0)
|
|
|
|
|
parser.add_argument("--quantization-aware", action="store_true")
|
|
|
|
|
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("--distillation-selection-weight", type=float, default=0.0)
|
|
|
|
|
parser.add_argument("--high-precision", type=float, default=0.98)
|
|
|
|
|
parser.add_argument("--accepted-precision", type=float, default=0.95)
|
|
|
|
|
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 main(argv: Sequence[str] | None = None) -> int:
|
|
|
|
|
parser = build_parser()
|
|
|
|
|
args = parser.parse_args(argv)
|
|
|
|
|
for name in (
|
|
|
|
|
"epochs",
|
|
|
|
|
"batch_size",
|
|
|
|
|
"eval_batch_size",
|
|
|
|
|
"early_stopping_patience",
|
|
|
|
|
):
|
|
|
|
|
if getattr(args, name) <= 0:
|
|
|
|
|
parser.error(f"--{name.replace('_', '-')} must be positive")
|
|
|
|
|
if args.progress_steps < 0:
|
|
|
|
|
parser.error("--progress-steps must be non-negative")
|
|
|
|
|
if args.learning_rate <= 0:
|
|
|
|
|
parser.error("--learning-rate must be positive")
|
|
|
|
|
if not 0 <= args.warmup_ratio < 1:
|
|
|
|
|
parser.error("--warmup-ratio must be in [0, 1)")
|
|
|
|
|
if args.boundary_weight <= 0:
|
|
|
|
|
parser.error("--boundary-weight must be positive")
|
|
|
|
|
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 not 0 <= args.distillation_selection_weight <= 1:
|
|
|
|
|
parser.error("--distillation-selection-weight must be in [0, 1]")
|
|
|
|
|
if (args.distillation_cache is None) != (args.distillation_weight == 0):
|
|
|
|
|
parser.error(
|
|
|
|
|
"--distillation-cache and a positive --distillation-weight "
|
|
|
|
|
"must be supplied together"
|
|
|
|
|
)
|
|
|
|
|
if args.distillation_selection_weight and args.distillation_cache is None:
|
|
|
|
|
parser.error(
|
|
|
|
|
"--distillation-selection-weight requires --distillation-cache"
|
|
|
|
|
)
|
|
|
|
|
try:
|
|
|
|
|
metrics = train(args)
|
|
|
|
|
except DataError as exc:
|
|
|
|
|
print(f"error: {exc}", file=sys.stderr)
|
|
|
|
|
return 2
|
|
|
|
|
print(
|
|
|
|
|
f"selected validation accuracy: {metrics['bestValidationAccuracy']:.4%}",
|
|
|
|
|
flush=True,
|
|
|
|
|
)
|
|
|
|
|
return 0
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
if __name__ == "__main__":
|
|
|
|
|
raise SystemExit(main())
|