From dbfb58ca1e6ff1bfdc39d5876043793d8802a4cb Mon Sep 17 00:00:00 2001 From: Davide Cappellini Date: Mon, 14 Sep 2026 23:00:22 +0200 Subject: [PATCH] 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 --- SNNBot_garage/src/SNNBot.nim | 127 ++++++++++++++++--------- SNNBot_garage/src/reservoir.nim | 164 ++++++++++++++++++++++++++++++++ 2 files changed, 245 insertions(+), 46 deletions(-) create mode 100644 SNNBot_garage/src/reservoir.nim diff --git a/SNNBot_garage/src/SNNBot.nim b/SNNBot_garage/src/SNNBot.nim index 0a7fac8..34cb810 100644 --- a/SNNBot_garage/src/SNNBot.nim +++ b/SNNBot_garage/src/SNNBot.nim @@ -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 diff --git a/SNNBot_garage/src/reservoir.nim b/SNNBot_garage/src/reservoir.nim new file mode 100644 index 0000000..06c5cb6 --- /dev/null +++ b/SNNBot_garage/src/reservoir.nim @@ -0,0 +1,164 @@ +# ponytail: binary reservoir aimer — prototype; if reservoir projection is poor, scale RESERVOIR_SIZE to 4096 + +import std/[bitops, math] + +# ── Constants ────────────────────────────────────────────────────────────────── + +const + RESERVOIR_SIZE* = 1024 + N_BINS* = 72 + BIN_WIDTH* = 5.0 + INPUT_BITS* = 80 + SPARSITY_IN = 0.1 + SPARSITY_REC = 0.05 + +# ── Types ────────────────────────────────────────────────────────────────────── + +type + BitVec80* = array[2, uint64] # 128 bits allocated, lower 80 used + BitVec* = array[16, uint64] # 1024 bits + + Reservoir* = object + wIn: array[RESERVOIR_SIZE, BitVec80] + wRec: array[RESERVOIR_SIZE, BitVec] + threshold: array[RESERVOIR_SIZE, int] + state: BitVec + readout: array[N_BINS, BitVec] + scores*: array[N_BINS, int] + rngState: uint64 + +# ── PRNG ─────────────────────────────────────────────────────────────────────── + +proc nextRand(r: var Reservoir): uint64 {.inline.} = + var x = r.rngState + x = x xor (x shl 13) + x = x xor (x shr 7) + x = x xor (x shl 17) + r.rngState = x + return x + +proc sparseBits(r: var Reservoir, density: float): uint64 = + ## uint64 with approximately density*64 bits set via repeated AND. + ## density=0.1 → ~1 AND needed for 1/2, but we use a direct Bernoulli approach. + ## ponytail: simple loop over 64 bits; replace with table-lookup if init is hot + result = 0'u64 + let thresh = uint64(density * float(high(uint64))) + for bit in 0 ..< 64: + if r.nextRand() < thresh: + result = result or (1'u64 shl bit) + +proc randomDecayMask(r: var Reservoir): uint64 = + ## ~1% bits set: AND 6 random words (1/2^6 = 1/64 ≈ 1.5% density). + result = r.nextRand() + for _ in 0 ..< 5: + result = result and r.nextRand() + +# ── Init ─────────────────────────────────────────────────────────────────────── + +proc initReservoir*(seed: int): Reservoir = + result.rngState = uint64(seed) or 1'u64 # avoid zero state + + for i in 0 ..< RESERVOIR_SIZE: + # Input masks: 80 bits across two uint64s (word 0: bits 0-63, word 1: bits 64-79) + result.wIn[i][0] = sparseBits(result, SPARSITY_IN) + # Only bits 0-15 of word 1 are meaningful (global bits 64-79) + result.wIn[i][1] = sparseBits(result, SPARSITY_IN) and 0x0000_0000_0000_FFFF'u64 + + for w in 0 ..< 16: + result.wRec[i][w] = sparseBits(result, SPARSITY_REC) + + let inPop = popcount(result.wIn[i][0]) + popcount(result.wIn[i][1]) + let recPop = block: + var s = 0 + for w in 0 ..< 16: s += popcount(result.wRec[i][w]) + s + # Threshold: ~50% of expected input votes + ~30% of expected recurrent votes + result.threshold[i] = max(1, int(float(inPop) * 0.5 + float(recPop) * 0.3)) + + # readout and state are zero-initialized by default + +# ── Forward ──────────────────────────────────────────────────────────────────── + +proc forward*(r: var Reservoir, input: BitVec80): int = + var newState: BitVec + + for i in 0 ..< RESERVOIR_SIZE: + let inScore = popcount(input[0] and r.wIn[i][0]) + + popcount(input[1] and r.wIn[i][1]) + var recScore = 0 + for w in 0 ..< 16: + recScore += popcount(r.state[w] and r.wRec[i][w]) + + if inScore + recScore > r.threshold[i]: + let wordIdx = i shr 6 # i div 64 + let bitIdx = i and 63 # i mod 64 + newState[wordIdx] = newState[wordIdx] or (1'u64 shl bitIdx) + + r.state = newState + + var bestBin = 0 + var bestScore = -1 + for k in 0 ..< N_BINS: + var score = 0 + for w in 0 ..< 16: + score += popcount(r.state[w] and r.readout[k][w]) + r.scores[k] = score + if score > bestScore: + bestScore = score + bestBin = k + + return bestBin + +# ── Angle helpers ────────────────────────────────────────────────────────────── + +proc binToAngle*(bin: int): float = + ## Bin 0 = -180°, Bin 36 = 0°, Bin 71 = +175° + result = -180.0 + float(bin) * BIN_WIDTH + +proc angleToBin*(angle: float): int = + ## angle in -180..+180, map to bin 0..71 + var a = angle + if a < -180.0: a += 360.0 + if a >= 180.0: a -= 360.0 + result = int((a + 180.0) / BIN_WIDTH) mod N_BINS + +proc interpolatedAngle*(r: Reservoir, winnerBin: int): float = + ## Weighted circular centroid of winner ± 1 bins for sub-5° precision. + let left = (winnerBin - 1 + N_BINS) mod N_BINS + let right = (winnerBin + 1) mod N_BINS + let sW = float(r.scores[winnerBin]) + let sL = float(r.scores[left]) + let sR = float(r.scores[right]) + let total = sW + sL + sR + if total == 0.0: return binToAngle(winnerBin) + let aW = degToRad(binToAngle(winnerBin)) + let aL = degToRad(binToAngle(left)) + let aR = degToRad(binToAngle(right)) + let sinAvg = (sW * sin(aW) + sL * sin(aL) + sR * sin(aR)) / total + let cosAvg = (sW * cos(aW) + sL * cos(aL) + sR * cos(aR)) / total + result = radToDeg(arctan2(sinAvg, cosAvg)) + +# ── Learning ─────────────────────────────────────────────────────────────────── + +proc learn*(r: var Reservoir, correctAngle: float) = + let correctBin = angleToBin(correctAngle) + + # Reinforce correct bin + for w in 0 ..< 16: + r.readout[correctBin][w] = r.readout[correctBin][w] or r.state[w] + + # Punish highest-scoring wrong bin + var worstBin = -1 + var worstScore = -1 + for k in 0 ..< N_BINS: + if k != correctBin and r.scores[k] > worstScore: + worstScore = r.scores[k] + worstBin = k + if worstBin >= 0: + for w in 0 ..< 16: + r.readout[worstBin][w] = r.readout[worstBin][w] and (not r.state[w]) + + # Decay: clear ~1.5% of bits per bin to prevent saturation + for k in 0 ..< N_BINS: + for w in 0 ..< 16: + r.readout[k][w] = r.readout[k][w] and (not randomDecayMask(r))