Decision-1.0-Nox-4B / code /predict.py
Xunzhuo's picture
Release measured Decision 1.0 decoder
43d6004 verified
Raw History Blame
3.16 kB
"""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))