File size: 4,465 Bytes
4397e12 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 | """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()
|