Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -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())
|
||||
Reference in New Issue
Block a user