tiny-agent-112m / code /scripts /rescore.py
darioooooo0o's picture
tiny-agent-112m: base + RL weights, tokenizer, code, model card
4397e12 verified
Raw History Blame Contribute Delete
1.66 kB
"""Score checkpoints on one common, larger validation set (same windows for every run).
Usage: source env.sh && $TA_PY scripts/rescore.py run_dir [run_dir ...] [--batches 12]"""
import argparse
import json
import os
import numpy as np
import torch
from tiny_agent.checkpoint import load_model
from tiny_agent.data import MixtureLoader
from tiny_agent.model import make_block_mask
@torch.no_grad()
def score(model, batches):
model.eval()
out = {}
for name, bs in batches.items():
ls = []
for inp, tgt, doc in bs:
inp, tgt, doc = inp.to("xpu"), tgt.to("xpu"), doc.to("xpu")
with torch.autocast("xpu", dtype=torch.bfloat16):
ls.append(model(inp, doc, make_block_mask(doc, model.cfg.swa_window), tgt).item())
out[name] = float(np.mean(ls))
return out
def main():
ap = argparse.ArgumentParser()
ap.add_argument("runs", nargs="+")
ap.add_argument("--batches", type=int, default=12)
a = ap.parse_args()
batches = MixtureLoader("val", "stable", 2048, 8, stream=False).fixed_batches(a.batches, seed=4321)
res = {}
for r in a.runs:
st = torch.load(os.path.join(r, "ckpt.pt"), map_location="cpu", weights_only=False)
model = load_model(os.path.join(r, "ckpt.pt"))
v = score(model, batches)
res[r] = {"tokens_M": round(st["tokens"] / 1e6), "step": st["step"], "mean_val": round(float(np.mean(list(v.values()))), 4),
**{k: round(x, 4) for k, x in sorted(v.items())}}
print(json.dumps({r: res[r]}), flush=True)
del model
torch.xpu.empty_cache()
if __name__ == "__main__":
main()