"""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) ap.add_argument("--bandlora", default=None, metavar="LORA_PT", help="load a lora_*_e*.pt loop-only band-LoRA checkpoint") ap.add_argument("--inner-iters", type=int, default=0, metavar="M", help="M in-place band iterations at the last pause") ap.add_argument("--kvmem", default=None, metavar="KVMEM_PT", help="KVMemoryAdapter checkpoint: burst states become " "band-layer KV prefix entries at generation") ap.add_argument("--symchain", default=None, metavar="PROJ_PT", help="item 32: discrete latent chain — zero-init " "projector checkpoint; generation snaps hard argmax") args = ap.parse_args() model, tok = load_model(dtype=torch.bfloat16) tok.padding_side = "left" looper = BandLooper(model) kvmem = None if args.kvmem: from kv_memory import install, KVMemoryAdapter from loop_common import BAND install(model) sd = torch.load(args.kvmem, map_location="cuda") code = sd["trunk.1.weight"].shape[0] kvmem = KVMemoryAdapter(model, band=BAND, code=code).cuda() kvmem.load_state_dict(sd) kvmem.eval() print(f"kv-memory loaded: {args.kvmem} (code={code})", flush=True) if args.bandlora: from lora_band import inject_band_lora ck = torch.load(args.bandlora, map_location="cuda") scales = {l: 1.0 for l in ck["band"]} ps = inject_band_lora(looper.tm, ck["band"][0], scales, rank=ck["rank"]) assert len(ps) == len(ck["tensors"]), (len(ps), len(ck["tensors"])) for pr, t in zip(ps, ck["tensors"]): pr.data = t.cuda() print(f"band-lora loaded: {args.bandlora} " f"(r={ck['rank']}, layers {ck['band'][0]}-{ck['band'][-1]})", flush=True) symchain = None if args.symchain: d_ = model.config.get_text_config().hidden_size proj = torch.nn.Linear(d_, d_).cuda() proj.load_state_dict(torch.load(args.symchain, map_location="cuda")) proj.eval() jbar = torch.load(Path(__file__).resolve().parent.parent / "results/jbar.pt", map_location="cuda")["Jbar"] from loop_common import BAND J30 = jbar[BAND[1]].float() tm = model.model.language_model softcap = model.config.get_text_config().final_logit_softcapping def lens_fn(h): x = tm.norm((h.float() @ J30.T).to(tm.norm.weight.dtype)) lg = model.lm_head(x) return softcap * torch.tanh(lg / softcap) if softcap else lg symchain = {"proj": proj, "lens_fn": lens_fn, "embed_w": model.get_input_embeddings().weight.detach(), "start_id": tok("\n", add_special_tokens=False)["input_ids"][0]} print(f"symchain loaded: {args.symchain}", flush=True) 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, inner_iters=args.inner_iters, kvmem=kvmem, symchain=symchain) if kvmem is not None: from kv_memory import arm_memory arm_memory(None) 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()