fix(SNNBot): store SNN state in aim snapshots for correct retroactive learning

The retroactive error signal matures many ticks after the snapshot was taken,
but superSpikeUpdate was using bot.lastSpikes/lastVSnap which had been
overwritten by subsequent DECIDE cycles. Learning was applied against wrong
neural activity — effectively random weight perturbations.

Fix: AimSnapshot now captures spikes, vSnap, and preTrace at DECIDE time.
superSpikeUpdate receives these saved values at maturation instead of the
stale current state.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-09-14 22:08:22 +02:00
parent 5740c56a9f
commit 0ee5c21ef5
+24 -6
View File
@@ -142,10 +142,11 @@ proc forward(snn: var SNN, inputs: array[N_IN, float],
proc superSpikeUpdate(snn: var SNN, proc superSpikeUpdate(snn: var SNN,
spikes: array[N_HID, float], spikes: array[N_HID, float],
vSnap: array[N_HID, float], vSnap: array[N_HID, float],
preTrace: array[N_IN, float],
targetAngle: float) = targetAngle: float) =
## SuperSpike three-factor weight update. ## SuperSpike three-factor weight update.
## Δw = η × pre_trace × σ'(U) × error ## Δw = η × pre_trace × σ'(U) × error
## spikes: accumulated counts over N_INFER ticks (0..N_INFER); normalized to rates. ## spikes/vSnap/preTrace: captured at DECIDE time for this inference window.
## Output error: target_rate − actual_rate (rate-coded target). ## Output error: target_rate − actual_rate (rate-coded target).
## Hidden error: projected via fixed random feedback weights B. ## Hidden error: projected via fixed random feedback weights B.
@@ -183,7 +184,7 @@ proc superSpikeUpdate(snn: var SNN,
let sg = surrogateDerivative(vSnap[h]) let sg = surrogateDerivative(vSnap[h])
let errHid = snn.bFb[h * 2 + 0] * errSin + snn.bFb[h * 2 + 1] * errCos let errHid = snn.bFb[h * 2 + 0] * errSin + snn.bFb[h * 2 + 1] * errCos
for i in 0 ..< N_IN: for i in 0 ..< N_IN:
snn.wih[i * N_HID + h] += ETA * snn.preTrace[i] * sg * errHid snn.wih[i * N_HID + h] += ETA * preTrace[i] * sg * errHid
snn.wih[i * N_HID + h] = snn.wih[i * N_HID + h].clamp(-W_CLAMP, W_CLAMP) snn.wih[i * N_HID + h] = snn.wih[i * N_HID + h].clamp(-W_CLAMP, W_CLAMP)
# ── Bot state machine ───────────────────────────────────────────────────────── # ── Bot state machine ─────────────────────────────────────────────────────────
@@ -198,6 +199,9 @@ type
botX: float # bot position at snapshot time botX: float # bot position at snapshot time
botY: float botY: float
gunHeading: float # absolute gun heading at snapshot time gunHeading: float # absolute gun heading at snapshot time
spikes: array[N_HID, float] # accumulated spike counts from this inference window
vSnap: array[N_HID, float] # hidden voltages from this inference window
preTrace: array[N_IN, float] # pre-synaptic traces from this inference window
SNNBot = ref object of Bot SNNBot = ref object of Bot
snn: SNN snn: SNN
@@ -427,14 +431,18 @@ method run*(bot: SNNBot) =
bot.snn.lastSinOut = totalSin; bot.snn.lastCosOut = totalCos bot.snn.lastSinOut = totalSin; bot.snn.lastCosOut = totalCos
bot.snn.lastSnnAngle = arctan2(totalSin, totalCos) * 180.0 / PI bot.snn.lastSnnAngle = arctan2(totalSin, totalCos) * 180.0 / PI
bot.targetAngle = gunDir + bot.snn.lastSnnAngle bot.targetAngle = gunDir + bot.snn.lastSnnAngle
# Record aim snapshot for retroactive error signal # Record aim snapshot for retroactive error signal; capture SNN state now
# before subsequent DECIDE cycles overwrite bot.lastSpikes / bot.lastVSnap.
bot.snapshots.add(AimSnapshot( bot.snapshots.add(AimSnapshot(
aimAngle: bot.snn.lastSnnAngle, aimAngle: bot.snn.lastSnnAngle,
distance: bot.enemyDist, distance: bot.enemyDist,
tick: bot.tick, tick: bot.tick,
botX: myX, botX: myX,
botY: myY, botY: myY,
gunHeading: gunDir)) gunHeading: gunDir,
spikes: bot.lastSpikes,
vSnap: bot.lastVSnap,
preTrace: bot.snn.preTrace))
# Log total spike count across inference window # Log total spike count across inference window
var spikeCount = 0 var spikeCount = 0
var maxV = 0.0 var maxV = 0.0
@@ -458,6 +466,10 @@ method run*(bot: SNNBot) =
var retroTarget = bot.lastRelBearing # fallback (unused if no matured snapshot) var retroTarget = bot.lastRelBearing # fallback (unused if no matured snapshot)
var hasMatured = false var hasMatured = false
var keepIdx = 0 # first non-matured snapshot to keep var keepIdx = 0 # first non-matured snapshot to keep
# matureSnap holds the last matured snapshot's SNN state for the weight update.
var matureSpikes: array[N_HID, float]
var matureVSnap: array[N_HID, float]
var maturePreTrace: array[N_IN, float]
for i in 0 ..< bot.snapshots.len: for i in 0 ..< bot.snapshots.len:
let snap = bot.snapshots[i] let snap = bot.snapshots[i]
let travelTicks = int(ceil(snap.distance / BULLET_SPEED)) let travelTicks = int(ceil(snap.distance / BULLET_SPEED))
@@ -475,15 +487,21 @@ method run*(bot: SNNBot) =
let ey = bot.posBuf[foundSlot].y let ey = bot.posBuf[foundSlot].y
let absBearing = directionTo(snap.botX, snap.botY, ex, ey) let absBearing = directionTo(snap.botX, snap.botY, ex, ey)
retroTarget = normalizeRelativeAngle(absBearing - snap.gunHeading) retroTarget = normalizeRelativeAngle(absBearing - snap.gunHeading)
matureSpikes = snap.spikes
matureVSnap = snap.vSnap
maturePreTrace = snap.preTrace
hasMatured = true hasMatured = true
keepIdx = i + 1 # discard matured snapshots up to and including this one keepIdx = i + 1 # discard matured snapshots up to and including this one
else: else:
break # snapshots are in order; stop at first non-matured break # snapshots are in order; stop at first non-matured
# Discard matured snapshots # Discard matured snapshots; immature ones are preserved automatically (keepIdx stays 0
# or points past the last matured entry; the rest of bot.snapshots is kept intact).
if keepIdx > 0: if keepIdx > 0:
bot.snapshots = bot.snapshots[keepIdx .. ^1] bot.snapshots = bot.snapshots[keepIdx .. ^1]
if hasMatured: if hasMatured:
bot.snn.superSpikeUpdate(bot.lastSpikes, bot.lastVSnap, retroTarget) # Use SNN state captured at DECIDE time — not the stale bot.lastSpikes/lastVSnap
# which have been overwritten by subsequent DECIDE cycles.
bot.snn.superSpikeUpdate(matureSpikes, matureVSnap, maturePreTrace, retroTarget)
# else: skip weight update — no matured snapshot yet (early game) # else: skip weight update — no matured snapshot yet (early game)
# Compute verbose logging metrics # Compute verbose logging metrics
var spikeCount = 0 var spikeCount = 0