Download code/scripts/ceiling_search.py from darioooooo0o/tiny-agent-112m: direct link, hf CLI and curl.
- Browser
- Download file 4.81 kB
-
https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/scripts/ceiling_search.py
- Command line
-
hf download hf://darioooooo0o/tiny-agent-112m/code/scripts/ceiling_search.py
-
curl -L -o ceiling_search.py https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/scripts/ceiling_search.py
4.81 kB
| """Search ceiling for the HotpotQA search agent: with BM25 over the 2017 Wikipedia abstracts | |
| (HotpotQA's full-wiki corpus, ~5M paragraphs), how often do perfect queries retrieve the gold | |
| paragraphs? Stages: | |
| hop1 BM25(question) top-k contains a gold paragraph / both gold paragraphs | |
| hop2 BM25(other gold title) top-k contains it (the agent copied the bridge entity perfectly) | |
| copyable the other gold title appears verbatim in the first gold paragraph (an open(title) | |
| tool would make hop 2 a pure copy) | |
| $TA_PY scripts/ceiling_search.py --build # once: parse the corpus, build the index (CPU) | |
| $TA_PY scripts/ceiling_search.py --n 1000 | |
| """ | |
| import argparse | |
| import bz2 | |
| import json | |
| import os | |
| import random | |
| import re | |
| import tarfile | |
| import bm25s | |
| import pyarrow.parquet as pq | |
| import Stemmer | |
| from tiny_agent.text import DATA | |
| WIKI = f"{DATA}/wiki" | |
| DEV = f"{DATA}/hotpot/distractor/validation-00000-of-00001.parquet" | |
| _TAG = re.compile(r"<[^>]+>") | |
| def corpus(): | |
| """(title, text) for every abstract paragraph; hyperlink markup stripped.""" | |
| with tarfile.open(f"{WIKI}/enwiki-abstracts.tar.bz2", "r:bz2") as tf: | |
| for m in tf: | |
| if not m.name.endswith(".bz2"): | |
| continue | |
| for line in bz2.decompress(tf.extractfile(m).read()).decode().splitlines(): | |
| d = json.loads(line) | |
| text = _TAG.sub("", "".join(d["text"])) | |
| if text.strip(): | |
| yield d["title"], text | |
| def build(): | |
| titles, docs = [], [] | |
| for t, x in corpus(): | |
| titles.append(t) | |
| docs.append(f"{t}. {x}") | |
| print("paragraphs", len(docs), flush=True) | |
| stem = Stemmer.Stemmer("english") | |
| toks = bm25s.tokenize(docs, stopwords="en", stemmer=stem, show_progress=False) | |
| ret = bm25s.BM25() | |
| ret.index(toks, show_progress=False) | |
| ret.save(f"{WIKI}/bm25") | |
| json.dump(titles, open(f"{WIKI}/titles.json", "w")) | |
| print("index saved", flush=True) | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--build", action="store_true") | |
| ap.add_argument("--n", type=int, default=1000) | |
| a = ap.parse_args() | |
| if a.build: | |
| build() | |
| return | |
| titles = json.load(open(f"{WIKI}/titles.json")) | |
| ret = bm25s.BM25.load(f"{WIKI}/bm25") | |
| stem = Stemmer.Stemmer("english") | |
| rows = pq.read_table(DEV).to_pylist() | |
| random.Random(0).shuffle(rows) | |
| rows = rows[:a.n] | |
| def search(qs, k=10): | |
| res, _ = ret.retrieve(bm25s.tokenize(qs, stopwords="en", stemmer=stem, show_progress=False), | |
| k=k, show_progress=False, n_threads=16) | |
| return [[titles[i] for i in r] for r in res] | |
| hop1 = search([r["question"] for r in rows]) | |
| stats = {k: 0 for k in ("any@5", "both@5", "any@10", "both@10", "bridge_n", "hop2@1", "hop2@5", "copyable", "chain@10")} | |
| for r, h in zip(rows, hop1): | |
| gold = list(dict.fromkeys(r["supporting_facts"]["title"])) | |
| stats["any@5"] += any(g in h[:5] for g in gold) | |
| stats["both@5"] += all(g in h[:5] for g in gold) | |
| stats["any@10"] += any(g in h for g in gold) | |
| stats["both@10"] += all(g in h for g in gold) | |
| # bridge questions: hop 2 = search the gold paragraph that hop 1 missed, by its exact title | |
| bridges = [(r, h) for r, h in zip(rows, hop1) if r["type"] == "bridge"] | |
| found, other = [], [] | |
| for r, h in bridges: | |
| gold = list(dict.fromkeys(r["supporting_facts"]["title"])) | |
| f = [g for g in gold if g in h] | |
| if len(f) == 1 and len(gold) == 2: | |
| found.append((r, f[0])) | |
| other.append([g for g in gold if g != f[0]][0]) | |
| stats["bridge_n"] = len(bridges) | |
| hop2 = search(other, k=5) if other else [] | |
| for (r, f), o, h2 in zip(found, other, hop2): | |
| stats["hop2@1"] += h2[0] == o | |
| stats["hop2@5"] += o in h2 | |
| ctx = {t: "".join(s) for t, s in zip(r["context"]["title"], r["context"]["sentences"])} | |
| stats["copyable"] += o.lower() in ctx.get(f, "").lower() | |
| stats["chain@10"] = sum(all(g in h for g in dict.fromkeys(r["supporting_facts"]["title"])) for r, h in bridges) \ | |
| + stats["hop2@5"] | |
| n, nb = len(rows), max(1, len(bridges)) | |
| print(json.dumps({"n": n, "any@5": stats["any@5"] / n, "both@5": stats["both@5"] / n, "any@10": stats["any@10"] / n, | |
| "both@10": stats["both@10"] / n, "bridge_n": len(bridges), "bridge_one_gold_found": len(found) / nb, | |
| "hop2@1 (of those)": stats["hop2@1"] / max(1, len(found)), "hop2@5 (of those)": stats["hop2@5"] / max(1, len(found)), | |
| "copyable (of those)": stats["copyable"] / max(1, len(found)), | |
| "bridge both gold within 2 perfect hops": stats["chain@10"] / nb}, indent=1)) | |
| if __name__ == "__main__": | |
| main() | |