File size: 2,991 Bytes
61b6fb9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Rebuild the typed-decisions card's two reference rows (Uniform, Prior) to pin metric definitions.

Card values (test, 2,000 decisions): Uniform acc 0.308 KL 0.444 Brier 0.238 ECE 0.169;
Prior (each question's train label frequencies) acc 0.470 KL 0.347 Brier 0.189 ECE 0.088.
Also writes the per-question teacher-agreement flag (label_agreement.argmax_agree) for the test.
"""
import collections
import json
import sys
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parent))
from metrics import summarize  # noqa: E402

REPO, REV = "LocalLLaMA/typed-decisions", "c76749ec58bd8c3d2ea706b31c333a9059c38f90"


def rows(split):
    import pyarrow.parquet as pq
    from huggingface_hub import hf_hub_download
    return pq.read_table(hf_hub_download(REPO, "all/%s-00000-of-00001.parquet" % split, repo_type="dataset",
                                         revision=REV)).to_pylist()


def keys_of(q):
    if q["type"] == "choice":
        return list(q["criteria"])
    if q["type"] == "score":
        return [str(i) for i in range(len(q["criteria"]))]
    return ["true", "false"]


def items(test, dist):
    out = []
    for r in test:
        qs, gold = json.loads(r["questions"]), json.loads(r["gold"])
        for qid, q in qs.items():
            k = keys_of(q)
            g = gold[qid]
            lab = str(g["label"]).lower() if q["type"] == "noul" else str(g["label"])
            soft = {(str(a).lower() if q["type"] == "noul" else str(a)): float(b) for a, b in g["probabilities"].items()}
            out.append({"id": r["id"], "qid": qid, "type": q["type"], "keys": k, "probs": dist(r, qid, k),
                        "gold": [lab], "soft": soft})
    return out


def main(out_path):
    test, train = rows("test"), rows("train")
    freq = collections.defaultdict(collections.Counter)
    for r in train:
        gold = json.loads(r["gold"])
        for qid, g in gold.items():
            lab = str(g["label"]).lower() if g.get("type") == "noul" else str(g["label"])
            freq[(r["workflow"], qid)][lab] += 1
    def uniform(r, qid, k):
        return [1 / len(k)] * len(k)
    def prior(r, qid, k):
        c = freq[(r["workflow"], qid)]
        t = sum(c[x] for x in k)
        return [c[x] / t for x in k]
    res = {}
    for name, f in (("uniform", uniform), ("prior", prior)):
        it = items(test, f)
        s = summarize(it)
        if name == "uniform":
            s["accuracy_expected_1_over_k"] = sum(1 / len(i["keys"]) for i in it) / len(it)
        res[name] = s
    agree = {}
    for r in test:
        la = json.loads(r["label_agreement"])
        for qid, v in la.items():
            agree["%s|%s" % (r["id"], qid)] = bool(v["argmax_agree"])
    res["argmax_agree_counts"] = dict(collections.Counter(agree.values()))
    Path(out_path).write_text(json.dumps({"reference_rows": res, "argmax_agree": agree}, indent=1) + "\n")
    print(json.dumps(res, indent=1))


if __name__ == "__main__":
    main(sys.argv[1])