187 lines
6.1 KiB
Python
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())
|