tiny-agent-112m / code /scripts /ceiling_read.py
darioooooo0o's picture
tiny-agent-112m: base + RL weights, tokenizer, code, model card
4397e12 verified
Raw History Blame Contribute Delete
4.47 kB
"""Reading ceiling for the HotpotQA search agent: given the right paragraphs (no search), how often
does the model produce the exact answer? Same prompt format as its synth_grounded training data
(READ_SYS, "Title: text" paragraphs, "Question: ..."). Greedy decoding, SQuAD-style EM/F1.
gold only the 2 supporting paragraphs -> the true reading ceiling
distractor all 10 paragraphs (training format) -> reading among noisy search results
$TA_PY scripts/ceiling_read.py --ckpt $TA_DATA/runs/main/final.pt --n 1000
"""
import argparse
import collections
import json
import random
import re
import string
import pyarrow.parquet as pq
import torch
from tokenizers import Tokenizer
from scripts.gen_synth import READ_SYS
from tiny_agent.chat import render
from tiny_agent.checkpoint import load_model
from tiny_agent.generate import KVCache, append
from tiny_agent.text import DATA
DEV = f"{DATA}/hotpot/distractor/validation-00000-of-00001.parquet"
def norm(s):
s = s.lower()
s = "".join(ch for ch in s if ch not in set(string.punctuation))
s = re.sub(r"\b(a|an|the)\b", " ", s)
return " ".join(s.split())
def f1(pred, gold):
p, g = norm(pred).split(), norm(gold).split()
common = collections.Counter(p) & collections.Counter(g)
k = sum(common.values())
if k == 0:
return 0.0
prec, rec = k / len(p), k / len(g)
return 2 * prec * rec / (prec + rec)
def prompt(r, setting):
titles, sents = r["context"]["title"], r["context"]["sentences"]
gold = set(r["supporting_facts"]["title"])
paras = [(t, s) for t, s in zip(titles, sents) if setting == "distractor" or t in gold]
ctx = "\n\n".join(f"{t}: {''.join(s)}" for t, s in paras)
return render([{"role": "system", "content": READ_SYS},
{"role": "user", "content": f"{ctx}\n\nQuestion: {r['question']}"}], add_generation_prompt=True)
@torch.no_grad()
def generate(model, tok, texts, max_new=48, B=64, device="xpu"):
end = tok.token_to_id("<|im_end|>")
outs = []
for i in range(0, len(texts), B):
ids = [tok.encode(t, add_special_tokens=False).ids[-3900:] for t in texts[i:i + B]]
b = len(ids)
cache = KVCache(model, b, 4096, device)
width = max(map(len, ids))
x = torch.zeros(b, width, dtype=torch.long)
for j, s in enumerate(ids):
x[j, :len(s)] = torch.tensor(s)
logits = append(model, cache, x.to(device), [len(s) for s in ids])
gen = [[] for _ in range(b)]
alive = [True] * b
for _ in range(max_new):
nxt = logits.argmax(-1)
nl = nxt.tolist()
for j in range(b):
if alive[j]:
if nl[j] == end:
alive[j] = False
else:
gen[j].append(nl[j])
if not any(alive):
break
logits = append(model, cache, nxt[:, None], [1 if a else 0 for a in alive])
outs += [tok.decode(g) for g in gen]
return outs
def answer_of(text):
return text.split("</think>")[-1].strip()
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--ckpt", required=True)
ap.add_argument("--n", type=int, default=1000)
ap.add_argument("--settings", default="gold,distractor")
ap.add_argument("--show", type=int, default=0)
a = ap.parse_args()
rows = pq.read_table(DEV).to_pylist()
random.Random(0).shuffle(rows)
rows = rows[:a.n]
tok = Tokenizer.from_file(f"{DATA}/tokenizer.json")
model = load_model(a.ckpt, dtype=torch.bfloat16).eval()
rep = {}
for setting in a.settings.split(","):
outs = generate(model, tok, [prompt(r, setting) for r in rows])
by = collections.defaultdict(lambda: [0, 0.0, 0])
for r, o in zip(rows, outs):
pred = answer_of(o)
for key in ("all", r["type"], "yes/no" if r["answer"].lower() in ("yes", "no") else "span"):
by[key][0] += norm(pred) == norm(r["answer"])
by[key][1] += f1(pred, r["answer"])
by[key][2] += 1
rep[setting] = {k: {"em": round(v[0] / v[2], 3), "f1": round(v[1] / v[2], 3), "n": v[2]} for k, v in by.items()}
for r, o in list(zip(rows, outs))[:a.show]:
print(f"[{setting}] Q: {r['question']} | gold: {r['answer']} | out: {o[:200]!r}")
print(json.dumps({"ckpt": a.ckpt, **rep}))
if __name__ == "__main__":
main()