#!/usr/bin/env python3 """Fine-tune the fixed-shape purpose-lite MiniLM classifier.""" from __future__ import annotations import argparse import math import random import shutil import sys import time from pathlib import Path from typing import Any, Sequence from purpose_data import ( LABELS, DataError, load_jsonl, normalize_prompt, prompt_hash, write_json, ) SCRIPT_DIR = Path(__file__).resolve().parent DEFAULT_DATASET_DIR = SCRIPT_DIR / ".artifacts" / "dataset-v1" DEFAULT_OUTPUT_DIR = SCRIPT_DIR / "outputs" / "purpose-lite-v1" DEFAULT_MODEL = "sentence-transformers/all-MiniLM-L6-v2" # Reproducibility requires a model commit, not a mutable `main` branch. DEFAULT_MODEL_REVISION = "1110a243fdf4706b3f48f1d95db1a4f5529b4d41" MAX_LENGTH = 128 HEAD_TAIL_SPECIAL_TOKENS = 3 HEAD_TOKENS = (MAX_LENGTH - HEAD_TAIL_SPECIAL_TOKENS + 1) // 2 TAIL_TOKENS = MAX_LENGTH - HEAD_TAIL_SPECIAL_TOKENS - HEAD_TOKENS def prepare_text(prompt: str) -> str: return normalize_prompt(prompt) def training_weight(record: dict[str, Any], boundary_weight: float) -> float: """Return the loss weight for one training record.""" return boundary_weight if record.get("slice") == "boundary" else 1.0 def knowledge_distillation_loss( torch: Any, student_logits: Any, teacher_logits: Any, *, temperature: float, ) -> Any: """Return per-record KL loss from a frozen float teacher to the student.""" student_log_probabilities = torch.nn.functional.log_softmax( student_logits / temperature, dim=-1, ) teacher_probabilities = torch.nn.functional.softmax( teacher_logits / temperature, dim=-1, ) return ( torch.nn.functional.kl_div( student_log_probabilities, teacher_probabilities, reduction="none", ).sum(dim=-1) * temperature * temperature ) def distillation_record_keys(records: Sequence[dict[str, Any]]) -> list[str]: """Bind cached teacher logits to both normalized prompt and expected label.""" return [ f"{prompt_hash(record['prompt'])}:{record['purpose']}" for record in records ] def enable_quantization_aware_training(torch: Any, model: Any) -> dict[str, int]: """Mirror the export graph's int8 policy with straight-through fake quantization. ONNX Runtime emits per-tensor uint8 embedding weights, per-channel symmetric int8 linear weights, and per-tensor uint8 activations. The replacement modules keep the original parameter names, so the selected checkpoint loads as an ordinary Transformers model for export after fake-quantization-aware fine-tuning. """ functional = torch.nn.functional def affine_parameters(value: Any) -> tuple[float, int]: detached = value.detach().float() # ONNX Runtime extends affine calibration ranges to include exact zero. minimum = min(0.0, float(detached.amin().item())) maximum = max(0.0, float(detached.amax().item())) scale = max((maximum - minimum) / 255.0, torch.finfo(torch.float32).eps) zero_point = max(0, min(255, round(-minimum / scale))) return scale, zero_point def fake_quantize_activation(value: Any) -> Any: scale, zero_point = affine_parameters(value) return torch.fake_quantize_per_tensor_affine( value, scale, zero_point, 0, 255, ) def fake_quantize_linear_weight(weight: Any) -> Any: detached = weight.detach().float() scales = detached.abs().amax(dim=1).div(127.0).clamp_min( torch.finfo(torch.float32).eps ) zero_points = torch.zeros_like(scales, dtype=torch.int32) return torch.fake_quantize_per_channel_affine( weight, scales, zero_points, 0, -127, 127, ) class QATLinear(torch.nn.Linear): def forward(self, value: Any) -> Any: result = functional.linear( fake_quantize_activation(value), fake_quantize_linear_weight(self.weight), self.bias, ) return fake_quantize_activation(result) class QATEmbedding(torch.nn.Embedding): def forward(self, indexes: Any) -> Any: embedded = functional.embedding( indexes, self.weight, self.padding_idx, self.max_norm, self.norm_type, self.scale_grad_by_freq, self.sparse, ) # Quantize only the selected rows using the full table's scale. This is # numerically equivalent to dequantizing the whole table before Gather but # avoids materializing a 30k x 384 fake-quantized embedding every batch. weight_scale, weight_zero_point = affine_parameters(self.weight) embedded = torch.fake_quantize_per_tensor_affine( embedded, weight_scale, weight_zero_point, 0, 255, ) return fake_quantize_activation(embedded) counts = {"linear": 0, "embedding": 0} def replace(parent: Any) -> None: for name, child in list(parent.named_children()): replacement = None if isinstance(child, torch.nn.Linear): replacement = QATLinear( child.in_features, child.out_features, bias=child.bias is not None, device=child.weight.device, dtype=child.weight.dtype, ) counts["linear"] += 1 elif isinstance(child, torch.nn.Embedding): replacement = QATEmbedding( child.num_embeddings, child.embedding_dim, padding_idx=child.padding_idx, max_norm=child.max_norm, norm_type=child.norm_type, scale_grad_by_freq=child.scale_grad_by_freq, sparse=child.sparse, device=child.weight.device, dtype=child.weight.dtype, ) counts["embedding"] += 1 if replacement is not None: replacement.weight = child.weight if isinstance(child, torch.nn.Linear): replacement.bias = child.bias setattr(parent, name, replacement) else: replace(child) replace(model) return counts def encode_fixed_shape( tokenizer: Any, texts: Sequence[str], torch: Any, ) -> dict[str, Any]: """Tokenize to 1x128 while retaining both context and a tail-buried request. Pasted logs and stack traces frequently put the actual ask after the context. Plain right truncation made generated boundary examples identical even when their final request — and therefore their label — differed. Long inputs use BERT's sentence-pair framing: [CLS] first 63 content tokens [SEP] last 62 content tokens [SEP]. """ 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": torch.tensor(input_rows, dtype=torch.long), "attention_mask": torch.tensor(mask_rows, dtype=torch.long), } if include_token_types: encoded["token_type_ids"] = torch.tensor(type_rows, dtype=torch.long) return encoded def classification_metrics( actual: Sequence[int], predicted: Sequence[int] ) -> dict[str, Any]: if len(actual) != len(predicted) or not actual: raise ValueError("metrics need equally sized, non-empty vectors") correct = sum(want == got for want, got in zip(actual, predicted)) recalls: dict[str, float] = {} confusion = [[0 for _ in LABELS] for _ in LABELS] for want, got in zip(actual, predicted): confusion[want][got] += 1 for index, label in enumerate(LABELS): total = sum(confusion[index]) recalls[label] = confusion[index][index] / total if total else 0.0 return { "records": len(actual), "accuracy": correct / len(actual), "macroRecall": sum(recalls.values()) / len(recalls), "perPurposeRecall": recalls, "confusionMatrix": { "labels": list(LABELS), "rows": confusion, }, } def confidence_score(top_probability: float, top_two_margin: float) -> float: """One monotonic score that keeps both calibration signals in the contract.""" return top_probability * (0.5 + 0.5 * top_two_margin) def _threshold_for_precision( scores: Sequence[float], correct: Sequence[bool], target_precision: float, ) -> tuple[float, float, float]: ranked = sorted(zip(scores, correct), key=lambda item: item[0], reverse=True) accepted = 0 accepted_correct = 0 best: tuple[float, float, float] | None = None index = 0 while index < len(ranked): score = ranked[index][0] while index < len(ranked) and ranked[index][0] == score: accepted += 1 accepted_correct += int(ranked[index][1]) index += 1 precision = accepted_correct / accepted if precision >= target_precision: best = (score, precision, accepted / len(ranked)) if best is None: return 1.000001, 1.0, 0.0 return best def choose_confidence_thresholds( top_probabilities: Sequence[float], top_two_margins: Sequence[float], correct: Sequence[bool], *, high_precision: float = 0.98, accepted_precision: float = 0.95, ) -> dict[str, Any]: if not ( len(top_probabilities) == len(top_two_margins) == len(correct) and top_probabilities ): raise ValueError("threshold calibration needs equally sized, non-empty vectors") scores = [ confidence_score(probability, margin) for probability, margin in zip(top_probabilities, top_two_margins) ] high = _threshold_for_precision(scores, correct, high_precision) medium = _threshold_for_precision(scores, correct, accepted_precision) # HIGH must always be a subset of the accepted MEDIUM-or-better population. high_threshold = max(high[0], medium[0]) return { "score": { "formula": "topProbability * (0.5 + 0.5 * topTwoMargin)", "probabilityWeight": 0.5, "marginInteractionWeight": 0.5, }, "high": { "minimumScore": high_threshold, "targetPrecision": high_precision, "validationPrecision": high[1], "validationCoverage": high[2] if high_threshold == high[0] else 0.0, }, "medium": { "minimumScore": medium[0], "targetAcceptedPrecision": accepted_precision, "validationAcceptedPrecision": medium[1], "validationAcceptedCoverage": medium[2], }, "low": {"minimumScore": 0.0}, } def expected_calibration_error( probabilities: Sequence[float], correct: Sequence[bool], bins: int = 15, ) -> float: if len(probabilities) != len(correct) or not probabilities: raise ValueError("ECE needs equally sized, non-empty vectors") total_error = 0.0 for lower_index in range(bins): lower = lower_index / bins upper = (lower_index + 1) / bins members = [ index for index, value in enumerate(probabilities) if lower <= value < upper or (upper == 1.0 and value == 1.0) ] if not members: continue confidence = sum(probabilities[index] for index in members) / len(members) accuracy = sum(correct[index] for index in members) / len(members) total_error += len(members) / len(probabilities) * abs(confidence - accuracy) return total_error def _validate_split(records: Sequence[dict[str, Any]], path: Path) -> None: if not records: raise DataError(f"{path}: split is empty") for index, record in enumerate(records, 1): if record.get("purpose") not in LABELS: raise DataError(f"{path}:{index}: invalid purpose") if not isinstance(record.get("prompt"), str) or not record["prompt"].strip(): raise DataError(f"{path}:{index}: invalid prompt") def _select_device(torch: Any, requested: str) -> Any: if requested != "auto": return torch.device(requested) if torch.cuda.is_available(): return torch.device("cuda") if hasattr(torch.backends, "mps") and torch.backends.mps.is_available(): return torch.device("mps") return torch.device("cpu") def _set_seeds(torch: Any, seed: int) -> None: random.seed(seed) torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) def _fit_temperature(torch: Any, logits: Any, labels: Any) -> float: log_temperature = torch.zeros(1, requires_grad=True) optimizer = torch.optim.LBFGS( [log_temperature], lr=0.05, max_iter=100, line_search_fn="strong_wolfe" ) def closure() -> Any: optimizer.zero_grad() temperature = log_temperature.exp().clamp(0.05, 20.0) loss = torch.nn.functional.cross_entropy(logits / temperature, labels) loss.backward() return loss optimizer.step(closure) return float(log_temperature.detach().exp().clamp(0.05, 20.0).item()) def _evaluate( torch: Any, model: Any, loader: Any, device: Any, *, progress_label: str | None = None, progress_steps: int = 0, ) -> tuple[Any, Any]: model.eval() all_logits = [] all_labels = [] started = time.perf_counter() with torch.inference_mode(): for step, batch in enumerate(loader, 1): labels = batch.pop("labels") batch.pop("sample_weights", None) batch.pop("teacher_logits", None) inputs = {key: value.to(device) for key, value in batch.items()} logits = model(**inputs).logits.cpu() all_logits.append(logits) all_labels.append(labels) if progress_label and progress_steps and ( step % progress_steps == 0 or step == len(loader) ): print( f"{progress_label} step {step}/{len(loader)} " f"elapsed={time.perf_counter() - started:.1f}s", flush=True, ) return torch.cat(all_logits), torch.cat(all_labels) def train(args: argparse.Namespace) -> dict[str, Any]: try: import torch from torch.utils.data import DataLoader, Dataset from transformers import ( AutoModelForSequenceClassification, AutoTokenizer, get_linear_schedule_with_warmup, ) except ImportError as exc: raise DataError( "training dependencies are missing; install requirements.txt in a virtualenv" ) from exc 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] output_dir: Path = args.output_dir local_model = Path(args.model).expanduser() distillation_cache = ( args.distillation_cache.expanduser() if args.distillation_cache is not None else None ) for option_name, local_path in ( ("--model", local_model if local_model.exists() else None), ("--distillation-cache", distillation_cache), ): if local_path is None: continue try: local_path.resolve().relative_to(output_dir.resolve()) except ValueError: pass else: raise DataError( f"local {option_name} must not be inside --output-dir; overwrite " "could destroy the continuation checkpoint" ) 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) _set_seeds(torch, args.seed) device = _select_device(torch, args.device) label_to_id = {label: index for index, label in enumerate(LABELS)} id_to_label = {index: label for label, index in label_to_id.items()} model_revision = None if local_model.exists() else args.model_revision pretrained_options = ( {"local_files_only": True} if local_model.exists() else {"revision": model_revision} ) tokenizer = AutoTokenizer.from_pretrained( args.model, use_fast=True, **pretrained_options ) model = AutoModelForSequenceClassification.from_pretrained( args.model, num_labels=len(LABELS), label2id=label_to_id, id2label=id_to_label, ignore_mismatched_sizes=True, **pretrained_options, ) config = model.config if getattr(config, "hidden_size", None) != 384 or getattr( config, "num_hidden_layers", None ) != 6: raise DataError( "purpose-lite must remain a 6-layer, 384-dimensional MiniLM encoder" ) config.purpose_classifier_version = "purpose-lite-v1" config.purpose_classifier_max_length = MAX_LENGTH config.purpose_classifier_fixed_shape = [1, MAX_LENGTH] config.purpose_classifier_truncation = { "strategy": "head-tail-pair", "headTokens": HEAD_TOKENS, "tailTokens": TAIL_TOKENS, } config.purpose_classifier_quantization_aware_training = bool( args.quantization_aware ) qat_modules = {"linear": 0, "embedding": 0} if args.quantization_aware: qat_modules = enable_quantization_aware_training(torch, model) model.to(device) class PromptDataset(Dataset): def __init__( self, records: Sequence[dict[str, Any]], teacher_logits: Any | None = None, ) -> None: self.records = records self.teacher_logits = teacher_logits def __len__(self) -> int: return len(self.records) def __getitem__(self, index: int) -> tuple[str, int, float, Any | None]: record = self.records[index] return ( prepare_text(record["prompt"]), label_to_id[record["purpose"]], training_weight(record, args.boundary_weight), ( self.teacher_logits[index] if self.teacher_logits is not None else None ), ) def collate( items: Sequence[tuple[str, int, float, Any | None]], ) -> dict[str, Any]: texts, labels, weights, teacher_rows = zip(*items) encoded = encode_fixed_shape(tokenizer, list(texts), torch) encoded["labels"] = torch.tensor(labels, dtype=torch.long) encoded["sample_weights"] = torch.tensor(weights, dtype=torch.float) if teacher_rows[0] is not None: encoded["teacher_logits"] = torch.stack(teacher_rows) return encoded teacher_train_logits = None teacher_validation_logits = None if distillation_cache is not None: if not distillation_cache.is_file(): raise DataError(f"{distillation_cache}: distillation cache is missing") try: cache = torch.load( distillation_cache, 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") teacher_train_logits = cache["trainLogits"][ : len(train_records) ].float().clone() teacher_validation_logits = cache["validationLogits"][ : len(validation_records) ].float().clone() except DataError: raise except (KeyError, TypeError, ValueError, RuntimeError) as exc: raise DataError( f"{distillation_cache}: cannot load distillation cache: {exc}" ) from exc expected_shape = (len(train_records), len(LABELS)) if tuple(teacher_train_logits.shape) != expected_shape: raise DataError("distillation training logits have the wrong shape") if tuple(teacher_validation_logits.shape) != ( len(validation_records), len(LABELS), ): raise DataError("distillation validation logits have the wrong shape") del cache generator = torch.Generator() generator.manual_seed(args.seed) train_loader = DataLoader( PromptDataset(train_records, teacher_train_logits), batch_size=args.batch_size, shuffle=True, generator=generator, collate_fn=collate, num_workers=args.workers, pin_memory=device.type == "cuda", ) validation_loader = DataLoader( PromptDataset(validation_records), batch_size=args.eval_batch_size, shuffle=False, collate_fn=collate, num_workers=args.workers, pin_memory=device.type == "cuda", ) validation_scorable = torch.tensor( [record.get("slice") != "vague-eval" for record in validation_records], dtype=torch.bool, ) initial_logits, initial_labels = _evaluate( torch, model, validation_loader, device, ) initial_predictions = initial_logits.argmax(dim=-1).tolist() initial_metrics = classification_metrics( initial_labels[validation_scorable].tolist(), [ prediction for prediction, scorable in zip( initial_predictions, validation_scorable.tolist(), ) if scorable ], ) teacher_validation_predictions = ( teacher_validation_logits.argmax(dim=-1).tolist() if teacher_validation_logits is not None else None ) def teacher_agreement(predictions: Sequence[int]) -> float | None: if teacher_validation_predictions is None: return None agreements = [ prediction == teacher_prediction for prediction, teacher_prediction, scorable in zip( predictions, teacher_validation_predictions, validation_scorable.tolist(), ) if scorable ] return sum(agreements) / len(agreements) 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 epochs_without_improvement = 0 stopped_early = False history = [] best_dir = output_dir / "model" model.save_pretrained(best_dir, safe_serialization=True) tokenizer.save_pretrained(best_dir) 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, ) optimizer = torch.optim.AdamW( model.parameters(), lr=args.learning_rate, weight_decay=args.weight_decay ) update_steps_per_epoch = math.ceil( len(train_loader) / args.gradient_accumulation_steps ) total_steps = update_steps_per_epoch * args.epochs scheduler = get_linear_schedule_with_warmup( optimizer, num_warmup_steps=round(total_steps * args.warmup_ratio), num_training_steps=total_steps, ) started = time.perf_counter() for epoch in range(1, args.epochs + 1): epoch_started = time.perf_counter() model.train() optimizer.zero_grad(set_to_none=True) running_loss = 0.0 running_label_loss = 0.0 running_distillation_loss = 0.0 for step, batch in enumerate(train_loader, 1): labels = batch.pop("labels").to(device) sample_weights = batch.pop("sample_weights").to(device) teacher_logits = batch.pop("teacher_logits", None) if teacher_logits is not None: teacher_logits = teacher_logits.to(device) inputs = {key: value.to(device) for key, value in batch.items()} student_logits = model(**inputs).logits label_loss = torch.nn.functional.cross_entropy( student_logits, labels, reduction="none", ) distillation_loss = torch.zeros_like(label_loss) if teacher_logits is not None: distillation_loss = knowledge_distillation_loss( torch, student_logits, teacher_logits, temperature=args.distillation_temperature, ) per_record_loss = ( (1.0 - args.distillation_weight) * label_loss + args.distillation_weight * distillation_loss ) loss = ( (per_record_loss * sample_weights).sum() / sample_weights.sum() ) / args.gradient_accumulation_steps loss.backward() running_loss += float(loss.item()) * args.gradient_accumulation_steps running_label_loss += float( (label_loss * sample_weights).sum().item() / sample_weights.sum().item() ) running_distillation_loss += float( (distillation_loss * sample_weights).sum().item() / sample_weights.sum().item() ) should_update = ( step % args.gradient_accumulation_steps == 0 or step == len(train_loader) ) if should_update: torch.nn.utils.clip_grad_norm_(model.parameters(), args.max_grad_norm) optimizer.step() scheduler.step() optimizer.zero_grad(set_to_none=True) if args.progress_steps and ( step % args.progress_steps == 0 or step == len(train_loader) ): print( f"epoch {epoch} step {step}/{len(train_loader)} " 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, labels = _evaluate(torch, model, validation_loader, device) predictions = logits.argmax(dim=-1).tolist() scored_labels = labels[validation_scorable].tolist() scored_predictions = [ prediction for prediction, scorable in zip( predictions, validation_scorable.tolist() ) if scorable ] metrics = classification_metrics(scored_labels, scored_predictions) 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 / len(train_loader) metrics["meanLabelLoss"] = running_label_loss / len(train_loader) metrics["meanDistillationLoss"] = ( running_distillation_loss / len(train_loader) ) 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 model.save_pretrained(best_dir, safe_serialization=True) tokenizer.save_pretrained(best_dir) 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 improvement " f"greater than {args.minimum_improvement:.4%} for " f"{args.early_stopping_patience} epoch(s)", flush=True, ) break model = AutoModelForSequenceClassification.from_pretrained(best_dir).to(device) logits, labels = _evaluate(torch, model, validation_loader, device) temperature = _fit_temperature( torch, logits[validation_scorable], labels[validation_scorable], ) calibrated = torch.softmax(logits / temperature, dim=-1) top = torch.topk(calibrated, k=2, dim=-1) top_probabilities = top.values[:, 0].tolist() margins = (top.values[:, 0] - top.values[:, 1]).tolist() predictions = top.indices[:, 0].tolist() correct = [ prediction == actual and scorable for prediction, actual, scorable in zip( predictions, labels.tolist(), validation_scorable.tolist(), ) ] thresholds = choose_confidence_thresholds( top_probabilities, margins, 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, correct), } metrics = { "modelVersion": "purpose-lite-v1", "baseModel": args.model, "baseModelRevision": model_revision or "local-checkpoint", "fixedInputShape": [1, MAX_LENGTH], "truncation": { "strategy": "head-tail-pair", "headTokens": HEAD_TOKENS, "tailTokens": TAIL_TOKENS, }, "device": str(device), "trainingSeconds": time.perf_counter() - started, "trainRecords": len(train_records), "boundaryTrainingWeight": args.boundary_weight, "quantizationAwareTraining": args.quantization_aware, "quantizationAwareModules": qat_modules, "distillation": { "cache": str(distillation_cache) if distillation_cache is not None else None, "weight": args.distillation_weight, "temperature": args.distillation_temperature, "selectionAgreementWeight": args.distillation_selection_weight, }, "validationRecords": len(validation_records), "scoredValidationRecords": int(validation_scorable.sum().item()), "vagueAbstentionValidationRecords": int( (~validation_scorable).sum().item() ), "bestValidationAccuracy": best_accuracy, "bestValidationSelectionScore": best_selection_score, "initialValidation": initial_metrics, "epochsCompleted": len(history), "stoppedEarly": stopped_early, "bestValidation": classification_metrics( labels[validation_scorable].tolist(), [ prediction for prediction, scorable in zip( predictions, validation_scorable.tolist() ) if scorable ], ), "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", default=DEFAULT_MODEL) parser.add_argument("--model-revision", default=DEFAULT_MODEL_REVISION) parser.add_argument("--device", default="auto") 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("--gradient-accumulation-steps", type=int, default=1) 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("--workers", type=int, default=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 _positive(parser: argparse.ArgumentParser, name: str, value: int) -> None: if value <= 0: parser.error(f"{name} 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", "gradient_accumulation_steps", "early_stopping_patience", ): _positive(parser, f"--{name.replace('_', '-')}", getattr(args, name)) if not 0.0 <= args.warmup_ratio < 1.0: parser.error("--warmup-ratio must be in [0, 1)") if args.minimum_improvement < 0.0: parser.error("--minimum-improvement must be non-negative") if args.progress_steps < 0: parser.error("--progress-steps must be non-negative") if args.boundary_weight <= 0.0: parser.error("--boundary-weight must be positive") if not 0.0 <= args.distillation_weight <= 1.0: parser.error("--distillation-weight must be in [0, 1]") if args.distillation_temperature <= 0.0: parser.error("--distillation-temperature must be positive") if not 0.0 <= args.distillation_selection_weight <= 1.0: parser.error("--distillation-selection-weight must be in [0, 1]") if (args.distillation_cache is None) != (args.distillation_weight == 0.0): parser.error( "--distillation-cache and a positive --distillation-weight " "must be supplied together" ) if ( args.distillation_cache is None and args.distillation_selection_weight != 0.0 ): parser.error( "--distillation-selection-weight requires --distillation-cache" ) if not 0.0 < args.accepted_precision <= args.high_precision <= 1.0: parser.error( "precision targets must satisfy 0 < accepted <= high <= 1" ) try: metrics = train(args) except (DataError, OSError, ValueError) as exc: print(f"error: {exc}", file=sys.stderr) return 1 print( f"Saved purpose-lite-v1; best validation accuracy " f"{metrics['bestValidationAccuracy']:.4%}." ) return 0 if __name__ == "__main__": raise SystemExit(main())