File size: 4,015 Bytes
de8702b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Run the offline baselines in-process and list shortcut-solvable scored cases.

A case solved by a surface baseline does not require understanding negation,
scope or temporal order. The report also checks the acceptance gate:
overlap, last_clause and first_clause must stay at or below GATE_OVERALL overall and
GATE_CATEGORY per category (categories with fewer than GATE_MIN_N cases are
reported but not gated).
"""
import argparse
import sys
from collections import defaultdict
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parent.parent))

from benchmark.adapters.baselines import BaselineAdapter  # noqa: E402
from benchmark.dataset import BENCHMARK_VERSION, VERSIONS, load_cases, load_tasks  # noqa: E402
from benchmark.metrics import group_consistency  # noqa: E402

GATE_OVERALL = 0.45
GATE_CATEGORY = 0.55
GATE_MIN_N = 10
GATED = ("overlap", "last_clause", "first_clause")


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--version", default=BENCHMARK_VERSION, choices=sorted(VERSIONS))
    parser.add_argument("--split", default="public")
    parser.add_argument("--dataset", type=Path, help="report on this file instead of the split")
    args = parser.parse_args()

    spec = VERSIONS[args.version]
    tasks = load_tasks(spec["tasks"])
    path = args.dataset or spec["splits"][args.split]
    cases = [c for c in load_cases(path) if c.get("scored", True)]
    if not cases:
        raise SystemExit("No scored cases.")

    baselines = list(BaselineAdapter.MODELS)
    adapters = {b: BaselineAdapter(b, cases=cases, tasks=tasks) for b in baselines}

    hits = {}
    for case in cases:
        question = tasks[case["task_id"]]
        hits[case["id"]] = {
            b: adapter.predict(case["state"], question)["choice"] == case["expected"]
            for b, adapter in adapters.items()
        }

    totals = defaultdict(int)
    correct = defaultdict(lambda: defaultdict(int))
    for case in cases:
        totals[case["category"]] += 1
        for b in baselines:
            correct[case["category"]][b] += hits[case["id"]][b]

    header = f"{'category':20} {'n':>3} " + " ".join(f"{b:>12}" for b in baselines)
    print(header)
    print("-" * len(header))
    for category in sorted(totals):
        n = totals[category]
        cells = " ".join(f"{correct[category][b] / n:>12.0%}" for b in baselines)
        print(f"{category:20} {n:>3} {cells}")
    n = len(cases)
    overall = {b: sum(h[b] for h in hits.values()) / n for b in baselines}
    print("-" * len(header))
    print(f"{'OVERALL':20} {n:>3} " + " ".join(f"{overall[b]:>12.0%}" for b in baselines))

    if any(c.get("group_id") for c in cases):
        line = []
        for b in baselines:
            records = [
                {**c, "predicted": c["expected"] if hits[c["id"]][b] else None}
                for c in cases
            ]
            consistency, groups = group_consistency(records)
            line.append(f"{consistency:>12.0%}")
        print(f"{'GROUP CONSISTENCY':20} {groups:>3} " + " ".join(line))

    shortcut_models = [b for b in baselines if b != "majority"]
    shortcut = [c for c in cases if any(hits[c["id"]][b] for b in shortcut_models)]
    hard = [c for c in cases if not any(hits[c["id"]].values())]
    print(f"\nShortcut-solvable ({' or '.join(shortcut_models)} correct): {len(shortcut)}/{n}")
    print(f"Solved by no baseline: {len(hard)}/{n}")

    failures = []
    for b in GATED:
        if overall[b] > GATE_OVERALL:
            failures.append(f"{b} overall {overall[b]:.0%} > {GATE_OVERALL:.0%}")
        for category, total in totals.items():
            rate = correct[category][b] / total
            if total >= GATE_MIN_N and rate > GATE_CATEGORY:
                failures.append(f"{b} on {category} {rate:.0%} > {GATE_CATEGORY:.0%}")
    print("\nACCEPTANCE GATE:", "PASS" if not failures else "FAIL")
    for f in failures:
        print(f"  - {f}")


if __name__ == "__main__":
    main()