420 lines
15 KiB
Python
420 lines
15 KiB
Python
#!/usr/bin/env python3
|
||
"""Gate A (OFFLINE VETO) for the OUTCOME-labelled learned movement (job j130).
|
||
|
||
Question: does labelling a resolved wave by its OUTCOME (would this wave have
|
||
hit me at candidate direction g?) — instead of by the GF bin the wave crossed
|
||
at — fix the measured inversion in `learned_surfer_gate.py` section F, where
|
||
`corr( P(arrival bin), P(hit | arrival bin) ) = -0.342` over the 31 bins, i.e.
|
||
minimising the resolved-position histogram steers INTO the bullets?
|
||
|
||
This is a VETO-ONLY harness (`docs/offline_harness_trust.md`): the live panel is
|
||
the decider, because the corpus was recorded while the enemy gun reacted to a
|
||
DIFFERENT mover.
|
||
|
||
--------------------------------------------------------------------------------
|
||
WHAT IS MEASURED
|
||
--------------------------------------------------------------------------------
|
||
Corpus `/tmp/tfil_ab2/out` (70 recorded battles, ~55k shots). Per shot fired by
|
||
side `s` at the dodging subject side `e` (the same instrument as
|
||
`learned_surfer_gate.py`):
|
||
* state at the fire tick: vlat, dist, room, turn (the module's 4 fields);
|
||
* `b_our` the 31-bin GF of the dodger's position at the NOMINAL arrival
|
||
tick (the histogram label the j128 module trains on);
|
||
* `b_bullet` the 31-bin GF of the bullet's own straight line (from the fire
|
||
event's `dir`) — the physically correct "arrival bin";
|
||
* `w` the body-width angular tolerance in bins: asin(R/d)/maxEA, R=18;
|
||
* `hit` the real server hit/miss outcome.
|
||
|
||
The label the MODULE can actually compute live (it cannot see bullet bodies) is
|
||
`hit(state, g) = hit and |g - b_our| <= w`
|
||
i.e. "this wave hit me at b; it would also have hit me at any g within a body
|
||
width of b". This is dense: every wave labels all 31 candidates.
|
||
|
||
Reported:
|
||
A. corr( LEARNED DANGER(g) , realised P(hit | b_our=g) ) over the 31 bins —
|
||
the exact metric that reads -0.342 for the histogram label. Computed for
|
||
the histogram danger, the module's outcome (hit-window) danger, and the
|
||
pure geometric bullet-line danger (the counterfactual if bullet bodies
|
||
were visible).
|
||
B. state-conditional information under the outcome label: held-out log-loss
|
||
of P(label | state, g) vs P(label | g) (state-free), paired per battle.
|
||
C. the open-loop decision counterfactual: if the mover picks argmin_g danger,
|
||
what fraction of held-out waves would still hit it (using b_bullet as the
|
||
ground truth)? Reported for the histogram and the outcome label, plus the
|
||
actual recorded trajectory as a floor.
|
||
|
||
Run:
|
||
python3 common_libs/tests/outcome_label_gate.py \
|
||
--corpus /tmp/tfil_ab2/out \
|
||
--report common_libs/tests/fixtures/outcome_label_gate_report.txt
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import math
|
||
import os
|
||
import random
|
||
import statistics
|
||
import sys
|
||
|
||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||
|
||
import analyze_drussgt_dodge_vs_power as adp # validated per-shot geometry
|
||
|
||
NBINS = 31
|
||
R = 18.0
|
||
FIELDS = ["vlat", "dist", "room", "turn"]
|
||
|
||
# The canonical, frozen edges hard-coded into `movements/learned_surfer.nim`
|
||
# (corpus quantiles). Using the deployed quantiser keeps the gate faithful to
|
||
# the module.
|
||
CANON = {
|
||
"vlat": [-6.736, 0.000, 6.753],
|
||
"dist": [431.321, 487.612, 552.670],
|
||
"room": [137.965, 206.589, 296.753],
|
||
"turn": [-0.142, 0.000, 0.105],
|
||
}
|
||
|
||
|
||
def wrap180(a):
|
||
return ((a + 180.0) % 360.0) - 180.0
|
||
|
||
|
||
def gf_to_bin(gf):
|
||
v = max(-1.0, min(1.0, gf))
|
||
return max(0, min(NBINS - 1, int(round((v + 1.0) * 0.5 * (NBINS - 1)))))
|
||
|
||
|
||
def bin_to_gf(i):
|
||
return i / (NBINS - 1) * 2.0 - 1.0
|
||
|
||
|
||
def code_of(sample, edges):
|
||
c = 0
|
||
for f in FIELDS:
|
||
cc = 0
|
||
for e in edges[f]:
|
||
if sample[f] > e:
|
||
cc += 1
|
||
c = c * 4 + cc
|
||
return c
|
||
|
||
|
||
# ------------------------------------------------------------------ extraction
|
||
def extract(runs):
|
||
recs = []
|
||
for run in runs:
|
||
bn = (os.path.basename(os.path.dirname(run.cap_path)) + "/" +
|
||
os.path.basename(run.cap_path).replace(".jsonl", ""))
|
||
for sh in run.shots():
|
||
t0, rnd = sh["tick"], sh["rnd"]
|
||
rs = run.start[rnd]
|
||
re = rs + run.count[rnd] - 1
|
||
r0 = run.by_tick.get(t0)
|
||
rp = run.by_tick.get(t0 - 1)
|
||
if r0 is None or rp is None or t0 <= rs:
|
||
continue
|
||
p0x, p0y, speed = sh["_x"], sh["_y"], sh["speed"]
|
||
tb0 = math.atan2(r0["ey"] - p0y, r0["ex"] - p0x)
|
||
ux, uy = math.cos(tb0), math.sin(tb0)
|
||
|
||
def lat_of(r):
|
||
dx, dy = r["ex"] - p0x, r["ey"] - p0y
|
||
return dx * (-uy) + dy * ux
|
||
|
||
lat = lat_of(r0)
|
||
vlat = lat - lat_of(rp)
|
||
dist = math.hypot(r0["ex"] - p0x, r0["ey"] - p0y)
|
||
sg = 1.0 if vlat >= 0 else -1.0
|
||
dxn, dyn = -uy * sg, ux * sg
|
||
t = float("inf")
|
||
for p, d, lo, hi in ((r0["ex"], dxn, R, 800 - R),
|
||
(r0["ey"], dyn, R, 600 - R)):
|
||
if abs(d) > 1e-9:
|
||
cand = (hi - p) / d if d > 0 else (lo - p) / d
|
||
if cand < t:
|
||
t = cand
|
||
room = 0.0 if t == float("inf") else max(0.0, t)
|
||
turn = wrap180(r0["eh"] - rp["eh"])
|
||
|
||
maxea = math.asin(min(8.0 / speed, 1.0))
|
||
kfix = max(1, int(math.ceil(dist / speed)))
|
||
r = run.by_tick.get(t0 + kfix)
|
||
if r is None or t0 + kfix > re:
|
||
continue
|
||
off = wrap180(math.atan2(r["ey"] - p0y, r["ex"] - p0x) - tb0)
|
||
b_our = gf_to_bin(max(-1.0, min(1.0, off / maxea)))
|
||
b_bullet = gf_to_bin(max(-1.0, min(1.0, wrap180(
|
||
sh["_dir"] - math.degrees(tb0)) / math.degrees(maxea))))
|
||
w = max(0.0, (math.asin(min(R / max(dist, 1e-6), 1.0)) / maxea) *
|
||
(NBINS - 1) / 2.0)
|
||
recs.append(dict(battle=bn, vlat=vlat, dist=dist, room=room,
|
||
turn=turn, b_our=b_our, b_bullet=b_bullet,
|
||
w=w, hit=sh["hit"]))
|
||
return recs
|
||
|
||
|
||
# ------------------------------------------------------------------- models
|
||
def split_battles(battles, seed, frac=0.7):
|
||
bs = sorted(battles)
|
||
random.Random(seed).shuffle(bs)
|
||
k = int(len(bs) * frac)
|
||
return set(bs[:k]), set(bs[k:])
|
||
|
||
|
||
class HistModel:
|
||
"""j128's counted SBC: state -> GF bin, per-cell posterior + prior blend."""
|
||
|
||
def __init__(self, decay=128, shift=1, alpha=5.0):
|
||
self.c = [[0] * NBINS for _ in range(256)]
|
||
self.g = [0] * NBINS
|
||
self.d, self.sh, self.a = decay, shift, alpha
|
||
self.lc = 0
|
||
|
||
def learn(self, st, label):
|
||
if self.c[st][label] < 255:
|
||
self.c[st][label] += 1
|
||
self.g[label] += 1
|
||
self.lc += 1
|
||
if self.d > 0 and self.lc >= self.d:
|
||
for row in self.c:
|
||
for k in range(NBINS):
|
||
row[k] -= row[k] >> self.sh
|
||
for k in range(NBINS):
|
||
self.g[k] -= self.g[k] >> self.sh
|
||
self.lc = 0
|
||
|
||
def prior(self):
|
||
tot = sum(self.g)
|
||
return [(self.g[k] + 1.0) / (tot + NBINS) for k in range(NBINS)]
|
||
|
||
def predict(self, st):
|
||
pr = self.prior()
|
||
n = sum(self.c[st])
|
||
if n == 0:
|
||
return pr
|
||
return [((self.c[st][k] / n) + self.a * pr[k]) / (1.0 + self.a)
|
||
for k in range(NBINS)]
|
||
|
||
|
||
class OutcomeModel:
|
||
"""The j130 module: counted 2-class SBC per (state, candidate bin).
|
||
|
||
`label(state, g) = hit and |g - b_our| <= w` (the live-computable dense
|
||
outcome label). P(hit | state, g) is the per-cell posterior blended with the
|
||
global hit rate. `state_free` forces every wave into one state cell (the
|
||
ablation)."""
|
||
|
||
def __init__(self, decay=128, shift=1, alpha=1.0, state_free=False):
|
||
self.c = [[[0, 0] for _ in range(NBINS)] for _ in range(256)]
|
||
self.tot = [0, 0]
|
||
self.d, self.sh, self.a = decay, shift, alpha
|
||
self.sf = state_free
|
||
self.lc = 0
|
||
|
||
def learn(self, st, g, lab):
|
||
c = self.c[0 if self.sf else st][g]
|
||
if c[lab] < 255:
|
||
c[lab] += 1
|
||
self.tot[lab] += 1
|
||
self.lc += 1
|
||
if self.d > 0 and self.lc >= self.d:
|
||
for row in self.c:
|
||
for cell in row:
|
||
for k in (0, 1):
|
||
cell[k] -= cell[k] >> self.sh
|
||
for k in (0, 1):
|
||
self.tot[k] -= self.tot[k] >> self.sh
|
||
self.lc = 0
|
||
|
||
def prior_hit(self):
|
||
t = self.tot[0] + self.tot[1]
|
||
return 0.5 if t == 0 else self.tot[1] / t
|
||
|
||
def predict_hit(self, st, g):
|
||
c = self.c[0 if self.sf else st][g]
|
||
n = c[0] + c[1]
|
||
pr = self.prior_hit()
|
||
if n == 0:
|
||
return pr
|
||
return (c[1] + self.a * pr) / (n + self.a)
|
||
|
||
|
||
def hitwin(r, g):
|
||
return 1 if (r["hit"] >= 0.5 and abs(g - r["b_our"]) <= r["w"]) else 0
|
||
|
||
|
||
def log2(x):
|
||
return math.log2(max(x, 1e-12))
|
||
|
||
|
||
# --------------------------------------------------------------------- gates
|
||
def gate_A_correlation(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 = [h[b] / n[b] for b in used]
|
||
out = {}
|
||
out["hist"] = statistics.correlation([n[b] / len(recs) for b in used], rate)
|
||
out["outcome"] = statistics.correlation(
|
||
[statistics.fmean([hitwin(r, b) for r in recs]) for b in used], rate)
|
||
out["bullet"] = statistics.correlation(
|
||
[statistics.fmean([1 if abs(b - r["b_bullet"]) <= r["w"] else 0
|
||
for r in recs]) for b in used], rate)
|
||
return out, used, rate
|
||
|
||
|
||
def run_split(recs, seed, decay=128, shift=1):
|
||
tr_b, te_b = 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(CANON)
|
||
|
||
hist = HistModel(decay, shift)
|
||
for r in tr:
|
||
hist.learn(code_of(r, edges), r["b_our"])
|
||
|
||
om = OutcomeModel(decay, shift, state_free=False)
|
||
om0 = OutcomeModel(decay, shift, state_free=True)
|
||
for r in tr:
|
||
st = code_of(r, edges)
|
||
for g in range(NBINS):
|
||
lab = hitwin(r, g)
|
||
om.learn(st, g, lab)
|
||
om0.learn(st, g, lab)
|
||
|
||
ll_s, ll_g = [], []
|
||
hit_hist, hit_out, hit_cur = [], [], []
|
||
for r in te:
|
||
st = code_of(r, edges)
|
||
ph = hist.predict(st)
|
||
gh = min(range(NBINS), key=lambda g: ph[g])
|
||
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_hist.append(real(gh))
|
||
hit_out.append(real(go))
|
||
hit_cur.append(r["hit"])
|
||
# held-out log-loss of the outcome label at every candidate g
|
||
for g in range(NBINS):
|
||
y = hitwin(r, g)
|
||
ll_s.append(-log2(om.predict_hit(st, g)) if y else
|
||
-log2(1.0 - om.predict_hit(st, g)))
|
||
ll_g.append(-log2(om0.predict_hit(st, g)) if y else
|
||
-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_hist=statistics.fmean(hit_hist),
|
||
hit_out=statistics.fmean(hit_out),
|
||
hit_cur=statistics.fmean(hit_cur),
|
||
)
|
||
|
||
|
||
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 = extract(runs)
|
||
|
||
lines = []
|
||
|
||
def out(s=""):
|
||
print(s)
|
||
lines.append(s)
|
||
|
||
out("# Outcome-label Gate A — offline sanity check (VETO ONLY)")
|
||
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 (the module's 4 fields, canonical edges)")
|
||
out("label : outcome hit(state,g) = hit and |g - b_our| <= w "
|
||
"(w = body width as an angle)")
|
||
out()
|
||
|
||
corr, used, rate = gate_A_correlation(recs)
|
||
out("## A. is the danger map the mover MINIMISES aligned with the realised "
|
||
"per-bin hit rate?")
|
||
out()
|
||
out("corr( danger(g) , P(hit | b_our = g) ) over the 31 bins:")
|
||
out()
|
||
out("| danger map | corr |")
|
||
out("|---|---:|")
|
||
out(f"| histogram label (j128) — P(arrival bin = g) | {corr['hist']:+.3f} |")
|
||
out(f"| **outcome label (j130, the module's live label)** | "
|
||
f"**{corr['outcome']:+.3f}** |")
|
||
out(f"| geometric bullet-line label (needs bullet bodies) | "
|
||
f"{corr['bullet']:+.3f} |")
|
||
out()
|
||
out("Negative = minimising the danger steers INTO the bullets (the j128 "
|
||
"defect). The histogram reproduces the ledger's -0.342.")
|
||
out()
|
||
out("| bin | P(hit) | P(arrival=bin) | outcome danger |")
|
||
out("|---:|---:|---:|---:|")
|
||
for b in used:
|
||
d = statistics.fmean([hitwin(r, b) for r in recs])
|
||
m = sum(1 for r in recs if r["b_our"] == b) / len(recs)
|
||
out(f"| {b} | {rate[used.index(b)]*100:.1f}% | {m*100:.1f}% | {d:.3f} |")
|
||
out()
|
||
|
||
per = [run_split(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 OUTCOME label")
|
||
out()
|
||
out("held-out per-candidate log-loss (bits) of the outcome label, "
|
||
"state-conditional vs state-free (same rows, same split):")
|
||
out()
|
||
out("| model | log-loss (bits) |")
|
||
out("|---|---:|")
|
||
out(f"| state-free P(hit | g) | {ll_g:.4f} |")
|
||
out(f"| state-conditional P(hit | 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("If the mover picks argmin_g danger, the fraction of held-out waves "
|
||
"whose bullet line would still pass within a body width of g "
|
||
"(ground truth = the recorded bullet line b_bullet).")
|
||
out()
|
||
out("| policy | held-out waves still hit |")
|
||
out("|---|---:|")
|
||
out(f"| histogram argmin (j128) | {statistics.fmean(p['hit_hist'] for p in per)*100:.2f}% |")
|
||
out(f"| outcome argmin (j130) | {statistics.fmean(p['hit_out'] for p in per)*100:.2f}% |")
|
||
out(f"| recorded trajectory (floor/ceiling) | {statistics.fmean(p['hit_cur'] for p in per)*100:.2f}% |")
|
||
out()
|
||
out("The counterfactual is OPEN LOOP: the recorded bullet lines were fired "
|
||
"at a different mover, so it cannot predict the live closed loop. It is "
|
||
"a veto, not a selection.")
|
||
out()
|
||
|
||
out("## MEASURED vs INFERRED")
|
||
out()
|
||
out("* MEASURED: every number above, on the recorded corpus.")
|
||
out("* INFERRED: that the offline alignment transfers live. It cannot — "
|
||
"see 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()
|