#!/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())