## ABLATION: does the radial TM's `bmPoint` advantage need the Tsetlin Machine? ## ## The committed radial-TM finding (589a230) is that its `bmPoint` win does NOT ## come from out-classifying a majority baseline (the radial head sits AT/BELOW ## the 58.2% majority rate) but from a NET-POSITIVE AVERAGE radial aim-distance ## shift. If so, a FIXED radial shift should reproduce most or all of the win ## with no learning at all. ## ## This sweep builds a constant-offset variant of the base gun ## (`guns/radial_offset.nim`: exact `forecastLinear` bearing, aim distance scaled ## and shifted by a fixed constant, same BotRadius clamp as the TM corrective ## path) and sweeps the constant. It compares against: ## 1. `Linear` — the unmodified base; ## 2. `TMRadial` — the learned radial head (current best config, defaults ## TM_RADIAL_RANGE=60, TM_RAD_MARGIN=0.25, 5 classes); ## 3. `TMRadialShuf` — the mandatory shuffled-label control; ## under BOTH metrics (`bmPoint`, and the SHIPPED `bmPath`). ## ## It also reports the raw radial-label distribution and per-adversary optima, ## so the reader can judge whether the net shift is a genuine surfer property or ## an artefact, and whether a single constant is fragile. ## ## Usage: ## nim c -r -d:release --path:common_libs \ ## common_libs/tests/sweep_radial_offset.nim \ ## --set=real --metric=point --seeds=3 ## Flags: --set=real|range|synthetic, --metric=point|path, --seeds=N, ## --maxrounds=N, --summaryOnly=0|1 ## ## EARLY = resolutions whose LOCAL tick is in the first 100 ticks of a round. ## OVERALL = every resolution in the round. Methodology (fixture set, round ## splitting, fresh gun per fixture) matches `sweep_tm_pattern.nim` exactly so ## the TMRadial numbers here are directly comparable to the committed ones. import std/[os, strformat, strutils, json, tables, math, random, algorithm, sequtils] import gun_harness/offline_range import gun_harness/gun_interface import gun_harness/virtual_bullets as vb import guns/linear import guns/tm_pattern import guns/lead_forecast import guns/radial_offset const repoRoot = currentSourcePath().parentDir.parentDir.parentDir fixturesDir = repoRoot / "tools" / "fixtures" type Adapt = object h100, n100, h300, n300, hall, nall, f100, m100: int rounds: int GunStats = object obs, labelMiss, traceMiss: int radLabel: array[TM_CLASSES, int] radChosen: array[TM_CLASSES, int] radCorrect, radTotal: int radDeltaSum, radDeltaAbsSum: float radDeltaN: int radOffsetSum: float radOffsetN: int ArmKind = enum akLinear, akOffset, akTMRad, akTMRadShuf Arm = object name: string kind: ArmKind scale, offsetPx: float RoundSpan = tuple[start, count: int] Row = object variant, fixture: string seed: int r: Adapt st: GunStats var gDropped = 0 ## unresolved bullets clobbered by the tracker ring (must be 0) # ── stats helpers ──────────────────────────────────────────────────────────── proc addAdapt(dst: var Adapt, src: Adapt) = inc dst.rounds, src.rounds dst.h100 += src.h100; dst.n100 += src.n100 dst.h300 += src.h300; dst.n300 += src.n300 dst.hall += src.hall; dst.nall += src.nall dst.f100 += src.f100; dst.m100 += src.m100 proc addStats(dst: var GunStats, src: GunStats) = dst.obs += src.obs; dst.labelMiss += src.labelMiss; dst.traceMiss += src.traceMiss dst.radCorrect += src.radCorrect; dst.radTotal += src.radTotal dst.radDeltaSum += src.radDeltaSum; dst.radDeltaAbsSum += src.radDeltaAbsSum dst.radDeltaN += src.radDeltaN dst.radOffsetSum += src.radOffsetSum; dst.radOffsetN += src.radOffsetN for c in 0.. n: return 0.0 var lg = 0.0 for i in 1..k: lg += ln(float(n - k + i)) - ln(float(i)) exp(lg - float(n) * ln(2.0)) proc signTestP(wins, n: int): float = if n == 0: return 1.0 let lo = min(wins, n - wins) var s = 0.0 for k in 0..lo: s += binomPmf(k, n) min(1.0, 2.0 * s) # ── fixture / round I/O ────────────────────────────────────────────────────── proc loadRounds(path: string): seq[RoundSpan] = let dir = path.parentDir let base = path.extractFilename var side = dir / "drussgt_meta" / (base & ".rounds.json") if not fileExists(side): side = dir / (base & ".rounds.json") if not fileExists(side): return @[] let node = parseJson(readFile(side)) if not node.hasKey("rounds"): return @[] for r in node["rounds"]: result.add (r["startTick"].getInt(), r["count"].getInt()) proc resolve(name: string): tuple[fx: Fixture, path: string] = let p = if fileExists(name): name else: fixturesDir / (name & ".jsonl") (loadFixture(p), p) proc fixtureSet(name: string): seq[string] = case name of "synthetic", "": for n in SyntheticFixtureNames: result.add n of "real": for n in ["drussgt_vs_crazy", "drussgt_vs_spinbot", "drussgt_vs_drussgt", "tr_drussgt_vs_crazy", "tr_drussgt_vs_spinbot", "tr_drussgt_vs_modularbot"]: result.add(fixturesDir / (n & ".jsonl")) of "range": for n in SyntheticFixtureNames: result.add n for n in ["drussgt_vs_crazy", "drussgt_vs_spinbot", "tr_drussgt_vs_crazy"]: result.add(fixturesDir / (n & ".jsonl")) else: discard # ── replay ─────────────────────────────────────────────────────────────────── proc runRound(states: seq[WorldState], lastSeen: seq[int], enemyId, baseTick: int, drivers: seq[GunDriver], metric: BulletMetric): seq[Adapt] = var tracker = initTracker(drivers.len, metric) var accum = newSeq[Adapt](drivers.len) for ad in accum.mitems: inc ad.rounds for si in 0..= 0: lst = lastSeen[si] if state.enemies.len > 0: for e in state.enemies: enemyPositions[e.id] = (x: e.x, y: e.y, lastSeenTick: lst, alive: true) else: enemyPositions[enemyId] = (x: state.enemyX, y: state.enemyY, lastSeenTick: lst, alive: true) let localTick = state.tick - baseTick let dref = drivers tracker.tickBullets(state, enemyPositions, proc(gunId: GunId, binIdx: int, e: FeedbackEvent) = inc accum[gunId].nall if e.hit: inc accum[gunId].hall if localTick < 100: inc accum[gunId].n100 if e.hit: inc accum[gunId].h100 if localTick < 300: inc accum[gunId].n300 if e.hit: inc accum[gunId].h300 let fireTick = e.fireTick - baseTick if fireTick < 100: inc accum[gunId].m100 if e.hit: inc accum[gunId].f100 dref[gunId].resultCb(e)) gDropped += tracker.droppedBullets result = accum proc runFixtureClean(fx: Fixture, path: string, drivers: seq[GunDriver], metric: BulletMetric, maxRounds = 0): seq[Adapt] = var spans = if fx.meta.source == "synthetic": @[(start: 0, count: fx.states.len)] else: loadRounds(path) if maxRounds > 0 and spans.len > maxRounds: spans.setLen(maxRounds) result = newSeq[Adapt](drivers.len) if spans.len == 0: result = runRound(fx.states, fx.lastSeen, fx.enemyId, 0, drivers, metric) return for sp in spans: var st: seq[WorldState] var ls: seq[int] for i in 0..= sp.start and t < sp.start + sp.count: st.add fx.states[i] ls.add(if i < fx.lastSeen.len: fx.lastSeen[i] else: -1) if st.len == 0: continue let rr = runRound(st, ls, fx.enemyId, sp.start, drivers, metric) for gi in 0..= 0: randomize(seed) result.gun = g result.driver = GunDriver( name: name, predictCb: proc(state: WorldState, bulletSpeed: float): GunPrediction = g[].predict(state, bulletSpeed), resultCb: proc(e: FeedbackEvent) = g[].onResult(e), readyCb: proc(): bool = g[].isWarmedUp()) proc detDriver(a: Arm): GunDriver = case a.kind of akLinear: makeDriver(a.name, LinearGun()) of akOffset: makeDriver(a.name, initRadialOffsetGun(a.scale, a.offsetPx)) else: raise newException(ValueError, "not a deterministic arm: " & a.name) # ── main ───────────────────────────────────────────────────────────────────── proc main() = var set = "real" var metricName = "point" var nSeeds = 3 var maxRounds = 0 for i in 1..paramCount(): let a = paramStr(i) if a.startsWith("--set="): set = a[6..^1] elif a.startsWith("--metric="): metricName = a[9..^1] elif a.startsWith("--seeds="): nSeeds = parseInt(a[8..^1]) elif a.startsWith("--maxrounds="): maxRounds = parseInt(a[12..^1]) else: discard let metric = if metricName == "point": bmPoint else: bmPath let names = fixtureSet(set) # ── arms ─────────────────────────────────────────────────────────────────── var arms: seq[Arm] arms.add Arm(name: "Linear", kind: akLinear) # multiplicative shrink (fraction of base fire distance) for s in [1.00, 0.98, 0.95, 0.90, 0.85, 0.80]: arms.add Arm(name: &"RO_s{s:.2f}", kind: akOffset, scale: s) # fixed px shift (negative = aim short); +30 = opposite-direction control for o in [-10.0, -20.0, -30.0, -40.0, -60.0, 30.0]: arms.add Arm(name: &"RO_o{int(o):+d}", kind: akOffset, scale: 1.0, offsetPx: o) arms.add Arm(name: "TMRadial", kind: akTMRad) arms.add Arm(name: "TMRadialShuf", kind: akTMRadShuf) let detArms = arms.filterIt(it.kind in {akLinear, akOffset}) let tmArms = arms.filterIt(it.kind in {akTMRad, akTMRadShuf}) echo "# set=", set, " metric=", metricName, " seeds=", nSeeds, " fixtureCount=", names.len var rows: seq[Row] echo "variant,fixture,seed,h100,n100,h300,n300,hall,nall,rounds,obs,labelMiss,traceMiss" for name in names: let (fx, path) = resolve(name) let fxName = path.extractFilename.replace(".jsonl", "") # Deterministic arms: one batched pass (stateless, order-independent). let dDrivers = detArms.mapIt(detDriver(it)) let dRes = runFixtureClean(fx, path, dDrivers, metric, maxRounds) for i, a in detArms: let row = Row(variant: a.name, fixture: fxName, seed: 1, r: dRes[i], st: GunStats()) rows.add row echo &"{row.variant},{fxName},1,{row.r.h100},{row.r.n100},{row.r.h300},{row.r.n300}," & &"{row.r.hall},{row.r.nall},{row.r.rounds},0,0,0" # TM arms: single-driver, fresh gun per fixture, seeded. for a in tmArms: for seed in 1..nSeeds: let pair = makeTMRadDriver(a.name, seed, shuffle = (a.kind == akTMRadShuf)) let rr = runFixtureClean(fx, path, @[pair.driver], metric, maxRounds)[0] let row = Row(variant: a.name, fixture: fxName, seed: seed, r: rr, st: tmStats(pair.gun)) rows.add row echo &"{row.variant},{fxName},{seed},{row.r.h100},{row.r.n100},{row.r.h300}," & &"{row.r.n300},{row.r.hall},{row.r.nall},{row.r.rounds},{row.st.obs}," & &"{row.st.labelMiss},{row.st.traceMiss}" if gDropped > 0: echo &"\n# WARNING: droppedBullets={gDropped} (ring clobbered unresolved bullets)" # ── pooled summary ───────────────────────────────────────────────────────── var pooled = initTable[string, Adapt]() var pooledSt = initTable[string, GunStats]() for a in arms: pooled[a.name] = Adapt() pooledSt[a.name] = GunStats() for row in rows: addAdapt(pooled[row.variant], row.r) addStats(pooledSt[row.variant], row.st) echo "\n# ── pooled summary ──" echo "variant,runs,early%,early_hits,early_n,overall%,overall_hits,overall_n" for a in arms: let p = pooled[a.name] echo &"{a.name},{p.rounds},{rateStr(p.h100, p.n100)},{p.h100},{p.n100}," & &"{rateStr(p.hall, p.nall)},{p.hall},{p.nall}" # ── per-run tables: deterministic arms replicated across seeds ───────────── type Key = tuple[fixture: string, seed: int] var byEarly = initTable[string, Table[Key, float]]() var byAll = initTable[string, Table[Key, float]]() var byEarlyHits = initTable[string, Table[Key, tuple[h, n: int]]]() var byAllHits = initTable[string, Table[Key, tuple[h, n: int]]]() for a in arms: byEarly[a.name] = initTable[Key, float]() byAll[a.name] = initTable[Key, float]() byEarlyHits[a.name] = initTable[Key, tuple[h, n: int]]() byAllHits[a.name] = initTable[Key, tuple[h, n: int]]() for row in rows: let k = (fixture: row.fixture, seed: row.seed) let e = if row.r.n100 > 0: row.r.h100.float / row.r.n100.float else: 0.0 let o = if row.r.nall > 0: row.r.hall.float / row.r.nall.float else: 0.0 byEarly[row.variant][k] = e byAll[row.variant][k] = o byEarlyHits[row.variant][k] = (row.r.h100, row.r.n100) byAllHits[row.variant][k] = (row.r.hall, row.r.nall) # replicate deterministic arms across the seed range for paired comparison for a in detArms: for row in rows: if row.variant == a.name: for s in 2..nSeeds: byEarly[a.name][(row.fixture, s)] = byEarly[a.name][(row.fixture, 1)] byAll[a.name][(row.fixture, s)] = byAll[a.name][(row.fixture, 1)] byEarlyHits[a.name][(row.fixture, s)] = byEarlyHits[a.name][(row.fixture, 1)] byAllHits[a.name][(row.fixture, s)] = byAllHits[a.name][(row.fixture, 1)] echo "\n# ── per-run distribution (mean / min / max, n runs) ──" echo "variant,earlyMean%,earlyMin%,earlyMax%,overallMean%,overallMin%,overallMax%,n" for a in arms: var es, osx: seq[float] for v in byEarly[a.name].values: es.add v for v in byAll[a.name].values: osx.add v if es.len == 0: continue es.sort(); osx.sort() echo &"{a.name},{es.sum/float(es.len)*100:.2f},{es[0]*100:.2f},{es[^1]*100:.2f}," & &"{osx.sum/float(osx.len)*100:.2f},{osx[0]*100:.2f},{osx[^1]*100:.2f},{es.len}" # ── sign tests ───────────────────────────────────────────────────────────── echo "\n# ── paired sign tests (rows = fixture x seed; exact two-sided binomial) ──" echo "A,B,metric,nA>B,nB>A,ties,p" proc signRow(aName, bName, label: string, tab: Table[string, Table[Key, float]]) = var winsA, winsB, ties, n = 0 for k, va in tab[aName].pairs: if k notin tab[bName]: continue let v = tab[bName][k] inc n if va > v: inc winsA elif v > va: inc winsB else: inc ties if n == 0: return echo &"{aName},{bName},{label},{n},{winsA},{winsB},{ties},{signTestP(winsA, n - ties):.4f}" for a in detArms: if a.name != "Linear": signRow(a.name, "Linear", "early", byEarly) signRow(a.name, "Linear", "overall", byAll) signRow("TMRadial", "Linear", "early", byEarly) signRow("TMRadial", "Linear", "overall", byAll) signRow("TMRadial", "TMRadialShuf", "early", byEarly) signRow("TMRadial", "TMRadialShuf", "overall", byAll) for a in detArms: if a.name != "Linear": signRow(a.name, "TMRadial", "early", byEarly) signRow(a.name, "TMRadial", "overall", byAll) # ── per-adversary optimum for the constant offset ────────────────────────── echo "\n# ── per-fixture best constant offset (pooled over seeds) ──" echo "fixture,bestEarlyVariant,bestEarly%,bestOverallVariant,bestOverall%,linearEarly%,linearOverall%,tmEarly%,tmOverall%" var fxNames: seq[string] for row in rows: if row.fixture notin fxNames: fxNames.add row.fixture fxNames.sort() for fxName in fxNames: var bestE = ("none", -1.0) var bestO = ("none", -1.0) for a in detArms: if a.name == "Linear": continue var eh, en, oh, on: int for row in rows: if row.variant == a.name and row.fixture == fxName: eh += row.r.h100; en += row.r.n100 oh += row.r.hall; on += row.r.nall let er = if en > 0: eh.float/en.float else: -1.0 let orr = if on > 0: oh.float/on.float else: -1.0 if er > bestE[1]: bestE = (a.name, er) if orr > bestO[1]: bestO = (a.name, orr) var lh, ln, loh, lon: int var th, tn, toh, ton: int for row in rows: if row.fixture != fxName: continue if row.variant == "Linear": lh += row.r.h100; ln += row.r.n100; loh += row.r.hall; lon += row.r.nall elif row.variant == "TMRadial": th += row.r.h100; tn += row.r.n100; toh += row.r.hall; ton += row.r.nall echo &"{fxName},{bestE[0]},{pctStr(bestE[1])},{bestO[0]},{pctStr(bestO[1])}," & &"{rateStr(lh,ln)},{rateStr(loh,lon)},{rateStr(th,tn)},{rateStr(toh,ton)}" # ── radial-label distribution + applied shift (from the real TM arm) ─────── echo "\n# ── radial-label distribution and applied shift (TMRadial, delta per fixture) ──" echo "scope,n,labelHist,labelMajority%,meanRadDeltaPx,meanAbsRadDeltaPx,chosenHist,meanAppliedShiftPx,onlineAcc" proc radReport(scope: string, st: GunStats) = var n = 0 var major = 0 var hs: string for c in 0.. 0: st.radDeltaSum / st.radDeltaN.float else: 0.0 let meanAbsD = if st.radDeltaN > 0: st.radDeltaAbsSum / st.radDeltaN.float else: 0.0 let meanShift = if st.radOffsetN > 0: st.radOffsetSum / st.radOffsetN.float else: 0.0 let acc = rateStr(st.radCorrect, st.radTotal) let majPct = if n > 0: &"{major.float/n.float*100.0:.1f}" else: "n/a" echo &"{scope},{n},{hs},{majPct},{meanD:.2f},{meanAbsD:.2f},{chs},{meanShift:.2f},{acc}" var poolSt: GunStats for row in rows: if row.variant == "TMRadial": addStats(poolSt, row.st) radReport(row.fixture, row.st) radReport("POOLED", poolSt) # ── raw radial DELTA distribution from the base forecast (metric-free) ───── echo "\n# ── raw base radial-error distribution (per fired-bullet arrival, from forecastLinear) ──" echo "scope,n,meanPx,meanAbsPx,fracNearer(short),fracFarther(long)" for name in names: let (fx, path) = resolve(name) let fxName = path.extractFilename.replace(".jsonl", "") # map tick -> enemy pose as the TM's own posRing would var pose = initTable[int, tuple[x, y: float]]() for s in fx.states: pose[s.tick] = (s.enemyX, s.enemyY) var sum, asum = 0.0 var n, near, far: int for s in fx.states: for b in 0.. BotRadius: inc far let meanD = if n > 0: sum/n.float else: 0.0 let meanA = if n > 0: asum/n.float else: 0.0 echo &"{fxName},{n},{meanD:.2f},{meanA:.2f},{near.float/max(1,n).float:.3f},{far.float/max(1,n).float:.3f}" when isMainModule: main()