"""E2 stage A eval (item 21): GSM8K accuracy for short-CoT-trained arms. Generates with generate_carry_c (prefill loop k + pause carry + per-token carry) or feedforward mode for the control arm; scores last_number vs gold, by STaR label. Grid over (k,p) cells. """ import argparse import json import os import sys import time from pathlib import Path import torch from carry_common import generate_carry_c from loop_common import (BandLooper, MergeAdapter, chat_prompt, DIRECT_SUFFIX, last_number, num_eq) sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) from jlens.core import load_model # noqa: E402 OUT = Path(os.environ.get("LOOP_OUT", Path(__file__).resolve().parent.parent / "results-loop")) @torch.no_grad() def main(): ap = argparse.ArgumentParser() ap.add_argument("--adapter", required=True) ap.add_argument("--tag", required=True) ap.add_argument("--grid", default="0:0,2:2,2:6", help="comma list of k:p cells") ap.add_argument("--n", type=int, default=256) ap.add_argument("--batch", type=int, default=8) ap.add_argument("--feedforward", action="store_true") ap.add_argument("--max-new", type=int, default=160) args = ap.parse_args() model, tok = load_model(dtype=torch.bfloat16) tok.padding_side = "left" looper = BandLooper(model) adapter = MergeAdapter( d=model.config.get_text_config().hidden_size).cuda() adapter.load_state_dict(torch.load(args.adapter, map_location="cuda")) adapter.eval() items = [it for it in json.load(open(OUT / "star_data.json")) if it["split"] == "test"][: args.n] print(f"[{args.tag}] GSM carry-cot eval on {len(items)}, " f"grid={args.grid} ff={args.feedforward}", flush=True) res = {"tag": args.tag, "grid": {}, "n": len(items)} for cell in args.grid.split(","): k, p = (int(x) for x in cell.split(":")) t0 = time.time() hits, per_label, per_item = 0, {}, [] for i in range(0, len(items), args.batch): chunk = items[i : i + args.batch] enc = tok([chat_prompt(tok, it["question"], DIRECT_SUFFIX) for it in chunk], return_tensors="pt", padding=True, add_special_tokens=False).to("cuda") if k == 0 and p == 0: gen = model.generate(**enc, max_new_tokens=args.max_new, do_sample=False) else: gen = generate_carry_c(looper, adapter, tok, enc["input_ids"], enc["attention_mask"], k, p, max_new_tokens=args.max_new, feedforward=args.feedforward) for j, it in enumerate(chunk): txt = tok.decode(gen[j, enc["input_ids"].shape[1]:], skip_special_tokens=True) ok = num_eq(last_number(txt), it["gold"]) hits += ok d = per_label.setdefault(it["label"], [0, 0]) d[0] += ok d[1] += 1 per_item.append({"idx": it["idx"], "ok": bool(ok)}) acc = hits / len(items) by = {l: c / n for l, (c, n) in per_label.items()} res["grid"][cell] = {"acc": acc, "by_label": by, "per_item": per_item} print(f"{cell}: acc={acc:.3f} " f"by_label={ {l: round(v,3) for l,v in by.items()} }" f" ({time.time()-t0:.0f}s)", flush=True) json.dump(res, open(OUT / f"eval_{args.tag}.json", "w"), indent=1) print("wrote", OUT / f"eval_{args.tag}.json") if __name__ == "__main__": main()