fix(SNNBot): tune SNN hyperparameters for 80-input layer

- Lower THRESH 0.2→0.08 to increase hidden firing rates (output layer was frozen due to rate≈0 in weight updates)
- Add ETA_IH=0.005 for input→hidden updates (10x smaller than output ETA=0.05 to prevent weight thrashing from large preTrace magnitudes)

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-09-14 22:16:51 +02:00
parent 0ee5c21ef5
commit 80dbf81479
+5 -4
View File
@@ -27,11 +27,12 @@ const
SPEED_BAND_WIDTH = 1.0 # units/tick per band SPEED_BAND_WIDTH = 1.0 # units/tick per band
MAX_SPEED = 8.0 MAX_SPEED = 8.0
LEAK = 0.9 # LIF membrane leak factor LEAK = 0.9 # LIF membrane leak factor
THRESH = 0.2 # ponytail: THRESH=0.2 — unitless system, must match weight scale; raise if neurons fire too much THRESH = 0.08 # ponytail: THRESH=0.08 tuned for 80-input layer; raise if firing saturates
MAX_GUN_TURN = 20.0 # max gun turn per tick (degrees) MAX_GUN_TURN = 20.0 # max gun turn per tick (degrees)
AIM_TOL = 2.0 # arrive tolerance (degrees) AIM_TOL = 2.0 # arrive tolerance (degrees)
ETA = 0.05 # SuperSpike learning rate (r_0 from paper) ETA = 0.05 # SuperSpike learning rate for hidden→output weights (r_0 from paper)
# ponytail: single learning rate, add RMaxProp optimizer if convergence unstable ETA_IH = 0.005 # ponytail: ETA_IH=0.005 scaled down for 80 inputs; raise if input→hidden converges too slowly
# ponytail: separate input→hidden rate; add RMaxProp optimizer if convergence still unstable
N_INFER = 10 # inference window ticks per DECIDE N_INFER = 10 # inference window ticks per DECIDE
# ponytail: N_INFER=10, increase if output still noisy; decrease if too slow per tick # 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) TRACE_DECAY = 0.9 # pre-synaptic trace decay (exponential low-pass)
@@ -184,7 +185,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 * preTrace[i] * sg * errHid snn.wih[i * N_HID + h] += ETA_IH * 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 ─────────────────────────────────────────────────────────