b509195ee9
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
350 lines
15 KiB
Nim
350 lines
15 KiB
Nim
## SAC_LSTM_Bot — Recurrent SAC-v2 bot (Gitea #48).
|
||
##
|
||
## Thread layout (decisions Q1–Q14, see integration.nim for plumbing):
|
||
## bot thread — this file's run(): inference only, <2ms/tick.
|
||
## training thread — permanent background SAC updates (integration.nim).
|
||
## I/O thread — atomic weight saves (integration.nim).
|
||
## No Arraymancer tensor ever crosses a thread boundary: the bot object and
|
||
## channels carry plain scalars / fixed arrays / plain seqs only.
|
||
|
||
import std/[os, math, random, algorithm]
|
||
import arraymancer except Linear
|
||
import tankroyale_botapi
|
||
import radar_lock
|
||
import SAC_LSTM_Bot/state
|
||
import SAC_LSTM_Bot/network
|
||
import SAC_LSTM_Bot/actions
|
||
import SAC_LSTM_Bot/rewards
|
||
import SAC_LSTM_Bot/integration
|
||
|
||
# Identity json: baked-in src json by default; SACLSTM_BOT_JSON lets a mirror
|
||
# twin (#54) boot the same binary under its own name (loadBotInfo gives the
|
||
# json total precedence over env, so the twin must point at its own file).
|
||
let botJsonPath = getEnv("SACLSTM_BOT_JSON",
|
||
currentSourcePath().parentDir / "SAC_LSTM_Bot.json")
|
||
|
||
# ── Colors (Recurrent Royalty palette) ───────────────────────────────────────
|
||
|
||
const
|
||
ColBody = fromHex("#7B2FBE")
|
||
ColTurret = fromHex("#FFD700")
|
||
ColGun = fromHex("#4A0E6B")
|
||
ColRadar = fromHex("#FFD700")
|
||
ColScan = fromHex("#FFB000")
|
||
ColBullet = fromHex("#FFC125")
|
||
ColTracks = fromHex("#2C2C34")
|
||
|
||
proc applyColors() =
|
||
setBodyColor(ColBody)
|
||
setTurretColor(ColTurret)
|
||
setGunColor(ColGun)
|
||
setRadarColor(ColRadar)
|
||
setScanColor(ColScan)
|
||
setBulletColor(ColBullet)
|
||
setTracksColor(ColTracks)
|
||
|
||
# ── Enemy bullet tracking ────────────────────────────────────────────────────
|
||
# PPO_Bot-proven pattern: dead-reckoned fixed buffer. getBulletStates() is NOT
|
||
# used from the bot thread — its seq refcount is shared with the main thread.
|
||
|
||
type InFlightBullet = object
|
||
x, y, vx, vy, power: float64
|
||
|
||
const MaxBotBullets = 4
|
||
const MaxHidden = 512 # ponytail: cap for the plain-array LSTM persistence; raise if SACLSTM_HIDDEN_SIZE > 512
|
||
|
||
# ── Bot type — PLAIN DATA ONLY on the shared object (no tensors/heap seqs:
|
||
# each round runs a fresh bot thread; heap blocks owned by the previous
|
||
# round's thread must not be freed from another thread) ─────────────────────
|
||
|
||
type SacBot = ref object of Bot
|
||
enemyBearing: float # last known absolute bearing to enemy
|
||
battleId: int # main thread bumps in onGameStarted; bot thread compares
|
||
seenBattle: int # bot-thread copy for battle-change detection
|
||
newBattleSent: bool # first scan of THIS battle emits NewBattle
|
||
hasContact: bool
|
||
enemy: EnemyData
|
||
ticksSinceScan: int
|
||
# per-step reward accumulators (consumed by the next tick's transition)
|
||
dmgDealt, dmgTaken, wastedPower: float64
|
||
wallHits, hits, ramTaken: int
|
||
# pending transition (episode spans the whole battle; round end is NOT a boundary)
|
||
hasLastTrans: bool
|
||
lastState: array[STATE_DIM, float32]
|
||
lastAction: array[ACTION_DIM, float32]
|
||
rn: RewardNormalizer # Welford running stats, persists across battles
|
||
bullets: array[MaxBotBullets, InFlightBullet]
|
||
bulletCount: int
|
||
hArr, cArr: array[MaxHidden, float32] # LSTM state across rounds; zeros at battle start
|
||
|
||
# ── Plain-array <-> tensor helpers (bot thread only) ──────────────────────────
|
||
|
||
proc stateToArr(t: Tensor[float32]): array[STATE_DIM, float32] =
|
||
for i in 0 ..< STATE_DIM: result[i] = t[i]
|
||
|
||
proc actionToArr(t: Tensor[float32]): array[ACTION_DIM, float32] =
|
||
for i in 0 ..< ACTION_DIM: result[i] = t[i]
|
||
|
||
proc hiddenToTensor(arr: array[MaxHidden, float32]; n: int): Tensor[float32] =
|
||
result = newTensor[float32](n)
|
||
for i in 0 ..< n: result[i] = arr[i]
|
||
|
||
proc tensorToHidden(t: Tensor[float32]; arr: var array[MaxHidden, float32]) =
|
||
for i in 0 ..< t.shape[0]: arr[i] = t[i]
|
||
|
||
# ── Reward ────────────────────────────────────────────────────────────────────
|
||
|
||
proc takeReward(bot: SacBot; win = false, loss = false): float32 =
|
||
## Consume accumulated step events -> Welford-normalized reward (#44).
|
||
## Lever 2 (#59): pass enemy distance (frac of arena diagonal) for the
|
||
## anti-charge term; sentinel 2.0 (> ChargeDistFrac) when no contact.
|
||
var distFrac = 2.0
|
||
if bot.hasContact:
|
||
let diag = hypot(getArenaWidth().float64, getArenaHeight().float64)
|
||
distFrac = hypot(bot.enemy.x - getX(), bot.enemy.y - getY()) / diag
|
||
let raw = computeReward(
|
||
damageInflicted = bot.dmgDealt,
|
||
damageReceived = bot.dmgTaken,
|
||
wallHitTicks = bot.wallHits,
|
||
wastedShotPower = bot.wastedPower,
|
||
hitCount = bot.hits,
|
||
ramTakenCount = bot.ramTaken,
|
||
enemyDistFrac = distFrac,
|
||
win = win, loss = loss)
|
||
# Lever-2 (#59) observability: env-gated one-liner for smoke/calibration
|
||
# greps — proves hit/ram/charge terms fire and shows raw magnitudes. File
|
||
# (not stderr): the battle runner swallows bot process streams.
|
||
# ponytail: grows unbounded if left on; keep off outside smokes.
|
||
if getEnv("SACLSTM_REWARD_DEBUG") == "1" and
|
||
(bot.hits > 0 or bot.ramTaken > 0 or (distFrac < ChargeDistFrac and bot.dmgDealt <= 0.0)):
|
||
try:
|
||
let f = open(getWeightsPath().parentDir.parentDir / "reward_debug.log", fmAppend)
|
||
f.writeLine("raw=" & $raw & " hits=" & $bot.hits & " ram=" & $bot.ramTaken &
|
||
" distFrac=" & $distFrac)
|
||
f.close()
|
||
except CatchableError:
|
||
discard
|
||
bot.dmgDealt = 0; bot.dmgTaken = 0; bot.wastedPower = 0; bot.wallHits = 0
|
||
bot.hits = 0; bot.ramTaken = 0
|
||
let norm = rewards.normalize(bot.rn, raw)
|
||
rewards.update(bot.rn, raw)
|
||
norm.float32
|
||
|
||
# ── Event handlers ────────────────────────────────────────────────────────────
|
||
# onGameStarted/onRoundStarted fire on the MAIN thread (bot thread not yet
|
||
# started or already joined) — plain-field writes only, no tensors here.
|
||
|
||
method onGameStarted*(bot: SacBot, e: GameStartedEventForBot) =
|
||
inc bot.battleId # bot thread zeroes LSTM + per-battle flags at next tick (Q4/Q12)
|
||
|
||
method onRoundStarted*(bot: SacBot, e: RoundStartedEvent) =
|
||
setAdjustRadarForBodyTurn(true)
|
||
setAdjustRadarForGunTurn(true)
|
||
radar_lock.init()
|
||
applyColors()
|
||
# Per-round reset ONLY. NOT the LSTM hidden state (persists across rounds, Q4);
|
||
# NOT hasLastTrans (the pending transition spans the round boundary — episode
|
||
# ends at battle end only).
|
||
bot.hasContact = false
|
||
bot.enemy = EnemyData()
|
||
bot.ticksSinceScan = 0
|
||
bot.bulletCount = 0
|
||
|
||
method onScannedBot*(bot: SacBot, e: ScannedBotEvent) =
|
||
bot.enemyBearing = directionTo(getX(), getY(), e.x, e.y)
|
||
# Same-tick radar lock: apply turn rate immediately so it takes effect this tick.
|
||
setRadarTurnRate(radar_lock.doRadar(getRadarDirection(), bot.enemyBearing))
|
||
# Fire detection: energy drop in [0.1, 3.0] between scans (PPO_Bot heuristic).
|
||
let prevE = if bot.hasContact: bot.enemy.energy else: e.energy
|
||
let drop = prevE - e.energy
|
||
bot.enemy.hasFired = bot.hasContact and drop >= 0.1 and drop <= 3.0
|
||
if bot.enemy.hasFired:
|
||
bot.enemy.lastFirePower = drop
|
||
# Keep previous-scan deltas before overwriting (state.nim derives accel/turn rate).
|
||
bot.enemy.prevSpeed = bot.enemy.speed
|
||
bot.enemy.prevDirection = bot.enemy.direction
|
||
bot.enemy.hasPrevScan = bot.hasContact
|
||
bot.enemy.x = e.x
|
||
bot.enemy.y = e.y
|
||
bot.enemy.direction = e.direction
|
||
bot.enemy.speed = e.speed
|
||
bot.enemy.energy = e.energy
|
||
bot.hasContact = true
|
||
bot.ticksSinceScan = 0
|
||
# Q12b/Q14+#49: one NewBattle per battle, numeric scannedBotId. Identity is
|
||
# keyed on getBotName(id) training-side (integration.opponentKey), numeric
|
||
# fallback in the pre-BotListUpdate window.
|
||
if not bot.newBattleSent:
|
||
bot.newBattleSent = true
|
||
discard sendTrainingMsg(TrainingMsg(kind: tmkNewBattle, enemyId: e.scannedBotId))
|
||
|
||
method onBulletHit*(bot: SacBot, e: BulletHitBotEvent) =
|
||
bot.dmgDealt += e.damage
|
||
inc bot.hits
|
||
|
||
method onHitByBullet*(bot: SacBot, e: HitByBulletEvent) =
|
||
bot.dmgTaken += e.damage
|
||
|
||
# Lever 2 (#59): anti-ram — every BotHitBotEvent receipt means a bot-bot
|
||
# collision happened and we ate RAM_DAMAGE (server deals 0.6 to both parties;
|
||
# only the hitter gets notified). Flat per-event penalty; being rammed without
|
||
# hitting back stays event-invisible.
|
||
# ponytail: enemy-initiated rams undetected — add energy-residual detection if
|
||
# v2 battle data shows ram-heavy losses.
|
||
method onHitBot*(bot: SacBot, e: BotHitBotEvent) =
|
||
inc bot.ramTaken
|
||
|
||
method onHitWall*(bot: SacBot, e: BotHitWallEvent) =
|
||
inc bot.wallHits
|
||
|
||
method onBulletHitWall*(bot: SacBot, e: BulletHitWallEvent) =
|
||
if e.bullet.ownerId == getMyId():
|
||
bot.wastedPower += e.bullet.power
|
||
|
||
method onGameAborted*(bot: SacBot) =
|
||
# Mid-round abort: drop the pending transition rather than leak it into the
|
||
# next battle's data.
|
||
bot.hasLastTrans = false
|
||
bot.dmgDealt = 0; bot.dmgTaken = 0; bot.wastedPower = 0; bot.wallHits = 0
|
||
bot.hits = 0; bot.ramTaken = 0
|
||
|
||
method onRoundEnded*(bot: SacBot, e: RoundEndedEventForBot) =
|
||
## Harness liveness signal (#49): RunTraining.java watches round_counter.txt
|
||
## and aborts the battle if it freezes (dead bot process). Main thread.
|
||
bumpRoundCounter()
|
||
|
||
method onGameEnded*(bot: SacBot, e: GameEndedEventForBot) =
|
||
## Battle end -> terminal transition with done=true. Main thread; the API has
|
||
## joined the bot thread before this fires, so these plain fields are quiescent.
|
||
if bot.hasLastTrans:
|
||
let r = bot.takeReward(win = e.results.rank == 1, loss = e.results.rank != 1)
|
||
discard sendTrainingMsg(TrainingMsg(kind: tmkTransition,
|
||
state: bot.lastState, action: bot.lastAction,
|
||
reward: r, nextState: bot.lastState, done: true))
|
||
bot.hasLastTrans = false
|
||
|
||
# ── Run loop (bot thread) ─────────────────────────────────────────────────────
|
||
|
||
method run(bot: SacBot) =
|
||
randomize()
|
||
# Per-thread locals: born and freed on THIS thread, every round. Nothing
|
||
# heap-owned survives the round boundary except the bot object's plain fields.
|
||
var actor: ActorNet
|
||
var actorReady = false
|
||
var myVersion = 0
|
||
var myHidden = 0
|
||
var flat: seq[float32]
|
||
|
||
while isRunning():
|
||
# Battle boundary (Q4/Q12): zero LSTM persistence + per-battle flags.
|
||
if bot.battleId != bot.seenBattle:
|
||
bot.seenBattle = bot.battleId
|
||
bot.newBattleSent = false
|
||
zeroMem(addr bot.hArr, sizeof(bot.hArr))
|
||
zeroMem(addr bot.cArr, sizeof(bot.cArr))
|
||
|
||
# Weight sync (Q7/Q11): always-latest; rebuild this thread's tensors on change.
|
||
if pullWeights(myVersion, myHidden, flat):
|
||
if myHidden > MaxHidden:
|
||
# hArr/cArr are fixed-capacity; a bigger SACLSTM_HIDDEN_SIZE would
|
||
# heap-overflow them in tensorToHidden. Loud misconfig beats corruption.
|
||
raise newException(ValueError, "SACLSTM_HIDDEN_SIZE=" & $myHidden &
|
||
" exceeds MaxHidden=" & $MaxHidden & " (bot-side LSTM persistence cap)")
|
||
var cur = 0
|
||
actor = actorFromFlat(flat, cur, myHidden)
|
||
actorReady = true
|
||
|
||
# Spawn an enemy bullet when a fresh scan shows they fired; then advance and
|
||
# prune the tracked bullets (positions feed state slots 22–33).
|
||
if bot.hasContact and bot.enemy.hasFired and bot.bulletCount < MaxBotBullets:
|
||
let p = bot.enemy.lastFirePower
|
||
let spd = 20.0 - 3.0 * p
|
||
let ang = arctan2(getY() - bot.enemy.y, getX() - bot.enemy.x)
|
||
bot.bullets[bot.bulletCount] = InFlightBullet(x: bot.enemy.x, y: bot.enemy.y,
|
||
vx: spd * cos(ang), vy: spd * sin(ang), power: p)
|
||
inc bot.bulletCount
|
||
let aW = float64(getArenaWidth())
|
||
let aH = float64(getArenaHeight())
|
||
var alive = 0
|
||
for i in 0 ..< bot.bulletCount:
|
||
let b = bot.bullets[i]
|
||
let nx = b.x + b.vx
|
||
let ny = b.y + b.vy
|
||
if nx >= 0.0 and nx <= aW and ny >= 0.0 and ny <= aH:
|
||
bot.bullets[alive] = InFlightBullet(x: nx, y: ny, vx: b.vx, vy: b.vy, power: b.power)
|
||
inc alive
|
||
bot.bulletCount = alive
|
||
|
||
# State build (35-dim, this thread's tensor).
|
||
inc bot.ticksSinceScan
|
||
var gs: GameState
|
||
gs.x = getX()
|
||
gs.y = getY()
|
||
gs.direction = getDirection()
|
||
gs.speed = getSpeed()
|
||
gs.energy = getEnergy()
|
||
gs.gunDirection = getGunDirection()
|
||
gs.gunHeat = getGunHeat()
|
||
gs.arenaWidth = aW
|
||
gs.arenaHeight = aH
|
||
gs.hasContact = bot.hasContact
|
||
gs.enemy = bot.enemy
|
||
gs.ticksSinceLastScan = bot.ticksSinceScan
|
||
var bd: array[MaxBotBullets, BulletData]
|
||
for i in 0 ..< bot.bulletCount:
|
||
bd[i] = BulletData(x: bot.bullets[i].x, y: bot.bullets[i].y, power: bot.bullets[i].power)
|
||
if bot.bulletCount > 1: # closest threats fill slots 0-2
|
||
bd.toOpenArray(0, bot.bulletCount - 1).sort(proc(a, b: BulletData): int =
|
||
cmp(hypot(a.x - gs.x, a.y - gs.y), hypot(b.x - gs.x, b.y - gs.y)))
|
||
gs.bulletCount = min(bot.bulletCount, 3)
|
||
for i in 0 ..< gs.bulletCount:
|
||
gs.bullets[i] = bd[i]
|
||
# Consume the fired pulse AFTER the state saw it (one shot -> one bullet).
|
||
bot.enemy.hasFired = false
|
||
|
||
let stateT = buildState(gs)
|
||
|
||
# Finalize the PREVIOUS transition: reward from events since the last tick,
|
||
# nextState is this tick's observation (PPO_Bot alignment).
|
||
if actorReady and bot.hasLastTrans:
|
||
discard sendTrainingMsg(TrainingMsg(kind: tmkTransition,
|
||
state: bot.lastState, action: bot.lastAction,
|
||
reward: bot.takeReward(), nextState: stateToArr(stateT), done: false))
|
||
bot.lastState = stateToArr(stateT)
|
||
|
||
if not actorReady:
|
||
setRadarTurnRate(45.0) # no weights yet (defensive; main pre-inits) — just sweep
|
||
go()
|
||
continue
|
||
|
||
# Inference: hidden state lives as plain arrays on the bot object (persists
|
||
# across rounds); tensors are rebuilt per tick on this thread.
|
||
let h = hiddenToTensor(bot.hArr, myHidden)
|
||
let c = hiddenToTensor(bot.cArr, myHidden)
|
||
let fwd = actor.actorForward(stateT, (h: h, c: c))
|
||
tensorToHidden(fwd.lstm.h, bot.hArr)
|
||
tensorToHidden(fwd.lstm.c, bot.cArr)
|
||
|
||
bot.lastAction = actionToArr(fwd.actions)
|
||
bot.hasLastTrans = true
|
||
|
||
# Actions -> intents (go() snapshots them at send time).
|
||
let mapped = mapActions(fwd.actions, getSpeed(), getGunHeat())
|
||
setTurnRate(mapped.turnRate)
|
||
setTargetSpeed(getSpeed() + mapped.acceleration) # actions.nim contract
|
||
setGunTurnRate(mapped.gunTurnRate)
|
||
if mapped.firePower > 0.0:
|
||
discard setFire(mapped.firePower)
|
||
if not bot.hasContact:
|
||
setRadarTurnRate(45.0) # sweep until first lock (onScannedBot overrides same-tick)
|
||
|
||
go()
|
||
|
||
# ── Entry point ───────────────────────────────────────────────────────────────
|
||
|
||
when isMainModule:
|
||
initIntegration() # spawn training + I/O threads, seed weight snapshot
|
||
var bot = SacBot()
|
||
start(bot, botJsonPath) # blocks until server disconnect
|
||
shutdownIntegration() # Shutdown msg -> final save -> joins
|