Files
SirRoboGarage/SAC_LSTM_Bot_garage/src/SAC_LSTM_Bot.nim
T

350 lines
15 KiB
Nim
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
## 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