## 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..