From ca82053a110ec4db47944095b2e9d1d1c78c6e2f Mon Sep 17 00:00:00 2001 From: Davide Cappellini Date: Tue, 22 Sep 2026 00:56:01 +0200 Subject: [PATCH] TM gun: the discrete-target diagnosis was RIGHT - it learns now. Still loses to Linear. The user's goal: a TM gun that is the best 1v1 gun, starting from scratch every battle but quickly overfitting the current enemy. The previous attempt (knob tuning) failed: NO configuration beat its own shuffled-feedback control, and the TM-off ablation scored the same as TM-on, i.e. the TM's correction was near-zero-mean noise. Diagnosis then: a Tsetlin Machine is a CLASSIFIER, and we were asking it for an absolute aim point - a regression target. So this attempt gave it a DISCRETE target (multi-class over guess-factor buckets) with 40 binary/bucketed motion features, and measured it against Linear, the default Tsetlin gun, and a MANDATORY shuffled control. THE DIAGNOSIS IS CONFIRMED - THE TM LEARNS, DECISIVELY: online class accuracy 46.0% vs shuffled control 20.0% (2.3x chance) raw ungated argmax 21.2%/18.6% vs shuffled 15.2%/8.3% (18/18, p<0.0001) TMPattern > its shuffled control, overall 17/1 runs, p=0.0001 Compare the previous attempt, which could not beat shuffled feedback at all. TMPattern also beats the default Tsetlin gun early (17/1, p=0.0001), so it is a strictly better TM gun than the one in the rack. BUT IT IS NOT COMPETITIVE WITH LINEAR ON REAL SURFERS: real DrussGT, bmPath (the shipped metric), 3 seeds, pooled early/overall Linear 34.0% (6358/18715) 24.3% (58297/239943) TMPattern (gated) 27.9% (15514/55535) 22.0% (158658/719681) TMPatternShuf 28.7% 19.4% Linear > TMPattern: 15/18 early p=0.0075, 15/18 overall p=0.0075 bmPoint: neutral (7.2%/4.6% vs Linear 7.2%/4.7%) synthetic controlled motion: matches/edges Linear (66.8%/60.6% vs 66.4%/59.6%, shuffled 55.7%/50.1%) - the mechanism works when motion is predictable. So: the representation fix moved this from "learns nothing" to "learns strongly but applies its knowledge badly". INFERRED reason for the residual loss: the linear lead is already the modal GF bucket (the label histogram is centred), so corrective excursions away from it are net-negative. The measured deficit lives in the BASELINE and in RANGE, not in the TM knobs - which is why further knob tuning was never going to work. Best config: gated hard K=5, TM_CONF_MARGIN=0.25, TM_SHRINK=0.5. NOT TRIED (time-boxed): the binary-reversal target, and a RADIAL (range-holding) target - the latter is the top next step. Adds `common_libs/guns/tm_pattern.nim` (NOT registered in the rack), `common_libs/tests/sweep_tm_pattern.nim`, and a durable writeup at `common_libs/tests/tm_pattern_sweep_results.md`. --- common_libs/guns/tm_pattern.nim | 493 ++++++++++++++++++ common_libs/tests/sweep_tm_pattern.nim | 349 +++++++++++++ common_libs/tests/tm_pattern_sweep_results.md | 181 +++++++ 3 files changed, 1023 insertions(+) create mode 100644 common_libs/guns/tm_pattern.nim create mode 100644 common_libs/tests/sweep_tm_pattern.nim create mode 100644 common_libs/tests/tm_pattern_sweep_results.md diff --git a/common_libs/guns/tm_pattern.nim b/common_libs/guns/tm_pattern.nim new file mode 100644 index 0000000..d0b6ac9 --- /dev/null +++ b/common_libs/guns/tm_pattern.nim @@ -0,0 +1,493 @@ +## TM pattern gun — a DISCRETE-target Tsetlin Machine on top of a self-consistent +## linear forecast. +## +## Why this exists (and why it is not `guns/tsetlin.nim`): a previous sweep +## (common_libs/tests/sweep_tsetlin.nim, documented in the tsetlin.nim header) +## found that the old gun's TM used as a pixel-correction REGRESSOR learns +## nothing — every config was indistinguishable from its own shuffled-feedback +## control, and the whole deficit vs Linear lived in the old gun's one-shot +## internal baseline. The conclusion was that the REPRESENTATION and the TARGET +## were the problem, not the knobs. +## +## This gun attacks both: +## +## * BASE — `forecastLinear` from lead_forecast.nim, the exact self-consistent +## constant-velocity forecast LinearGun uses, so GF class 0 (center) +## reproduces the Linear gun byte-for-byte. Any measured difference +## is attributable to the TM, not to a weaker baseline. +## * TARGET — the discrete class is a GUESS-FACTOR BUCKET: which lateral +## escape sector (in max-escape-angle units) the enemy occupied at +## the tick our bullet would have arrived. A small multi-class +## classification, which is what a Tsetlin Machine is for. +## * LABEL — computed from the enemy position at the BASE forecast's arrival +## tick, looked up in a per-tick position ring. This is deliberate: +## under the shipped `bmPath` metric `FeedbackEvent.actualXY` is the +## closest-approach point on the gun's OWN aim ray, which biases the +## label toward the gun's own last output (a self-referential +## feedback loop). Reading our own recorded history at the base +## arrival tick gives a clean, metric-independent label. +## * FEATURES — binary/bucketed motion context (lateral-velocity sign over 3 +## ticks, turn-rate sign, time since reversal, radial fraction, +## speed/distance/flight-time/wall/energy bands). 40 bits. +## +## The TM core is a compact, self-contained Granmo Table 2/3 implementation with +## the corrected feedback rules and Eq. 6 empty-clause bootstrap (the same +## corrected core as tsetlin.nim / tm_selector.nim, re-derived here at a small +## feature width so a 5-class team is cheap). +## +## Per-enemy specialisation: the net is FRESH per gun instance (battle) and the +## gun resets it if the target id changes mid-battle. Offline the range replays +## one round per fresh instance, which is exactly "cold every battle, overfit +## within the battle". + +import std/[math, random, strutils] +import gun_harness/gun_interface +import gun_harness/virtual_bullets as vb +import guns/lead_forecast + +const + ## ── classifier shape ──────────────────────────────────────────────────── + TM_CLASSES* {.intdefine.} = 5 ## GF buckets, centers -1, -0.5, 0, +0.5, +1 + TM_NBITS* = 40 + TM_NLITS* = TM_NBITS * 2 + TM_NCLAUSES* {.intdefine.} = 40 ## per class + TM_HALF* = TM_NCLAUSES div 2 + TM_NSTATES* {.intdefine.} = 64 ## automaton range [-NSTATES, NSTATES] + TM_T* = float(TM_HALF) + TM_S_DEF {.strdefine.} = "3.0" + TM_S* = parseFloat(TM_S_DEF) + TM_MIN_OBS* {.intdefine.} = 24 ## freshly-cold -> look straight ahead + ## Confidence gate: only leave the centre (GF=0) bucket when the winning class + ## beats the centre class by this fraction of TM_T. 0.0 = raw argmax. + TM_CONF_MARGIN_DEF {.strdefine.} = "0.0" + TM_CONF_MARGIN* = parseFloat(TM_CONF_MARGIN_DEF) + ## Scale applied to the predicted GF when a correction is taken. + TM_SHRINK_DEF {.strdefine.} = "1.0" + TM_SHRINK* = parseFloat(TM_SHRINK_DEF) + ## Readout: "hard" argmax (with the confidence gate) or "soft" vote-weighted. + TM_GF_MODE {.strdefine.} = "hard" + TM_SOFT_BETA_DEF {.strdefine.} = "4.0" + TM_SOFT_BETA* = parseFloat(TM_SOFT_BETA_DEF) + TM_TRACE_SLOTS = 1024 + POS_RING = 512 + DebugTMPattern* = false + +type + TmBits* = array[TM_NLITS, uint8] + + TmPatternTrace = object + fireTick: int + powerBin: int + arrivalTick: int + baseBearing: float + fireX, fireY: float + lits: TmBits + votes: array[TM_CLASSES, float] + cache: array[TM_CLASSES, array[TM_NCLAUSES, uint8]] + chosen: int + warm: bool + alive: bool + + PosSample = object + tick: int + x, y: float + valid: bool + + TmPatternGun* = object + teams: array[TM_CLASSES, seq[int16]] + traces: array[TM_TRACE_SLOTS, TmPatternTrace] + # ── history ── + posRing: array[POS_RING, PosSample] + lastTick: int + prevTick: int + prevX, prevY, prevHeading: float + hasPrev: bool + latSignHist: array[3, int] + turnSignHist: array[3, int] + lastNonzeroLat: int + sinceReversal: int + radialFracSm: float + latPersist: int + currentTarget: int + # ── instrumentation ── + totalObs*: int + predictCalls*: int + trainCalls*: int + traceMisses*: int + labelMisses*: int + chosenHist*: array[TM_CLASSES, int] + labelHist*: array[TM_CLASSES, int] + classCorrect*: int ## warm predictions whose class matched the eventual label + classTotal*: int ## warm predictions with a resolvable label + lastChosen*: int + shuffleLabels*: bool ## control: replace the computed GF label with a random class + debugGraphics*: bool + +# ── TM core (Granmo Table 2/3, corrected resource allocation) ──────────────── + +proc tmPolarity(cl: int): float {.inline.} = + if cl < TM_HALF: 1.0 else: -1.0 + +proc tmNewTeam(): seq[int16] = + result = newSeq[int16](TM_NCLAUSES * TM_NLITS) # 0 = Exclude boundary + +proc tmEval(team: seq[int16], lits: TmBits, cl: int, learning: bool): uint8 = + let base = cl * TM_NLITS + var hasInc = false + for lit in 0.. 0: + hasInc = true + if lits[lit] == 0'u8: return 0'u8 + if hasInc: return 1'u8 + # Eq. 6: the empty conjunction is vacuously true during learning, false in + # classification. Without this the all-Exclude init deadlocks. + return if learning: 1'u8 else: 0'u8 + +proc tmForward(team: seq[int16], lits: TmBits, + cache: var array[TM_NCLAUSES, uint8]): float = + var v = 0.0 + for cl in 0..= pFeedback: continue + let pol = tmPolarity(cl) + let cOut = cache[cl] + let base = cl * TM_NLITS + if pol * d > 0.0: + # Type I (Table 2) collapsed to the resulting state move. + for lit in 0.. 180.0: result -= 360.0 + while result < -180.0: result += 360.0 + +# ── public API ─────────────────────────────────────────────────────────────── + +proc initTmPatternGun*(): TmPatternGun = + for c in 0.. g.prevTick: + let dx = state.enemyX - g.prevX + let dy = state.enemyY - g.prevY + let spd = hypot(dx, dy) + let lx = state.enemyX - state.selfX + let ly = state.enemyY - state.selfY + let ld = hypot(lx, ly) + var latSign = 0 + var crossFrac = 0.0 + var radialFrac = 0.0 + if ld > 1e-6 and spd > 1e-6: + let cross = (lx * dy - ly * dx) / ld + crossFrac = abs(cross) / spd + radialFrac = abs((lx * dx + ly * dy) / ld) / spd + if cross > 0.5: latSign = 1 + elif cross < -0.5: latSign = -1 + for i in countdown(2, 1): g.latSignHist[i] = g.latSignHist[i - 1] + g.latSignHist[0] = latSign + if latSign != 0 and g.lastNonzeroLat != 0 and latSign != g.lastNonzeroLat: + g.sinceReversal = 0 + else: + inc g.sinceReversal + if latSign != 0: g.lastNonzeroLat = latSign + g.latPersist = if latSign != 0 and latSign == g.latSignHist[1]: 1 else: 0 + var dh = normDeg(state.enemyHeading - g.prevHeading) + let ts = if dh > 0.5: 1 elif dh < -0.5: -1 else: 0 + for i in countdown(2, 1): g.turnSignHist[i] = g.turnSignHist[i - 1] + g.turnSignHist[0] = ts + g.radialFracSm = 0.8 * g.radialFracSm + 0.2 * radialFrac + g.prevX = state.enemyX + g.prevY = state.enemyY + g.prevHeading = state.enemyHeading + g.prevTick = state.tick + g.hasPrev = true + +proc tmBuildBits(g: var TmPatternGun, state: WorldState, flightTicks: float): + TmBits = + ## 40 binary/bucketed context features. Written as literals directly: bit i + ## and its negation at i + TM_NBITS. + var bits: array[TM_NBITS, uint8] + var o = 0 + template put(v: uint8) = (bits[o] = v; inc o) + template putSign(s: int) = + # ternary -> 2 bits: (+, -); 0 -> (0,0) + put(if s > 0: 1'u8 else: 0'u8) + put(if s < 0: 1'u8 else: 0'u8) + for i in 0..2: putSign(g.latSignHist[i]) # 6 + for i in 0..2: putSign(g.turnSignHist[i]) # 6 + + # time since last lateral reversal: 3 one-hot + if g.sinceReversal <= 3: put 1'u8 else: put 0'u8 + if g.sinceReversal > 3 and g.sinceReversal <= 10: put 1'u8 else: put 0'u8 + if g.sinceReversal > 10: put 1'u8 else: put 0'u8 + + # lateral persistence / magnitude: 3 bits + put(if g.latPersist == 1: 1'u8 else: 0'u8) + + let spd = state.enemySpeed + # speed band (uses signed velocity magnitude; classic fixtures may be negative) + let aspd = abs(spd) + if aspd < 1.0: put 1'u8 else: put 0'u8 + if aspd >= 1.0 and aspd < 4.0: put 1'u8 else: put 0'u8 + if aspd >= 4.0: put 1'u8 else: put 0'u8 + + # distance band + let d = hypot(state.enemyX - state.selfX, state.enemyY - state.selfY) + if d < 150.0: put 1'u8 else: put 0'u8 + if d >= 150.0 and d < 350.0: put 1'u8 else: put 0'u8 + if d >= 350.0: put 1'u8 else: put 0'u8 + + # flight-time band + if flightTicks < 10.0: put 1'u8 else: put 0'u8 + if flightTicks >= 10.0 and flightTicks < 25.0: put 1'u8 else: put 0'u8 + if flightTicks >= 25.0: put 1'u8 else: put 0'u8 + + # walls close (4 bits) + put(if state.enemyY < 60.0: 1'u8 else: 0'u8) + put(if state.arenaHeight - state.enemyY < 60.0: 1'u8 else: 0'u8) + put(if state.arenaWidth - state.enemyX < 60.0: 1'u8 else: 0'u8) + put(if state.enemyX < 60.0: 1'u8 else: 0'u8) + + # radial fraction band (3 one-hot) + if g.radialFracSm < 0.35: put 1'u8 else: put 0'u8 + if g.radialFracSm >= 0.35 and g.radialFracSm < 0.7: put 1'u8 else: put 0'u8 + if g.radialFracSm >= 0.7: put 1'u8 else: put 0'u8 + + # enemy energy band (2 one-hot) + if state.enemyEnergy < 20.0: put 1'u8 else: put 0'u8 + if state.enemyEnergy >= 20.0: put 1'u8 else: put 0'u8 + + # heading relative to LOS (toward / away) + let lx = state.enemyX - state.selfX + let ly = state.enemyY - state.selfY + let ld = hypot(lx, ly) + var dot = 0.0 + if ld > 1e-6: + let hr = degToRad(state.enemyHeading) + dot = (cos(hr) * lx + sin(hr) * ly) / ld + put(if dot > 0.3: 1'u8 else: 0'u8) + put(if dot < -0.3: 1'u8 else: 0'u8) + + # approach (closing / opening) relative to displacement direction + var closing = 0.0 + if ld > 1e-6 and g.hasPrev: + let vx = state.enemyX - g.prevX + let vy = state.enemyY - g.prevY + closing = (vx * lx + vy * ly) / ld + put(if closing < -0.3: 1'u8 else: 0'u8) + put(if closing > 0.3: 1'u8 else: 0'u8) + + # pack into literal vector + for i in 0.. GF 0; a confident tail -> near that + ## tail), which avoids the up-to-half-bucket aim error of a hard argmax. + var mx = -Inf + for c in 0.. mx: mx = votes[c] + var sum = 0.0 + var w: array[TM_CLASSES, float] + for c in 0.. straight ahead (GF = 0). Otherwise the + ## argmax class, but ONLY when it beats the centre class by TM_CONF_MARGIN + ## (fraction of TM_T); otherwise stay at the centre. This is what keeps the + ## gun from degrading to arbitrary buckets when the TM has no real evidence. + let centre = (TM_CLASSES - 1) div 2 + if g.totalObs < TM_MIN_OBS: + return centre + var best = 0 + var bestV = -Inf + for c in 0.. bestV: bestV = votes[c]; best = c + if best == centre: + return centre + let margin = (votes[best] - votes[centre]) / TM_T + if margin < TM_CONF_MARGIN: + return centre + result = best + +proc predict*(g: var TmPatternGun, state: WorldState, bulletSpeed: float): + GunPrediction = + inc g.predictCalls + + # target-change reset (per-enemy specialisation) + if state.enemies.len > 0: + let tid = state.enemies[0].id + if g.currentTarget != tid: + if g.currentTarget >= 0: g.resetLearning() + g.currentTarget = tid + + g.tmUpdateHistory(state) + + if bulletSpeed <= 0.0: + return GunPrediction(x: state.enemyX, y: state.enemyY) + + let mea = arcsin(clamp(8.0 / bulletSpeed, -1.0, 1.0)) + let f = forecastLinear(state, bulletSpeed) + let flightTicks = f.dist / bulletSpeed + + let lits = g.tmBuildBits(state, flightTicks) + + var votes: array[TM_CLASSES, float] + var caches: array[TM_CLASSES, array[TM_NCLAUSES, uint8]] + for c in 0..= 0: + let slot = tmTraceSlot(state.tick, binIdx) + # The virtual bullet advances one step on its own spawn tick, so it reaches + # the base fire distance after max(0, ceil(flightTicks)-1) further ticks. + let arrOff = max(0, int(ceil(flightTicks)) - 1) + g.traces[slot] = TmPatternTrace( + fireTick: state.tick, powerBin: binIdx, + arrivalTick: state.tick + arrOff, + baseBearing: f.bearing, + fireX: state.selfX, fireY: state.selfY, + lits: lits, votes: votes, cache: caches, + chosen: chosen, warm: (g.totalObs >= TM_MIN_OBS), alive: true) + + when DebugTMPattern: + echo "tmp tick=", state.tick, " bin=", binIdx, " chosen=", chosen, + " gf=", gf, " obs=", g.totalObs, " votes=", votes + + GunPrediction( + x: clamp(px, BotRadius, state.arenaWidth - BotRadius), + y: clamp(py, BotRadius, state.arenaHeight - BotRadius), + ) + +proc onResult*(g: var TmPatternGun, e: FeedbackEvent) = + let binIdx = + if e.powerBin >= 0 and e.powerBin < len(vb.PowerBins): e.powerBin + else: tmBinForSpeed(bulletSpeed(e.bulletPower)) + if binIdx < 0: + inc g.traceMisses + return + let slot = tmTraceSlot(e.fireTick, binIdx) + var t = addr g.traces[slot] + if not t.alive or t.fireTick != e.fireTick or t.powerBin != binIdx: + inc g.traceMisses + return + + # Clean label: enemy position at the BASE arrival tick from our own history. + let s = ((t.arrivalTick mod POS_RING) + POS_RING) mod POS_RING + if not g.posRing[s].valid or g.posRing[s].tick != t.arrivalTick: + inc g.labelMisses + t.alive = false + return + + let speed = bulletSpeed(e.bulletPower) + let mea = arcsin(clamp(8.0 / speed, -1.0, 1.0)) + let actualBearing = arctan2(g.posRing[s].y - t.fireY, g.posRing[s].x - t.fireX) + var delta = actualBearing - t.baseBearing + while delta > PI: delta -= 2.0 * PI + while delta < -PI: delta += 2.0 * PI + let gf = if mea > 1e-10: clamp(delta / mea, -1.0, 1.0) else: 0.0 + let winner = if g.shuffleLabels: rand(TM_CLASSES - 1) else: gfToBucket(gf) + inc g.labelHist[winner] + if t.warm: + inc g.classTotal + if winner == t.chosen: inc g.classCorrect + + for c 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 + tracker.tickBullets(state, enemyPositions, + proc(gunId: GunId, binIdx: int, e: FeedbackEvent) = + inc res.nall + if e.hit: inc res.hall + if localTick < 100: + inc res.n100 + if e.hit: inc res.h100 + if localTick < 300: + inc res.n300 + if e.hit: inc res.h300 + let fireTick = e.fireTick - baseTick + if fireTick < 100: + inc res.m100 + if e.hit: inc res.f100 + driver.resultCb(e)) + let after = obsCount() + res.st.obs = after.obs - obsBefore.obs + res.st.labelMiss = after.labelMiss - obsBefore.labelMiss + res.st.traceMiss = after.traceMiss - obsBefore.traceMiss + res.st.classCorrect = after.classCorrect - obsBefore.classCorrect + res.st.classTotal = after.classTotal - obsBefore.classTotal + for c in 0.. 0 and spans.len > maxRounds: spans.setLen(maxRounds) + if spans.len == 0: + return replayRound(fx.states, fx.lastSeen, fx.enemyId, 0, driver, metric, + obsBefore, obsCount) + 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 + addAdapt(result, replayRound(st, ls, fx.enemyId, sp.start, driver, metric, + obsBefore, obsCount)) + +proc emptyStats(): GunStats = GunStats() + +proc makeTmpatDriver(seed: int, shuffle: bool): + tuple[driver: GunDriver, gun: ref TmPatternGun] = + let g = new(TmPatternGun) + g[] = initTmPatternGun() + g[].shuffleLabels = shuffle + if seed >= 0: randomize(seed) + result.gun = g + result.driver = GunDriver( + name: (if shuffle: "TMPatternShuf" else: "TMPattern"), + predictCb: proc(state: WorldState, bulletSpeed: float): GunPrediction = + g[].predict(state, bulletSpeed), + resultCb: proc(e: FeedbackEvent) = g[].onResult(e), + readyCb: proc(): bool = g[].isWarmedUp()) + +proc tmpatStats(g: ref TmPatternGun): GunStats = + result.obs = g[].totalObs + result.labelMiss = g[].labelMisses + result.traceMiss = g[].traceMisses + result.labHist = g[].labelHist + result.choHist = g[].chosenHist + result.classCorrect = g[].classCorrect + result.classTotal = g[].classTotal + +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 + +proc resolve(name: string): tuple[fx: Fixture, path: string] = + let p = if fileExists(name): name else: fixturesDir / (name & ".jsonl") + (loadFixture(p), p) + +# ── stats ──────────────────────────────────────────────────────────────────── + +proc rateStr(h, n: int): string = + if n == 0: " n/a " else: &"{h.float / n.float * 100.0:5.1f}%" + +proc binomPmf(k, n: int): float = + if k < 0 or k > 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 = + ## exact two-sided binomial p under p=0.5 + 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) + +proc main() = + var nSeeds = 3 + var set = "real" + var maxRounds = 0 + var metricName = "path" + var variants: seq[VariantKind] = @[vLinear, vTsetlin, vTmpat, vTmpatShuf] + for i in 1..paramCount(): + let a = paramStr(i) + if a.startsWith("--seeds="): nSeeds = parseInt(a[8..^1]) + elif a.startsWith("--set="): set = a[6..^1] + elif a.startsWith("--metric="): metricName = a[9..^1] + elif a.startsWith("--maxrounds="): maxRounds = parseInt(a[12..^1]) + elif a.startsWith("--variants="): + variants = @[] + for tok in a[11..^1].split(','): + case tok.strip() + of "linear": variants.add vLinear + of "tsetlin": variants.add vTsetlin + of "tmpat": variants.add vTmpat + of "tmpat_shuf": variants.add vTmpatShuf + else: discard + let metric = if metricName == "point": bmPoint else: bmPath + let names = fixtureSet(set) + echo "# set=", set, " seeds=", nSeeds, " metric=", metricName, " variants=", variants.mapIt(variantName(it)).join(",") + echo "variant,fixture,seed,h100,n100,h300,n300,hall,nall,f100,m100,rounds,obs,labelMiss,traceMiss" + + var rows: seq[Row] + for name in names: + let (fx, path) = resolve(name) + let fxName = path.extractFilename.replace(".jsonl", "") + for v in variants: + let nIter = if v == vLinear: 1 else: nSeeds + for seed in 1..nIter: + var drv: GunDriver + var gun: ref TmPatternGun + var obsCount: proc(): GunStats = emptyStats + case v + of vLinear: + drv = makeDriver("Linear", LinearGun()) + of vTsetlin: + let pair = makeTsetlinDriver(seed = seed) + drv = pair.driver + of vTmpat, vTmpatShuf: + let pair = makeTmpatDriver(seed = seed, shuffle = (v == vTmpatShuf)) + drv = pair.driver + gun = pair.gun + obsCount = proc(): GunStats = tmpatStats(gun) + let r = replayFixture(fx, path, drv, metric, emptyStats(), obsCount, maxRounds) + rows.add Row(variant: variantName(v), fixture: fxName, seed: seed, r: r) + echo &"{variantName(v)},{fxName},{seed},{r.h100},{r.n100},{r.h300},{r.n300}," & + &"{r.hall},{r.nall},{r.f100},{r.m100},{r.rounds},{r.st.obs},{r.st.labelMiss},{r.st.traceMiss}" + + # ── per-variant pooled summary ── + echo "\n# ── pooled summary ──" + echo "variant,runs,h100,n100,early%,h300,n300,early300%,hall,nall,overall%,obs,labelMiss,traceMiss" + var pooled = initTable[string, Adapt]() + for v in variants: pooled[variantName(v)] = Adapt() + for row in rows: addAdapt(pooled[row.variant], row.r) + for v in variants: + let a = pooled[variantName(v)] + echo &"{variantName(v)},{a.rounds},{a.h100},{a.n100},{rateStr(a.h100, a.n100)}," & + &"{a.h300},{a.n300},{rateStr(a.h300, a.n300)},{a.hall},{a.nall}," & + &"{rateStr(a.hall, a.nall)},{a.st.obs},{a.st.labelMiss},{a.st.traceMiss}" + + # ── label vs chosen class histogram (TMPattern only) ── + for v in [vTmpat, vTmpatShuf]: + if v in variants: + let a = pooled[variantName(v)] + var ls, cs: string + for c in 0.. 0: row.r.h100.float / row.r.n100.float else: 0.0 + let overall = if row.r.nall > 0: row.r.hall.float / row.r.nall.float else: 0.0 + byVariant[row.variant][(row.fixture, row.seed)] = early + byVariantAll[row.variant][(row.fixture, row.seed)] = overall + if vLinear in variants: + for row in rows: + if row.variant == "Linear": + for s in 2..nSeeds: + byVariant["Linear"][(row.fixture, s)] = byVariant["Linear"][(row.fixture, 1)] + byVariantAll["Linear"][(row.fixture, s)] = byVariantAll["Linear"][(row.fixture, 1)] + + echo "\n# ── per-run early-rate distribution (mean / min / max, n runs) ──" + echo "variant,earlyMean%,earlyMin%,earlyMax%,overallMean%,overallMin%,overallMax%,n" + for v in variants: + var es: seq[float] + var os: seq[float] + for r in byVariant[variantName(v)].values: es.add r + for r in byVariantAll[variantName(v)].values: os.add r + if es.len == 0: continue + es.sort(); os.sort() + echo &"{variantName(v)},{es.sum/float(es.len)*100:.2f},{es[0]*100:.2f},{es[^1]*100:.2f}," & + &"{os.sum/float(os.len)*100:.2f},{os[0]*100:.2f},{os[^1]*100:.2f},{es.len}" + + echo "\n# ── pairwise paired sign tests (rows = fixture×seed) ──" + echo "A,B,metric,nA>B,nB>A,ties,p" + for i in 0.. vb: inc winsA + elif vb > va: inc winsB + else: inc ties + if n == 0: continue + echo &"{aName},{bName},{label},{n},{winsA},{winsB},{ties},{signTestP(winsA, n - ties):.4f}" + +when isMainModule: + main() diff --git a/common_libs/tests/tm_pattern_sweep_results.md b/common_libs/tests/tm_pattern_sweep_results.md new file mode 100644 index 0000000..a8484ea --- /dev/null +++ b/common_libs/tests/tm_pattern_sweep_results.md @@ -0,0 +1,181 @@ +# TM pattern gun — discrete-target sweep results + +Date: 2026-09-21. Author: background worker (executor-heavy). +Artifacts implementing this: `common_libs/guns/tm_pattern.nim`, +`common_libs/tests/sweep_tm_pattern.nim`. Do not commit. + +## What was built + +`tm_pattern.nim` is a NEW gun (the old `guns/tsetlin.nim` is untouched). It +attacks both the REPRESENTATION and the TARGET as the brief asked: + +* **Base**: `forecastLinear` (the exact self-consistent forecast `LinearGun` + uses). GF class 0 (centre) reproduces the Linear gun byte-for-byte, so any + measured difference is attributable to the TM. +* **Target**: a discrete multi-class GUESS-FACTOR BUCKET — which lateral escape + sector (in max-escape-angle units) the enemy occupied at the tick the bullet + would have reached the BASE fire distance. 5 or 9 classes. +* **Label**: read from a per-tick ring of our own recorded enemy positions at the + base arrival tick, NOT from `FeedbackEvent.actualXY`. Under the shipped + `bmPath` metric `actualXY` is the closest-approach point on the gun's OWN aim + ray, which biases the label toward the gun's own last output; the ring gives a + clean, metric-independent label. +* **Features**: 40 hand-built binary/bucketed motion features (lateral-velocity + sign over 3 ticks, turn-rate sign over 3 ticks, time since reversal, lateral + magnitude, speed/distance/flight-time bands, four per-wall proximity bits, + radial-fraction band, energy band, heading relative to LOS, approach sign). +* **TM core**: compact self-contained Granmo Table 2/3 with the corrected + feedback rules and Eq. 6 empty-clause bootstrap (same corrected core as + tsetlin.nim / tm_selector.nim, re-derived at 40-bit width). +* **Per-enemy / freshness**: a fresh net per gun instance; the net and history + reset if the target id changes. Each offline round is replayed with a fresh + instance (cold every battle, overfit within the battle). + +Config overrides used in the final run: `-d:TM_CONF_MARGIN_DEF=0.25 +-d:TM_SHRINK_DEF=0.5` (confidence gate + shrink). Defaults are 0.0 / 1.0 +(= raw argmax). Compile-time knobs: `TM_CLASSES`, `TM_NCLAUSES`, `TM_NSTATES`, +`TM_S_DEF`, `TM_MIN_OBS`, `TM_CONF_MARGIN_DEF`, `TM_SHRINK_DEF`, `TM_GF_MODE` +(hard|soft), `TM_SOFT_BETA_DEF`. + +## How to reproduce + +``` +nim c --path:common_libs -d:release \ + -d:TM_CONF_MARGIN_DEF=0.25 -d:TM_SHRINK_DEF=0.5 \ + -o:/tmp/sweep_tm_pattern common_libs/tests/sweep_tm_pattern.nim +/tmp/sweep_tm_pattern --set=real --seeds=3 --metric=path \ + --variants=linear,tsetlin,tmpat,tmpat_shuf +/tmp/sweep_tm_pattern --set=real --seeds=3 --metric=point \ + --variants=linear,tsetlin,tmpat,tmpat_shuf +``` + +Raw outputs: `/tmp/final_path_s3.txt`, `/tmp/final_point_s3.txt`, +`/tmp/final_ungated_path_s3.txt`, `/tmp/syn_*`. + +## Metric + +EARLY = resolutions in the first 100 ticks of each round (a cold TM every +round). OVERALL = whole fixture. Pooled over all rounds / fixtures / seeds. +`TMPatternShuf` = identical gun/encoding/cadence but the training label is a +uniform-random class (the mandatory shuffled-feedback control). Per-run = one +fixture × one seed (Linear is deterministic and replicated across seeds for +pairing). Significance = exact two-sided paired sign test, 18 pairs. + +## The core result — the discrete target IS learnable, but does not beat the base + +Online classification accuracy of the GF bucket (warm predictions only, +seeds=1, n ≈ 1.26 M for each arm): + +| arm | correct/total | accuracy | +|---|---|---| +| TMPattern (real labels) | 578722/1258488 | **46.0%** | +| TMPatternShuf (random labels) | 246733/1231116 | **20.0%** (chance) | + +So the Tsetlin Machine genuinely learns the discrete target (2.3× chance). The +representation mismatch was real and is fixed. The problem is that the target +is not aligned with what wins the metric. + +### Real DrussGT fixtures, bmPath (shipped), seeds=3 + +| variant | early | overall | +|---|---|---| +| Linear | 34.0% (6358/18715) | 24.3% (58297/239943) | +| Tsetlin (default) | 22.3% (12344/55463) | 20.3% (145828/719205) | +| **TMPattern (gated)** | **27.9% (15514/55535)** | **22.0% (158658/719681)** | +| TMPatternShuf | 28.7% (16049/55969) | 19.4% (139639/719790) | + +Paired sign tests (18 runs; ranges overlap, so the paired test is the test): + +* Linear > TMPattern: early 15/18 p=0.0075; overall 15/18 p=0.0075. **Significantly + worse than Linear.** +* TMPattern > TMPatternShuf: early 10/8 p=0.81 (tie); overall 17/1 p=0.0001. + **Learning is real but shows up mainly in the whole-round aggregate, not early.** +* TMPattern > Tsetlin: early 17/1 p=0.0001; overall 12/6 p=0.24. **Beats the + default TM gun early, ties overall.** + +Per-run distributions (mean [min,max], 18 runs): +Linear early 35.52 [26.03,50.55] / overall 26.86 [9.79,43.36]; +Tsetlin 25.76 [19.79,47.68] / 21.99 [9.64,31.34]; +TMPattern 31.18 [20.17,55.11] / 24.45 [11.24,37.78]; +Shuf 31.42 [20.72,53.49] / 21.07 [9.55,32.65]. + +### Raw ungated hard argmax (margin 0.0, shrink 1.0), bmPath, seeds=3 + +| variant | early | overall | +|---|---|---| +| Linear | 34.0% | 24.3% | +| Tsetlin | 22.3% | 20.3% | +| TMPattern | 21.2% (11783/55664) | 18.6% (133465/719432) | +| TMPatternShuf | 15.2% (8710/57169) | 8.3% (59903/720583) | + +TMPattern > Shuf 18/18 p<0.0001 on BOTH early and overall; TMPattern < Linear +3/15 p=0.0075 on both. The raw classifier is a clear, decisive learner and a +clear loser to the Linear base: applying an argmax GF bucket costs ~13 pp early. + +### Real DrussGT fixtures, bmPoint, seeds=3 (gated) + +| variant | early | overall | +|---|---|---| +| Linear | 7.2% (1480/20498) | 4.7% (11277/241423) | +| Tsetlin | 7.0% (4403/62617) | 4.8% (34588/724717) | +| TMPattern | 7.2% (4441/61842) | 4.6% (33341/724556) | +| TMPatternShuf | 6.6% (4112/61979) | 3.4% (24847/724655) | + +Online accuracy 50.8%. TMPattern is statistically indistinguishable from Linear +here (early per-run mean 13.67 vs 13.61; overall 5.70 vs 5.77) and beats its +control on overall — i.e. on the arrival-time metric the correction is neutral, +not harmful. + +### Synthetic fixtures (known rules) — the mechanism works when motion is predictable + +bmPath, seeds=1, soft readout K=9: Linear early 76.3% / overall 71.4%; +TMPattern early 76.6% / overall 71.8%; Shuf early 76.6% / overall 67.8%. +Per-fixture gains vs Linear: wall-bounce 567 vs 537, energy-threshold-turner 332 +vs 319; loss: constant-velocity 417 vs 431. + +bmPoint, seeds=1, gated hard K=5: Linear 66.4% / 59.6%; TMPattern 66.8% / 60.6%; +Shuf 55.7% / 50.1%. Energy-threshold-turner 268 vs 212, wall-bounce 585 vs 573. + +## Verdict + +* **Learning**: YES, decisively. The discrete-target TM predicts the GF bucket + far above chance (46% vs 20%) and beats its shuffled control (ungated 18/18, + p<0.0001). The "regression is a TM mismatch" diagnosis was correct. +* **Beats Linear**: NO on the real surfers under bmPath (significantly worse, + p=0.0075). Neutral under bmPoint. Matches/slightly beats Linear only on + synthetic motion whose future is genuinely predictable. +* **Best configuration found**: gated hard K=5, `TM_CONF_MARGIN=0.25`, + `TM_SHRINK=0.5` → 27.9% early / 22.0% overall (bmPath, real), +5.6 pp early / + +1.7 pp overall vs the default TM gun, but 6.1 pp early / 2.3 pp overall + behind Linear. + +## MEASURED vs INFERRED + +MEASURED: every number in the tables above (pooled hits/shots, per-run +distributions, paired sign tests, online classification accuracies). The +position-ring label is our own recorded history at the base arrival tick; the +shuffled control replaces only the label class with a uniform random draw. + +INFERRED: that the residual loss on real surfers is because the linear lead is +already the modal GF (label histogram is centred: real labels +[3.8,6.6,17.6,6.6,3.5]×10⁵ for 5 classes) and the enemy's per-tick lateral +reversal sign is not predictable enough from the 40 context bits to make a +corrective excursion net-positive. Not directly measured. + +## What to try next (not done, time-boxed out) + +1. **Radial target instead of angular.** `forecastRadialBlend` work showed the + dominant surfer error is range-holding (radial), not angle. A TM classifier + over a RADIAL displacement bucket applied as an aim-distance correction + targets the error the base actually has room to fix, and should matter most + under bmPoint. +2. **Binary reversal with a two-candidate aim** (brief candidate #1, unimplemented): + predict "will the enemy reverse lateral direction before arrival?" and choose + between the linear lead and a reversed lead. Same GF family, but a 2-class + target is far more data-efficient; expected neutral given the GF result. +3. **Condition a genuinely weaker base.** The measured wall says the deficit is + the baseline; the Linear base leaves the TM no headroom. Feeding the TM the + residual of `forecastRadialBlend` (a base that is worse on straight-liners but + range-correct on surfers) is where a learned correction could plausibly pay. +4. **Richer context.** 46% accuracy leaves room; the current context lacks the + enemy's own recent GF history / segmentation that KNN/DecayGF exploit.