tiny-agent-112m / code /scripts /ceiling_search.py
darioooooo0o's picture
tiny-agent-112m: base + RL weights, tokenizer, code, model card
4397e12 verified
Raw History Blame Contribute Delete
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()