## GATE 2 — does a learned Tsetlin model SHRINK THE MISS, and does that ## translate into hits? OFFLINE ONLY. ## ## Pipeline (all offline, fixtures READ-ONLY): ## 1. Extract the DRAFT 49-bit TM feature spec (walls / us / motion / bullets) ## from the committed DrussGT fixtures, plus a 4-bit one-hot horizon block ## (h in {15,20,25,30}) -> 53 raw bits. ## 2. Build FACT labels: for a sample at tick t and horizon h, look up where ## the enemy ACTUALLY was at t+h (never across a round boundary; the last h ## ticks of each round are dropped). Two binaries: ## (a) side: enemy LEFT / RIGHT of the naive straight-line guess, ## (b) magnitude: |angular error| bigger / smaller than the TRAIN median. ## Four quadrants: left-small / left-big / right-small / right-big. ## 3. Train a real Tsetlin Machine with the validated core at ## `common_libs/tm_diag/tm_core.nim` (one fresh model per round = the ## intended "fresh every round, overfit the current enemy" semantics). ## 4. Map each predicted quadrant to a representative signed angular offset ## (median signed error of the TRAINING samples in that quadrant), apply it ## to the naive aim, and measure the residual angular error. ## ## Four arms + floor: ## naive : straight-line guess (baseline) ## TM : trained model ## shuffled : same pipeline, labels randomised (pipeline-integrity control) ## turn-only : uses ONLY the enemy's current turn direction (critical arm) ## majority : constant majority-quadrant offset (floor) ## ## Metric: median / p90 residual |angular error| (deg) and the estimated hit ## fraction (|residual| < atan(18px / range)). The absolute hit fraction is ## OPTIMISTIC (perfect arrival knowledge every tick = bmPoint-style); only the ## DELTA before/after is meaningful. ## ## Protocol: within-round split, train on the EARLY portion, evaluate on the ## LATER portion. Focus horizons h = 15,20,25 (+30). ## ## Run: nim c -r -d:release --path:common_libs \ ## common_libs/tests/measure_tm_miss_shrink.nim ## [--fixtures=a,b] [--epochs=N] [--trainfrac=F] [--clauses=N] import std/[json, os, strformat, strutils, math, algorithm, tables] import tm_diag/tm_core # ── configuration ──────────────────────────────────────────────────────────── const repoRoot* = currentSourcePath().parentDir.parentDir.parentDir const fixturesDir* = repoRoot / "tools" / "fixtures" const metaDir* = fixturesDir / "drussgt_meta" const HORIZONS = [15, 20, 25, 30] const NH = 4 const N_BASE = 49 # draftTMSpec() bit count const N_BITS = N_BASE + 4 # + 4-bit horizon one-hot const N_CLASSES = 4 # left-small, left-big, right-small, right-big const MIN_I = 12 # need 10 ticks of history for the motion features const BOT_RADIUS = 18.0 var cfgClauses = 50 cfgStates = 64 cfgS = 3.0 cfgEpochs = 5 cfgTrainFrac = 0.70 const PRIMARY = ["tr_drussgt_vs_modularbot.jsonl", "tr_drussgt_vs_modularbot_shield.jsonl"] # ── small helpers ──────────────────────────────────────────────────────────── proc wrap180(a: float): float {.inline.} = var r = a while r > 180.0: r -= 360.0 while r <= -180.0: r += 360.0 r proc signf(x: float): int {.inline.} = if x > 1e-9: 1 elif x < -1e-9: -1 else: 0 proc jf(d: JsonNode, k: string): float = let n = d[k] case n.kind of JFloat: n.getFloat of JInt: float(n.getInt) else: parseFloat(n.getStr) # ── data model ─────────────────────────────────────────────────────────────── type Tick* = object tick*: int ex*, ey*, eh*, es*, ee*: float sx*, sy*, sh*, ss*, se*: float Rnd* = object roundNo*: int st*: seq[Tick] BulletSeries* = object ## Per-tick proxy for OUR in-flight bullets (INFERRED from self-energy ## drops; the fixture records no gun heading/power). tta*: seq[int] ## ticks until the nearest in-flight bullet arrives, -1 none lat*: seq[float] ## lateral offset of the enemy from the fired path (px) Arm = enum aNaive, aTM, aShuf, aTurn, aMaj Hist = object errs: array[Arm, array[NH, seq[float]]] hits: array[Arm, array[NH, int]] n: array[NH, int] # diagnostics tmQCor: array[NH, int] ## TM predicted the true quadrant tmQTot: array[NH, int] tmHitCor: array[NH, int] ## TM predicted the true sign tmHitTot: array[NH, int] turnSignCor: array[NH, int] ## sign(turn) == sign(err) turnSignTot: array[NH, int] shufQCor: array[NH, int] ## shuffled-model predicted the true quadrant shufQTot: array[NH, int] predCounts: array[NH, array[N_CLASSES, int]] trueCounts: array[NH, array[N_CLASSES, int]] # ── fixture loading ────────────────────────────────────────────────────────── proc loadTicks(path: string): seq[Tick] = for line in lines(path): let ln = line.strip() if ln.len == 0: continue let d = parseJson(ln) if not d.hasKey("tick"): continue result.add Tick(tick: d["tick"].getInt, ex: jf(d, "ex"), ey: jf(d, "ey"), eh: jf(d, "eh"), es: jf(d, "es"), ee: jf(d, "ee"), sx: jf(d, "sx"), sy: jf(d, "sy"), sh: jf(d, "sh"), ss: jf(d, "ss"), se: jf(d, "se")) proc loadRounds(path: string, ticks: seq[Tick]): seq[Rnd] = let rp = metaDir / (extractFilename(path) & ".rounds.json") var spans: seq[(int, int)] if fileExists(rp): let j = parseFile(rp) for r in j["rounds"]: spans.add (r["startTick"].getInt, r["count"].getInt) elif ticks.len > 0: spans.add (ticks[0].tick, ticks.len) var idxByTick = initTable[int, int]() for i, t in ticks: idxByTick[t.tick] = i for sp in spans: let (s0, c) = sp if not idxByTick.hasKey(s0): continue let i0 = idxByTick[s0] var st: seq[Tick] for k in 0.. 0: result.add Rnd(roundNo: result.len + 1, st: st) # ── bullet proxy (INFERRED) ────────────────────────────────────────────────── proc buildBulletSeries(r: Rnd): BulletSeries = let L = r.st.len result.tta = newSeq[int](L) for i in 0.. 3.1: continue # damage, not a fire var power = drop if power < 0.1: power = 0.1 if power > 3.0: power = 3.0 let speed = 20.0 - 3.0 * power let rng = hypot(r.st[t0].ex - r.st[t0].sx, r.st[t0].ey - r.st[t0].sy) let flight = int(ceil(rng / speed)) let dx = r.st[t0].ex - r.st[t0].sx let dy = r.st[t0].ey - r.st[t0].sy let nrm = max(1e-6, hypot(dx, dy)) let ux = dx / nrm let uy = dy / nrm for k in 0..flight: let t = t0 + k if t >= L: break let ta = t0 + flight - t if result.tta[t] < 0 or ta < result.tta[t]: result.tta[t] = ta let vx = r.st[t].ex - r.st[t0].sx let vy = r.st[t].ey - r.st[t0].sy result.lat[t] = ux * vy - uy * vx # ── feature extraction: the 49 draft bits (causal, at tick i) ──────────────── # # Block layout (mirrors draftTMSpec()): # 0..3 dist-to-nearest-wall (4 one-hot) # 4..7 which-wall-nearest (4 one-hot) # 8..13 dist-from-us (6 one-hot) # 14..16 enemy-heading-vs-line-to-us (3 one-hot) # 17..19 turn-direction t, t-1, t-2 (3 boolean "was turning left") # 20..24 ticks-since-reversal (5 one-hot) # 25..27 turn-consistency-10 (3 one-hot) # 28..30 distance-moved-10 (3 one-hot) # 31..33 speed-trend-10 (3 one-hot) # 34..36 turn-rate-change-5 (3 one-hot) # 37..41 time-until-bullet (5 one-hot) # 42..48 bullet-lateral-offset (7 one-hot) proc buildBase(r: Rnd, bi: BulletSeries, i: int, sinceRev: seq[int]): array[N_BASE, int] = let s = r.st let cur = s[i] # walls let dL = cur.ex let dR = 800.0 - cur.ex let dT = 600.0 - cur.ey let dBottom = cur.ey let dmin = min(min(dL, dR), min(dT, dBottom)) var wallBin = 3 if dmin < 50.0: wallBin = 0 elif dmin < 100.0: wallBin = 1 elif dmin < 200.0: wallBin = 2 result[wallBin] = 1 var wb = 0 let walls = [dL, dR, dT, dBottom] for w in 1..3: if walls[w] < walls[wb]: wb = w result[4 + wb] = 1 # us let rng = hypot(cur.ex - cur.sx, cur.ey - cur.sy) var ub = 5 if rng < 100.0: ub = 0 elif rng < 200.0: ub = 1 elif rng < 300.0: ub = 2 elif rng < 400.0: ub = 3 elif rng < 600.0: ub = 4 result[8 + ub] = 1 let lane = arctan2(cur.sy - cur.ey, cur.sx - cur.ex) let hdg = cur.eh * PI / 180.0 let perp = abs(sin(hdg - lane)) var hb = 1 if perp < 0.5: hb = 2 elif perp > 0.866: hb = 0 result[14 + hb] = 1 # motion: turn direction for k in 0..2: if i - 1 - k >= 0: let d = wrap180(s[i - k].eh - s[i - 1 - k].eh) if d > 1e-6: result[17 + k] = 1 # ticks since reversal var rb = 4 let sr = sinceRev[i] if sr < 5: rb = 0 elif sr < 10: rb = 1 elif sr < 20: rb = 2 elif sr < 40: rb = 3 result[20 + rb] = 1 # turn consistency over last 10 var pos = 0 var neg = 0 for k in 0..9: if i - 1 - k < 0: break let d = wrap180(s[i - k].eh - s[i - 1 - k].eh) if d > 1e-6: inc pos elif d < -1e-6: inc neg let tot = pos + neg let cons = if tot > 0: max(pos, neg).float / tot.float else: 0.0 var cb = 0 if cons > 0.8: cb = 2 elif cons >= 0.5: cb = 1 result[25 + cb] = 1 # distance moved over 10 let j0 = max(0, i - 10) let dm = hypot(cur.ex - s[j0].ex, cur.ey - s[j0].ey) var mb = 1 if dm < 20.0: mb = 0 elif dm > 50.0: mb = 2 result[28 + mb] = 1 # speed trend over 10 let st10 = abs(s[max(0, i - 10)].es) let spdDiff = abs(cur.es) - st10 var sb = 1 if spdDiff < -0.5: sb = 0 elif spdDiff > 0.5: sb = 2 result[31 + sb] = 1 # turn-rate change: last 5 deltas vs previous 5 var r1 = 0.0 var n1 = 0 for k in 0..4: if i - 1 - k >= 0: r1 += abs(wrap180(s[i - k].eh - s[i - 1 - k].eh)); inc n1 var r2 = 0.0 var n2 = 0 for k in 5..9: if i - 1 - k >= 0: r2 += abs(wrap180(s[i - k].eh - s[i - 1 - k].eh)); inc n2 let m1 = if n1 > 0: r1 / float(n1) else: 0.0 let m2 = if n2 > 0: r2 / float(n2) else: 0.0 let dtr = m1 - m2 var tb = 1 if dtr < -0.3: tb = 0 elif dtr > 0.3: tb = 2 result[34 + tb] = 1 # bullets let tta = bi.tta[i] var b1 = 0 if tta >= 0: if tta < 5: b1 = 1 elif tta < 10: b1 = 2 elif tta < 20: b1 = 3 else: b1 = 4 result[37 + b1] = 1 let lat = bi.lat[i] var lb = 3 if lat < -72.0: lb = 0 elif lat < -36.0: lb = 1 elif lat < -18.0: lb = 2 elif lat <= 18.0: lb = 3 elif lat <= 36.0: lb = 4 elif lat <= 72.0: lb = 5 else: lb = 6 result[42 + lb] = 1 proc toLits(base: array[N_BASE, int], h: int): seq[uint8] = var raw: array[N_BITS, int] for i in 0.. bestV: bestV = v result = c # ── statistics ─────────────────────────────────────────────────────────────── proc medOf(v: seq[float]): float = if v.len == 0: return NaN var s = v s.sort() s[s.len div 2] proc qOf(v: seq[float], q: float): float = if v.len == 0: return NaN var s = v s.sort() s[min(s.len - 1, max(0, int(q * float(s.len - 1) + 0.5)))] proc medOfInts(v: seq[int]): float = if v.len == 0: return NaN var s = v s.sort() float(s[s.len div 2]) proc mergeHist(dst: var Hist, src: Hist) = for a in Arm: for hi in 0..offset mapping is # calibrated OUT-OF-SAMPLE so an overfit in-sample median cannot leak. let fitEnd = max(MIN_I + 1, int(0.50 * float(L))) let calEnd = max(fitEnd + 1, int(cfgTrainFrac * float(L))) # local collector: samples with i in [lo,stop) and the label j = i+h < stop # (so NOTHING here reads past `stop` — no cross-region label leakage). proc collect(lo, stop, hidx: int): tuple[hs: seq[int], es: seq[float], ls: seq[seq[uint8]], ts: seq[float]] = let h = HORIZONS[hidx] for i in lo..= stop: continue let cur = s[i] let gx = cur.ex + cur.es * cos(cur.eh * PI / 180.0) * float(h) let gy = cur.ey + cur.es * sin(cur.eh * PI / 180.0) * float(h) let ba = arctan2(s[j].ey - cur.sy, s[j].ex - cur.sx) let bg = arctan2(gy - cur.sy, gx - cur.sx) let err = radToDeg(arctan2(sin(ba - bg), cos(ba - bg))) result.hs.add hidx result.es.add err result.ls.add toLits(base[i], h) result.ts.add(if i - 1 >= 0: wrap180(cur.eh - s[i - 1].eh) else: 0.0) var trH: seq[int] var trErr: seq[float] var trLits: seq[seq[uint8]] var trTurn: seq[float] var calH: seq[int] var calErr: seq[float] var calLits: seq[seq[uint8]] var calTurn: seq[float] var evH: seq[int] var evIdx: seq[int] var evErr: seq[float] var evRange: seq[float] var evTurn: seq[float] for hi in 0..= L: continue let cur = s[i] let gx = cur.ex + cur.es * cos(cur.eh * PI / 180.0) * float(h) let gy = cur.ey + cur.es * sin(cur.eh * PI / 180.0) * float(h) let ba = arctan2(s[j].ey - cur.sy, s[j].ex - cur.sx) let bg = arctan2(gy - cur.sy, gx - cur.sx) let err = radToDeg(arctan2(sin(ba - bg), cos(ba - bg))) evH.add hi evIdx.add i evErr.add err evRange.add hypot(s[j].ex - cur.sx, s[j].ey - cur.sy) evTurn.add(if i - 1 >= 0: wrap180(cur.eh - s[i - 1].eh) else: 0.0) # ── labels (true quadrants); magnitude threshold = FIT median |err| ── var medAbs: array[NH, float] for hi in 0.. 0.0: 0 else: 2) + (if abs(trErr[k]) > medAbs[hi]: 1 else: 0)) tlits.add trLits[k] terr.add trErr[k] th.add hi tturn.add trTurn[k] tlab.add cls proc sideMedian(errs: seq[float]): array[2, float] = var l, r: seq[float] for e in errs: if e > 1e-9: l.add e elif e < -1e-9: r.add e result[0] = medOf(l) result[1] = medOf(r) # ── train TM (fresh per round) ── var tm = newMachine(N_BITS, N_CLASSES, cfgClauses, cfgStates, cfgS, seed = 12345'u64 + uint64(r.roundNo)) trainMachine(tm, tlits, tlab, cfgEpochs, seed = 999'u64 + uint64(r.roundNo)) # Calibrate OUT-OF-SAMPLE on the calibration slice: # offset[c] = median(err | model predicts c) on calibration ticks. # This is the L1-optimal correction for a model-dependent partition. var repOff: array[NH, array[N_CLASSES, float]] block: var clsE: array[NH, array[N_CLASSES, seq[float]]] var cache = newSeq[uint8](tm.nClauses) for idx in 0.. 0: medOf(clsE[hi][c]) else: (if c < 2: sm2[0] else: sm2[1]) # constant floor: the majority true quadrant (from the fit portion) var majClass: array[NH, int] block: var cnt: array[NH, array[N_CLASSES, int]] for idx in 0.. cnt[hi][best]: best = c majClass[hi] = best # ── shuffled-label control ── # Permute the fit labels within each horizon, retrain, and calibrate the # SAME way (offset = median err | shuffled model predicts c). var slab = tlab block: var rng = seedRng(4242'u64 + uint64(r.roundNo)) for hi in 0.. 0: medOf(clsE[hi][c]) else: (if c < 2: sm2[0] else: sm2[1]) # ── turn-direction-only rule ── # LEARN the direction of the association per horizon on the CALIBRATION # slice (the headroom study shows it flips sign with horizon), then # calibrate the offset on the rule's own predictions. var turnPred: array[NH, array[2, int]] # [hi][turn>0 ? 0 : 1] -> side var turnOff: array[NH, array[2, float]] block: for hi in 0.. 1e-6: if calErr[idx] > 0.0: inc posL else: inc posR elif calTurn[idx] < -1e-6: if calErr[idx] > 0.0: inc negL else: inc negR turnPred[hi][0] = if posL >= posR: 0 else: 1 turnPred[hi][1] = if negL >= negR: 0 else: 1 var offE: array[NH, array[2, seq[float]]] for idx in 0.. 1e-6: turnPred[hi][0] elif calTurn[idx] < -1e-6: turnPred[hi][1] else: majClass[hi] div 2 offE[hi][ps].add calErr[idx] for hi in 0.. 0: medOf(offE[hi][sd]) else: sm2[sd] # ── evaluate ── var sc = newSeq[uint8](sm.nClauses) var tc = newSeq[uint8](tm.nClauses) for k in 0.. 1e-6: turnPred[hi][0] elif evTurn[k] < -1e-6: turnPred[hi][1] else: majClass[hi] div 2 let offs: array[Arm, float] = [ aNaive: 0.0, aTM: repOff[hi][predTM], aShuf: shufOff[hi][predShuf], aTurn: turnOff[hi][predSide], aMaj: repOff[hi][majClass[hi]]] # diagnostics on the SAME eval tick let trueCls = ((if err > 0.0: 0 else: 2) + (if abs(err) > medAbs[hi]: 1 else: 0)) inc result.trueCounts[hi][trueCls] inc result.predCounts[hi][predTM] inc result.tmQTot[hi] if predTM == trueCls: inc result.tmQCor[hi] inc result.shufQTot[hi] if predShuf == trueCls: inc result.shufQCor[hi] inc result.tmHitTot[hi] if (predTM <= 1) == (trueCls <= 1): inc result.tmHitCor[hi] inc result.turnSignTot[hi] if (predSide == 0) == (err > 0.0): inc result.turnSignCor[hi] for a in Arm: let res = err - offs[a] result.errs[a][hi].add abs(res) if abs(res) < half: inc result.hits[a][hi] # ── per-round table ── if emit: echo &" round {r.roundNo:>2} (L={L:>5}) " & "nEval/h = " & $[result.n[0], result.n[1], result.n[2], result.n[3]] echo " h arm med p90 hit%" for hi in 0.. 0: 100.0 * result.hits[a][hi].float / result.n[hi].float else: NaN echo &" {HORIZONS[hi]:>3} {ArmName[a]:<9} {m:>6.2f} {p:>6.2f} {hit:>6.1f}" # ── table printing ─────────────────────────────────────────────────────────── proc printPooled(h: Hist, title: string) = echo "\n" & title echo " h arm N med|err| p90|err| hit% miss%" for hi in 0.. 0: 100.0 * h.hits[a][hi].float / n.float else: NaN echo &"{HORIZONS[hi]:>3} {ArmName[a]:<9} {n:>7} {m:>9.2f} {p:>9.2f} {hit:>7.1f} {100.0-hit:>7.1f}" proc printDeltas(h: Hist, title: string) = echo "\n" & title echo " h d(med) TM-naive d(p90) TM-naive d(hit) TM-naive d(hit) TM-turn d(hit) shuf-naive" for hi in 0..3} {dm:>+16.2f} {dp:>+16.2f} " & &"{hitTM-hitNaive:>+16.1f} {hitTM-hitTurn:>+15.1f} {hitShuf-hitNaive:>+18.1f}" proc printDiagnostics(h: Hist, title: string) = ## Is the learner actually learning? (quadrant / side accuracy vs chance) echo "\n" & title echo " h N TMquad% shufquad% TMside% turnSide% majQuad% pred[Ls Lb Rs Rb] / true[Ls Lb Rs Rb]" for hi in 0..3} {n:>6} {tmq:>8.1f} {shufq:>10.1f} {tms:>8.1f} {ts:>10.1f} {maj:>9.1f} [{pcs}] / [{tcs}]" # ── main ───────────────────────────────────────────────────────────────────── proc main() = var names = @PRIMARY for i in 1..paramCount(): let a = paramStr(i) if a.startsWith("--fixtures="): names = a[11..^1].split(',') elif a.startsWith("--epochs="): cfgEpochs = parseInt(a[9..^1]) elif a.startsWith("--trainfrac="): cfgTrainFrac = parseFloat(a[12..^1]) elif a.startsWith("--clauses="): cfgClauses = parseInt(a[10..^1]) echo "=" .repeat(90) echo "GATE 2 - does a learned Tsetlin model SHRINK THE MISS (and turn it into hits)?" echo "=" .repeat(90) echo &"fixtures : {names.join(\", \")}" echo &"bits : {N_BITS} ({N_BASE} draft + 4 horizon one-hot)" echo &"TM : {cfgClauses} clauses, {cfgStates} states, s={cfgS}, " & &"{cfgEpochs} epochs, fresh per round" echo &"split : within-round, train first {cfgTrainFrac*100:.0f}%, eval later portion" echo &"label : side=sign(err); mag=|err| > train median (per round,h)" echo &"estimated hit : |residual| < atan(18px / range) -- ABSOLUTE IS OPTIMISTIC" echo &"bullet block : INFERRED from self-energy drops (no gun heading recorded)" echo "=" .repeat(90) var pooled = Hist() var pooledByFix = initTable[string, Hist]() var roundHists = initTable[string, seq[Hist]]() for name in names: let path = if name.endsWith(".jsonl"): fixturesDir / name else: fixturesDir / (name & ".jsonl") if not fileExists(path): echo &"# SKIP missing fixture {path}" continue let ticks = loadTicks(path) let rounds = loadRounds(path, ticks) echo &"\n## FIXTURE {name}: {rounds.len} rounds, {ticks.len} ticks" var fhist = Hist() var rlist: seq[Hist] for r in rounds: let bi = buildBulletSeries(r) var cache = newSeq[uint8](cfgClauses) let rh = runRound(r, bi, cache, emit = true) mergeHist(fhist, rh) mergeHist(pooled, rh) rlist.add rh echo "" pooledByFix[name] = fhist roundHists[name] = rlist printPooled(fhist, &"## PER-FIXTURE pooled ({name})") printDiagnostics(fhist, &"## PER-FIXTURE learnability ({name})") echo "\n" & "=".repeat(90) printPooled(pooled, "## POOLED PRIMARY (tr_drussgt_vs_modularbot*): four-arm + floor table") printDeltas(pooled, "## DELTAS (percentage points of estimated hit fraction)") printDiagnostics(pooled, "## POOLED learnability diagnostics (is the learner learning?)") echo "\n" & "=" .repeat(90) echo "## SHUFFLED-CONTROL INTEGRITY (per round, should stay ~flat)" echo "fixture round h TM-naive shuf-naive turn-naive" for name in names: if not roundHists.hasKey(name): continue for rn, rh in roundHists[name]: for hi in 0..5} {HORIZONS[hi]:>3} {hitTM-hitNaive:>+9.1f} " & &"{hitShuf-hitNaive:>+11.1f} {hitTurn-hitNaive:>+11.1f}" when isMainModule: main()