#!/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, 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 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) -> tuple[Any, Any]: model.eval() all_logits = [] all_labels = [] with torch.inference_mode(): for batch in loader: labels = batch.pop("labels") inputs = {key: value.to(device) for key, value in batch.items()} logits = model(**inputs).logits.cpu() all_logits.append(logits) all_labels.append(labels) 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 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()} tokenizer = AutoTokenizer.from_pretrained( args.model, revision=args.model_revision, use_fast=True ) model = AutoModelForSequenceClassification.from_pretrained( args.model, revision=args.model_revision, num_labels=len(LABELS), label2id=label_to_id, id2label=id_to_label, ignore_mismatched_sizes=True, ) 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, } model.to(device) class PromptDataset(Dataset): def __init__(self, records: Sequence[dict[str, Any]]) -> None: self.records = records def __len__(self) -> int: return len(self.records) def __getitem__(self, index: int) -> tuple[str, int]: record = self.records[index] return prepare_text(record["prompt"]), label_to_id[record["purpose"]] def collate(items: Sequence[tuple[str, int]]) -> dict[str, Any]: texts, labels = zip(*items) encoded = encode_fixed_shape(tokenizer, list(texts), torch) encoded["labels"] = torch.tensor(labels, dtype=torch.long) return encoded generator = torch.Generator() generator.manual_seed(args.seed) train_loader = DataLoader( PromptDataset(train_records), 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, ) 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, ) best_accuracy = -1.0 history = [] best_dir = output_dir / "model" started = time.perf_counter() for epoch in range(1, args.epochs + 1): model.train() optimizer.zero_grad(set_to_none=True) running_loss = 0.0 for step, batch in enumerate(train_loader, 1): batch = {key: value.to(device) for key, value in batch.items()} loss = model(**batch).loss / args.gradient_accumulation_steps loss.backward() running_loss += float(loss.item()) * args.gradient_accumulation_steps 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) 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) metrics["epoch"] = epoch metrics["meanTrainingLoss"] = running_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%}", flush=True, ) if metrics["accuracy"] > best_accuracy: best_accuracy = metrics["accuracy"] model.save_pretrained(best_dir, safe_serialization=True) tokenizer.save_pretrained(best_dir) 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": args.model_revision, "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), "validationRecords": len(validation_records), "scoredValidationRecords": int(validation_scorable.sum().item()), "vagueAbstentionValidationRecords": int( (~validation_scorable).sum().item() ), "bestValidationAccuracy": best_accuracy, "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("--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", ): _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 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())