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:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user