Files
SirRoboGarage/common_libs/guns/tm_pattern.nim
T
SirStone 9cd6e9b8ce Ablation: the radial TM is replaceable by a CONSTANT, and its avenue is dead on the
shipped metric

The radial TM beats Linear on bmPoint, but its head never beat the majority
baseline after the label bias was fixed - suggesting the win is a constant lean
rather than learning. So: sweep a stateless constant short-range offset (new
`common_libs/guns/radial_offset.nim`, no learning at all) against the learned TM.

VERDICT (measured, offline range, seeds=3, 18 paired runs, 231 TM rounds):
1. **bmPoint - REPLACE the TM with a constant.** `RO_s0.95` (aim distance x0.95)
   TIES it early (9/9, p=1.0) and BEATS it overall (15/3, p=0.0075; 7.47% vs
   6.89% per-run mean). A fixed -20px does the same. The head never beats its
   majority baseline (56.2% vs 57.2%).
2. **bmPath (the SHIPPED metric) - the radial avenue is a DEAD END.** TMRadial is
   a systematic LOSS there (2/16, p=0.0013); every constant is within +-0.2pp; the
   only real bmPath effect is the BotRadius clamp. So the radial shift cannot help
   the shipped configuration.
3. The per-adversary optimum DOES vary (fixed -10 for crazy, -30 for tr_crazy,
   scale 0.95 for three others) - but ONE GLOBAL CONSTANT still beats the
   adaptively-trained head, so the "fragility justifies learning" argument FAILS.

THE REAL FINDING UNDERNEATH, and it generalises beyond this gun: the base linear
prediction systematically OVERSHOOTS. Measured raw per-tick base radial error has
mean -71 to -100 px; the enemy is NEARER than the prediction in 63-81% of shots
and farther in only 4-14%, CONSISTENT ACROSS ALL SIX CAPTURES. Radial label
histogram [415166,126461,120325,44694,19351] = 57.2% majority class, mean label
-82.3 px, mean applied shift -37.6 px. So the net-short bias is a GENUINE property
of these range-holders against a constant-velocity extrapolation (they decelerate
and turn, so the true position is closer than the straight-line guess) - NOT a
fixture artefact. That is worth chasing for the guns that actually ship.

Caveat: bmPoint is not the shipped metric (bmPath won the real-hit-rate A/B for
SELECTION), so a bmPoint win is not yet evidence of a real win. That needs a live
test - and the natural target is Pattern, which is now the default and best gun.

Adds radial_offset.nim + sweep_radial_offset.nim; tm_pattern.nim gains additive
instrumentation only (radial label mean and applied-shift mean; no behaviour
change, and test_tm_pattern_registration still passes all 20 checks).
2026-09-22 02:08:38 +02:00

709 lines
28 KiB
Nim

## 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
## Radial-target instrumentation (ablation support): raw label statistics
## and the mean APPLIED aim-distance shift, so a constant-offset control can
## be compared against the actual average readout the TM produces.
radDeltaSum*, radDeltaAbsSum*: float
radDeltaN*: int
radOffsetSum*: float
radOffsetN*: 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..<TM_NLITS:
if team[base + lit] > 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..<TM_NCLAUSES:
let o = tmEval(team, lits, cl, learning = false)
cache[cl] = tmEval(team, lits, cl, learning = true)
v += tmPolarity(cl) * float(o)
clamp(v, -TM_T, TM_T)
proc tmLearnDir(team: var seq[int16], lits: TmBits,
cache: array[TM_NCLAUSES, uint8], vote, d: float) =
## One Granmo update of one class team with desired vote direction `d`.
let pFeedback = (TM_T - d * vote) / (2.0 * TM_T)
if pFeedback <= 0.0: return
for cl in 0..<TM_NCLAUSES:
if rand(1.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..<TM_NLITS:
var st = int(team[base + lit])
if lits[lit] == 1'u8:
if cOut == 1'u8:
if rand(1.0) < (TM_S - 1.0) / TM_S: st = min(st + 1, TM_NSTATES)
else:
if rand(1.0) < 1.0 / TM_S: st = max(st - 1, -TM_NSTATES)
else:
if rand(1.0) < 1.0 / TM_S: st = max(st - 1, -TM_NSTATES)
team[base + lit] = int16(st)
else:
# Type II (Table 3): penalise exclusion of a zero literal when firing.
if cOut == 1'u8:
for lit in 0..<TM_NLITS:
if lits[lit] == 0'u8:
if team[base + lit] <= 0:
team[base + lit] = int16(min(int(team[base + lit]) + 1, TM_NSTATES))
# ── geometry helpers ─────────────────────────────────────────────────────────
proc tmBinForSpeed(spd: float): int {.inline.} =
for i in 0..<len(vb.PowerBins):
if abs(spd - bulletSpeed(vb.PowerBins[i])) < 1e-6: return i
-1
proc tmTraceSlot(fireTick, binIdx: int): int {.inline.} =
((fireTick * len(vb.PowerBins)) + binIdx) mod TM_TRACE_SLOTS
proc gfToBucket(gf: float): int {.inline.} =
clamp(int(round((gf + 1.0) * 0.5 * float(TM_CLASSES - 1))), 0, TM_CLASSES - 1)
proc bucketToGF(c: int): float {.inline.} =
if TM_CLASSES <= 1: 0.0
else: float(c) / float(TM_CLASSES - 1) * 2.0 - 1.0
proc radToBucket(delta: float): int {.inline.} =
let u = clamp(delta / TM_RADIAL_RANGE, -1.0, 1.0)
clamp(int(round((u + 1.0) * 0.5 * float(TM_CLASSES - 1))), 0, TM_CLASSES - 1)
proc bucketToRadial(c: int): float {.inline.} =
if TM_CLASSES <= 1: 0.0
else: (float(c) / float(TM_CLASSES - 1) * 2.0 - 1.0) * TM_RADIAL_RANGE
proc normDeg(d: float): float {.inline.} =
result = d
while result > 180.0: result -= 360.0
while result < -180.0: result += 360.0
# ── public API ───────────────────────────────────────────────────────────────
proc initTmPatternGun*(): TmPatternGun =
for c in 0..<TM_CLASSES: result.teams[c] = tmNewTeam()
for c in 0..<TM_CLASSES: result.radTeams[c] = tmNewTeam()
for c in 0..<2: result.revTeams[c] = tmNewTeam()
result.targetMode = tmGF
result.lastTick = -1
result.prevTick = -1
result.currentTarget = -1
result.lastChosen = (TM_CLASSES - 1) div 2
randomize()
result.debugGraphics = false
proc initTmRadialGun*(): TmPatternGun =
## The RACK-REGISTERED instance: the RADIAL target mode, which is the
## control-validated winner under `bmPoint` (see tm_pattern_sweep_results.md,
## Round 2 Task 2). The gun type carries all three heads; the live rack only
## ever selects this radial-mode instance.
result = initTmPatternGun()
result.targetMode = tmRadial
proc isWarmedUp*(g: TmPatternGun): bool {.inline.} = true
proc resetLearning*(g: var TmPatternGun) =
## Fresh concept: wipe every clause team and the motion history. Called when
## the target id changes so a new opponent starts from a cold net.
for c in 0..<TM_CLASSES: g.teams[c] = tmNewTeam()
for c in 0..<TM_CLASSES: g.radTeams[c] = tmNewTeam()
for c in 0..<2: g.revTeams[c] = tmNewTeam()
g.totalObs = 0
g.hasPrev = false
for i in 0..<3:
g.latSignHist[i] = 0
g.turnSignHist[i] = 0
g.sinceReversal = 0
g.radialFracSm = 0.0
g.latPersist = 0
g.pendingCount = 0
proc tmUpdateHistory(g: var TmPatternGun, state: WorldState) =
if state.tick == g.lastTick: return
g.lastTick = state.tick
let slot = ((state.tick mod POS_RING) + POS_RING) mod POS_RING
g.posRing[slot] = PosSample(tick: state.tick, x: state.enemyX, y: state.enemyY,
heading: state.enemyHeading, valid: true)
if g.hasPrev and state.tick > 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..<TM_NBITS:
result[i] = bits[i]
result[i + TM_NBITS] = 1'u8 - bits[i]
proc tmSoftGF(votes: array[TM_CLASSES, float]): float =
## Vote-weighted expectation of the GF bucket centres. The multiclass votes are
## not calibrated probabilities, but a softmax over them gives a smooth,
## self-shrinking readout (flat votes -> 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..<TM_CLASSES:
if votes[c] > mx: mx = votes[c]
var sum = 0.0
var w: array[TM_CLASSES, float]
for c in 0..<TM_CLASSES:
w[c] = exp((votes[c] - mx) * TM_SOFT_BETA / TM_T)
sum += w[c]
if sum <= 0.0: return 0.0
for c in 0..<TM_CLASSES:
result += (w[c] / sum) * bucketToGF(c)
proc tmChooseAt(votes: openArray[float], centre: int, margin: float,
nObs: int): int =
## Cold gun or no class beating the centre by `margin` (fraction of TM_T)
## -> 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..<votes.len:
if votes[c] > 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
inc g.radDeltaN
g.radDeltaSum += radDelta
g.radDeltaAbsSum += abs(radDelta)
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..<TM_CLASSES:
let d = if c == winner: 1.0 else: -1.0
g.teams[c].tmLearnDir(t.lits, t.cache[c], t.votes[c], d)
of tmRadial:
for c in 0..<TM_CLASSES:
let d = if c == radWinner: 1.0 else: -1.0
g.radTeams[c].tmLearnDir(t.lits, t.radCache[c], t.radVotes[c], d)
of tmReversal:
for c in 0..<TM_CLASSES:
let d = if c == winner: 1.0 else: -1.0
g.teams[c].tmLearnDir(t.lits, t.cache[c], t.votes[c], d)
for c in 0..<2:
let d = if c == revWinner: 1.0 else: -1.0
g.revTeams[c].tmLearnDir(t.lits, t.revCache[c], t.revVotes[c], d)
inc g.totalObs
inc g.trainCalls
proc tmFlushPending(g: var TmPatternGun) =
## Resolve every deferred trace whose BASE arrival tick is now recorded in
## `posRing`. Called once per `predict` right after `tmUpdateHistory`, so the
## just-written current tick is visible. Entries are compacted in place.
if g.pendingCount == 0: return
var w = 0
for i in 0..<g.pendingCount:
let p = addr g.pending[i]
let s = ((p.arrivalTick mod POS_RING) + POS_RING) mod POS_RING
if g.posRing[s].valid and g.posRing[s].tick == p.arrivalTick:
g.tmResolveTrace(p.trace, p.power)
else:
if w != i: g.pending[w] = g.pending[i]
inc w
g.pendingCount = w
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)
# 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..<TM_CLASSES:
votes[c] = tmForward(g.teams[c], lits, caches[c])
var radVotes: array[TM_CLASSES, float]
var radCaches: array[TM_CLASSES, array[TM_NCLAUSES, uint8]]
for c in 0..<TM_CLASSES:
radVotes[c] = tmForward(g.radTeams[c], lits, radCaches[c])
var revVotes: array[2, float]
var revCaches: array[2, array[TM_NCLAUSES, uint8]]
for c in 0..<2:
revVotes[c] = tmForward(g.revTeams[c], lits, revCaches[c])
let chosen = g.tmChooseClass(votes)
let radChosen = tmChooseAt(radVotes, (TM_CLASSES - 1) div 2, TM_RAD_MARGIN,
g.totalObs)
let revChosen = tmChooseAt(revVotes, 0, TM_REV_MARGIN, g.totalObs)
g.lastChosen = chosen
inc g.chosenHist[chosen]
inc g.radChosenHist[radChosen]
inc g.revChosenHist[revChosen]
var gf = 0.0
var radOffset = 0.0
if not (g.forceBase or g.totalObs < TM_MIN_OBS):
case g.targetMode
of tmGF:
if TM_GF_MODE == "soft": gf = TM_SHRINK * tmSoftGF(votes)
else: gf = TM_SHRINK * bucketToGF(chosen)
of tmRadial:
radOffset = bucketToRadial(radChosen)
g.radOffsetSum += radOffset
inc g.radOffsetN
of tmReversal:
if TM_GF_MODE == "soft": gf = TM_SHRINK * tmSoftGF(votes)
else: gf = TM_SHRINK * bucketToGF(chosen)
if revChosen == 1:
gf = -TM_REV_GAIN * gf
# An all-zero correction must reproduce `LinearGun` BYTE-FOR-BYTE, so use its
# exact aim point and its exact [0, arena] clamp rather than the
# BotRadius-inset clamp the corrective excursions use.
let usedBase = (gf == 0.0 and radOffset == 0.0)
var px, py: float
if usedBase:
px = f.x
py = f.y
else:
let aimAngle = f.bearing + gf * mea
let aimDist = f.dist + radOffset
px = state.selfX + cos(aimAngle) * aimDist
py = state.selfY + sin(aimAngle) * aimDist
let binIdx = tmBinForSpeed(bulletSpeed)
if binIdx >= 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