From 0ee5c21ef5b8b234fdff69ae4c6917abf463d4f2 Mon Sep 17 00:00:00 2001 From: Davide Cappellini Date: Mon, 14 Sep 2026 22:08:22 +0200 Subject: [PATCH] fix(SNNBot): store SNN state in aim snapshots for correct retroactive learning MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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 --- SNNBot_garage/src/SNNBot.nim | 34 ++++++++++++++++++++++++++-------- 1 file changed, 26 insertions(+), 8 deletions(-) diff --git a/SNNBot_garage/src/SNNBot.nim b/SNNBot_garage/src/SNNBot.nim index 74ca77f..085e58c 100644 --- a/SNNBot_garage/src/SNNBot.nim +++ b/SNNBot_garage/src/SNNBot.nim @@ -142,10 +142,11 @@ proc forward(snn: var SNN, inputs: array[N_IN, float], proc superSpikeUpdate(snn: var SNN, spikes: array[N_HID, float], vSnap: array[N_HID, float], + preTrace: array[N_IN, float], targetAngle: float) = ## SuperSpike three-factor weight update. ## Δ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). ## Hidden error: projected via fixed random feedback weights B. @@ -183,7 +184,7 @@ proc superSpikeUpdate(snn: var SNN, let sg = surrogateDerivative(vSnap[h]) let errHid = snn.bFb[h * 2 + 0] * errSin + snn.bFb[h * 2 + 1] * errCos 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) # ── Bot state machine ───────────────────────────────────────────────────────── @@ -198,6 +199,9 @@ type botX: float # bot position at snapshot time botY: float 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 snn: SNN @@ -427,14 +431,18 @@ method run*(bot: SNNBot) = bot.snn.lastSinOut = totalSin; bot.snn.lastCosOut = totalCos bot.snn.lastSnnAngle = arctan2(totalSin, totalCos) * 180.0 / PI 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( aimAngle: bot.snn.lastSnnAngle, distance: bot.enemyDist, tick: bot.tick, botX: myX, botY: myY, - gunHeading: gunDir)) + gunHeading: gunDir, + spikes: bot.lastSpikes, + vSnap: bot.lastVSnap, + preTrace: bot.snn.preTrace)) # Log total spike count across inference window var spikeCount = 0 var maxV = 0.0 @@ -458,6 +466,10 @@ method run*(bot: SNNBot) = var retroTarget = bot.lastRelBearing # fallback (unused if no matured snapshot) var hasMatured = false 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: let snap = bot.snapshots[i] let travelTicks = int(ceil(snap.distance / BULLET_SPEED)) @@ -474,16 +486,22 @@ method run*(bot: SNNBot) = let ex = bot.posBuf[foundSlot].x let ey = bot.posBuf[foundSlot].y let absBearing = directionTo(snap.botX, snap.botY, ex, ey) - retroTarget = normalizeRelativeAngle(absBearing - snap.gunHeading) - hasMatured = true + retroTarget = normalizeRelativeAngle(absBearing - snap.gunHeading) + matureSpikes = snap.spikes + matureVSnap = snap.vSnap + maturePreTrace = snap.preTrace + hasMatured = true keepIdx = i + 1 # discard matured snapshots up to and including this one else: 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: bot.snapshots = bot.snapshots[keepIdx .. ^1] 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) # Compute verbose logging metrics var spikeCount = 0