File size: 4,465 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
116
117
118
119
120
121
"""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)


@torch.no_grad()
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()