File size: 4,807 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
"""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()