Download code/scripts/ceiling_read.py from darioooooo0o/tiny-agent-112m: direct link, hf CLI and curl.
- Browser
- Download file 4.47 kB
-
https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/scripts/ceiling_read.py
- Command line
-
hf download hf://darioooooo0o/tiny-agent-112m/code/scripts/ceiling_read.py
-
curl -L -o ceiling_read.py https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/scripts/ceiling_read.py
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) | |
| 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() | |