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