77 lines
2.8 KiB
Python
77 lines
2.8 KiB
Python
"""Item 20 phase 1: record the E1c halting head's probabilities per item.
|
|
|
|
One pass over the test set with the gatehead adapter: p0 (pre-loop, on
|
|
s_0) and p_1..p_kmax (after each iteration). With these, k*(threshold)
|
|
is computable offline for any threshold — no more GPU passes.
|
|
"""
|
|
|
|
import argparse
|
|
import json
|
|
import os
|
|
import sys
|
|
from pathlib import Path
|
|
|
|
import torch
|
|
|
|
from halting_common import HaltingMergeAdapter
|
|
from loop_common import BandLooper
|
|
from prep_mbpp import DIRECT_SUFFIX, mbpp_prompt
|
|
|
|
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"))
|
|
|
|
|
|
def main():
|
|
ap = argparse.ArgumentParser()
|
|
ap.add_argument("--adapter", required=True)
|
|
ap.add_argument("--kmax", type=int, default=4)
|
|
ap.add_argument("--n", type=int, default=250)
|
|
ap.add_argument("--batch", type=int, default=8)
|
|
args = ap.parse_args()
|
|
|
|
model, tok = load_model(dtype=torch.bfloat16)
|
|
tok.padding_side = "left"
|
|
looper = BandLooper(model)
|
|
adapter = HaltingMergeAdapter(
|
|
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 / "mbpp_data.json"))
|
|
if it["split"] == "test"][: args.n]
|
|
rows = []
|
|
with torch.no_grad():
|
|
for i in range(0, len(items), args.batch):
|
|
chunk = items[i : i + args.batch]
|
|
enc = tok([mbpp_prompt(tok, it, DIRECT_SUFFIX) for it in chunk],
|
|
return_tensors="pt", padding=True,
|
|
add_special_tokens=False).to("cuda")
|
|
ids, msk = enc["input_ids"], enc["attention_mask"]
|
|
calls, _ = looper.capture(ids, msk, logits_to_keep=1)
|
|
e = looper._hin[looper.l0].detach()
|
|
B = e.shape[0]
|
|
bidx = torch.arange(B, device=e.device)
|
|
pl = torch.full((B,), ids.shape[1] - 1, device=e.device,
|
|
dtype=torch.long)
|
|
s = looper.band(e, calls)
|
|
probs = [adapter.halt_prob(e[bidx, pl], s[bidx, pl])]
|
|
for _ in range(args.kmax):
|
|
x = adapter(e, s)
|
|
s = looper.band(x, calls)
|
|
probs.append(adapter.halt_prob(e[bidx, pl], s[bidx, pl]))
|
|
P = torch.stack(probs, -1).float().cpu() # (B, kmax+1)
|
|
for j, it in enumerate(chunk):
|
|
rows.append({"task_id": it["task_id"], "label": it["label"],
|
|
"p": [round(v, 5) for v in P[j].tolist()]})
|
|
if i % 40 == 0:
|
|
print(f"[{i+len(chunk)}/{len(items)}]", flush=True)
|
|
json.dump(rows, open(OUT / "gate_probs.json", "w"), indent=1)
|
|
print("wrote", OUT / "gate_probs.json")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|