"""Portable one-GPU JSONL inference for frozen dynamic-option checkpoints.""" import argparse import hashlib import json from pathlib import Path import time import torch from decision_model import DecisionModel, encode, collate p = argparse.ArgumentParser() p.add_argument("--model", required=True) p.add_argument("--base", action="store_true") p.add_argument("--revision", default="local-checkpoint") p.add_argument("--input", required=True) p.add_argument("--output", required=True) p.add_argument("--batch-size", type=int, default=8) p.add_argument("--max-length", type=int, default=4096) p.add_argument("--temperature", type=float, default=1.0) a = p.parse_args() assert a.temperature > 0 torch.cuda.set_device(0) model, tokenizer = DecisionModel.from_base(a.model, a.revision) if a.base else DecisionModel.from_checkpoint(a.model) model = model.cuda().eval() rows = [json.loads(line) for line in Path(a.input).read_text().splitlines() if line.strip()] encoded = [encode(row, tokenizer, a.max_length) for row in rows] assert len({row["id"] for row in rows}) == len(rows) output = Path(a.output); output.parent.mkdir(parents=True, exist_ok=True) if output.exists(): raise RuntimeError("Refusing to overwrite predictions") with torch.inference_mode(), output.open("w") as f: for start in range(0, len(encoded), a.batch_size): examples = encoded[start:start + a.batch_size] batch = {key: value.cuda() if torch.is_tensor(value) else value for key, value in collate(examples, tokenizer.pad_token_id if tokenizer.pad_token_id is not None else tokenizer.eos_token_id).items()} torch.cuda.synchronize(); tick = time.perf_counter() with torch.autocast("cuda", dtype=torch.bfloat16): logits = model(**batch) torch.cuda.synchronize(); elapsed = time.perf_counter() - tick probabilities = (logits / a.temperature).softmax(-1).cpu().tolist(); scores = logits.cpu().tolist() for row, example, prob, score in zip(rows[start:start + a.batch_size], examples, probabilities, scores): k = example["nopts"]; pred = max(range(k), key=lambda i: prob[i]) rec = {"id": row["id"], "family": row.get("family"), "label": row.get("label"), "prediction": pred, "prediction_key": row["options"][pred]["key"], "probabilities": prob[:k], "logits": score[:k], "temperature": a.temperature, "input_tokens": len(example["ids"]), "prompt_sha256": example["prompt_sha256"], "batch_elapsed_seconds": elapsed, "batch_size": len(examples)} if "score_values" in row: rec["expected_score"] = sum(value * probability for value, probability in zip(row["score_values"], prob[:k])) f.write(json.dumps(rec, ensure_ascii=False) + "\n") f.flush(); print(json.dumps({"completed": min(start + a.batch_size, len(rows)), "total": len(rows)}), flush=True) Path(str(output) + ".metadata.json").write_text(json.dumps({"model": a.model, "revision": a.revision, "temperature": a.temperature, "input_sha256": hashlib.sha256(Path(a.input).read_bytes()).hexdigest(), "predictions_sha256": hashlib.sha256(output.read_bytes()).hexdigest(), "rows": len(rows), "model_metadata": model.metadata}, indent=2))