Files

187 lines
6.1 KiB
Python

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