225 lines
8.8 KiB
Python
225 lines
8.8 KiB
Python
#!/usr/bin/env python3
|
||
"""Exact-geometry Gate A/B for the learned movement (job j131).
|
||
|
||
Question: does labelling/resolving a wave by the REAL bullet endpoint (the
|
||
exact origin->endpoint straight line, available live from `onHitByBullet` and
|
||
from a bullet-vs-bullet intercept) fix the danger-map inversion that j128
|
||
measured with the histogram label (`corr = -0.342`) and that j130 replaced with
|
||
a suspicious live-computable proxy (`corr = +0.566`)?
|
||
|
||
This is the SAME corpus, SAME per-shot extraction and SAME metric as
|
||
`outcome_label_gate.py` (which job j130 used), so the three danger maps are
|
||
computed under ONE consistent computation and are directly comparable:
|
||
|
||
(a) histogram label danger(g) = P(arrival bin = g) (j128)
|
||
(b) outcome proxy label danger(g) = P(hit and |g - b_our| <= w) (j130 live)
|
||
(c) EXACT bullet line danger(g) = P(|g - b_bullet| <= w) (this job)
|
||
|
||
`b_our` is the GF of OUR position at the nominal arrival tick; `b_bullet` is
|
||
the GF of the bullet's own straight line (from the recorded fire direction -
|
||
exactly the line the real endpoint would give). `w` is the body half-width as an
|
||
angle, in bins. Both correlations use the SAME realised per-bin hit rate
|
||
`P(hit | b_our = g)` (the j128 metric), and a second, bullet-conditioned target
|
||
is printed as a cross-check.
|
||
|
||
Gate B: held-out per-candidate log-loss of the EXACT (bullet-line) label,
|
||
state-conditional vs state-free, the same measurement j130 ran for its proxy.
|
||
|
||
Run:
|
||
python3 common_libs/tests/exact_geometry_gate.py \
|
||
--corpus /tmp/tfil_ab2/out \
|
||
--report common_libs/tests/fixtures/exact_geometry_gate_report.txt
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import os
|
||
import statistics
|
||
import sys
|
||
|
||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||
|
||
import outcome_label_gate as olg # validated extraction + metric
|
||
import analyze_drussgt_dodge_vs_power as adp
|
||
|
||
NBINS = olg.NBINS
|
||
|
||
|
||
def p_hist(recs, b):
|
||
return sum(1 for r in recs if r["b_our"] == b) / len(recs)
|
||
|
||
|
||
def p_proxy(recs, b):
|
||
"""j130 live label: P(hit and |b - b_our| <= w)."""
|
||
return statistics.fmean(1 if (r["hit"] >= 0.5 and abs(b - r["b_our"]) <= r["w"])
|
||
else 0 for r in recs)
|
||
|
||
|
||
def p_exact(recs, b):
|
||
"""Exact bullet-line label: P(|b - b_bullet| <= w)."""
|
||
return statistics.fmean(1 if abs(b - r["b_bullet"]) <= r["w"] else 0
|
||
for r in recs)
|
||
|
||
|
||
def p_exact_hit(recs, b):
|
||
"""Exact bullet-line AND hit: P(hit and |b - b_bullet| <= w)."""
|
||
return statistics.fmean(1 if (r["hit"] >= 0.5 and abs(b - r["b_bullet"]) <= r["w"])
|
||
else 0 for r in recs)
|
||
|
||
|
||
def correlations(recs):
|
||
n = [0] * NBINS
|
||
h = [0] * NBINS
|
||
for r in recs:
|
||
n[r["b_our"]] += 1
|
||
h[r["b_our"]] += r["hit"]
|
||
used = [b for b in range(NBINS) if n[b] > 0]
|
||
rate_our = [h[b] / n[b] for b in used]
|
||
|
||
nb = [0] * NBINS
|
||
hb = [0] * NBINS
|
||
for r in recs:
|
||
nb[r["b_bullet"]] += 1
|
||
hb[r["b_bullet"]] += r["hit"]
|
||
usedb = [b for b in range(NBINS) if nb[b] > 0]
|
||
rate_bullet = [hb[b] / nb[b] for b in usedb]
|
||
|
||
maps = {
|
||
"histogram (j128): P(arrival = g)": [p_hist(recs, b) for b in used],
|
||
"outcome proxy (j130 live): P(hit & |g-b_our|<=w)": [p_proxy(recs, b) for b in used],
|
||
"EXACT bullet line: P(|g-b_bullet|<=w)": [p_exact(recs, b) for b in used],
|
||
"EXACT bullet line & hit: P(hit & |g-b_bullet|<=w)": [p_exact_hit(recs, b) for b in used],
|
||
}
|
||
out = {}
|
||
for name, d in maps.items():
|
||
out[name] = (statistics.correlation(d, rate_our),
|
||
statistics.correlation([d[used.index(b)] if b in used else 0.0
|
||
for b in usedb], rate_bullet))
|
||
return out, used, rate_our, usedb, rate_bullet
|
||
|
||
|
||
def run_split_exact(recs, seed, decay=128, shift=1):
|
||
tr_b, te_b = olg.split_battles({r["battle"] for r in recs}, seed)
|
||
tr = [r for r in recs if r["battle"] in tr_b]
|
||
te = [r for r in recs if r["battle"] in te_b]
|
||
edges = dict(olg.CANON)
|
||
om = olg.OutcomeModel(decay, shift, state_free=False)
|
||
om0 = olg.OutcomeModel(decay, shift, state_free=True)
|
||
for r in tr:
|
||
st = olg.code_of(r, edges)
|
||
for g in range(NBINS):
|
||
lab = 1 if (r["hit"] >= 0.5 and abs(g - r["b_bullet"]) <= r["w"]) else 0
|
||
om.learn(st, g, lab)
|
||
om0.learn(st, g, lab)
|
||
ll_s, ll_g, hit_out = [], [], []
|
||
for r in te:
|
||
st = olg.code_of(r, edges)
|
||
go = min(range(NBINS), key=lambda g: om.predict_hit(st, g))
|
||
real = lambda g: 1 if abs(g - r["b_bullet"]) <= r["w"] else 0
|
||
hit_out.append(real(go))
|
||
for g in range(NBINS):
|
||
y = 1 if (r["hit"] >= 0.5 and abs(g - r["b_bullet"]) <= r["w"]) else 0
|
||
ll_s.append(-olg.log2(om.predict_hit(st, g)) if y
|
||
else -olg.log2(1.0 - om.predict_hit(st, g)))
|
||
ll_g.append(-olg.log2(om0.predict_hit(st, g)) if y
|
||
else -olg.log2(1.0 - om0.predict_hit(st, g)))
|
||
return dict(seed=seed,
|
||
ll_state=statistics.fmean(ll_s),
|
||
ll_statefree=statistics.fmean(ll_g),
|
||
delta=statistics.fmean([a - b for a, b in zip(ll_s, ll_g)]),
|
||
hit_out=statistics.fmean(hit_out))
|
||
|
||
|
||
def main():
|
||
ap = argparse.ArgumentParser()
|
||
ap.add_argument("--corpus", default="/tmp/tfil_ab2/out")
|
||
ap.add_argument("--report", default=None)
|
||
ap.add_argument("--seeds", type=int, default=3)
|
||
args = ap.parse_args()
|
||
|
||
runs = adp.discover_tfil(args.corpus)
|
||
recs = olg.extract(runs)
|
||
lines = []
|
||
|
||
def out(s=""):
|
||
print(s)
|
||
lines.append(s)
|
||
|
||
out("# Exact-geometry Gate A/B — learned movement (job j131)")
|
||
out()
|
||
out(f"corpus : {args.corpus}")
|
||
out(f"battles : {len(runs)}")
|
||
out(f"shots : {len(recs)}")
|
||
out(f"base hit : {statistics.fmean(r['hit'] for r in recs)*100:.2f}%")
|
||
out("state : vlat, dist, room, turn (module's 4 fields, canonical edges)")
|
||
out()
|
||
|
||
corr, used, rate_our, usedb, rate_bullet = correlations(recs)
|
||
out("## A. danger-map alignment (ONE consistent computation)")
|
||
out()
|
||
out("corr( danger(g) , P(hit | b_our = g) ) [the j128 metric, = -0.342 hist]")
|
||
out("corr( danger(g) , P(hit | b_bullet = g) ) [same danger, bullet-conditioned target]")
|
||
out()
|
||
out("| danger map | corr vs P(hit\\|b_our=g) | corr vs P(hit\\|b_bullet=g) |")
|
||
out("|---|---:|---:|")
|
||
for name, (c_our, c_bul) in corr.items():
|
||
out(f"| {name} | {c_our:+.3f} | {c_bul:+.3f} |")
|
||
out()
|
||
out("Negative = minimising the danger steers INTO where the observed hits")
|
||
out("happen (the j128 defect). The exact bullet line is the physically")
|
||
out("correct 'would this wave hit me at g' map; if its correlation is still")
|
||
out("negative, exact geometry does NOT fix the inversion.")
|
||
out()
|
||
|
||
out("| bin | P(hit\\|b_our) | P(hit\\|b_bullet) | hist danger | proxy danger | exact danger |")
|
||
out("|---:|---:|---:|---:|---:|---:|")
|
||
rb = {b: rate_bullet[usedb.index(b)] for b in usedb}
|
||
for b in used:
|
||
rb_str = f"{rb[b]*100:.1f}%" if b in rb else "—"
|
||
out(f"| {b} | {rate_our[used.index(b)]*100:.1f}% | "
|
||
f"{rb_str} | "
|
||
f"{p_hist(recs, b):.3f} | {p_proxy(recs, b):.3f} | {p_exact(recs, b):.3f} |")
|
||
out()
|
||
|
||
per = [run_split_exact(recs, s) for s in range(args.seeds)]
|
||
ll_s = statistics.fmean(p["ll_state"] for p in per)
|
||
ll_g = statistics.fmean(p["ll_statefree"] for p in per)
|
||
out("## B. state-conditional information under the EXACT bullet-line label")
|
||
out()
|
||
out("held-out per-candidate log-loss (bits) of the exact label, "
|
||
"state-conditional vs state-free (same rows, same split):")
|
||
out()
|
||
out("| model | log-loss (bits) |")
|
||
out("|---|---:|")
|
||
out(f"| state-free P(label | g) | {ll_g:.4f} |")
|
||
out(f"| state-conditional P(label | state, g) | {ll_s:.4f} |")
|
||
out(f"| Δ (state − state-free) | {ll_s - ll_g:+.4f} |")
|
||
out()
|
||
neg = sum(1 for p in per if p["delta"] < 0)
|
||
out(f"state conditioning is better in {neg}/{len(per)} splits "
|
||
f"(negative Δ = better).")
|
||
out()
|
||
|
||
out("## C. open-loop decision counterfactual (VETO ONLY)")
|
||
out()
|
||
out("argmin_g danger with the recorded bullet line as ground truth:")
|
||
out()
|
||
out(f"| exact-label argmin (j131) | "
|
||
f"{statistics.fmean(p['hit_out'] for p in per)*100:.2f}% |")
|
||
out()
|
||
|
||
out("## MEASURED vs INFERRED")
|
||
out()
|
||
out("* MEASURED: every number above, on the recorded corpus.")
|
||
out("* INFERRED: that an offline alignment transfers live — it cannot, the")
|
||
out(" corpus is open loop (`docs/offline_harness_trust.md`).")
|
||
|
||
if args.report:
|
||
os.makedirs(os.path.dirname(args.report), exist_ok=True)
|
||
with open(args.report, "w") as f:
|
||
f.write("\n".join(lines) + "\n")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|