#!/usr/bin/env python3 """TMComposites GATE — analyse the per-sample confidence dump. Reads the JSONL produced by ``common_libs/tests/measure_tmcomposites.nim`` (one row per resolved virtual bullet: fixture, split, gun, tick, bin, conf, hit, relDeg, range, missPx) and answers the three questions of ``docs/tmcomposites_gate.md``: A. Is each gun's intrinsic confidence FAITHFUL? Rank samples by the gun's own confidence; report the accuracy-vs-confidence curve, Spearman(confidence, hit), and top-half vs bottom-half accuracy with a two-proportion p-value. B. Are the faithful guns COMPLEMENTARY SPECIALISTS? For each pair, on the samples where A's alpha-normalised confidence beats B's, is A the more accurate one? Report the two slices and the win rates. C. Does the Eq-8 alpha-normalised confidence-weighted composite beat the best single gun on held-out battles, and is the gain attributable to competence? The recorder already ran ``Composite`` and its within-sample confidence-shuffle control ``CompositeShuf`` through the SAME virtual-bullet geometry, so the comparison is paired (McNemar). Pure stdlib (no numpy/scipy on this machine). MEASURED = every number printed; INFERRED = the causal reading in the doc. """ from __future__ import annotations import argparse import collections import json import math import random import sys DETERMINISTIC = ["HeadOn", "Linear", "Circular", "WallBounce", "Accel", "StopShot", "Displace", "AvgLead"] # ───────────────────────────── stats (stdlib) ────────────────────────────── def mean(xs): return sum(xs) / len(xs) if xs else float("nan") def spearman(xs, ys): """Spearman rho with average ranks for ties, and a normal-approx p.""" n = len(xs) if n < 3: return float("nan"), float("nan") def ranks(v): order = sorted(range(n), key=lambda i: v[i]) r = [0.0] * n i = 0 while i < n: j = i while j + 1 < n and v[order[j + 1]] == v[order[i]]: j += 1 avg = (i + j) / 2.0 + 1.0 for k in range(i, j + 1): r[order[k]] = avg i = j + 1 return r rx, ry = ranks(xs), ranks(ys) mx, my = mean(rx), mean(ry) num = sum((a - mx) * (b - my) for a, b in zip(rx, ry)) den = math.sqrt(sum((a - mx) ** 2 for a in rx) * sum((b - my) ** 2 for b in ry)) if den == 0: return 0.0, 1.0 rho = num / den rho = max(-1.0, min(1.0, rho)) z = rho * math.sqrt(n - 1) p = math.erfc(abs(z) / math.sqrt(2.0)) return rho, p def norm_two_prop(z): return math.erfc(abs(z) / math.sqrt(2.0)) def two_prop_p(h1, n1, h2, n2): if n1 == 0 or n2 == 0: return float("nan"), float("nan") p1, p2 = h1 / n1, h2 / n2 p = (h1 + h2) / (n1 + n2) se = math.sqrt(p * (1 - p) * (1 / n1 + 1 / n2)) if se == 0: return float("nan"), float("nan") z = (p1 - p2) / se return z, norm_two_prop(z) def binom_two_sided(k, n, p=0.5): """Exact two-sided binomial p (used for McNemar's discordant pairs).""" if n == 0: return 1.0 def pmf(i): return math.comb(n, i) * p ** i * (1 - p) ** (n - i) obs = pmf(k) tot = 0.0 for i in range(n + 1): if pmf(i) <= obs + 1e-12: tot += pmf(i) return min(1.0, tot) def mcnemar_p(a_hit_b_miss, a_miss_b_hit): """McNemar: exact binomial for small discordant counts, normal approx for large (the exact branch overflows math.comb for n in the hundred-thousands).""" b, c = a_hit_b_miss, a_miss_b_hit n = b + c if n == 0: return 1.0 if n < 500: return binom_two_sided(b, n) z = (b - c) / math.sqrt(n) return math.erfc(abs(z) / math.sqrt(2.0)) # ───────────────────────────── data loading ──────────────────────────────── class Dump: def __init__(self, path): self.rows = [] # list of dicts self.by_gun = collections.defaultdict(list) # key -> {gun: (conf, hit)} self.by_sample = collections.defaultdict(dict) with open(path) as f: for line in f: line = line.strip() if not line: continue o = json.loads(line) self.rows.append(o) self.by_gun[o["gun"]].append(o) self.by_sample[(o["fixture"], o["tick"], o["bin"])][o["gun"]] = ( o["conf"], bool(o["hit"])) def guns(self): return sorted(self.by_gun.keys()) def split(self, gun, split): return [r for r in self.by_gun[gun] if r["split"] == split] def faithful_curve(rows, bins=10): rows = sorted(rows, key=lambda r: r["conf"]) n = len(rows) if n == 0: return [] out = [] for b in range(bins): lo = b * n // bins hi = (b + 1) * n // bins chunk = rows[lo:hi] if not chunk: continue out.append(dict( lo=chunk[0]["conf"], hi=chunk[-1]["conf"], n=len(chunk), hit=mean([1.0 if r["hit"] else 0.0 for r in chunk]), )) return out def faithfulness_report(dump): """Question A: per-gun faithfulness on the pooled test samples.""" out = {} for gun in dump.guns(): rows = dump.split(gun, "test") if not rows: continue confs = [r["conf"] for r in rows] hits = [1.0 if r["hit"] else 0.0 for r in rows] nz = sum(1 for c in confs if c > 1e-12) rho, p = spearman(confs, hits) order = sorted(range(len(rows)), key=lambda i: confs[i]) half = len(order) // 2 lo_idx, hi_idx = order[:half], order[half:] h_lo = sum(hits[i] for i in lo_idx) h_hi = sum(hits[i] for i in hi_idx) z, phalf = two_prop_p(h_hi, len(hi_idx), h_lo, len(lo_idx)) out[gun] = dict( n=len(rows), nonzero_conf=nz, base=mean(hits), rho=rho, rho_p=p, bottom_half_acc=(h_lo / len(lo_idx)) if lo_idx else float("nan"), top_half_acc=(h_hi / len(hi_idx)) if hi_idx else float("nan"), half_z=z, half_p=phalf, curve=faithful_curve(rows), ) return out def alphas(dump): """alpha_t = max-min of the gun's confidence over the TRAIN samples (Eq 7).""" out = {} for gun in dump.guns(): rows = dump.split(gun, "train") if not rows: continue cs = [r["conf"] for r in rows] out[gun] = max(1e-12, max(cs) - min(min(cs), 0.0)) return out def normalised(conf, gun, alpha): return conf / alpha.get(gun, 1.0) def complementarity(dump, alphas_, faithful): """Question B: pairwise complementary slices over the test samples. For each pair we split the samples by which gun has the higher alpha-normalised confidence and, ON EACH SLICE, measure BOTH guns' accuracy. A pair is complementary when each gun is the more accurate one on its own winning slice, and each slice is a substantial (>=10%) share. """ tfs = test_fixtures(dump) pairs = [] names = [g for g in faithful if faithful[g]["nonzero_conf"] > 0 and not g.startswith("Composite")] for i in range(len(names)): for j in range(i + 1, len(names)): A, B = names[i], names[j] a_win = b_win = 0 # hits ON the A-winning slice, and ON the B-winning slice aA = bA = aB = bB = 0 aA_only = bA_only = aB_only = bB_only = 0 for key, guns in dump.by_sample.items(): if key[0] not in tfs: continue if A not in guns or B not in guns: continue ca = normalised(guns[A][0], A, alphas_) cb = normalised(guns[B][0], B, alphas_) if ca <= 0 and cb <= 0: continue ha, hb = int(guns[A][1]), int(guns[B][1]) if ca > cb: a_win += 1; aA += ha; bA += hb if ha and not hb: aA_only += 1 elif hb and not ha: bA_only += 1 elif cb > ca: b_win += 1; aB += ha; bB += hb if ha and not hb: aB_only += 1 elif hb and not ha: bB_only += 1 both = a_win + b_win def frac(h, n): return h / n if n else float("nan") accA_on_A, accB_on_A = frac(aA, a_win), frac(bA, a_win) accA_on_B, accB_on_B = frac(aB, b_win), frac(bB, b_win) comp = (both > 0 and a_win >= 0.10 * both and b_win >= 0.10 * both and accA_on_A > accB_on_A and accB_on_B > accA_on_B) pairs.append(dict( A=A, B=B, a_win=a_win, b_win=b_win, both=both, A_acc_on_A_slice=accA_on_A, B_acc_on_A_slice=accB_on_A, A_acc_on_B_slice=accA_on_B, B_acc_on_B_slice=accB_on_B, p_on_A_slice=mcnemar_p(aA_only, bA_only), p_on_B_slice=mcnemar_p(bB_only, aB_only), complementary=comp)) return pairs def compl_ok(a_win, b_win, both, aa, ba): """Retained for backwards compatibility; superseded by the per-slice test in `complementarity`.""" if both == 0 or a_win == 0 or b_win == 0: return False return a_win >= 0.10 * both and b_win >= 0.10 * both # ───────────────────────── composite comparison ──────────────────────────── def paired(dump, gun_a, gun_b, tfs): """Return (a_hit_b_miss, a_miss_b_hit, a_hits, b_hits) over the test fixtures.""" ab = ba = na = nb = 0 for key, guns in dump.by_sample.items(): if key[0] not in tfs: continue if gun_a in guns and gun_b in guns: a = guns[gun_a][1] b = guns[gun_b][1] if a and not b: ab += 1 elif b and not a: ba += 1 if a: na += 1 if b: nb += 1 return ab, ba, na, nb def test_fixtures(dump): out = set() for r in dump.rows: if r["split"] == "test": out.add(r["fixture"]) return out def composite_report(dump): """Question C: composite vs best single vs shuffle control on test.""" tfs = test_fixtures(dump) acc = {} n = {} for gun in dump.guns(): rows = [r for r in dump.by_gun[gun] if r["split"] == "test"] if not rows: continue acc[gun] = mean([1.0 if r["hit"] else 0.0 for r in rows]) n[gun] = len(rows) members = [g for g in dump.guns() if not g.startswith("Composite") and g not in DETERMINISTIC] best_member = max(members, key=lambda g: acc[g]) if members else None res = dict(acc=acc, n=n, best_member=best_member) # Oracle ceiling: if a perfect per-sample selector could pick ANY member, # how often would it hit? This bounds what a member-selection composite # could ever reach (the vote can do worse but not better than this). oracle = 0 oracle_n = 0 for key, guns in dump.by_sample.items(): if key[0] not in tfs: continue hits = [guns[m][1] for m in members if m in guns] if hits: oracle += int(any(hits)) oracle_n += 1 res["member_oracle"] = (oracle / oracle_n) if oracle_n else float("nan") res["member_oracle_n"] = oracle_n comps = [g for g in dump.guns() if g.startswith("Composite") and not g.endswith("Shuf")] res["composites"] = {} for comp in comps: entry = {} if best_member: ab, ba, na, nb = paired(dump, comp, best_member, tfs) entry["vs_best"] = dict(best=best_member, comp_hit=na, best_hit=nb, total=n[comp], mcnemar_ab=ab, mcnemar_ba=ba, p=mcnemar_p(ab, ba)) if "Pattern" in acc: ab, ba, na, nb = paired(dump, comp, "Pattern", tfs) entry["vs_pattern"] = dict(comp_hit=na, pattern_hit=nb, total=n[comp], mcnemar_ab=ab, mcnemar_ba=ba, p=mcnemar_p(ab, ba)) shuf = comp + "Shuf" if shuf in acc: ab, ba, na, nb = paired(dump, comp, shuf, tfs) entry["vs_shuffle"] = dict(comp_hit=na, shuffle_hit=nb, total=n[comp], mcnemar_ab=ab, mcnemar_ba=ba, p=mcnemar_p(ab, ba)) res["composites"][comp] = entry # per-fixture composite vs best single vs shuffle per = {} for fx in sorted(tfs): row = {} for gun in ([best_member] if best_member else []) + comps: rows = [r for r in dump.by_gun[gun] if r["split"] == "test" and r["fixture"] == fx] if rows: row[gun] = dict(n=len(rows), acc=mean([1.0 if r["hit"] else 0.0 for r in rows])) per[fx] = row res["per_fixture"] = per return res # ───────────────────────────────── main ──────────────────────────────────── def main(): ap = argparse.ArgumentParser() ap.add_argument("--input", default="/tmp/tmc_full.jsonl") ap.add_argument("--json", default=None) args = ap.parse_args() dump = Dump(args.input) print(f"rows={len(dump.rows)} guns={len(dump.guns())} " f"test fixtures={sorted(test_fixtures(dump))}") a = faithfulness_report(dump) print("\n=== A. FAITHFULNESS (test samples; rank by own confidence) ===") print(f"{'gun':<14}{'n':>7}{'nonzero':>8}{'base%':>7}{'rho':>8}{'rho_p':>9}" f"{'bot%':>7}{'top%':>7}{'z':>7}{'p':>9} verdict") verdicts = {} for gun in sorted(a, key=lambda g: -a[g]["base"]): r = a[gun] if r["nonzero_conf"] == 0: v = "NO SIGNAL" elif r["rho_p"] < 0.01 and r["rho"] > 0.05: v = "FAITHFUL" elif r["rho_p"] < 0.01 and r["rho"] < -0.05: v = "ANTI-FAITHFUL" else: v = "USELESS" verdicts[gun] = v print(f"{gun:<14}{r['n']:>7}{r['nonzero_conf']:>8}{100*r['base']:>7.2f}" f"{r['rho']:>8.3f}{r['rho_p']:>9.2g}{100*r['bottom_half_acc']:>7.2f}" f"{100*r['top_half_acc']:>7.2f}{r['half_z']:>7.2f}{r['half_p']:>9.2g} {v}") print("\ncurves (deciles, low->high confidence):") for gun in sorted(a, key=lambda g: -a[g]["base"]): if verdicts[gun] == "NO SIGNAL": continue cur = " ".join(f"{100*c['hit']:.0f}%" for c in a[gun]["curve"]) print(f" {gun:<14} {cur}") al = alphas(dump) pairs = complementarity(dump, al, a) print("\n=== B. PAIRWISE COMPLEMENTARITY (test; alpha-normalised confidence) ===") print(f"{'A':<14}{'B':<14}{'A_wins':>8}{'A|A':>7}{'B|A':>7}{'p_A':>9}| {'B_wins':>8}{'A|B':>7}{'B|B':>7}{'p_B':>9} comp") for p in pairs: print(f"{p['A']:<14}{p['B']:<14}{p['a_win']:>8}" f"{100*p['A_acc_on_A_slice']:>7.1f}{100*p['B_acc_on_A_slice']:>7.1f}{p['p_on_A_slice']:>9.2g}| " f"{p['b_win']:>8}{100*p['A_acc_on_B_slice']:>7.1f}" f"{100*p['B_acc_on_B_slice']:>7.1f}{p['p_on_B_slice']:>9.2g} {p['complementary']}") print(" (A|A = A's accuracy on the slice A wins; B|A = B's accuracy on that same slice; etc.)") c = composite_report(dump) print("\n=== C. COMPOSITE vs BEST SINGLE vs SHUFFLE CONTROL (test) ===") for gun in sorted(c["acc"], key=lambda g: -c["acc"][g]): print(f" {gun:<14} {100*c['acc'][gun]:>6.2f}% n={c['n'][gun]}") if "member_oracle" in c: print(f" member_oracle (any member hits, per sample) : {100*c['member_oracle']:.2f}%") for comp, entry in c.get("composites", {}).items(): print(f" {comp}:") for k, r in entry.items(): print(f" {k}: {r}") print("\nper-fixture:") for fx, row in c["per_fixture"].items(): parts = " ".join(f"{g}={100*v['acc']:.1f}%({v['n']})" for g, v in row.items()) print(f" {fx:<32} {parts}") if args.json: blob = dict( faithfulness=a, alphas=al, verdicts=verdicts, complementarity=pairs, composite=c, test_fixtures=sorted(test_fixtures(dump)), ) with open(args.json, "w") as f: json.dump(blob, f, indent=2) print(f"\n[json] wrote {args.json}") if __name__ == "__main__": main()