86 lines
3.1 KiB
Python
86 lines
3.1 KiB
Python
"""Blocksworld eval: pass@1 vs k (prompt-only loops, fast path)."""
|
|
|
|
import argparse
|
|
import json
|
|
import os
|
|
import sys
|
|
import time
|
|
from pathlib import Path
|
|
|
|
import torch
|
|
|
|
from bw_common import verify_plan
|
|
from bw_prep import DIRECT_SUFFIX, chat
|
|
from loop_common import AdaptiveMergeAdapter, BandLooper, MergeAdapter
|
|
|
|
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 acc_at_k(looper, adapter, tok, items, k, batch=12, max_new=200):
|
|
oks, per_item = [], []
|
|
for i in range(0, len(items), batch):
|
|
torch.cuda.empty_cache()
|
|
chunk = items[i : i + batch]
|
|
enc = tok([chat(tok, it, DIRECT_SUFFIX) for it in chunk],
|
|
return_tensors="pt", padding=True,
|
|
add_special_tokens=False).to("cuda")
|
|
gen = looper.generate_frozen_prompt(adapter, tok, enc["input_ids"], k,
|
|
max_new_tokens=max_new,
|
|
attention_mask=enc["attention_mask"])
|
|
for j, it in enumerate(chunk):
|
|
txt = tok.decode(gen[j, enc["input_ids"].shape[1]:],
|
|
skip_special_tokens=True)
|
|
ok = verify_plan(it, txt)
|
|
oks.append(ok)
|
|
per_item.append({"task_id": it["task_id"], "ok": bool(ok)})
|
|
by = {}
|
|
for it, ok in zip(items, oks):
|
|
d = by.setdefault(it["label"], [0, 0])
|
|
d[0] += ok
|
|
d[1] += 1
|
|
return (sum(oks) / len(items),
|
|
{l: c / n for l, (c, n) in by.items()}, per_item)
|
|
|
|
|
|
def main():
|
|
ap = argparse.ArgumentParser()
|
|
ap.add_argument("--adapter", default=None)
|
|
ap.add_argument("--tag", default="bw")
|
|
ap.add_argument("--ks", default="0,2,4")
|
|
ap.add_argument("--adaptive", action="store_true")
|
|
args = ap.parse_args()
|
|
ks = [int(x) for x in args.ks.split(",")]
|
|
|
|
model, tok = load_model(dtype=torch.bfloat16)
|
|
tok.padding_side = "left"
|
|
looper = BandLooper(model)
|
|
cls = AdaptiveMergeAdapter if args.adaptive else MergeAdapter
|
|
adapter = cls(
|
|
d=model.config.get_text_config().hidden_size).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 / "bw_data.json"))
|
|
if it["split"] == "test"]
|
|
print(f"[{args.tag}] Blocksworld eval on {len(items)} items", flush=True)
|
|
res = {"tag": args.tag, "ks": {}, "n": len(items)}
|
|
for k in ks:
|
|
t0 = time.time()
|
|
acc, by, per_item = acc_at_k(looper, adapter, tok, items, k)
|
|
res["ks"][k] = {"acc": acc, "by_label": by, "per_item": per_item}
|
|
print(f"k={k}: 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_bw_{args.tag}.json", "w"), indent=1)
|
|
print("wrote", OUT / f"eval_bw_{args.tag}.json")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|