feat(SNNBot): add binary reservoir aimer as alternative to SuperSpike (#159)
New architecture: 1024 binary neurons in fixed random reservoir, 72-bin population-coded output, WTA Hebbian learning with binary ops. Forward pass: AND + popcount. Learning: OR (reinforce) / AND NOT (punish). No backprop, no floats in hot path. Toggle via USE_RESERVOIR const. Forecast: ~200-400 ticks to learn stationary target aiming. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -2,15 +2,19 @@
|
||||
# 80-input population-coded → 12 LIF hidden → polar-coded (sin/cos) output decoded via atan2.
|
||||
# Inputs: bearing [0-35], velocity direction [36-71], speed [72-79].
|
||||
# No movement, no firing. SuperSpike three-factor rule on all weights (issue #158).
|
||||
# Binary reservoir alternative toggled via USE_RESERVOIR const (issue #159).
|
||||
#
|
||||
# State machine:
|
||||
# DECIDE → feed bearing+velocity into SNN, store targetAngle, → WAITING
|
||||
# DECIDE → feed bearing+velocity into SNN/reservoir, store targetAngle, → WAITING
|
||||
# WAITING → aimTo() each tick; when error < 2° → EVALUATE
|
||||
# EVALUATE → measure error, compute SuperSpike update, log, → DECIDE
|
||||
# EVALUATE → measure error, compute SuperSpike/reservoir update, log, → DECIDE
|
||||
|
||||
import std/[math, random, os, strutils]
|
||||
import robocode_tankroyale_botapi
|
||||
import radar_lock/radar_lock as radar_lock
|
||||
import reservoir
|
||||
|
||||
const USE_RESERVOIR* = true
|
||||
|
||||
# ── Constants ──────────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -194,6 +198,7 @@ type
|
||||
|
||||
SNNBot = ref object of Bot
|
||||
snn: SNN
|
||||
res: Reservoir
|
||||
phase: Phase
|
||||
targetAngle: float # SNN output (absolute bearing)
|
||||
enemyBearing: float # last known enemy bearing
|
||||
@@ -370,6 +375,17 @@ method onRoundStarted*(bot: SNNBot, e: RoundStartedEvent) =
|
||||
|
||||
method onGameStarted*(bot: SNNBot, e: GameStartedEventForBot) =
|
||||
initSNN(bot.snn)
|
||||
bot.res = initReservoir(42)
|
||||
|
||||
# ── Reservoir helpers ─────────────────────────────────────────────────────────
|
||||
|
||||
proc toBitVec80(inputs: array[N_IN, float]): BitVec80 =
|
||||
result = [0'u64, 0'u64]
|
||||
for i in 0 ..< N_IN:
|
||||
if inputs[i] > 0.5:
|
||||
let word = i div 64
|
||||
let bit = i mod 64
|
||||
result[word] = result[word] or (1'u64 shl bit)
|
||||
|
||||
# ── Main loop ─────────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -393,30 +409,38 @@ method run*(bot: SNNBot) =
|
||||
let relBearing = normalizeRelativeAngle(bot.enemyBearing - gunDir)
|
||||
bot.lastRelBearing = relBearing
|
||||
let inputs = encodeInputFull(relBearing, bot.velDirDeg, bot.velSpeed, bot.hasLastPos)
|
||||
# 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
|
||||
when USE_RESERVOIR:
|
||||
let binInput = toBitVec80(inputs)
|
||||
let winBin = bot.res.forward(binInput)
|
||||
let aimAngle = bot.res.interpolatedAngle(winBin)
|
||||
bot.targetAngle = gunDir + aimAngle
|
||||
bot.lastDecideGunDir = gunDir
|
||||
echo "RES tick=" & $bot.tick & " bin=" & $winBin & " aim=" & formatFloat(aimAngle, ffDecimal, 1)
|
||||
else:
|
||||
# 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
|
||||
bot.lastDecideGunDir = gunDir
|
||||
# Log total spike count across inference window
|
||||
var spikeCount = 0
|
||||
var maxV = 0.0
|
||||
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
|
||||
bot.lastDecideGunDir = gunDir
|
||||
# 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)
|
||||
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:
|
||||
@@ -427,33 +451,44 @@ method run*(bot: SNNBot) =
|
||||
bot.phase = EVALUATE
|
||||
|
||||
of EVALUATE:
|
||||
let err = abs(normalizeRelativeAngle(gunDir - bot.enemyBearing))
|
||||
# Predictive error signal: extrapolate enemy position at bullet impact time.
|
||||
if bot.hasLastPos:
|
||||
let travelTime = bot.enemyDist / BULLET_SPEED
|
||||
let velRad = degToRad(bot.velDirDeg)
|
||||
let futureX = bot.lastEnemyX + cos(velRad) * bot.velSpeed * travelTime
|
||||
let futureY = bot.lastEnemyY + sin(velRad) * bot.velSpeed * travelTime
|
||||
let absBearing = directionTo(myX, myY, futureX, futureY)
|
||||
let targetAngle = normalizeRelativeAngle(absBearing - bot.lastDecideGunDir)
|
||||
bot.snn.superSpikeUpdate(bot.lastSpikes, bot.lastVSnap, bot.snn.preTrace, targetAngle)
|
||||
# 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)
|
||||
let absBearing = directionTo(myX, myY, futureX, futureY)
|
||||
let targetRel = normalizeRelativeAngle(absBearing - bot.lastDecideGunDir)
|
||||
when USE_RESERVOIR:
|
||||
let winBin = bot.res.forward(toBitVec80(encodeInputFull(
|
||||
normalizeRelativeAngle(bot.enemyBearing - bot.lastDecideGunDir),
|
||||
bot.velDirDeg, bot.velSpeed, bot.hasLastPos)))
|
||||
let aimAngle = bot.res.interpolatedAngle(winBin)
|
||||
bot.res.learn(targetRel)
|
||||
echo "RES tick=" & $bot.tick & " bin=" & $winBin &
|
||||
" aim=" & formatFloat(aimAngle, ffDecimal, 1) &
|
||||
" target=" & formatFloat(targetRel, ffDecimal, 1) &
|
||||
" err=" & formatFloat(abs(aimAngle - targetRel), ffDecimal, 1)
|
||||
else:
|
||||
let err = abs(normalizeRelativeAngle(gunDir - bot.enemyBearing))
|
||||
bot.snn.superSpikeUpdate(bot.lastSpikes, bot.lastVSnap, bot.snn.preTrace, targetRel)
|
||||
# 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
|
||||
|
||||
Reference in New Issue
Block a user