## 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 ## Deferred-label queue (Task 2): a virtual bullet whose radial correction ## aimed SHORT resolves BEFORE its BASE arrival tick, when the arrival-tick ## position is not yet in `posRing`. Instead of dropping the sample ## (`labelMisses`), the trace is copied here and resolved on the first later ## `predict` tick at which the base arrival tick's position exists, so every ## fired virtual bullet contributes an unbiased training sample. TM_PENDING_SLOTS = 1024 DebugTMPattern* = false ## ── radial head (Task 2) ──────────────────────────────────────────────── ## Radial label = (enemy radius at the BASE arrival tick) - (base fire ## distance), bucketed over +/-TM_RADIAL_RANGE px. Readout advances/retards ## the aim distance along the base bearing. TM_RADIAL_RANGE_DEF {.strdefine.} = "60.0" TM_RADIAL_RANGE* = parseFloat(TM_RADIAL_RANGE_DEF) TM_RAD_MARGIN_DEF {.strdefine.} = "0.25" TM_RAD_MARGIN* = parseFloat(TM_RAD_MARGIN_DEF) ## ── reversal head (Task 3) ────────────────────────────────────────────── ## Binary: did the enemy's heading turn direction over the flight oppose the ## direction it was turning at fire time? TM_REV_TURN_DEG_DEF {.strdefine.} = "10.0" TM_REV_TURN_DEG* = parseFloat(TM_REV_TURN_DEG_DEF) TM_REV_MARGIN_DEF {.strdefine.} = "0.0" TM_REV_MARGIN* = parseFloat(TM_REV_MARGIN_DEF) TM_REV_GAIN_DEF {.strdefine.} = "1.0" TM_REV_GAIN* = parseFloat(TM_REV_GAIN_DEF) type TmTargetMode* = enum tmGF ## round-1: lateral guess-factor bucket tmRadial ## Task 2: radial displacement bucket (aim-distance correction) tmReversal ## Task 3: binary turn reversal; flips the GF correction sign TmBits* = array[TM_NLITS, uint8] TmPatternTrace = object fireTick: int powerBin: int arrivalTick: int baseBearing: float fireX, fireY: float fireHeading: float fireTurn: int fireDist: float lits: TmBits votes: array[TM_CLASSES, float] cache: array[TM_CLASSES, array[TM_NCLAUSES, uint8]] chosen: int radVotes: array[TM_CLASSES, float] radCache: array[TM_CLASSES, array[TM_NCLAUSES, uint8]] radChosen: int revVotes: array[2, float] revCache: array[2, array[TM_NCLAUSES, uint8]] revChosen: int warm: bool alive: bool PosSample = object tick: int x, y: float heading: float valid: bool PendingResolve = object ## A fired virtual bullet whose label was not yet resolvable at resolution ## time. `trace` is a COPY of the fire-time trace (features + clause ## caches), so the deferred training update is identical to an immediate ## one, just later. arrivalTick: int powerBin: int power: float trace: TmPatternTrace TmPatternGun* = object teams: array[TM_CLASSES, seq[int16]] radTeams: array[TM_CLASSES, seq[int16]] revTeams: array[2, seq[int16]] targetMode*: TmTargetMode traces: array[TM_TRACE_SLOTS, TmPatternTrace] # ── deferred labels (Task 2) ── pending: array[TM_PENDING_SLOTS, PendingResolve] pendingCount: int pendingDropped*: int # ── 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] radChosenHist*: array[TM_CLASSES, int] radLabelHist*: array[TM_CLASSES, int] revChosenHist*: array[2, int] revLabelHist*: array[2, int] radCorrect*, radTotal*: int revCorrect*, revTotal*: 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 forceBase*: bool ## measurement: ignore the TM, emit the pure LinearGun base 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.. the centre. Shared by the GF, radial and reversal heads. if nObs < 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 if (votes[best] - votes[centre]) / TM_T < margin: return centre result = best proc tmChooseClass(g: var TmPatternGun, votes: array[TM_CLASSES, float]): int = tmChooseAt(votes, (TM_CLASSES - 1) div 2, TM_CONF_MARGIN, g.totalObs) proc tmResolveTrace(g: var TmPatternGun, t: TmPatternTrace, power: float) = ## One label + one TM update for a fired virtual bullet, using the enemy ## position recorded at the BASE arrival tick. `t` is a value copy of the ## fire-time trace, so this is safe to call either from `onResult` (the label ## is already resolvable) or from `tmFlushPending` (the label was deferred ## because the bullet resolved before its base arrival tick). let s = ((t.arrivalTick mod POS_RING) + POS_RING) mod POS_RING let speed = bulletSpeed(power) 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 # The shuffled control randomises ONLY the head the current mode is claiming. let shuffleGF = g.shuffleLabels and g.targetMode == tmGF let shuffleRad = g.shuffleLabels and g.targetMode == tmRadial let shuffleRev = g.shuffleLabels and g.targetMode == tmReversal let winner = if shuffleGF: rand(TM_CLASSES - 1) else: gfToBucket(gf) inc g.labelHist[winner] if t.warm: inc g.classTotal if winner == t.chosen: inc g.classCorrect # Radial label: enemy radius at the base arrival tick minus the base fire # distance. Independent of our own aim, so it is a clean target. let actualRadius = hypot(g.posRing[s].x - t.fireX, g.posRing[s].y - t.fireY) let radDelta = actualRadius - t.fireDist let radWinner = if shuffleRad: rand(TM_CLASSES - 1) else: radToBucket(radDelta) inc g.radLabelHist[radWinner] if t.warm: inc g.radTotal if radWinner == t.radChosen: inc g.radCorrect # Reversal label: net heading turn over the flight, opposite to the direction # the enemy was turning at fire time. let dh = normDeg(g.posRing[s].heading - t.fireHeading) let netTurn = if dh > TM_REV_TURN_DEG: 1 elif dh < -TM_REV_TURN_DEG: -1 else: 0 let revWinner = if shuffleRev: rand(1) elif t.fireTurn != 0 and netTurn != 0 and netTurn != t.fireTurn: 1 else: 0 inc g.revLabelHist[revWinner] if t.warm: inc g.revTotal if revWinner == t.revChosen: inc g.revCorrect case g.targetMode of tmGF: for c in 0.. 0: let tid = state.enemies[0].id if g.currentTarget != tid: if g.currentTarget >= 0: g.resetLearning() g.currentTarget = tid g.tmUpdateHistory(state) # Deferred-label flush (Task 2): resolve any fired bullet whose BASE arrival # tick is now in the ring, before the cold-start gate reads totalObs. g.tmFlushPending() 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, fireHeading: state.enemyHeading, fireTurn: g.turnSignHist[0], fireDist: f.dist, lits: lits, votes: votes, cache: caches, chosen: chosen, radVotes: radVotes, radCache: radCaches, radChosen: radChosen, revVotes: revVotes, revCache: revCaches, revChosen: revChosen, warm: (g.totalObs >= TM_MIN_OBS), alive: true) when DebugTMPattern: echo "tmp tick=", state.tick, " bin=", binIdx, " chosen=", chosen, " radChosen=", radChosen, " revChosen=", revChosen, " gf=", gf, " roff=", radOffset, " obs=", g.totalObs, " votes=", votes GunPrediction( x: if usedBase: clamp(px, 0.0, state.arenaWidth) else: clamp(px, BotRadius, state.arenaWidth - BotRadius), y: if usedBase: clamp(py, 0.0, state.arenaHeight) else: 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 g.posRing[s].valid and g.posRing[s].tick == t.arrivalTick: g.tmResolveTrace(t[], e.bulletPower) else: # DEFER (Task 2): the bullet resolved BEFORE its BASE arrival tick, which # happens whenever the radial correction aimed SHORT. The arrival-tick # position is not recorded yet, so keep a COPY of the trace and train on it # once that tick is in the ring (`tmFlushPending`). Dropping it here is what # biased the training set toward only the resolvable (long/centre) aims. if g.pendingCount < TM_PENDING_SLOTS: g.pending[g.pendingCount] = PendingResolve( arrivalTick: t.arrivalTick, powerBin: binIdx, power: e.bulletPower, trace: t[]) inc g.pendingCount else: inc g.pendingDropped inc g.labelMisses t.alive = false