diff --git a/README.md b/README.md index a3913ac..a9ad53f 100644 --- a/README.md +++ b/README.md @@ -72,7 +72,23 @@ platform selector, then install `requirements-base.txt`. Training writes a local checkpoint, `calibration.json`, and `metrics.json` under `outputs/purpose-lite-v1/`. It selects checkpoints and fits temperature on label-scorable validation records. When deriving nested HIGH/MEDIUM/LOW cutoffs, every `vague-eval` -record counts as an abstention miss even if its synthetic label happens to match. +record counts as an abstention miss even if its synthetic label happens to match. The +incoming checkpoint is scored and retained as epoch zero, so a continuation run cannot +silently replace it with a regression. Validation early stopping defaults to two epochs +without an improvement greater than 0.05 points. + +Continuation training accepts a local checkpoint. `--boundary-weight` is an opt-in, +validation-selected loss weight for the measured weakest slice; it does not add held-out +fixtures to training: + +```bash +ml/purpose-classifier/.venv/bin/python ml/purpose-classifier/train.py \ + --model ml/purpose-classifier/outputs/purpose-lite-v1/model \ + --epochs 3 --learning-rate 3e-6 --warmup-ratio 0 \ + --boundary-weight 2 \ + --output-dir ml/purpose-classifier/outputs/purpose-lite-v1-boundary-tune \ + --overwrite-output +``` For a wiring smoke test, use a small deterministic prefix: @@ -104,8 +120,21 @@ ml/purpose-classifier/.venv/bin/python ml/purpose-classifier/eval.py \ `export.py` emits fixed-shape opset-17 fp16 and int8-QDQ graphs, a tokenizer/ normalization contract, golden tokenizations, shared calibration config, graph checks, -artifact hashes, and a size report. The int8 graph is the ≤25 MiB shipping candidate; -the fp16 graph remains the accelerator-oriented conversion input. +artifact hashes, and a size report. Its default 256-record quantization calibration sample +is deterministic and stratified by purpose, slice, and primary language; the export report +records the seed, distribution, and prompt hashes. The int8 graph is the ≤25 MiB shipping +candidate; the fp16 graph remains the accelerator-oriented conversion input. + +When scoring ONNX, add `--compare-pytorch` to measure artifact drift against +`--model-dir`. The report then includes overall, label-scorable, and per-slice label +agreement plus every correct→incorrect, incorrect→correct, and changed-wrong-label +transition: + +```bash +ml/purpose-classifier/.venv/bin/python ml/purpose-classifier/eval.py \ + --onnx-model ml/purpose-classifier/outputs/purpose-lite-v1/export/purpose-lite-v1-int8-qdq.onnx \ + --compare-pytorch --no-gate +``` ## Audit curation diff --git a/eval.py b/eval.py index 0c2aea6..54dc13f 100644 --- a/eval.py +++ b/eval.py @@ -140,6 +140,80 @@ def routing_tier_drift( } +def prediction_agreement( + records: Sequence[dict[str, Any]], + actual: Sequence[int], + reference: Sequence[int], + candidate: Sequence[int], +) -> dict[str, Any]: + """Report label drift from a reference runtime to a candidate artifact.""" + + if not ( + len(records) == len(actual) == len(reference) == len(candidate) + and records + ): + raise ValueError("prediction agreement needs equally sized, non-empty vectors") + agreements = [ + expected == got for expected, got in zip(reference, candidate) + ] + scorable_indexes = [ + index + for index, record in enumerate(records) + if record.get("slice") != "vague-eval" + ] + transitions = Counter() + disagreements = [] + for index, agrees in enumerate(agreements): + if agrees: + continue + reference_correct = reference[index] == actual[index] + candidate_correct = candidate[index] == actual[index] + if reference_correct and not candidate_correct: + transition = "correctToIncorrect" + elif not reference_correct and candidate_correct: + transition = "incorrectToCorrect" + else: + transition = "differentIncorrectLabel" + transitions[transition] += 1 + disagreements.append( + { + "promptHash": prompt_hash(records[index]["prompt"]), + "slice": records[index].get("slice", "unknown"), + "expected": LABELS[actual[index]], + "reference": LABELS[reference[index]], + "candidate": LABELS[candidate[index]], + "transition": transition, + } + ) + + by_slice = {} + for slice_name in sorted( + {record.get("slice", "unknown") for record in records} + ): + indexes = [ + index + for index, record in enumerate(records) + if record.get("slice", "unknown") == slice_name + ] + by_slice[slice_name] = { + "records": len(indexes), + "labelAgreement": ( + sum(agreements[index] for index in indexes) / len(indexes) + ), + } + return { + "records": len(records), + "labelAgreement": sum(agreements) / len(agreements), + "scoredLabelAgreement": ( + sum(agreements[index] for index in scorable_indexes) + / len(scorable_indexes) + ), + "transitionCounts": dict(sorted(transitions.items())), + "bySlice": by_slice, + "disagreements": disagreements, + } + + def evaluate(args: argparse.Namespace) -> dict[str, Any]: try: import torch @@ -171,6 +245,7 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]: label_to_id = {label: index for index, label in enumerate(LABELS)} tokenizer = AutoTokenizer.from_pretrained(args.model_dir, local_files_only=True) onnx_session = None + reference_model = None if args.onnx_model is not None: try: import onnxruntime as ort @@ -191,7 +266,15 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]: device = torch.device("cpu") model = None runtime_name = "onnxruntime-cpu" + if args.compare_pytorch: + reference_model = AutoModelForSequenceClassification.from_pretrained( + args.model_dir, + local_files_only=True, + ).to(device) + reference_model.eval() else: + if args.compare_pytorch: + raise DataError("--compare-pytorch requires --onnx-model") device = _device(torch, args.device) model = AutoModelForSequenceClassification.from_pretrained( args.model_dir, local_files_only=True @@ -216,6 +299,7 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]: probabilities: list[float] = [] confidences: list[str] = [] margins: list[float] = [] + reference_predictions: list[int] = [] inference_batch_size = 1 if onnx_session is not None else args.batch_size with torch.inference_mode(): for start in range(0, len(records), inference_batch_size): @@ -226,6 +310,11 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]: torch, ) logits = predict_logits(encoded) / temperature + if reference_model is not None: + reference_logits = reference_model(**encoded).logits + reference_predictions.extend( + reference_logits.argmax(dim=-1).tolist() + ) distribution = torch.softmax(logits, dim=-1) top = torch.topk(distribution, k=2, dim=-1) batch_probabilities = top.values[:, 0].tolist() @@ -419,6 +508,13 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]: }, "misclassifications": misclassifications, } + if reference_model is not None: + report["pytorchParity"] = prediction_agreement( + records, + actual, + reference_predictions, + predicted, + ) write_json(args.report, report) return report @@ -431,6 +527,11 @@ def build_parser() -> argparse.ArgumentParser: type=Path, help="score a fixed-shape ONNX artifact instead of the PyTorch checkpoint", ) + parser.add_argument( + "--compare-pytorch", + action="store_true", + help="include label-level drift from --model-dir when scoring ONNX", + ) parser.add_argument("--calibration", type=Path, default=DEFAULT_CALIBRATION) parser.add_argument("--test", type=Path, default=DEFAULT_TEST) parser.add_argument("--fixtures", type=Path, default=DEFAULT_FIXTURES) diff --git a/export.py b/export.py index 2c68351..1b86886 100644 --- a/export.py +++ b/export.py @@ -5,15 +5,23 @@ from __future__ import annotations import argparse import hashlib +import math import shutil import sys import tempfile +from collections import Counter, defaultdict from pathlib import Path from typing import Any, Sequence import numpy as np -from purpose_data import DataError, load_jsonl, normalize_prompt, write_json +from purpose_data import ( + DataError, + load_jsonl, + normalize_prompt, + prompt_hash, + write_json, +) from train import ( HEAD_TOKENS, MAX_LENGTH, @@ -41,6 +49,59 @@ GOLDEN_PROMPTS = ( ) +def stratified_calibration_sample( + records: Sequence[dict[str, Any]], + count: int, + *, + seed: int, +) -> list[dict[str, Any]]: + """Select an exact, deterministic purpose/slice/language calibration sample.""" + + if not records: + raise DataError("cannot calibrate quantization from an empty validation split") + if count <= 0: + raise DataError("calibration record count must be positive") + target = min(count, len(records)) + groups: dict[tuple[str, str, str], list[dict[str, Any]]] = defaultdict(list) + for record in records: + language = str(record.get("lang", "unknown")).split("-", 1)[0].casefold() + key = ( + str(record.get("purpose", "unknown")), + str(record.get("slice", "unknown")), + language, + ) + groups[key].append(record) + + allocations = {} + remainders = [] + allocated = 0 + for key in sorted(groups): + quota = len(groups[key]) * target / len(records) + base = math.floor(quota) + allocations[key] = base + allocated += base + tie_break = hashlib.sha256(f"{seed}\0{key}".encode("utf-8")).hexdigest() + remainders.append((quota - base, tie_break, key)) + for _, _, key in sorted(remainders, reverse=True)[: target - allocated]: + allocations[key] += 1 + + selected = [] + for key in sorted(groups): + ranked = sorted( + groups[key], + key=lambda record: hashlib.sha256( + f"{seed}\0{prompt_hash(record['prompt'])}".encode("utf-8") + ).hexdigest(), + ) + selected.extend(ranked[: allocations[key]]) + return sorted( + selected, + key=lambda record: hashlib.sha256( + f"{seed + 1}\0{prompt_hash(record['prompt'])}".encode("utf-8") + ).hexdigest(), + ) + + def _sha256(path: Path) -> str: digest = hashlib.sha256() with path.open("rb") as handle: @@ -216,11 +277,16 @@ def export(args: argparse.Namespace) -> dict[str, Any]: onnx.save(fp16_model, fp16_path) validation = load_jsonl(args.validation) + calibration_samples = stratified_calibration_sample( + validation, + args.calibration_records, + seed=args.calibration_seed, + ) class Reader(CalibrationDataReader): def __init__(self) -> None: self.index = 0 - self.samples = validation[: args.calibration_records] + self.samples = calibration_samples def get_next(self) -> dict[str, np.ndarray] | None: if self.index >= len(self.samples): @@ -281,7 +347,20 @@ def export(args: argparse.Namespace) -> dict[str, Any]: "opset": args.opset, "fixedInputShape": [1, MAX_LENGTH], "inputNames": input_names, - "calibrationRecords": min(args.calibration_records, len(validation)), + "calibrationRecords": len(calibration_samples), + "calibrationSeed": args.calibration_seed, + "calibrationSample": { + "strategy": "stratified by purpose, slice, and primary language", + "purposeCounts": dict( + sorted(Counter(item["purpose"] for item in calibration_samples).items()) + ), + "sliceCounts": dict( + sorted(Counter(item["slice"] for item in calibration_samples).items()) + ), + "promptHashes": sorted( + prompt_hash(item["prompt"]) for item in calibration_samples + ), + }, "shippingArtifact": "int8QDQ", "shippingBudgetBytes": args.shipping_budget_bytes, "shippingBudgetPassed": int8_size <= args.shipping_budget_bytes, @@ -303,6 +382,7 @@ def build_parser() -> argparse.ArgumentParser: parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR) parser.add_argument("--opset", type=int, default=17) parser.add_argument("--calibration-records", type=int, default=256) + parser.add_argument("--calibration-seed", type=int, default=20260730) parser.add_argument( "--shipping-budget-bytes", type=int, default=SHIPPING_BUDGET_BYTES ) diff --git a/tests/test_eval.py b/tests/test_eval.py index f7a12b2..9bf9ed6 100644 --- a/tests/test_eval.py +++ b/tests/test_eval.py @@ -23,5 +23,33 @@ class TierDriftTests(unittest.TestCase): self.assertLessEqual(report["maximumTierDrift"], 1) +class PredictionAgreementTests(unittest.TestCase): + def test_reports_accuracy_transitions_and_scorable_agreement(self): + records = [ + {"prompt": "one", "slice": "core"}, + {"prompt": "two", "slice": "boundary"}, + {"prompt": "three", "slice": "vague-eval"}, + {"prompt": "four", "slice": "core"}, + ] + report = purpose_eval.prediction_agreement( + records, + actual=[0, 1, 2, 3], + reference=[0, 0, 3, 4], + candidate=[1, 1, 4, 5], + ) + self.assertEqual(0.0, report["labelAgreement"]) + self.assertEqual(0.0, report["scoredLabelAgreement"]) + self.assertEqual( + { + "correctToIncorrect": 1, + "differentIncorrectLabel": 2, + "incorrectToCorrect": 1, + }, + report["transitionCounts"], + ) + self.assertEqual(2, report["bySlice"]["core"]["records"]) + self.assertEqual(4, len(report["disagreements"])) + + if __name__ == "__main__": unittest.main() diff --git a/tests/test_export.py b/tests/test_export.py new file mode 100644 index 0000000..2e253bf --- /dev/null +++ b/tests/test_export.py @@ -0,0 +1,51 @@ +import sys +import unittest +from collections import Counter +from pathlib import Path + + +MODULE_DIR = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(MODULE_DIR)) + +import export + + +class CalibrationSampleTests(unittest.TestCase): + def test_sample_is_exact_deterministic_and_stratified(self): + records = [] + for index in range(100): + records.append( + { + "prompt": f"prompt {index}", + "purpose": "planning" if index < 80 else "writing", + "slice": "core" if index % 2 else "boundary", + "lang": "en" if index % 5 else "fr", + } + ) + first = export.stratified_calibration_sample(records, 25, seed=42) + second = export.stratified_calibration_sample(records, 25, seed=42) + self.assertEqual(25, len(first)) + self.assertEqual( + [item["prompt"] for item in first], + [item["prompt"] for item in second], + ) + purposes = Counter(item["purpose"] for item in first) + self.assertEqual({"planning": 20, "writing": 5}, dict(purposes)) + + def test_sample_caps_at_population(self): + records = [ + { + "prompt": "one", + "purpose": "planning", + "slice": "core", + "lang": "en", + } + ] + self.assertEqual( + records, + export.stratified_calibration_sample(records, 10, seed=1), + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_train.py b/tests/test_train.py index 73665a9..1f6958f 100644 --- a/tests/test_train.py +++ b/tests/test_train.py @@ -12,6 +12,16 @@ import train class MetricsTests(unittest.TestCase): + def test_boundary_training_weight_is_opt_in(self): + self.assertEqual( + 2.0, + train.training_weight({"slice": "boundary"}, boundary_weight=2.0), + ) + self.assertEqual( + 1.0, + train.training_weight({"slice": "core"}, boundary_weight=2.0), + ) + def test_classification_metrics_include_every_label(self): actual = list(range(8)) predicted = [0, 1, 2, 3, 4, 5, 6, 0] diff --git a/train.py b/train.py index d6ab54f..81e8b00 100644 --- a/train.py +++ b/train.py @@ -31,6 +31,12 @@ 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 encode_fixed_shape( tokenizer: Any, texts: Sequence[str], @@ -284,6 +290,7 @@ def _evaluate(torch: Any, model: Any, loader: Any, device: Any) -> tuple[Any, An with torch.inference_mode(): for batch in loader: labels = batch.pop("labels") + batch.pop("sample_weights", None) inputs = {key: value.to(device) for key, value in batch.items()} logits = model(**inputs).logits.cpu() all_logits.append(logits) @@ -317,6 +324,17 @@ def train(args: argparse.Namespace) -> dict[str, Any]: validation_records = validation_records[: args.max_validation_records] output_dir: Path = args.output_dir + local_model = Path(args.model).expanduser() + if local_model.exists(): + try: + local_model.resolve().relative_to(output_dir.resolve()) + except ValueError: + pass + else: + raise DataError( + "local --model 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( @@ -329,16 +347,22 @@ def train(args: argparse.Namespace) -> dict[str, Any]: 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, revision=args.model_revision, use_fast=True + args.model, use_fast=True, **pretrained_options ) 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, + **pretrained_options, ) config = model.config if getattr(config, "hidden_size", None) != 384 or getattr( @@ -364,14 +388,19 @@ def train(args: argparse.Namespace) -> dict[str, Any]: def __len__(self) -> int: return len(self.records) - def __getitem__(self, index: int) -> tuple[str, int]: + def __getitem__(self, index: int) -> tuple[str, int, float]: record = self.records[index] - return prepare_text(record["prompt"]), label_to_id[record["purpose"]] + return ( + prepare_text(record["prompt"]), + label_to_id[record["purpose"]], + training_weight(record, args.boundary_weight), + ) - def collate(items: Sequence[tuple[str, int]]) -> dict[str, Any]: - texts, labels = zip(*items) + def collate(items: Sequence[tuple[str, int, float]]) -> dict[str, Any]: + texts, labels, weights = 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) return encoded generator = torch.Generator() @@ -397,6 +426,36 @@ def train(args: argparse.Namespace) -> dict[str, Any]: [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 + ], + ) + best_accuracy = initial_metrics["accuracy"] + 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%}", + flush=True, + ) optimizer = torch.optim.AdamW( model.parameters(), lr=args.learning_rate, weight_decay=args.weight_decay @@ -411,17 +470,23 @@ def train(args: argparse.Namespace) -> dict[str, Any]: 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 + labels = batch.pop("labels").to(device) + sample_weights = batch.pop("sample_weights").to(device) + inputs = {key: value.to(device) for key, value in batch.items()} + per_record_loss = torch.nn.functional.cross_entropy( + model(**inputs).logits, + labels, + reduction="none", + ) + 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 should_update = ( @@ -454,10 +519,23 @@ def train(args: argparse.Namespace) -> dict[str, Any]: f"macro_recall={metrics['macroRecall']:.4%}", flush=True, ) - if metrics["accuracy"] > best_accuracy: + improvement = metrics["accuracy"] - best_accuracy + if improvement > args.minimum_improvement: best_accuracy = metrics["accuracy"] + 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 validation 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) @@ -497,7 +575,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]: metrics = { "modelVersion": "purpose-lite-v1", "baseModel": args.model, - "baseModelRevision": args.model_revision, + "baseModelRevision": model_revision or "local-checkpoint", "fixedInputShape": [1, MAX_LENGTH], "truncation": { "strategy": "head-tail-pair", @@ -507,12 +585,16 @@ def train(args: argparse.Namespace) -> dict[str, Any]: "device": str(device), "trainingSeconds": time.perf_counter() - started, "trainRecords": len(train_records), + "boundaryTrainingWeight": args.boundary_weight, "validationRecords": len(validation_records), "scoredValidationRecords": int(validation_scorable.sum().item()), "vagueAbstentionValidationRecords": int( (~validation_scorable).sum().item() ), "bestValidationAccuracy": best_accuracy, + "initialValidation": initial_metrics, + "epochsCompleted": len(history), + "stoppedEarly": stopped_early, "bestValidation": classification_metrics( labels[validation_scorable].tolist(), [ @@ -555,6 +637,9 @@ def build_parser() -> argparse.ArgumentParser: 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("--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("--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) @@ -576,10 +661,15 @@ def main(argv: Sequence[str] | None = None) -> int: "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.boundary_weight <= 0.0: + parser.error("--boundary-weight must be positive") if not 0.0 < args.accepted_precision <= args.high_precision <= 1.0: parser.error( "precision targets must satisfy 0 < accepted <= high <= 1"