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