diff --git a/common_libs/tests/measure_horizon_headroom.py b/common_libs/tests/measure_horizon_headroom.py new file mode 100644 index 0000000..35ba73f --- /dev/null +++ b/common_libs/tests/measure_horizon_headroom.py @@ -0,0 +1,581 @@ +#!/usr/bin/env python3 +"""OFFLINE MEASUREMENT ONLY - no live battles, no GUI, no training, no source edits. + +"Find the hard questions": how learnable is the *horizon* prediction problem — +where will the enemy be h ticks from now — as a function of h? + +For every horizon h = 1..50 and every tick t with t+h inside the SAME round, and +using the shooter (s*) position at t as the observer: + + 1. NAIVE GUESS = straight-line extrapolation of the enemy's current velocity + (signed speed * heading) for h ticks. Never bounced off walls. + 2. ANGULAR ERROR = bearing(observer -> actual enemy at t+h) + - bearing(observer -> naive guess), wrapped to (-180,180]. + The angle is what aim cares about; distance-only error cannot change a shot. + 3. Per horizon: sign balance, signed/absolute error spread, the fraction of + ticks where the naive guess points OUTSIDE the enemy body (18 px radius, + half-angle atan(18/dist)), trivial-predictor accuracy for the sign, and the + LEFT/RIGHT distribution shift conditioned on state features. + 4. A recommended horizon range. + +Data: tools/fixtures/*drussgt*.jsonl (READ-ONLY) + their *.rounds.json sidecars. +PRIMARY = the tr-bridge captures where s* is ModularBot (our own movement). +STABILITY= every *drussgt* fixture, each using its own s* as observer. + +Run: + python3 common_libs/tests/measure_horizon_headroom.py +""" + +import json +import math +import os +import sys +from collections import defaultdict + +ROOT = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +FIX = os.path.join(ROOT, "tools", "fixtures") +META = os.path.join(FIX, "drussgt_meta") + +BOT_RADIUS = 18.0 # enemy body radius (px) +WALL_NEAR = 60.0 # "near a wall" threshold (px) +FAST_SPEED = 4.0 # "fast" threshold (px/tick) +FIRE_DROP_MAX = 3.1 # se drop above this is damage, not a fire (max power 3) +H_MAX = 50 + +PRIMARY = ["tr_drussgt_vs_modularbot.jsonl", + "tr_drussgt_vs_modularbot_shield.jsonl"] + + +def wrap180(a): + return math.degrees(math.atan2(math.sin(math.radians(a)), + math.cos(math.radians(a)))) + + +def load_fixture(path): + """Return (states_by_tick, rounds, meta). states are keyed by global tick.""" + states = {} + meta = {} + rounds = None + with open(path) as f: + for line in f: + line = line.strip() + if not line: + continue + d = json.loads(line) + if "meta" in d: + meta = d["meta"] + elif "rounds" in d: + rounds = d["rounds"] + elif "end" in d: + pass + elif "tick" in d: + states[d["tick"]] = d + rp = os.path.join(META, os.path.basename(path) + ".rounds.json") + if os.path.exists(rp): + rounds = json.load(open(rp))["rounds"] + if rounds is None: + # fall back to one round spanning everything + ticks = sorted(states) + rounds = [{"round": 1, "startTick": ticks[0], + "count": ticks[-1] - ticks[0] + 1}] + return states, rounds, meta + + +def bullet_in_flight_series(states, rounds): + """Per-tick bool: is one of OUR bullets plausibly still flying? + + Fires are detected as self-energy drops in (0, 3.1] px (fire cost = power; + larger drops are enemy damage). Bullet speed = 20 - 3*power; flight ticks = + range_at_fire / speed. This is a PROXY (no gun/power/heading is recorded), + labelled INFERRED in the report. + """ + active = defaultdict(list) # tick -> list of remaining-ticks + for r in rounds: + s0, n = int(r["startTick"]), int(r["count"]) + for t in range(s0 + 1, s0 + n): + prev = states.get(t - 1) + cur = states.get(t) + if prev is None or cur is None: + continue + drop = prev["se"] - cur["se"] + if 0.05 < drop <= FIRE_DROP_MAX: + power = min(3.0, max(0.1, drop)) + speed = 20.0 - 3.0 * power + rng = math.hypot(cur["ex"] - cur["sx"], cur["ey"] - cur["sy"]) + flight = int(math.ceil(rng / speed)) + for k in range(1, flight + 1): + active[t + k].append(1) + return {t: True for t, v in active.items() if v} + + +def build_round_arrays(states, rounds): + """Contiguous per-round arrays of the fields we need.""" + out = [] + all_ticks = set(states) + for r in rounds: + s0, n = int(r["startTick"]), int(r["count"]) + arr = [] + ok = True + for t in range(s0, s0 + n): + if t not in all_ticks: + ok = False + break + arr.append(states[t]) + if ok and arr: + out.append(arr) + return out + + +def turn_sign(arr, i): + """Heading change (deg, CCW+) from i-1 to i within the round; None at i=0.""" + if i <= 0: + return None + return wrap180(arr[i]["eh"] - arr[i - 1]["eh"]) + + +def feature_wall(st): + return min(st["ex"], st["ey"], 800.0 - st["ex"], 600.0 - st["ey"]) < WALL_NEAR + + +def analyze_fixture(path, horizons=range(1, H_MAX + 1)): + states, rounds, meta = load_fixture(path) + ra = build_round_arrays(states, rounds) + flying = bullet_in_flight_series(states, rounds) + + # accumulators per horizon + acc = {h: dict(n=0, left=0, right=0, zero=0, naive_miss=0, + med_signed=[], abs_err=[], + turn_n=0, turn_ok=0, + pers1_n=0, pers1_ok=0, + persist_n=0, persist_ok=0, + feat=defaultdict(lambda: [0, 0])) for h in horizons} + # feat keys: (name, group) -> [n, left] + + for arr in ra: + L = len(arr) + for i in range(L): + selfx, selfy = arr[i]["sx"], arr[i]["sy"] + ex, ey = arr[i]["ex"], arr[i]["ey"] + eh = arr[i]["eh"] + es = arr[i]["es"] + vx = es * math.cos(math.radians(eh)) + vy = es * math.sin(math.radians(eh)) + turn = turn_sign(arr, i) + near_wall = feature_wall(arr[i]) + fast = abs(es) >= FAST_SPEED + rev = es < 0.0 + inf = flying.get(arr[i]["tick"], False) + rng = math.hypot(ex - selfx, ey - selfy) + if i >= 1: + psx0, psy0 = arr[i - 1]["sx"], arr[i - 1]["sy"] + prng = math.hypot(arr[i - 1]["ex"] - psx0, arr[i - 1]["ey"] - psy0) + closing = rng < prng + else: + closing = None + + # causal persistence predictor: the OBSERVED 1-step angular error + # sign at time t (naive 1-step guess made at t-1 vs actual at t). + err1_sign = 0 + if i - 1 >= 0: + psx, psy = arr[i - 1]["sx"], arr[i - 1]["sy"] + pex, pey = arr[i - 1]["ex"], arr[i - 1]["ey"] + peh, pes = arr[i - 1]["eh"], arr[i - 1]["es"] + pgx = pex + pes * math.cos(math.radians(peh)) + pgy = pey + pes * math.sin(math.radians(peh)) + b_a = math.atan2(ey - psy, ex - psx) + b_g = math.atan2(pgy - psy, pgx - psx) + e1 = math.degrees(math.atan2(math.sin(b_a - b_g), + math.cos(b_a - b_g))) + if abs(e1) > 1e-9: + err1_sign = 1 if e1 > 0 else -1 + + for h in horizons: + j = i + h + if j >= L: + break + a = acc[h] + ax, ay = arr[j]["ex"], arr[j]["ey"] + gx, gy = ex + vx * h, ey + vy * h + + ba = math.atan2(ay - selfy, ax - selfx) + bg = math.atan2(gy - selfy, gx - selfx) + err = math.degrees(math.atan2(math.sin(ba - bg), + math.cos(ba - bg))) + + a["n"] += 1 + if err > 1e-9: + sg = 1 + elif err < -1e-9: + sg = -1 + else: + sg = 0 + if sg > 0: + a["left"] += 1 + elif sg < 0: + a["right"] += 1 + else: + a["zero"] += 1 + a["med_signed"].append(err) + a["abs_err"].append(abs(err)) + + dist_actual = math.hypot(ax - selfx, ay - selfy) + if dist_actual > 1e-6: + half = math.degrees(math.atan2(BOT_RADIUS, dist_actual)) + if abs(err) > half: + a["naive_miss"] += 1 + + # (a) turn-direction predictor + if turn is not None and abs(turn) > 1e-6 and sg != 0: + a["turn_n"] += 1 + pred = 1 if turn > 0 else -1 + if pred == sg: + a["turn_ok"] += 1 + + # (b) causal persistence: side of the last OBSERVED 1-step error + if err1_sign != 0 and sg != 0: + a["pers1_n"] += 1 + if err1_sign == sg: + a["pers1_ok"] += 1 + + # (b') NON-CAUSAL diagnostic: same sign as the h-step error at + # t-1. NOT usable online for h>1 (needs t-1+h > t). + if i - 1 >= 0 and sg != 0: + pex, pey = arr[i - 1]["ex"], arr[i - 1]["ey"] + psx, psy = arr[i - 1]["sx"], arr[i - 1]["sy"] + pax, pay = arr[j - 1]["ex"], arr[j - 1]["ey"] + pb = math.atan2(pay - psy, pax - psx) + peh, pes = arr[i - 1]["eh"], arr[i - 1]["es"] + pgx = pex + pes * math.cos(math.radians(peh)) * h + pgy = pey + pes * math.sin(math.radians(peh)) * h + pbg = math.atan2(pgy - psy, pgx - psx) + perr = math.degrees(math.atan2(math.sin(pb - pbg), + math.cos(pb - pbg))) + if abs(perr) > 1e-9: + a["persist_n"] += 1 + if (perr > 0) == (sg > 0): + a["persist_ok"] += 1 + + # conditional-distribution groups (only defined when sg != 0) + if sg != 0: + for name, grp in (("wall", near_wall), + ("turn", None if turn is None else turn > 0), + ("speed", fast), + ("bullet", inf), + ("rev", rev), + ("closing", closing)): + if grp is None or (name == "turn" and turn is not None + and abs(turn) < 1e-6): + continue + key = (name, grp) + a["feat"][key][0] += 1 + if sg > 0: + a["feat"][key][1] += 1 + + return acc, meta, len(ra), sum(len(x) for x in ra) + + +def pct(x): + return 100.0 * x + + +def median(v): + return sorted(v)[len(v) // 2] if v else float("nan") + + +def quantile(v, q): + if not v: + return float("nan") + s = sorted(v) + i = min(len(s) - 1, max(0, int(q * (len(s) - 1)))) + return s[i] + + +def z2prop(l1, n1, l2, n2): + if n1 == 0 or n2 == 0: + return 0.0 + p1, p2 = l1 / n1, l2 / n2 + p = (l1 + l2) / (n1 + n2) + var = max(0.0, p * (1 - p) * (1 / n1 + 1 / n2)) + se = math.sqrt(var) + return 0.0 if se <= 0 else (p1 - p2) / se + + +def group_shift(a, name): + n1, l1 = a["feat"].get((name, True), [0, 0]) + n2, l2 = a["feat"].get((name, False), [0, 0]) + if n1 == 0 or n2 == 0: + return float("nan"), 0.0, n1, n2 + d = (l1 / n1) - (l2 / n2) + return d, z2prop(l1, n1, l2, n2), n1, n2 + + +def merged(accs): + """Merge several per-fixture accumulators (add counters, concat lists).""" + out = {} + for acc in accs: + for h, a in acc.items(): + if h not in out: + out[h] = dict(n=0, left=0, right=0, zero=0, naive_miss=0, + med_signed=[], abs_err=[], + turn_n=0, turn_ok=0, + pers1_n=0, pers1_ok=0, + persist_n=0, persist_ok=0, + feat=defaultdict(lambda: [0, 0])) + o = out[h] + for k in ("n", "left", "right", "zero", "naive_miss", + "turn_n", "turn_ok", "pers1_n", "pers1_ok", + "persist_n", "persist_ok"): + o[k] += a[k] + o["med_signed"].extend(a["med_signed"]) + o["abs_err"].extend(a["abs_err"]) + for key, val in a["feat"].items(): + o["feat"][key][0] += val[0] + o["feat"][key][1] += val[1] + return out + + +def main(): + args = [a for a in sys.argv[1:] if not a.startswith("-")] + if args: + files = args + primary = args + else: + files = sorted(f for f in os.listdir(FIX) + if "drussgt" in f and f.endswith(".jsonl")) + primary = PRIMARY + + print("=" * 100) + print("HORIZON HEADROOM - offline angular measurement on DrussGT fixtures") + print("=" * 100) + print(f"fixtures ({len(files)}): {', '.join(files)}") + print(f"naive guess : straight-line extrapolation of (signed es * heading), not bounced") + print(f"angular error : bearing(us->actual) - bearing(us->naive), deg; LEFT = err>0 (CCW)") + print(f"naive_miss : |err| > atan(18px / actual range) -> aim-at-naive lands outside body") + print(f" (body half-angle at 300px = {math.degrees(math.atan2(18,300)):.2f} deg)") + print() + + per_fixture = {} + for path in files: + full = os.path.join(FIX, path) + acc, meta, nrounds, nticks = analyze_fixture(full) + per_fixture[path] = (acc, meta, nrounds, nticks) + print(f" {path:<42} rounds={nrounds:<3} ticks={nticks}") + + print() + prim_accs = [per_fixture[p][0] for p in primary] + P = merged(prim_accs) + + # ── Table A: sign balance + error size + naive miss ── + print("=" * 100) + print("TABLE A - POOLED PRIMARY (s* = ModularBot): sign balance, angular error, naive miss") + print("=" * 100) + hdr = (f"{'h':>3} {'N':>7} {'%L':>5} {'%R':>5} {'%0':>4} " + f"{'medSig':>7} {'|p10':>6} {'|p50':>6} {'|p90':>6} " + f"{'naiveMiss%':>10} {'9}") + print(hdr) + for h in range(1, H_MAX + 1): + a = P[h] + n = a["n"] + if n == 0: + continue + fl = a["left"] / n + fr = a["right"] / n + fz = a["zero"] / n + ms = median(a["med_signed"]) + p10 = quantile(a["abs_err"], 0.10) + p50 = quantile(a["abs_err"], 0.50) + p90 = quantile(a["abs_err"], 0.90) + nm = a["naive_miss"] / n + # fraction of samples with an error big enough to matter at 300px (3.43deg) + big = sum(1 for e in a["abs_err"] if e > math.degrees(math.atan2(18, 300))) + print(f"{h:>3} {n:>7} {pct(fl):>5.1f} {pct(fr):>5.1f} {pct(fz):>4.1f} " + f"{ms:>7.2f} {p10:>6.2f} {p50:>6.2f} {p90:>6.2f} " + f"{pct(nm):>10.1f} {pct(big/max(1,n)):>9.1f}") + + # ── Table B: trivial predictors + feature shifts ── + print() + print("=" * 100) + print("TABLE B - POOLED PRIMARY: trivial sign predictors and conditional-distribution shift") + print("=" * 100) + print(" turnAcc = sign(enemy heading change at t) == sign(err) (n shown)") + print(" pers1Acc = sign(observed 1-step angular error at t) == sign(err) [CAUSAL, available at t]") + print(" persNcAcc= sign(err at t-1) == sign(err at t) [NON-CAUSAL diagnostic; needs t-1+h>t]") + print(" majAcc = always predict the majority side (in-sample ceiling)") + print(" dWall/dTurn/dSpeed/dBullet = %L(group=True) - %L(group=False); z in ()") + print(" wall True=enemy <60px from a wall") + print(" turn True=enemy heading changing CCW at t") + print(" speed True=|es| >= 4") + print(" bullet True=one of OUR bullets plausibly in flight (proxy, INFERRED)") + print(" rev True=enemy signed speed es < 0 (reversing)") + print(" close True=range to enemy shrank since t-1 (moving toward us)") + print() + hdr = (f"{'h':>3} {'turnAcc':>7} {'pers1Acc':>8} {'persNcAcc':>9} {'majAcc':>7} " + f"{'dWall':>12} {'dTurn':>12} {'dSpeed':>12} {'dBullet':>13}") + print(hdr) + for h in range(1, H_MAX + 1): + a = P[h] + n = a["n"] + if n == 0: + continue + ta = pct(a["turn_ok"] / a["turn_n"]) if a["turn_n"] else float("nan") + p1 = pct(a["pers1_ok"] / a["pers1_n"]) if a["pers1_n"] else float("nan") + pn = pct(a["persist_ok"] / a["persist_n"]) if a["persist_n"] else float("nan") + maj = pct(max(a["left"], a["right"]) / max(1, a["left"] + a["right"])) + dw, zw, _, _ = group_shift(a, "wall") + dt, zt, _, _ = group_shift(a, "turn") + ds, zs, _, _ = group_shift(a, "speed") + db, zb, _, _ = group_shift(a, "bullet") + + def cell(d, z): + if d != d: + return f"{'--':>12}" + return f"{pct(d):+6.1f}({z:+4.1f})" + + print(f"{h:>3} {ta:>7.1f} {p1:>8.1f} {pn:>9.1f} {maj:>7.1f} " + f"{cell(dw,zw):>12} {cell(dt,zt):>12} {cell(ds,zs):>12} {cell(db,zb):>13}") + + # ── Feature group detail at selected horizons ── + print() + print("=" * 100) + print("FEATURE GROUP DETAIL (pooled primary): %L in each split") + print("=" * 100) + sel = [2, 3, 5, 8, 10, 15, 20, 25, 30, 40, 50] + print(f"{'h':>3} {'%L all':>7} | {'wall':>16} {'turn':>16} {'speed':>16} " + f"{'bullet':>16} {'rev':>16} {'closing':>16}") + print(f"{'':>3} {'':>7} | {'near':>7} {'open':>7} {'left':>7} {'right':>7} " + f"{'fast':>7} {'slow':>7} {'infl':>7} {'none':>7} {'rev':>7} {'fwd':>7} " + f"{'close':>7} {'open':>7}") + for h in sel: + a = P[h] + nl = a["left"] + a["right"] + if nl == 0: + continue + row = f"{h:>3} {pct(a['left']/nl):>7.1f} |" + for name in ("wall", "turn", "speed", "bullet", "rev", "closing"): + nt, lt = a["feat"].get((name, True), [0, 0]) + nf, lf = a["feat"].get((name, False), [0, 0]) + row += f" {pct(lt/nt) if nt else float('nan'):>7.1f} {pct(lf/nf) if nf else float('nan'):>7.1f} " + print(row) + + # ── Recommendation logic ── + print() + print("=" * 100) + print("RECOMMENDATION") + print("=" * 100) + print("A horizon counts as HARD/LEARNABLE when:") + print(" (a) sign balance near 50/50 ........ 42% <= %L <= 58%") + print(" (b) trivial predictors do NOT solve . max(turnAcc, pers1Acc) < 75%") + print(" (majority is implied by balance: min(%L,%R) small => majAcc high)") + print(" (c) a state feature shifts the split some |d| >= 5pp with |z| >= 3") + print(" (d) naive guess actually misses ...... naiveMiss >= 15%") + print() + hard = [] + hard_any = [] + for h in range(1, H_MAX + 1): + a = P[h] + n = a["n"] + if n == 0: + continue + fl = a["left"] / max(1, a["left"] + a["right"]) + bal = 0.42 <= fl <= 0.58 + ta = (a["turn_ok"] / a["turn_n"]) if a["turn_n"] else 0.0 + p1 = (a["pers1_ok"] / a["pers1_n"]) if a["pers1_n"] else 0.0 + nontrivial = max(ta, p1) < 0.75 + sig = False + best = ("", 0.0, 0.0) + for name in ("wall", "turn", "speed", "bullet"): + d, z, _, _ = group_shift(a, name) + if d == d and abs(d) >= 0.05 and abs(z) >= 3.0: + sig = True + if abs(d) > abs(best[1]): + best = (name, d, z) + sig_any = sig + best_any = best + for name in ("rev", "closing"): + d, z, _, _ = group_shift(a, name) + if d == d and abs(d) >= 0.05 and abs(z) >= 3.0: + if not sig_any or abs(d) > abs(best_any[1]): + best_any = (name, d, z) + sig_any = True + nm = a["naive_miss"] / n + ok = bal and nontrivial and sig and nm >= 0.15 + ok_any = bal and nontrivial and sig_any and nm >= 0.15 + if ok: + hard.append(h) + if ok_any: + hard_any.append(h) + print(f" h={h:>2} bal={'Y' if bal else 'n'} ({pct(fl):4.1f}%) " + f"nontriv={'Y' if nontrivial else 'n'} (turn={pct(ta):4.1f} pers1={pct(p1):4.1f}) " + f"signal4={'Y' if sig else 'n'} ({best[0]}:{pct(best[1]):+.1f}pp z={best[2]:+.1f}) " + f"signalAll={'Y' if sig_any else 'n'} ({best_any[0]}:{pct(best_any[1]):+.1f}pp) " + f"miss={'Y' if nm>=0.15 else 'n'} ({pct(nm):4.1f}%) => {'HARD' if ok else '-'}") + + print() + if hard: + print(f"RECOMMENDED HORIZON RANGE (4 requested features): {hard[0]}..{hard[-1]} ticks " + f"({len(hard)} horizons, contiguous={hard == list(range(hard[0], hard[-1]+1))})") + else: + print("RECOMMENDED HORIZON RANGE (4 requested features): NONE.") + if hard_any: + print(f"RECOMMENDED HORIZON RANGE (+ reversing/closing): {hard_any[0]}..{hard_any[-1]} ticks " + f"({len(hard_any)} horizons, contiguous={hard_any == list(range(hard_any[0], hard_any[-1]+1))})") + else: + print("RECOMMENDED HORIZON RANGE (+ reversing/closing): NONE.") + print("(4-feature hard horizons:", hard, ")") + print("(all-feature hard horizons:", hard_any, ")") + + # ── Stability across fixtures ── + print() + print("=" * 100) + print("STABILITY ACROSS FIXTURES: %L and naiveMiss% per fixture at chosen horizons") + print("=" * 100) + print("(each fixture uses its OWN s* as observer; only the two modularbot files have") + print(" s* = ModularBot, so only those isolate OUR movement's effect on DrussGT)") + show = [1, 2, 3, 5, 8, 10, 15, 20, 30, 40, 50] + short = {p: p.replace("tr_drussgt_vs_", "tr:").replace("_vs_", "~") + .replace(".jsonl", "") for p in files} + print(f"{'fixture':<22} " + " ".join(f"h{h:<2}" for h in show)) + for path in files: + acc = per_fixture[path][0] + cells = [] + for h in show: + a = acc[h] + n = a["n"] + if n == 0: + cells.append(" -- ") + continue + fl = pct(a["left"] / n) + nm = pct(a["naive_miss"] / n) + cells.append(f"{fl:4.0f}") + print(f"{short[path]:<22} " + " ".join(f"{c:>4}" for c in cells) + " (%L)") + print() + print(f"{'fixture':<22} " + " ".join(f"h{h:<2}" for h in show) + " (naiveMiss%)") + for path in files: + acc = per_fixture[path][0] + cells = [] + for h in show: + a = acc[h] + n = a["n"] + cells.append(" -- " if n == 0 else f"{pct(a['naive_miss']/n):4.0f}") + print(f"{short[path]:<22} " + " ".join(f"{c:>4}" for c in cells)) + + # Cross-fixture spread of %L at each horizon + print() + print("cross-fixture spread of %L (min..max, pp) at chosen horizons:") + for h in show: + vals = [] + for path in files: + a = per_fixture[path][0][h] + n = a["n"] + if n: + vals.append(pct(a["left"] / n)) + if vals: + print(f" h={h:>2}: min={min(vals):5.1f} max={max(vals):5.1f} " + f"spread={max(vals)-min(vals):5.1f}pp (n_fixtures={len(vals)})") + print() + print("SAMPLE SIZES: see N per horizon in Table A (pooled primary);") + print("per-fixture ticks are listed at the top.") + + +if __name__ == "__main__": + main()