feat(SNNBot): add velocity input encoding and retroactive error signal (#159)
- Expand input layer from 36 to 80 neurons: bearing (0-35) + velocity direction (36-71) + speed (72-79) - Population-code velocity direction (same 10° band scheme as bearing) and speed (8 bands, 1 unit/tick) - Compute 1-tick velocity from position deltas in onScannedBot - Add 100-slot ring buffer with per-slot tick tracking for enemy positions - Replace instantaneous bearing error with retroactive would-have-hit signal - Skip learning until first aim snapshot matures (buffer fill period) - Delayed-target SuperSpike: weight updates use most recently matured retroactive bearing Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
+124
-24
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user