feat(SNNBot): replace R-STDP with SuperSpike learning rule (#158)

Three-factor local learning: Δw = η × pre_trace × σ'(U) × error
- Surrogate derivative σ'(U) = 1/(1+|β(U-ϑ)|)² gives directional gradient
- Random feedback weights project output error to hidden layer
- 10-tick inference window for rate-coded sin/cos output
- Fixed stale learning target bug (EVALUATE used wrong gun direction)

Hyperparams: ETA=0.05, THRESH=0.2, W_CLAMP=1.0, N_INFER=10, BETA=1.0
Removed: R-STDP eligibility traces, STDP timing window, exploration noise

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-09-13 20:07:31 +02:00
parent 7e1f1483a4
commit dbae449273
+178 -89
View File
@@ -1,13 +1,13 @@
# SNNBot — SNN core + aiming loop prototype (issue #152).
# 36-input population-coded → 6 LIF hidden → 1 membrane-readout output.
# No movement, no firing, no learning. Proves plumbing.
# 36-input population-coded → 12 LIF hidden → polar-coded (sin/cos) output decoded via atan2.
# No movement, no firing. SuperSpike three-factor rule on all weights (issue #158).
#
# State machine:
# DECIDE → feed bearing into SNN, store targetAngle, → WAITING
# WAITING → aimTo() each tick; when error < 2° → EVALUATE
# EVALUATE → measure error, compute reward, log, → DECIDE
# EVALUATE → measure error, compute SuperSpike update, log, → DECIDE
import std/[math, random, os]
import std/[math, random, os, strutils]
import robocode_tankroyale_botapi
import radar_lock/radar_lock as radar_lock
@@ -17,38 +17,46 @@ const botJsonPath = currentSourcePath().parentDir / "SNNBot.json"
const
N_IN = 36 # input neurons (10°-wide bands, -180..+180)
N_HID = 6 # hidden LIF neurons
N_HID = 12 # hidden LIF neurons
BAND_DEG = 10.0 # degrees per input band
LEAK = 0.9 # LIF membrane leak factor
THRESH = 1.0 # LIF spike threshold
THRESH = 0.2 # ponytail: THRESH=0.2 — unitless system, must match weight scale; raise if neurons fire too much
MAX_GUN_TURN = 20.0 # max gun turn per tick (degrees)
AIM_TOL = 2.0 # arrive tolerance (degrees)
STDP_WIN = 20 # STDP timing window (ticks)
ETA = 0.01 # learning rate
# ponytail: global learning rate, no per-synapse adaptation; add when performance plateaus
ELG_DECAY = 0.95 # eligibility trace decay per tick
W_CLAMP = 2.0 # weight magnitude clamp
ETA = 0.05 # SuperSpike learning rate (r_0 from paper)
# ponytail: single learning rate, add RMaxProp optimizer if convergence unstable
N_INFER = 10 # inference window ticks per DECIDE
# ponytail: N_INFER=10, increase if output still noisy; decrease if too slow per tick
TRACE_DECAY = 0.9 # pre-synaptic trace decay (exponential low-pass)
W_CLAMP = 1.0 # ponytail: W_CLAMP=1.0, paper says ±0.1 but that's for mV-scale voltages; our unitless THRESH=1.0 needs larger weights
BETA = 1.0 # surrogate steepness (unitless; potentials are unitless)
NOISE_AMP = 0.3 # exploration noise amplitude
# ponytail: uniform exploration noise; upgrade to annealed Gaussian if convergence needs tuning
# ── SNN types ─────────────────────────────────────────────────────────────────
type
SNN = object
wih: array[N_IN * N_HID, float] # 36×6 input→hidden weights
who: array[N_HID, float] # 6×1 hidden→output weights
vHid: array[N_HID, float] # hidden membrane potentials
vOut: float # output membrane potential (readout)
lastSpikeIn: array[N_IN, int] # tick of last input spike (-1 = never)
lastSpikeHid: array[N_HID, int] # tick of last hidden spike (-1 = never)
eligibility: array[N_IN * N_HID, float] # per-synapse eligibility traces
tick: int # internal tick counter for spike timing
wih: array[N_IN * N_HID, float] # 36×12 input→hidden weights
wSin: array[N_HID, float] # hidden→sin-channel weights
wCos: array[N_HID, float] # hidden→cos-channel weights
vHid: array[N_HID, float] # hidden membrane potentials
preTrace: array[N_IN, float] # pre-synaptic traces (low-pass of input spikes)
# Random feedback weights B[h, o] for hidden error projection; fixed, never updated.
# Indexed as bFb[h * 2 + o], o=0 → sin, o=1 → cos.
bFb: array[N_HID * 2, float]
tick: int
lastSinOut: float # last forward pass sin output (weighted sum)
lastCosOut: float # last forward pass cos output (weighted sum)
lastSnnAngle: float # last atan2 result (degrees)
proc initSNN(snn: var SNN) =
randomize()
for w in snn.wih.mitems: w = rand(1.0) - 0.5
for w in snn.who.mitems: w = rand(1.0) - 0.5
for t in snn.lastSpikeIn.mitems: t = -STDP_WIN - 1
for t in snn.lastSpikeHid.mitems: t = -STDP_WIN - 1
for e in snn.eligibility.mitems: e = 0.0
for w in snn.wih.mitems: w = rand(0.2) - 0.1 # init within clamp ±0.1, room to grow to ±1.0
for w in snn.wSin.mitems: w = rand(0.2) - 0.1
for w in snn.wCos.mitems: w = rand(0.2) - 0.1
for b in snn.bFb.mitems: b = rand(2.0) - 1.0 # N(0,1)-ish; fixed forever
for t in snn.preTrace.mitems: t = 0.0
snn.tick = 0
proc encodeInput(bearing: float): array[N_IN, float] =
@@ -63,51 +71,93 @@ proc encodeInput(bearing: float): array[N_IN, float] =
result[lo] = 1.0 - frac
result[hi] = frac
proc forward(snn: var SNN, inputs: array[N_IN, float]): float =
## One SNN tick. Returns target angle in degrees (-180..+180).
proc surrogateDerivative(v: float): float {.inline.} =
## σ'(U) = (1 + |β(U − ϑ)|)^{-2} — peaks at threshold, gives gradient direction.
let x = BETA * (v - THRESH)
result = 1.0 / ((1.0 + abs(x)) * (1.0 + abs(x)))
proc forward(snn: var SNN, inputs: array[N_IN, float],
spikesOut: var array[N_HID, float],
vHidSnap: var array[N_HID, float],
sinOut: var float, cosOut: var float) =
## One SNN tick. Writes hidden spikes, pre-spike voltages, and raw sin/cos outputs.
## Caller accumulates sinOut/cosOut across N_INFER ticks, then calls atan2.
inc snn.tick
# Decay eligibility traces each tick
for e in snn.eligibility.mitems: e *= ELG_DECAY
# Record input spikes
# Pre-synaptic trace: low-pass of input spikes
for i in 0 ..< N_IN:
if inputs[i] > 0.0:
snn.lastSpikeIn[i] = snn.tick
snn.preTrace[i] = TRACE_DECAY * snn.preTrace[i] + inputs[i]
# Hidden layer: LIF update + STDP trace
var spikes: array[N_HID, float]
# Hidden layer: LIF update
for h in 0 ..< N_HID:
var wsum = 0.0
for i in 0 ..< N_IN:
wsum += inputs[i] * snn.wih[i * N_HID + h]
# ponytail: noise removed — SuperSpike surrogate derivative provides gradient direction, no exploration needed
snn.vHid[h] = LEAK * snn.vHid[h] + wsum
vHidSnap[h] = snn.vHid[h] # snapshot voltage before reset (for surrogate)
if snn.vHid[h] >= THRESH:
spikes[h] = 1.0
spikesOut[h] = 1.0
snn.vHid[h] = 0.0
snn.lastSpikeHid[h] = snn.tick
# STDP: post fires → check recent pre-spikes (potentiation)
for i in 0 ..< N_IN:
let dt = snn.tick - snn.lastSpikeIn[i]
if dt >= 0 and dt <= STDP_WIN:
snn.eligibility[i * N_HID + h] += exp(-float(dt) / float(STDP_WIN))
else:
spikes[h] = 0.0
# STDP: pre fires after post → depression for synapses where post spiked recently
for i in 0 ..< N_IN:
if inputs[i] > 0.0:
let dt = snn.tick - snn.lastSpikeHid[h]
if dt >= 0 and dt <= STDP_WIN:
snn.eligibility[i * N_HID + h] -= exp(-float(dt) / float(STDP_WIN))
spikesOut[h] = 0.0
# Output layer: membrane readout (no threshold)
var osum = 0.0
# Output: polar-coded via sin/cos channels (raw; caller does atan2)
sinOut = 0.0
cosOut = 0.0
for h in 0 ..< N_HID:
osum += spikes[h] * snn.who[h]
snn.vOut = LEAK * snn.vOut + osum
sinOut += spikesOut[h] * snn.wSin[h]
cosOut += spikesOut[h] * snn.wCos[h]
snn.lastSinOut = sinOut
snn.lastCosOut = cosOut
snn.lastSnnAngle = arctan2(sinOut, cosOut) * 180.0 / PI
# Scale vOut to [-180, +180]
result = tanh(snn.vOut) * 180.0
proc superSpikeUpdate(snn: var SNN,
spikes: array[N_HID, float],
vSnap: array[N_HID, 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.
## Output error: target_rate − actual_rate (rate-coded target).
## Hidden error: projected via fixed random feedback weights B.
# Target rates for sin/cos channels: map [-1,1] → [0,1]
let tSin = (sin(degToRad(targetAngle)) + 1.0) / 2.0
let tCos = (cos(degToRad(targetAngle)) + 1.0) / 2.0
# Normalize accumulated spike counts to rates in [0,1]
# Actual output activity using spike rates
var sinAct = 0.0; var cosAct = 0.0
for h in 0 ..< N_HID:
let rate = spikes[h] / float(N_INFER)
sinAct += rate * snn.wSin[h]
cosAct += rate * snn.wCos[h]
# Map actual output to [0,1] for rate comparison
let sinActNorm = (sinAct.clamp(-1.0, 1.0) + 1.0) / 2.0
let cosActNorm = (cosAct.clamp(-1.0, 1.0) + 1.0) / 2.0
# Output error (target_rate − actual_rate)
let errSin = tSin - sinActNorm
let errCos = tCos - cosActNorm
# Update hidden→output weights: Δw = η × rate_h × σ'(U_h) × error_o
for h in 0 ..< N_HID:
let sg = surrogateDerivative(vSnap[h])
let rate = spikes[h] / float(N_INFER)
snn.wSin[h] += ETA * rate * sg * errSin
snn.wSin[h] = snn.wSin[h].clamp(-W_CLAMP, W_CLAMP)
snn.wCos[h] += ETA * rate * sg * errCos
snn.wCos[h] = snn.wCos[h].clamp(-W_CLAMP, W_CLAMP)
# Update input→hidden weights: Δw = η × preTrace_j × σ'(U_h) × error_h
# Hidden error projected via fixed random feedback: error_h = Σ_o B[h,o] × error_o
for h in 0 ..< N_HID:
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] = snn.wih[i * N_HID + h].clamp(-W_CLAMP, W_CLAMP)
# ── Bot state machine ─────────────────────────────────────────────────────────
@@ -122,6 +172,9 @@ type
enemyDist: float # last known distance to enemy
hasContact: bool
tick: int
lastSpikes: array[N_HID, float] # accumulated spike counts over N_INFER ticks
lastVSnap: array[N_HID, float] # hidden voltages before reset from last tick (for SuperSpike)
lastRelBearing: float # relative bearing at DECIDE time (fixed for learning)
# ── aimTo helper ──────────────────────────────────────────────────────────────
@@ -161,12 +214,11 @@ proc drawOverlay(bot: SNNBot, myX, myY, gunDir, enemyBearing, targetAngle: float
let
panelW = getArenaWidth().float / 2.0
panelH = getArenaHeight().float / 2.0
CH = (panelH - PY - 20.0) / 7.0 # row height (same formula as before)
CH = (panelH - PY - 20.0) / 7.0 # row height
LAYER_V = CH * 2.0 # inter-layer gap
IN_W = panelW # input layer spans full panel width
IN_STEP = IN_W / float(N_IN) # spacing between input lines
HID_STEP = IN_W / float(N_HID) # spacing between hidden lines
# Anchor x of each hidden line — centred within IN_W
HID_STEP = IN_W / float(N_HID) # spacing between hidden lines (scales with N_HID)
HID_OFF = 0.0
# Input layer: vertical line per neuron, visible when activation > 0
@@ -179,7 +231,7 @@ proc drawOverlay(bot: SNNBot, myX, myY, gunDir, enemyBearing, targetAngle: float
setStrokeColor(fromRgba(0, 200, 255, alpha))
drawLine(lx, PY, lx, PY + CH)
# Hidden layer: vertical line per neuron, visible when spiking
# Hidden layer: vertical line per neuron, visible when membrane > 0
let hidY = PY + CH + LAYER_V
setStrokeWidth(1.5)
for h in 0 ..< N_HID:
@@ -190,8 +242,7 @@ proc drawOverlay(bot: SNNBot, myX, myY, gunDir, enemyBearing, targetAngle: float
setStrokeColor(fromRgba(255, 140, 0, alpha))
drawLine(lx, hidY, lx, hidY + CH)
# Weight lines: input → hidden (sample every 4th input to avoid clutter)
# Fixed light-blue color; thickness proportional to normalized weight magnitude.
# Weight lines: input → hidden
var maxAbsWih = 0.0
for w in bot.snn.wih: maxAbsWih = max(maxAbsWih, abs(w))
setStrokeColor(fromRgba(180, 220, 255, 180))
@@ -206,47 +257,49 @@ proc drawOverlay(bot: SNNBot, myX, myY, gunDir, enemyBearing, targetAngle: float
setStrokeWidth(0.5 + norm * 2.5)
drawLine(ix, iy, hx, hy)
# Output neuron: single vertical line centred in panel, visible when active
# Output: two small vertical lines for sin/cos channels side-by-side
let outY = hidY + CH + LAYER_V
let outV = (tanh(bot.snn.vOut) + 1.0) / 2.0 # 0..1 for display
if outV > 0.0:
let lx = PX + IN_W / 2.0
let alpha = uint8(outV * 255.0)
setStrokeColor(fromRgba(200, 0, 255, alpha))
setStrokeWidth(2.0)
drawLine(lx, outY, lx, outY + CH)
var sinOut = 0.0; var cosOut = 0.0
for h in 0 ..< N_HID:
let rate = bot.lastSpikes[h] / float(N_INFER)
sinOut += rate * bot.snn.wSin[h]
cosOut += rate * bot.snn.wCos[h]
let sinV = (sinOut.clamp(-1.0, 1.0) + 1.0) / 2.0
let cosV = (cosOut.clamp(-1.0, 1.0) + 1.0) / 2.0
let cx = PX + IN_W / 2.0
setStrokeWidth(2.0)
if sinV > 0.0:
setStrokeColor(fromRgba(200, 0, 255, uint8(sinV * 255.0)))
drawLine(cx - 4.0, outY, cx - 4.0, outY + CH)
if cosV > 0.0:
setStrokeColor(fromRgba(0, 200, 100, uint8(cosV * 255.0)))
drawLine(cx + 4.0, outY, cx + 4.0, outY + CH)
# ── Aiming-line legend ────────────────────────────────────────────────────
# Arena Y=0 is bottom; top = getArenaHeight(). LEG_Y is the bottom edge of
# the legend block so it sits near the top of the screen.
const
LEG_X = 10.0 # left margin
LEG_SQ = 8.0 # coloured square side
LEG_GAP = 4.0 # gap between square and text
LEG_ROW = 14.0 # row height
LEG_PAD = 6.0 # inner padding of background rect
let LEG_Y = getArenaHeight().float - 60.0 # near top of arena
LEG_X = 10.0
LEG_SQ = 8.0
LEG_GAP = 4.0
LEG_ROW = 14.0
LEG_PAD = 6.0
let LEG_Y = getArenaHeight().float - 60.0
# Semi-transparent background
setFillColor(fromRgba(0, 0, 0, 160))
fillRectangle(LEG_X - LEG_PAD,
LEG_Y - LEG_PAD,
LEG_SQ + LEG_GAP + 80.0 + LEG_PAD,
3.0 * LEG_ROW + LEG_PAD)
# Row 0 — Green: Enemy bearing
setFillColor(GREEN)
fillRectangle(LEG_X, LEG_Y, LEG_SQ, LEG_SQ)
setFillColor(WHITE)
drawText("Enemy bearing", LEG_X + LEG_SQ + LEG_GAP, LEG_Y + LEG_SQ)
# Row 1 — Red: Gun direction
setFillColor(RED)
fillRectangle(LEG_X, LEG_Y + LEG_ROW, LEG_SQ, LEG_SQ)
setFillColor(WHITE)
drawText("Gun direction", LEG_X + LEG_SQ + LEG_GAP, LEG_Y + LEG_ROW + LEG_SQ)
# Row 2 — Yellow: SNN target
setFillColor(YELLOW)
fillRectangle(LEG_X, LEG_Y + 2.0 * LEG_ROW, LEG_SQ, LEG_SQ)
setFillColor(WHITE)
@@ -294,8 +347,31 @@ method run*(bot: SNNBot) =
case bot.phase
of DECIDE:
let relBearing = normalizeRelativeAngle(bot.enemyBearing - gunDir)
bot.lastRelBearing = relBearing
let inputs = encodeInput(relBearing)
bot.targetAngle = bot.snn.forward(inputs)
# Multi-tick inference: accumulate sin/cos and spike counts over N_INFER ticks
var totalSin = 0.0; var totalCos = 0.0
var tickSpikes: array[N_HID, float]
var tickVSnap: array[N_HID, float]
for h in 0 ..< N_HID: bot.lastSpikes[h] = 0.0
for _ in 0 ..< N_INFER:
var sinT, cosT: float
bot.snn.forward(inputs, tickSpikes, tickVSnap, sinT, cosT)
totalSin += sinT; totalCos += cosT
for h in 0 ..< N_HID:
bot.lastSpikes[h] += tickSpikes[h] # accumulate counts
# Store last-tick voltages for learning
bot.lastVSnap = tickVSnap
bot.snn.lastSinOut = totalSin; bot.snn.lastCosOut = totalCos
bot.snn.lastSnnAngle = arctan2(totalSin, totalCos) * 180.0 / PI
bot.targetAngle = gunDir + bot.snn.lastSnnAngle
# Log total spike count across inference window
var spikeCount = 0
var maxV = 0.0
for h in 0 ..< N_HID:
spikeCount += int(bot.lastSpikes[h])
maxV = max(maxV, bot.snn.vHid[h])
echo "tick=" & $bot.tick & " infer_spk=" & $spikeCount & "/" & $(N_INFER * N_HID) & " maxV=" & formatFloat(maxV, ffDecimal, 3)
bot.phase = WAITING
of WAITING:
@@ -305,14 +381,27 @@ method run*(bot: SNNBot) =
bot.phase = EVALUATE
of EVALUATE:
let err = abs(normalizeRelativeAngle(gunDir - bot.enemyBearing))
let reward = 1.0 / (1.0 + err)
# R-STDP weight update: w += η * R * e; then clamp and decay traces
for idx in 0 ..< N_IN * N_HID:
bot.snn.wih[idx] += ETA * reward * bot.snn.eligibility[idx]
bot.snn.wih[idx] = bot.snn.wih[idx].clamp(-W_CLAMP, W_CLAMP)
bot.snn.eligibility[idx] *= ELG_DECAY
echo "tick=" & $bot.tick & " error=" & $err & "° reward=" & $reward
let err = abs(normalizeRelativeAngle(gunDir - bot.enemyBearing))
# SuperSpike update: use relative bearing (what SNN should have learned to output) as target
let relTarget = bot.lastRelBearing
bot.snn.superSpikeUpdate(bot.lastSpikes, bot.lastVSnap, relTarget)
# Compute verbose logging metrics
var spikeCount = 0
var maxV = 0.0
for h in 0 ..< N_HID:
spikeCount += int(bot.lastSpikes[h])
maxV = max(maxV, bot.snn.vHid[h])
var meanWih = 0.0
for w in bot.snn.wih:
meanWih += abs(w)
meanWih /= float(N_IN * N_HID)
var meanWout = 0.0
for w in bot.snn.wSin:
meanWout += abs(w)
for w in bot.snn.wCos:
meanWout += abs(w)
meanWout /= float(2 * N_HID)
echo "tick=" & $bot.tick & " err=" & formatFloat(err, ffDecimal, 1) & "° infer_spk=" & $spikeCount & "/" & $(N_INFER * N_HID) & " maxV=" & formatFloat(maxV, ffDecimal, 3) & " |wih|=" & formatFloat(meanWih, ffDecimal, 4) & " |wOut|=" & formatFloat(meanWout, ffDecimal, 4) & " sin=" & formatFloat(bot.snn.lastSinOut, ffDecimal, 3) & " cos=" & formatFloat(bot.snn.lastCosOut, ffDecimal, 3) & " snnAngle=" & formatFloat(bot.snn.lastSnnAngle, ffDecimal, 1)
bot.phase = DECIDE
# Radar lock