File size: 3,773 Bytes
b2ebc95
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Compare two eval runs: solve rates, Wilson intervals, and whether the difference survives.

Usage: compare.py runs/A runs/B [...]
Prints a pairwise table. The p-value is a two-sided Fisher exact test on the 2x2 table of
solved/not — the right test for "did this actually move" when n is a hundred episodes and the
rate is under 30%. Overlapping Wilson intervals are reported too, since that is the check the
task brief asks for.
"""

import itertools
import json
import math
import os
import sys


def wilson(k, n, z=1.96):
    if n == 0:
        return 0.0, 0.0, 0.0
    p = k / n
    d = 1 + z * z / n
    c = (p + z * z / (2 * n)) / d
    h = z * math.sqrt(p * (1 - p) / n + z * z / (4 * n * n)) / d
    return p, max(0.0, c - h), min(1.0, c + h)


def load(path):
    """-> (solved_count, n, {task_name: solved_bool})"""
    f = os.path.join(path, "traces.jsonl")
    k = n = 0
    per = {}
    if not os.path.exists(f):
        return 0, 0, per
    for line in open(f):
        rec = json.loads(line)
        for t in rec.get("traces", []):
            rw = {k: v for k, v in (t.get("rewards") or {}).items()
                  if isinstance(v, dict) and v.get("score") is not None}
            if not rw:
                continue
            tot = sum(v["score"] * v.get("weight", 1.0) for v in rw.values())
            wsum = sum(v.get("weight", 1.0) for v in rw.values()) or 1.0
            solved = tot / wsum >= 0.999
            n += 1
            k += 1 if solved else 0
            name = ((t.get("task") or {}).get("data") or {}).get("name")
            if name is not None:
                per[name] = solved or per.get(name, False)
    return k, n, per


def mcnemar(pa, pb):
    """Paired test on the tasks both runs attempted. Returns (b, c, p, n_paired)."""
    shared = set(pa) & set(pb)
    b = sum(1 for t in shared if pa[t] and not pb[t])   # only A solved
    c = sum(1 for t in shared if pb[t] and not pa[t])   # only B solved
    m = b + c
    if m == 0:
        return b, c, 1.0, len(shared)
    from math import comb

    # exact two-sided binomial on the discordant pairs
    tail = sum(comb(m, i) for i in range(0, min(b, c) + 1)) / (2 ** m)
    return b, c, min(1.0, 2 * tail), len(shared)


def fisher(a, b, c, d):
    """Two-sided Fisher exact p for [[a,b],[c,d]]."""
    from math import comb

    def hyp(x, r1, r2, c1):
        return comb(r1, x) * comb(r2, c1 - x) / comb(r1 + r2, c1)

    r1, r2, c1 = a + b, c + d, a + c
    obs = hyp(a, r1, r2, c1)
    lo = max(0, c1 - r2)
    hi = min(r1, c1)
    return min(1.0, sum(hyp(x, r1, r2, c1) for x in range(lo, hi + 1) if hyp(x, r1, r2, c1) <= obs * (1 + 1e-9)))


def main():
    runs = sys.argv[1:]
    if len(runs) < 1:
        sys.exit(__doc__)
    stats = {}
    for r in runs:
        k, n, per = load(r)
        p, lo, hi = wilson(k, n)
        stats[r] = (k, n, p, lo, hi, per)
        print(f"{os.path.basename(r.rstrip('/')):<28} {k:>4}/{n:<4} {p:.3f}  ci95=[{lo:.3f},{hi:.3f}]")
    print()
    for a, b in itertools.combinations(runs, 2):
        ka, na, pa, loa, hia, pera = stats[a]
        kb, nb, pb, lob, hib, perb = stats[b]
        if not na or not nb:
            continue
        p = fisher(ka, na - ka, kb, nb - kb)
        onlya, onlyb, pm, npair = mcnemar(pera, perb)
        overlap = not (hia < lob or hib < loa)
        print(
            f"{os.path.basename(a.rstrip('/'))} vs {os.path.basename(b.rstrip('/'))}: "
            f"{pa:.3f} vs {pb:.3f}  diff={pb - pa:+.3f}  fisher_p={p:.3f}\n"
            f"    paired on {npair} tasks: only-A {onlya}, only-B {onlyb}, mcnemar_p={pm:.3f}  "
            f"| {'unpaired intervals OVERLAP' if overlap else 'unpaired intervals separate'}"
        )


if __name__ == "__main__":
    main()