diff --git a/SNNBot_garage/src/SNNBot.nim b/SNNBot_garage/src/SNNBot.nim index b58fe43..74ca77f 100644 --- a/SNNBot_garage/src/SNNBot.nim +++ b/SNNBot_garage/src/SNNBot.nim @@ -1,9 +1,10 @@ # SNNBot — SNN core + aiming loop prototype (issue #152). -# 36-input population-coded → 12 LIF hidden → polar-coded (sin/cos) output decoded via atan2. +# 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). # # State machine: -# DECIDE → feed bearing into SNN, store targetAngle, → WAITING +# DECIDE → feed bearing+velocity into SNN, store targetAngle, → WAITING # WAITING → aimTo() each tick; when error < 2° → EVALUATE # EVALUATE → measure error, compute SuperSpike update, log, → DECIDE @@ -16,9 +17,15 @@ import radar_lock/radar_lock as radar_lock const botJsonPath = currentSourcePath().parentDir / "SNNBot.json" const - N_IN = 36 # input neurons (10°-wide bands, -180..+180) - N_HID = 12 # hidden LIF neurons - BAND_DEG = 10.0 # degrees per input band + N_IN = 80 # input neurons: bearing(36) + vel_dir(36) + speed(8) + N_HID = 12 # hidden LIF neurons + BAND_DEG = 10.0 # degrees per input band + BEARING_OFFSET = 0 # neurons 0-35: relative bearing + VEL_DIR_OFFSET = 36 # neurons 36-71: velocity direction + SPEED_OFFSET = 72 # neurons 72-79: speed bands + N_SPEED_BANDS = 8 + SPEED_BAND_WIDTH = 1.0 # units/tick per band + MAX_SPEED = 8.0 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 MAX_GUN_TURN = 20.0 # max gun turn per tick (degrees) @@ -32,12 +39,15 @@ const 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 + FIRE_POWER = 1.0 # fixed firing power; bullet speed = 20 - 3*FIRE_POWER + BULLET_SPEED = 20.0 - 3.0 * FIRE_POWER + POS_BUF_LEN = 100 # ring buffer length for enemy positions (~100 ticks) # ── SNN types ───────────────────────────────────────────────────────────────── type SNN = object - wih: array[N_IN * N_HID, float] # 36×12 input→hidden weights + wih: array[N_IN * N_HID, float] # 80×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 @@ -59,17 +69,34 @@ proc initSNN(snn: var SNN) = for t in snn.preTrace.mitems: t = 0.0 snn.tick = 0 -proc encodeInput(bearing: float): array[N_IN, float] = - ## Population-code bearing into N_IN neurons. - ## Each neuron owns a 10° band centred at -175, -165, …, +175. - ## The band that contains the bearing fires 1.0; both neighbours get - ## linear interpolation for smoother encoding. +proc encodeBearing(inputs: var array[N_IN, float], bearing: float, offset: int) = + ## Population-code a -180..+180 angle into 36 neurons starting at offset. + ## Triangular interpolation between lo and hi band. let norm = ((bearing + 180.0) / BAND_DEG) # 0..36 - let lo = int(norm) mod N_IN - let hi = (lo + 1) mod N_IN - let frac = norm - float(int(norm)) # fractional position in band - result[lo] = 1.0 - frac - result[hi] = frac + let lo = int(norm) mod 36 + let hi = (lo + 1) mod 36 + let frac = norm - float(int(norm)) + inputs[offset + lo] = 1.0 - frac + inputs[offset + hi] = frac + +proc encodeInput(bearing: float): array[N_IN, float] = + ## Bearing-only encode for overlay (velocity channels stay 0). + encodeBearing(result, bearing, BEARING_OFFSET) + +proc encodeInputFull(bearing: float, velDirDeg: float, speed: float, + hasVel: bool): array[N_IN, float] = + ## Full 80-neuron encode: bearing + velocity direction + speed. + encodeBearing(result, bearing, BEARING_OFFSET) + if hasVel: + encodeBearing(result, velDirDeg, VEL_DIR_OFFSET) + # Speed: triangular interpolation over N_SPEED_BANDS bands, clamp to [0, MAX_SPEED] + let s = speed.clamp(0.0, MAX_SPEED) + let norm = s / SPEED_BAND_WIDTH + let lo = min(int(norm), N_SPEED_BANDS - 1) + let hi = min(lo + 1, N_SPEED_BANDS - 1) + let frac = norm - float(int(norm)) + result[SPEED_OFFSET + lo] += 1.0 - frac + result[SPEED_OFFSET + hi] += frac proc surrogateDerivative(v: float): float {.inline.} = ## σ'(U) = (1 + |β(U − ϑ)|)^{-2} — peaks at threshold, gives gradient direction. @@ -164,6 +191,14 @@ proc superSpikeUpdate(snn: var SNN, type Phase = enum DECIDE, WAITING, EVALUATE + AimSnapshot = object + aimAngle: float # SNN output angle (relative, degrees) + distance: float # distance to enemy at snapshot time + tick: int # tick when snapshot was taken + botX: float # bot position at snapshot time + botY: float + gunHeading: float # absolute gun heading at snapshot time + SNNBot = ref object of Bot snn: SNN phase: Phase @@ -175,6 +210,16 @@ type 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) + lastEnemyX: float64 # previous tick enemy position (for velocity) + lastEnemyY: float64 + hasLastPos: bool + velDirDeg: float64 # velocity direction (degrees) from last scan delta + velSpeed: float64 # speed (units/tick) from last scan delta + # Retroactive would-have-hit error signal + posBuf: array[POS_BUF_LEN, tuple[x, y: float, tick: int]] # ring buffer of enemy (x,y,tick) + posBufWrite: int # next write index into posBuf + posBufTick: int # game tick of the most recently written slot + snapshots: seq[AimSnapshot] # pending aim snapshots awaiting maturation # ── aimTo helper ────────────────────────────────────────────────────────────── @@ -312,15 +357,32 @@ method onScannedBot*(bot: SNNBot, e: ScannedBotEvent) = bot.enemyBearing = directionTo(bx, by, e.x, e.y) bot.enemyDist = distanceTo(bx, by, e.x, e.y) bot.hasContact = true + if bot.hasLastPos: + let dx = e.x - bot.lastEnemyX + let dy = e.y - bot.lastEnemyY + bot.velDirDeg = arctan2(dy, dx) * 180.0 / PI + bot.velSpeed = sqrt(dx * dx + dy * dy) + bot.lastEnemyX = e.x; bot.lastEnemyY = e.y + bot.hasLastPos = true + # Record enemy position in ring buffer + bot.posBuf[bot.posBufWrite] = (x: e.x, y: e.y, tick: bot.tick) + bot.posBufTick = bot.tick + bot.posBufWrite = (bot.posBufWrite + 1) mod POS_BUF_LEN method onRoundStarted*(bot: SNNBot, e: RoundStartedEvent) = setAdjustGunForBodyTurn(true) setAdjustRadarForBodyTurn(true) setAdjustRadarForGunTurn(true) radar_lock.init() - bot.hasContact = false - bot.phase = DECIDE - bot.tick = 0 + bot.hasContact = false + bot.hasLastPos = false + bot.velDirDeg = 0.0 + bot.velSpeed = 0.0 + bot.phase = DECIDE + bot.tick = 0 + bot.posBufWrite = 0 + bot.posBufTick = 0 + bot.snapshots = @[] setTargetSpeed(0.0) setTurnRate(0.0) @@ -348,7 +410,7 @@ method run*(bot: SNNBot) = of DECIDE: let relBearing = normalizeRelativeAngle(bot.enemyBearing - gunDir) bot.lastRelBearing = relBearing - let inputs = encodeInput(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] @@ -365,6 +427,14 @@ method run*(bot: SNNBot) = bot.snn.lastSinOut = totalSin; bot.snn.lastCosOut = totalCos bot.snn.lastSnnAngle = arctan2(totalSin, totalCos) * 180.0 / PI bot.targetAngle = gunDir + bot.snn.lastSnnAngle + # Record aim snapshot for retroactive error signal + bot.snapshots.add(AimSnapshot( + aimAngle: bot.snn.lastSnnAngle, + distance: bot.enemyDist, + tick: bot.tick, + botX: myX, + botY: myY, + gunHeading: gunDir)) # Log total spike count across inference window var spikeCount = 0 var maxV = 0.0 @@ -382,9 +452,39 @@ method run*(bot: SNNBot) = of EVALUATE: 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) + # Retroactive would-have-hit error signal: find the most recently matured snapshot. + # A snapshot matures when currentTick >= snapshotTick + ceil(distance / BULLET_SPEED). + # Look up enemy position at impact tick from ring buffer (offset from most-recent write). + var retroTarget = bot.lastRelBearing # fallback (unused if no matured snapshot) + var hasMatured = false + var keepIdx = 0 # first non-matured snapshot to keep + for i in 0 ..< bot.snapshots.len: + let snap = bot.snapshots[i] + let travelTicks = int(ceil(snap.distance / BULLET_SPEED)) + let impactTick = snap.tick + travelTicks + if bot.tick >= impactTick: + # Scan backward from most-recent write slot to find closest tick <= impactTick + var foundSlot = -1 + for j in 0 ..< POS_BUF_LEN: + let s = ((bot.posBufWrite - 1) - j + POS_BUF_LEN) mod POS_BUF_LEN + if bot.posBuf[s].tick > 0 and bot.posBuf[s].tick <= impactTick: + foundSlot = s + break + if foundSlot >= 0: + let ex = bot.posBuf[foundSlot].x + let ey = bot.posBuf[foundSlot].y + let absBearing = directionTo(snap.botX, snap.botY, ex, ey) + retroTarget = normalizeRelativeAngle(absBearing - snap.gunHeading) + hasMatured = true + keepIdx = i + 1 # discard matured snapshots up to and including this one + else: + break # snapshots are in order; stop at first non-matured + # Discard matured snapshots + if keepIdx > 0: + bot.snapshots = bot.snapshots[keepIdx .. ^1] + if hasMatured: + bot.snn.superSpikeUpdate(bot.lastSpikes, bot.lastVSnap, retroTarget) + # else: skip weight update — no matured snapshot yet (early game) # Compute verbose logging metrics var spikeCount = 0 var maxV = 0.0 @@ -401,7 +501,7 @@ method run*(bot: SNNBot) = 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) + echo "tick=" & $bot.tick & " err=" & formatFloat(err, ffDecimal, 1) & "° matured=" & $hasMatured & " retroAngle=" & formatFloat(retroTarget, 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