"""Eval design C (prefill k + carry through p pauses) on GSM8K test. Usage: eval_carry.py --adapter ../results-loop/adapter_carry_e600.pt \ --tag carry --grid "0:0,2:0,2:2,2:6" Grid entries are k:p pairs; k=0,p=0 is the plain baseline (same harness). """ import argparse import json import sys import time from pathlib import Path import torch from carry_common import generate_carry_c from loop_common import (DIRECT_SUFFIX, BandLooper, MergeAdapter, chat_prompt, last_number, num_eq) sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) from jlens.core import load_model # noqa: E402 OUT = Path(__file__).resolve().parent.parent / "results-loop" @torch.no_grad() def acc_at(looper, adapter, tok, items, k, p, batch=16, feedforward=False): hits, per_label = 0, {} for i in range(0, len(items), batch): torch.cuda.empty_cache() chunk = items[i : i + 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") gen = generate_carry_c(looper, adapter, tok, enc["input_ids"], enc["attention_mask"], k=k, p=p, max_new_tokens=10, feedforward=feedforward) n0 = enc["input_ids"].shape[1] + p for j, it in enumerate(chunk): txt = tok.decode(gen[j, n0:], 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 return hits / len(items), {l: c / n for l, (c, n) in per_label.items()} def main(): ap = argparse.ArgumentParser() ap.add_argument("--adapter", default=None) ap.add_argument("--tag", default="carry") ap.add_argument("--grid", default="0:0,2:0,2:2,2:6") ap.add_argument("--n", type=int, default=0) ap.add_argument("--feedforward", action="store_true") args = ap.parse_args() model, tok = load_model(dtype=torch.bfloat16) tok.padding_side = "left" looper = BandLooper(model) adapter = MergeAdapter().cuda() if args.adapter: 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"] if args.n: items = items[: args.n] print(f"[{args.tag}] GSM8K carry eval on {len(items)} items, " f"grid={args.grid}", flush=True) res = {"tag": args.tag, "grid": {}, "n": len(items)} for kp in args.grid.split(","): k, p = (int(x) for x in kp.split(":")) t0 = time.time() acc, by_label = acc_at(looper, adapter, tok, items, k, p, feedforward=args.feedforward) res["grid"][kp] = {"acc": acc, "by_label": by_label} print(f"k={k} p={p}: acc={acc:.3f} " f"by_label={ {l: round(v,3) for l,v in by_label.items()} }" f" ({time.time()-t0:.0f}s)", flush=True) with open(OUT / f"eval_{args.tag}.json", "w") as f: json.dump(res, f, indent=1) print("wrote", OUT / f"eval_{args.tag}.json") if __name__ == "__main__": main()