diff --git a/SNNBot_garage/src/SNNBot.nim b/SNNBot_garage/src/SNNBot.nim index daf4b98..b58fe43 100644 --- a/SNNBot_garage/src/SNNBot.nim +++ b/SNNBot_garage/src/SNNBot.nim @@ -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