"""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()