#!/usr/bin/env python3 """Fine-tune purpose-lite with MLX, using Metal by default.""" 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" 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]: try: import mlx.core as mx import mlx.nn as nn import mlx.optimizers as optim except ImportError as exc: raise DataError( "MLX training requires requirements-mlx.txt" ) from exc _configure_mlx_device(mx, device) 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]: mx, nn, optim = _load_mlx(args.device) 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", "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), "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) parser.add_argument( "--device", choices=("metal", "cpu"), default="metal", help="MLX execution device (default: metal; cpu is a diagnostic fallback)", ) 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())