diff --git a/README.md b/README.md index 23c8613..5614468 100644 --- a/README.md +++ b/README.md @@ -117,6 +117,34 @@ it is rejected. Do not continue optimizer-only QAT sweeps on this split. The nex iteration should incorporate reviewed boundary data and be selected on a revised validation/frozen dataset version. +To target only the remaining float→int8 decision drift, cache the float teacher in a +separate inference process and use its logits for QAT distillation. Keeping teacher and +student models out of the same process avoids doubling peak resident memory: + +```bash +ml/purpose-classifier/.venv/bin/python ml/purpose-classifier/cache_teacher.py \ + --model ml/purpose-classifier/outputs/purpose-lite-v1-boundary-tune/model \ + --output ml/purpose-classifier/outputs/purpose-lite-v1-boundary-teacher.pt \ + --overwrite-output +ml/purpose-classifier/.venv/bin/python ml/purpose-classifier/train.py \ + --model ml/purpose-classifier/outputs/purpose-lite-v1-boundary-tune/model \ + --distillation-cache \ + ml/purpose-classifier/outputs/purpose-lite-v1-boundary-teacher.pt \ + --distillation-weight 0.9 --distillation-temperature 2 \ + --distillation-selection-weight 0.5 --quantization-aware \ + --epochs 2 --learning-rate 1e-6 --warmup-ratio 0 --boundary-weight 1 \ + --output-dir ml/purpose-classifier/outputs/purpose-lite-v1-distilled-qat \ + --overwrite-output +``` + +The cache binds each logit row to normalized prompt hash plus expected label. Training +fails closed if either split changes. Selection combines label accuracy with float-teacher +agreement, retains the incoming checkpoint as epoch zero, and logs label/distillation loss +separately. A 64-record wiring run exercised cache loading, shuffled row alignment, +backpropagation, selection, and ordinary checkpoint reload. The current shared CPU runtime +then showed severe post-batch throttling, so no full candidate result is claimed from that +canary. + For a wiring smoke test, use a small deterministic prefix: ```bash @@ -178,10 +206,12 @@ the frozen split automatically. The 18 word-trigram exclusions and the human-rev completion rule are recorded in `data/curation-review-v1.json`; the semantic report is versioned as `data/semantic-audit-v1.json`. -## Complete the human review +## Optional human review -The deterministic CSV currently contains 1,219 blank review rows. Check progress without -running the embedding audit again: +The dataset owner accepted the curated generated labels and difficulty metadata as-is on +2026-07-31, so the blank 1,219-row review sample is not a training or rollout blocker. It +remains available as an optional future audit. Check its progress without running the +embedding audit again: ```bash ml/purpose-classifier/.venv/bin/python ml/purpose-classifier/review_data.py diff --git a/cache_teacher.py b/cache_teacher.py new file mode 100644 index 0000000..a634e10 --- /dev/null +++ b/cache_teacher.py @@ -0,0 +1,186 @@ +#!/usr/bin/env python3 +"""Cache float-teacher logits for memory-bounded QAT distillation.""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path +from typing import Any, Sequence + +from purpose_data import LABELS, DataError, load_jsonl +from train import ( + DEFAULT_DATASET_DIR, + DEFAULT_MODEL_REVISION, + _select_device, + _set_seeds, + _validate_split, + distillation_record_keys, + encode_fixed_shape, + prepare_text, +) + + +SCRIPT_DIR = Path(__file__).resolve().parent +DEFAULT_OUTPUT = SCRIPT_DIR / "outputs" / "purpose-lite-v1-teacher-logits.pt" + + +def _predict( + torch: Any, + model: Any, + tokenizer: Any, + records: Sequence[dict[str, Any]], + *, + device: Any, + batch_size: int, + progress_steps: int, + label: str, +) -> Any: + rows = [] + batches = (len(records) + batch_size - 1) // batch_size + started = time.perf_counter() + model.eval() + with torch.inference_mode(): + for batch_index, start in enumerate(range(0, len(records), batch_size), 1): + batch = records[start : start + batch_size] + encoded = encode_fixed_shape( + tokenizer, + [prepare_text(record["prompt"]) for record in batch], + torch, + ) + inputs = {key: value.to(device) for key, value in encoded.items()} + rows.append(model(**inputs).logits.cpu()) + if progress_steps and ( + batch_index % progress_steps == 0 or batch_index == batches + ): + print( + f"teacher {label} step {batch_index}/{batches} " + f"elapsed={time.perf_counter() - started:.1f}s", + flush=True, + ) + return torch.cat(rows) + + +def cache_teacher(args: argparse.Namespace) -> dict[str, Any]: + try: + import torch + from transformers import AutoModelForSequenceClassification, AutoTokenizer + except ImportError as exc: + raise DataError( + "teacher-cache dependencies are missing; install requirements.txt" + ) 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) + + output = args.output.resolve() + local_model = Path(args.model).expanduser() + if local_model.is_dir(): + try: + output.relative_to(local_model.resolve()) + except ValueError: + pass + else: + raise DataError("teacher cache output must not overwrite the model directory") + if output.exists() and not args.overwrite_output: + raise DataError(f"{output}: cache exists; pass --overwrite-output intentionally") + if output.exists() and output.is_dir(): + raise DataError(f"{output}: cache output must be a file path") + output.parent.mkdir(parents=True, exist_ok=True) + + _set_seeds(torch, args.seed) + device = _select_device(torch, args.device) + options = ( + {"local_files_only": True} + if local_model.exists() + else {"revision": args.model_revision} + ) + tokenizer = AutoTokenizer.from_pretrained(args.model, use_fast=True, **options) + model = AutoModelForSequenceClassification.from_pretrained( + args.model, + **options, + ).to(device) + teacher_labels = [ + model.config.id2label.get(index, model.config.id2label.get(str(index))) + for index in range(len(LABELS)) + ] + if teacher_labels != list(LABELS): + raise DataError("teacher label order does not match purpose-lite") + + train_logits = _predict( + torch, + model, + tokenizer, + train_records, + device=device, + batch_size=args.batch_size, + progress_steps=args.progress_steps, + label="train", + ) + validation_logits = _predict( + torch, + model, + tokenizer, + validation_records, + device=device, + batch_size=args.batch_size, + progress_steps=args.progress_steps, + label="validation", + ) + artifact = { + "schemaVersion": 1, + "labels": list(LABELS), + "teacher": str(args.model), + "trainRecordKeys": distillation_record_keys(train_records), + "validationRecordKeys": distillation_record_keys(validation_records), + "trainLogits": train_logits, + "validationLogits": validation_logits, + } + torch.save(artifact, output) + return { + "trainRecords": len(train_records), + "validationRecords": len(validation_records), + "output": str(output), + } + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--dataset-dir", type=Path, default=DEFAULT_DATASET_DIR) + parser.add_argument("--model", required=True) + parser.add_argument("--model-revision", default=DEFAULT_MODEL_REVISION) + parser.add_argument("--output", type=Path, default=DEFAULT_OUTPUT) + parser.add_argument("--device", default="auto") + parser.add_argument("--seed", type=int, default=20260730) + parser.add_argument("--batch-size", type=int, default=64) + parser.add_argument("--progress-steps", type=int, default=50) + 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) + if args.batch_size <= 0: + parser.error("--batch-size must be positive") + if args.progress_steps < 0: + parser.error("--progress-steps must be non-negative") + try: + metrics = cache_teacher(args) + except (DataError, OSError, ValueError, RuntimeError) as exc: + print(f"error: {exc}", file=sys.stderr) + return 1 + print( + f"Cached teacher logits for {metrics['trainRecords']} train and " + f"{metrics['validationRecords']} validation records at {metrics['output']}." + ) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/data/curation-review-v1.json b/data/curation-review-v1.json index 1823bcd..901ad5c 100644 --- a/data/curation-review-v1.json +++ b/data/curation-review-v1.json @@ -61,7 +61,7 @@ ] }, "humanLabelAndDifficultyReview": { - "status": "planned", + "status": "accepted-as-generated", "populationRecords": 12193, "sampleFraction": 0.1, "sampleRecords": 1219, @@ -80,6 +80,8 @@ "notes" ], "generatedArtifact": "ml/purpose-classifier/.artifacts/human-review-v1.csv", - "completionRule": "Every sampled row must be marked accept, relabel, or reject by a human reviewer. review_data.py validates the exact sample and writes a versioned complete ledger; prepare_data.py applies relabel/reject decisions before splitting. The revised split and frozen evaluation must then be intentionally reviewed and versioned before the dataset can be called fully curated." + "decisionDate": "2026-07-31", + "decisionBasis": "The dataset owner explicitly directed the project to assume the generated labels, secondary purposes, difficulties, and slices are correct and validated without completing the row-by-row sample.", + "decision": "Accept the curated generated population as-is. No relabel or reject decisions are inferred, and the blank CSV remains an optional future audit artifact rather than a rollout blocker." } } diff --git a/data/dataset-v1-manifest.json b/data/dataset-v1-manifest.json index ff40fde..f016379 100644 --- a/data/dataset-v1-manifest.json +++ b/data/dataset-v1-manifest.json @@ -9,7 +9,7 @@ "nearDuplicateThreshold": 0.92, "retainedRecords": 12193, "reviewPath": "ml/purpose-classifier/data/curation-review-v1.json", - "reviewSha256": "625251a98bcdcde0bee074e3bba3c564af4347eb7955497a9c3e018aae00b30f", + "reviewSha256": "1721ab77a0f723a4b03f42b14850b5c351ead2cd8fa72bb2d1fe3f0d2bb568da", "vagueEvalPolicy": "validation/test only" }, "datasetVersion": "purpose-dataset-v1", diff --git a/review_data.py b/review_data.py index 0e21bb3..924e7c0 100644 --- a/review_data.py +++ b/review_data.py @@ -65,6 +65,17 @@ def reviewed_population(args: argparse.Namespace) -> list[SourceRecord]: ).records +def review_policy_status(path: Path) -> str: + try: + value = json.loads(path.read_text(encoding="utf-8")) + status = value["humanLabelAndDifficultyReview"]["status"] + except (OSError, UnicodeError, json.JSONDecodeError, KeyError, TypeError) as exc: + raise DataError(f"{path}: cannot read human-review policy: {exc}") from exc + if not isinstance(status, str) or not status: + raise DataError(f"{path}: invalid human-review policy status") + return status + + def build_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--source", action="append", type=Path) @@ -98,6 +109,7 @@ def main(argv: Sequence[str] | None = None) -> int: args = build_parser().parse_args(argv) try: population = reviewed_population(args) + policy_status = review_policy_status(args.curation_review.resolve()) sample = stratified_review_sample( population, fraction=args.review_fraction, @@ -152,7 +164,8 @@ def main(argv: Sequence[str] | None = None) -> int: print( f"Human review: {progress.completed}/{progress.records} complete " f"({progress.accepted} accept, {progress.relabeled} relabel, " - f"{progress.rejected} reject, {progress.incomplete} remaining)." + f"{progress.rejected} reject, {progress.incomplete} remaining). " + f"Dataset policy: {policy_status}." ) return 0 diff --git a/tests/test_train.py b/tests/test_train.py index 2eb0e20..b82272b 100644 --- a/tests/test_train.py +++ b/tests/test_train.py @@ -12,6 +12,19 @@ import train class MetricsTests(unittest.TestCase): + def test_distillation_loss_is_zero_for_matching_logits_and_backpropagates(self): + teacher = torch.tensor([[2.0, 0.0, -1.0]]) + student = teacher.clone().requires_grad_(True) + loss = train.knowledge_distillation_loss( + torch, + student, + teacher, + temperature=2.0, + ).mean() + self.assertAlmostEqual(0.0, loss.item(), places=6) + loss.backward() + self.assertIsNotNone(student.grad) + def test_qat_replacements_keep_checkpoint_keys_and_gradients(self): model = torch.nn.Sequential( torch.nn.Embedding(16, 8), diff --git a/train.py b/train.py index c9564e7..ba25147 100644 --- a/train.py +++ b/train.py @@ -12,7 +12,14 @@ import time from pathlib import Path from typing import Any, Sequence -from purpose_data import LABELS, DataError, load_jsonl, normalize_prompt, write_json +from purpose_data import ( + LABELS, + DataError, + load_jsonl, + normalize_prompt, + prompt_hash, + write_json, +) SCRIPT_DIR = Path(__file__).resolve().parent @@ -37,6 +44,42 @@ def training_weight(record: dict[str, Any], boundary_weight: float) -> float: 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. @@ -400,18 +443,36 @@ def _fit_temperature(torch: Any, logits: Any, labels: Any) -> float: 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]: +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 batch in loader: + 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) @@ -442,15 +503,25 @@ def train(args: argparse.Namespace) -> dict[str, Any]: output_dir: Path = args.output_dir local_model = Path(args.model).expanduser() - if local_model.exists(): + 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_model.resolve().relative_to(output_dir.resolve()) + local_path.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" + 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: @@ -505,31 +576,88 @@ def train(args: argparse.Namespace) -> dict[str, Any]: model.to(device) class PromptDataset(Dataset): - def __init__(self, records: Sequence[dict[str, Any]]) -> None: + 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]: + 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]]) -> dict[str, Any]: - texts, labels, weights = zip(*items) + 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), + PromptDataset(train_records, teacher_train_logits), batch_size=args.batch_size, shuffle=True, generator=generator, @@ -567,7 +695,42 @@ def train(args: argparse.Namespace) -> dict[str, Any]: 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 = [] @@ -576,7 +739,13 @@ def train(args: argparse.Namespace) -> dict[str, Any]: tokenizer.save_pretrained(best_dir) print( f"epoch 0: validation_accuracy={initial_metrics['accuracy']:.4%} " - f"macro_recall={initial_metrics['macroRecall']:.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, ) @@ -599,20 +768,46 @@ def train(args: argparse.Namespace) -> dict[str, Any]: 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()} - per_record_loss = torch.nn.functional.cross_entropy( - model(**inputs).logits, + 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) @@ -628,6 +823,8 @@ def train(args: argparse.Namespace) -> dict[str, Any]: 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, ) @@ -643,18 +840,34 @@ def train(args: argparse.Namespace) -> dict[str, Any]: 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"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 = metrics["accuracy"] - best_accuracy + 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) @@ -663,7 +876,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]: if epochs_without_improvement >= args.early_stopping_patience: stopped_early = True print( - f"early stopping after epoch {epoch}: no validation improvement " + 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, @@ -721,12 +934,19 @@ def train(args: argparse.Namespace) -> dict[str, Any]: "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, @@ -777,6 +997,10 @@ def build_parser() -> argparse.ArgumentParser: 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) @@ -809,6 +1033,24 @@ def main(argv: Sequence[str] | None = None) -> int: 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"