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