refactor(SNNBot): replace retroactive error with predictive lead target
Delete ring buffer, aim snapshots, maturation loop (~80 lines). Replace with simple velocity extrapolation: predict enemy position at bullet impact time, use bearing to predicted position as SuperSpike target. Immediate error signal every EVALUATE — no delay, no snapshot state storage. Mathematically equivalent to retroactive for linear movers.
This commit is contained in:
@@ -42,7 +42,6 @@ const
|
|||||||
# ponytail: uniform exploration noise; upgrade to annealed Gaussian if convergence needs tuning
|
# 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
|
FIRE_POWER = 1.0 # fixed firing power; bullet speed = 20 - 3*FIRE_POWER
|
||||||
BULLET_SPEED = 20.0 - 3.0 * FIRE_POWER
|
BULLET_SPEED = 20.0 - 3.0 * FIRE_POWER
|
||||||
POS_BUF_LEN = 100 # ring buffer length for enemy positions (~100 ticks)
|
|
||||||
|
|
||||||
# ── SNN types ─────────────────────────────────────────────────────────────────
|
# ── SNN types ─────────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
@@ -193,17 +192,6 @@ proc superSpikeUpdate(snn: var SNN,
|
|||||||
type
|
type
|
||||||
Phase = enum DECIDE, WAITING, EVALUATE
|
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
|
|
||||||
spikes: array[N_HID, float] # accumulated spike counts from this inference window
|
|
||||||
vSnap: array[N_HID, float] # hidden voltages from this inference window
|
|
||||||
preTrace: array[N_IN, float] # pre-synaptic traces from this inference window
|
|
||||||
|
|
||||||
SNNBot = ref object of Bot
|
SNNBot = ref object of Bot
|
||||||
snn: SNN
|
snn: SNN
|
||||||
phase: Phase
|
phase: Phase
|
||||||
@@ -220,11 +208,7 @@ type
|
|||||||
hasLastPos: bool
|
hasLastPos: bool
|
||||||
velDirDeg: float64 # velocity direction (degrees) from last scan delta
|
velDirDeg: float64 # velocity direction (degrees) from last scan delta
|
||||||
velSpeed: float64 # speed (units/tick) from last scan delta
|
velSpeed: float64 # speed (units/tick) from last scan delta
|
||||||
# Retroactive would-have-hit error signal
|
lastDecideGunDir: float # gun heading captured at DECIDE time for EVALUATE
|
||||||
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 ──────────────────────────────────────────────────────────────
|
# ── aimTo helper ──────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
@@ -369,10 +353,6 @@ method onScannedBot*(bot: SNNBot, e: ScannedBotEvent) =
|
|||||||
bot.velSpeed = sqrt(dx * dx + dy * dy)
|
bot.velSpeed = sqrt(dx * dx + dy * dy)
|
||||||
bot.lastEnemyX = e.x; bot.lastEnemyY = e.y
|
bot.lastEnemyX = e.x; bot.lastEnemyY = e.y
|
||||||
bot.hasLastPos = true
|
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) =
|
method onRoundStarted*(bot: SNNBot, e: RoundStartedEvent) =
|
||||||
setAdjustGunForBodyTurn(true)
|
setAdjustGunForBodyTurn(true)
|
||||||
@@ -385,9 +365,6 @@ method onRoundStarted*(bot: SNNBot, e: RoundStartedEvent) =
|
|||||||
bot.velSpeed = 0.0
|
bot.velSpeed = 0.0
|
||||||
bot.phase = DECIDE
|
bot.phase = DECIDE
|
||||||
bot.tick = 0
|
bot.tick = 0
|
||||||
bot.posBufWrite = 0
|
|
||||||
bot.posBufTick = 0
|
|
||||||
bot.snapshots = @[]
|
|
||||||
setTargetSpeed(0.0)
|
setTargetSpeed(0.0)
|
||||||
setTurnRate(0.0)
|
setTurnRate(0.0)
|
||||||
|
|
||||||
@@ -432,18 +409,7 @@ method run*(bot: SNNBot) =
|
|||||||
bot.snn.lastSinOut = totalSin; bot.snn.lastCosOut = totalCos
|
bot.snn.lastSinOut = totalSin; bot.snn.lastCosOut = totalCos
|
||||||
bot.snn.lastSnnAngle = arctan2(totalSin, totalCos) * 180.0 / PI
|
bot.snn.lastSnnAngle = arctan2(totalSin, totalCos) * 180.0 / PI
|
||||||
bot.targetAngle = gunDir + bot.snn.lastSnnAngle
|
bot.targetAngle = gunDir + bot.snn.lastSnnAngle
|
||||||
# Record aim snapshot for retroactive error signal; capture SNN state now
|
bot.lastDecideGunDir = gunDir
|
||||||
# before subsequent DECIDE cycles overwrite bot.lastSpikes / bot.lastVSnap.
|
|
||||||
bot.snapshots.add(AimSnapshot(
|
|
||||||
aimAngle: bot.snn.lastSnnAngle,
|
|
||||||
distance: bot.enemyDist,
|
|
||||||
tick: bot.tick,
|
|
||||||
botX: myX,
|
|
||||||
botY: myY,
|
|
||||||
gunHeading: gunDir,
|
|
||||||
spikes: bot.lastSpikes,
|
|
||||||
vSnap: bot.lastVSnap,
|
|
||||||
preTrace: bot.snn.preTrace))
|
|
||||||
# Log total spike count across inference window
|
# Log total spike count across inference window
|
||||||
var spikeCount = 0
|
var spikeCount = 0
|
||||||
var maxV = 0.0
|
var maxV = 0.0
|
||||||
@@ -462,57 +428,15 @@ method run*(bot: SNNBot) =
|
|||||||
|
|
||||||
of EVALUATE:
|
of EVALUATE:
|
||||||
let err = abs(normalizeRelativeAngle(gunDir - bot.enemyBearing))
|
let err = abs(normalizeRelativeAngle(gunDir - bot.enemyBearing))
|
||||||
# Retroactive would-have-hit error signal: find the most recently matured snapshot.
|
# Predictive error signal: extrapolate enemy position at bullet impact time.
|
||||||
# A snapshot matures when currentTick >= snapshotTick + ceil(distance / BULLET_SPEED).
|
if bot.hasLastPos:
|
||||||
# Look up enemy position at impact tick from ring buffer (offset from most-recent write).
|
let travelTime = bot.enemyDist / BULLET_SPEED
|
||||||
var retroTarget = bot.lastRelBearing # fallback; overwritten after averaging (used for logging)
|
let velRad = degToRad(bot.velDirDeg)
|
||||||
var hasMatured = false
|
let futureX = bot.lastEnemyX + cos(velRad) * bot.velSpeed * travelTime
|
||||||
var keepIdx = 0 # first non-matured snapshot to keep
|
let futureY = bot.lastEnemyY + sin(velRad) * bot.velSpeed * travelTime
|
||||||
# Accumulate retroTarget circular components across all matured snapshots.
|
let absBearing = directionTo(myX, myY, futureX, futureY)
|
||||||
var sumSin = 0.0; var sumCos = 0.0
|
let targetAngle = normalizeRelativeAngle(absBearing - bot.lastDecideGunDir)
|
||||||
var matureCount = 0
|
bot.snn.superSpikeUpdate(bot.lastSpikes, bot.lastVSnap, bot.snn.preTrace, targetAngle)
|
||||||
# Most-recent matured snapshot's SNN state (most relevant for weight update).
|
|
||||||
var bestSpikes: array[N_HID, float]
|
|
||||||
var bestVSnap: array[N_HID, float]
|
|
||||||
var bestPreTrace: array[N_IN, float]
|
|
||||||
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)
|
|
||||||
let rt = normalizeRelativeAngle(absBearing - snap.gunHeading)
|
|
||||||
# Accumulate circular mean components
|
|
||||||
sumSin += sin(degToRad(rt))
|
|
||||||
sumCos += cos(degToRad(rt))
|
|
||||||
inc matureCount
|
|
||||||
# Keep most-recent (highest index) snapshot's SNN state
|
|
||||||
bestSpikes = snap.spikes
|
|
||||||
bestVSnap = snap.vSnap
|
|
||||||
bestPreTrace = snap.preTrace
|
|
||||||
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
|
|
||||||
# One update per EVALUATE using circular-mean target — stable learning rate regardless of snapshot count.
|
|
||||||
if matureCount > 0:
|
|
||||||
retroTarget = arctan2(sumSin, sumCos) * 180.0 / PI
|
|
||||||
bot.snn.superSpikeUpdate(bestSpikes, bestVSnap, bestPreTrace, retroTarget)
|
|
||||||
# Discard matured snapshots; immature ones are preserved automatically (keepIdx stays 0
|
|
||||||
# or points past the last matured entry; the rest of bot.snapshots is kept intact).
|
|
||||||
if keepIdx > 0:
|
|
||||||
bot.snapshots = bot.snapshots[keepIdx .. ^1]
|
|
||||||
# else: skip weight update — no matured snapshot yet (early game)
|
|
||||||
# Compute verbose logging metrics
|
# Compute verbose logging metrics
|
||||||
var spikeCount = 0
|
var spikeCount = 0
|
||||||
var maxV = 0.0
|
var maxV = 0.0
|
||||||
@@ -529,7 +453,7 @@ method run*(bot: SNNBot) =
|
|||||||
for w in bot.snn.wCos:
|
for w in bot.snn.wCos:
|
||||||
meanWout += abs(w)
|
meanWout += abs(w)
|
||||||
meanWout /= float(2 * N_HID)
|
meanWout /= float(2 * N_HID)
|
||||||
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)
|
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
|
bot.phase = DECIDE
|
||||||
|
|
||||||
# Radar lock
|
# Radar lock
|
||||||
|
|||||||
Reference in New Issue
Block a user