Files
jspace/scripts/eval_carry_cot.py
T

156 lines
6.7 KiB
Python

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