feat(SAC_LSTM_Bot): main bot integration (#48)
This commit is contained in:
@@ -5,6 +5,9 @@ switch("passL", "-L/nix/store/v07svn2y92bvzjl51aj7c9ca1cwg7rw7-openblas-0.3.32/l
|
||||
# ponytail: nix store path; adjust per machine, or use pkg-config
|
||||
switch("passL", "-L/nix/store/wqvz31s598bvj3zb747943xhl38hjc6h-libzip-1.11.4/lib -lzip")
|
||||
switch("threads", "on")
|
||||
# Submodules import each other as SAC_LSTM_Bot/<mod>; make that resolvable for
|
||||
# the binary build too (tests already add ../src via tests/config.nims).
|
||||
switch("path", thisDir() & "/src")
|
||||
# begin Nimble config (version 2)
|
||||
when withDir(thisDir(), system.fileExists("nimble.paths")):
|
||||
include "nimble.paths"
|
||||
|
||||
@@ -1,13 +1,26 @@
|
||||
## SAC_LSTM_Bot — skeleton: radar lock + "Recurrent Royalty" color scheme.
|
||||
## No RL yet. Connects, sets colors, locks radar onto enemy.
|
||||
## 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
|
||||
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
|
||||
|
||||
const botJsonPath = currentSourcePath().parentDir / "SAC_LSTM_Bot.json"
|
||||
|
||||
# ── Colors (Recurrent Royalty palette) ───────────────────────────────────────
|
||||
|
||||
const
|
||||
ColBody = fromHex("#7B2FBE")
|
||||
ColTurret = fromHex("#FFD700")
|
||||
@@ -26,35 +39,262 @@ proc applyColors() =
|
||||
setBulletColor(ColBullet)
|
||||
setTracksColor(ColTracks)
|
||||
|
||||
# ── Bot type ──────────────────────────────────────────────────────────────────
|
||||
# ── 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
|
||||
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: 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).
|
||||
let raw = computeReward(
|
||||
damageInflicted = bot.dmgDealt,
|
||||
damageReceived = bot.dmgTaken,
|
||||
wallHitTicks = bot.wallHits,
|
||||
wastedShotPower = bot.wastedPower,
|
||||
win = win, loss = loss)
|
||||
bot.dmgDealt = 0; bot.dmgTaken = 0; bot.wastedPower = 0; bot.wallHits = 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: one NewBattle per battle, numeric scannedBotId
|
||||
# ponytail: name-based opponent identity deferred to #49 (protocol has no names).
|
||||
if not bot.newBattleSent:
|
||||
bot.newBattleSent = true
|
||||
discard sendTrainingMsg(TrainingMsg(kind: tmkNewBattle, enemyId: e.scannedBotId))
|
||||
|
||||
# ── Run loop ──────────────────────────────────────────────────────────────────
|
||||
method onBulletHit*(bot: SacBot, e: BulletHitBotEvent) =
|
||||
bot.dmgDealt += e.damage
|
||||
|
||||
method onHitByBullet*(bot: SacBot, e: HitByBulletEvent) =
|
||||
bot.dmgTaken += e.damage
|
||||
|
||||
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
|
||||
|
||||
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():
|
||||
# Spin radar when no enemy is visible (full sweep).
|
||||
if bot.enemyBearing == 0.0:
|
||||
setRadarTurnRate(45.0)
|
||||
# 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):
|
||||
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)
|
||||
start(bot, botJsonPath) # blocks until server disconnect
|
||||
shutdownIntegration() # Shutdown msg -> final save -> joins
|
||||
|
||||
@@ -0,0 +1,369 @@
|
||||
## integration.nim — #48 thread plumbing for SAC_LSTM_Bot.
|
||||
##
|
||||
## Threads added by the bot (on top of the bot API's main/bot/sender):
|
||||
## training thread — permanent, drain-then-train loop, owns SACTrainer +
|
||||
## ReplayBuffer. Q1/Q2/Q10/Q12.
|
||||
## I/O thread — cap-1 channel of weight snapshots, atomic zip saves. Q5.
|
||||
##
|
||||
## Cross-thread payloads are plain arrays/seqs ONLY. No Arraymancer tensor ever
|
||||
## crosses a thread boundary: Tensor is a ref type and ORC refcounts are
|
||||
## non-atomic — sharing them across threads is SIGSEGV territory (PPO_Bot,
|
||||
## empirically confirmed). Each thread builds its own tensors from plain data.
|
||||
## Decisions Q1–Q14: Gitea #48.
|
||||
|
||||
import arraymancer except Linear
|
||||
import std/[locks, os, math, random, strutils]
|
||||
import SAC_LSTM_Bot/network
|
||||
import SAC_LSTM_Bot/state # STATE_DIM
|
||||
import SAC_LSTM_Bot/actions # ACTION_DIM
|
||||
import SAC_LSTM_Bot/replay_buffer
|
||||
import SAC_LSTM_Bot/training
|
||||
import SAC_LSTM_Bot/weights
|
||||
|
||||
# ── Config ────────────────────────────────────────────────────────────────────
|
||||
|
||||
proc getUtdRatio*(): int =
|
||||
## Gradient steps per drained transition (Q2). 200 ticks -> 200 steps at 1.
|
||||
parseInt(getEnv("SACLSTM_UTD_RATIO", "1"))
|
||||
|
||||
proc getBatchSize*(): int =
|
||||
# ponytail: 16 is an unprofiled guess sized so UTD=1 keeps up with 30tps;
|
||||
# env knob is the calibration point if steps/sec falls short.
|
||||
parseInt(getEnv("SACLSTM_BATCH_SIZE", "16"))
|
||||
|
||||
proc getSaveInterval*(): int =
|
||||
parseInt(getEnv("SACLSTM_SAVE_INTERVAL", "500"))
|
||||
|
||||
proc getWeightsPath*(): string =
|
||||
getEnv("SACLSTM_WEIGHTS_PATH",
|
||||
currentSourcePath().parentDir / "weights" / "sac_latest.zip")
|
||||
|
||||
# ── Channel message (plain data only) ────────────────────────────────────────
|
||||
|
||||
type
|
||||
TrainingMsgKind* = enum tmkTransition, tmkNewBattle, tmkShutdown
|
||||
|
||||
TrainingMsg* = object
|
||||
case kind*: TrainingMsgKind
|
||||
of tmkTransition:
|
||||
state*: array[STATE_DIM, float32] # Q13 plain arrays, no tensors
|
||||
action*: array[ACTION_DIM, float32]
|
||||
reward*: float32
|
||||
nextState*: array[STATE_DIM, float32]
|
||||
done*: bool # true only at battle end
|
||||
of tmkNewBattle:
|
||||
enemyId*: int # Q14 numeric scannedBotId
|
||||
of tmkShutdown:
|
||||
discard
|
||||
|
||||
proc arrToTensor*[N: static int](arr: array[N, float32]): Tensor[float32] =
|
||||
result = newTensor[float32](N)
|
||||
for i in 0 ..< N: result[i] = arr[i]
|
||||
|
||||
# ── Flat weight snapshots ─────────────────────────────────────────────────────
|
||||
# Layout mirrors network.nim's fixed architecture (fc1 -> hiddenDim-wide LSTM
|
||||
# with [4h, 2h] combined weights, 128-wide fc2, 4-out heads). The asserts catch
|
||||
# layout drift if network.nim shapes ever change.
|
||||
|
||||
proc actorSize*(h: int): int = 8*h*h + 168*h + 1160
|
||||
proc criticSize*(h: int): int = 8*h*h + 172*h + 257
|
||||
|
||||
proc putT(t: Tensor[float32]; dst: var seq[float32]; c: var int) =
|
||||
for v in t:
|
||||
dst[c] = v
|
||||
inc c
|
||||
|
||||
proc takeT(src: seq[float32]; c: var int; rows, cols: int): Tensor[float32] =
|
||||
# seq slice copies, then toTensor copies again: result owns its memory —
|
||||
# never a view into src (src may be a cross-thread buffer).
|
||||
let n = rows * cols
|
||||
result = src[c ..< c + n].toTensor().reshape(rows, cols)
|
||||
c += n
|
||||
|
||||
proc takeV(src: seq[float32]; c: var int; n: int): Tensor[float32] =
|
||||
## Rank-1 vector (biases) — reshape(n) keeps rank 1.
|
||||
result = src[c ..< c + n].toTensor().reshape(n)
|
||||
c += n
|
||||
|
||||
proc packActor*(a: ActorNet; dst: var seq[float32]; c: var int) =
|
||||
putT(a.fc1.w, dst, c); putT(a.fc1.b, dst, c)
|
||||
putT(a.lstm.wCombined, dst, c); putT(a.lstm.bCombined, dst, c)
|
||||
putT(a.fc2.w, dst, c); putT(a.fc2.b, dst, c)
|
||||
putT(a.muHead.w, dst, c); putT(a.muHead.b, dst, c)
|
||||
putT(a.logStdHead.w, dst, c); putT(a.logStdHead.b, dst, c)
|
||||
|
||||
proc packCritic*(net: CriticNet; dst: var seq[float32]; c: var int) =
|
||||
putT(net.fc1.w, dst, c); putT(net.fc1.b, dst, c)
|
||||
putT(net.lstm.wCombined, dst, c); putT(net.lstm.bCombined, dst, c)
|
||||
putT(net.fc2.w, dst, c); putT(net.fc2.b, dst, c)
|
||||
putT(net.fc3.w, dst, c); putT(net.fc3.b, dst, c)
|
||||
|
||||
proc actorFromFlat*(src: seq[float32]; c: var int; h: int): ActorNet =
|
||||
result.fc1.w = takeT(src, c, h, 35)
|
||||
result.fc1.b = takeV(src, c, h)
|
||||
result.lstm.wCombined = takeT(src, c, 4*h, 2*h)
|
||||
result.lstm.bCombined = takeV(src, c, 4*h)
|
||||
result.fc2.w = takeT(src, c, 128, h)
|
||||
result.fc2.b = takeV(src, c, 128)
|
||||
result.muHead.w = takeT(src, c, 4, 128)
|
||||
result.muHead.b = takeV(src, c, 4)
|
||||
result.logStdHead.w = takeT(src, c, 4, 128)
|
||||
result.logStdHead.b = takeV(src, c, 4)
|
||||
result.lstm.hiddenDim = result.lstm.bCombined.size div 4 # same as weights.nim loadLSTMCell
|
||||
result.hiddenDim = h
|
||||
assert c == actorSize(h), "actor flat layout drift"
|
||||
|
||||
proc criticFromFlat*(src: seq[float32]; c: var int; h: int): CriticNet =
|
||||
result.fc1.w = takeT(src, c, h, 39)
|
||||
result.fc1.b = takeV(src, c, h)
|
||||
result.lstm.wCombined = takeT(src, c, 4*h, 2*h)
|
||||
result.lstm.bCombined = takeV(src, c, 4*h)
|
||||
result.fc2.w = takeT(src, c, 128, h)
|
||||
result.fc2.b = takeV(src, c, 128)
|
||||
result.fc3.w = takeT(src, c, 1, 128)
|
||||
result.fc3.b = takeV(src, c, 1)
|
||||
result.lstm.hiddenDim = result.lstm.bCombined.size div 4
|
||||
result.hiddenDim = h
|
||||
|
||||
type
|
||||
FullSnap* = object
|
||||
hiddenDim*: int
|
||||
data*: seq[float32] # actor | critic1 | critic2 | targetCritic1 | targetCritic2 | alpha
|
||||
|
||||
proc packFull*(t: SACTrainer): FullSnap =
|
||||
let h = t.actor.hiddenDim
|
||||
result.hiddenDim = h
|
||||
result.data = newSeq[float32](actorSize(h) + 4 * criticSize(h) + 1)
|
||||
var c = 0
|
||||
packActor(t.actor, result.data, c)
|
||||
packCritic(t.critic1, result.data, c)
|
||||
packCritic(t.critic2, result.data, c)
|
||||
packCritic(t.targetCritic1, result.data, c)
|
||||
packCritic(t.targetCritic2, result.data, c)
|
||||
assert abs(t.alpha().float64 - exp(t.logAlpha.float64)) < 1e-6
|
||||
result.data[c] = t.alpha()
|
||||
inc c
|
||||
assert c == result.data.len, "full snapshot layout drift"
|
||||
|
||||
proc unpackFull*(fs: FullSnap):
|
||||
tuple[a: ActorNet, c1, c2, t1, t2: CriticNet, alpha: float32] =
|
||||
var c = 0
|
||||
result.a = actorFromFlat(fs.data, c, fs.hiddenDim)
|
||||
result.c1 = criticFromFlat(fs.data, c, fs.hiddenDim)
|
||||
result.c2 = criticFromFlat(fs.data, c, fs.hiddenDim)
|
||||
result.t1 = criticFromFlat(fs.data, c, fs.hiddenDim)
|
||||
result.t2 = criticFromFlat(fs.data, c, fs.hiddenDim)
|
||||
result.alpha = fs.data[c]
|
||||
|
||||
# ── Shared weight snapshot (training thread writes, bot thread copies out) ────
|
||||
# Q7/Q11: Lock + always-latest semantics. `data` is allocated ONCE and written
|
||||
# IN PLACE under gWeightLock — it is never reassigned, so the shared heap block
|
||||
# never sees cross-thread refcount traffic. Readers copy element-wise out.
|
||||
|
||||
type
|
||||
WeightSnapshot* = object
|
||||
hiddenDim*: int
|
||||
version*: int
|
||||
data*: seq[float32]
|
||||
|
||||
var gWeightLock: Lock
|
||||
var gSharedSnap: WeightSnapshot
|
||||
var gTrainChan: Channel[TrainingMsg]
|
||||
var gSaveChan: Channel[FullSnap]
|
||||
var gTrainingThread: Thread[void]
|
||||
var gIoThread: Thread[void]
|
||||
var gInitialFull: FullSnap # built on main before spawn; read-once after (happens-before)
|
||||
var gWeightsPath: string
|
||||
|
||||
proc pullWeights*(myVersion: var int; hidden: var int;
|
||||
flat: var seq[float32]): bool =
|
||||
## Copy the latest actor snapshot out under the lock. Returns true when a new
|
||||
## version arrived (caller rebuilds its tensors on ITS OWN thread).
|
||||
withLock(gWeightLock):
|
||||
if gSharedSnap.version == myVersion:
|
||||
return false
|
||||
if flat.len != gSharedSnap.data.len:
|
||||
flat = newSeq[float32](gSharedSnap.data.len) # caller-thread-owned buffer
|
||||
for i in 0 ..< flat.len:
|
||||
flat[i] = gSharedSnap.data[i]
|
||||
myVersion = gSharedSnap.version
|
||||
hidden = gSharedSnap.hiddenDim
|
||||
true
|
||||
|
||||
proc sendTrainingMsg*(msg: TrainingMsg): bool {.inline.} =
|
||||
## Bot-side enqueue (cap-256, drops on overflow per Q10). Thread-safe.
|
||||
gTrainChan.trySend(msg)
|
||||
|
||||
# ── Training state (testable without threads) ─────────────────────────────────
|
||||
|
||||
type
|
||||
TrainState* = object
|
||||
trainer*: SACTrainer
|
||||
buf*: ReplayBuffer
|
||||
lastEnemyId*: int
|
||||
stepCount*: int
|
||||
nextSave*: int
|
||||
|
||||
proc trainerFromFull*(initial: FullSnap): SACTrainer =
|
||||
let (a, c1, c2, t1, t2, alpha) = unpackFull(initial)
|
||||
result.actor = a
|
||||
result.critic1 = c1
|
||||
result.critic2 = c2
|
||||
result.targetCritic1 = t1
|
||||
result.targetCritic2 = t2
|
||||
assert alpha > 0.0'f32, "checkpoint alpha must be positive"
|
||||
result.logAlpha = ln(alpha)
|
||||
result.targetEntropy = getTargetEntropy()
|
||||
result.tau = getTau()
|
||||
result.lrActor = getLrActor()
|
||||
result.lrCritic = getLrCritic()
|
||||
result.lrAlpha = getLrAlpha()
|
||||
result.gamma = getGamma()
|
||||
result.adam = initSACAdamStates(result.actor, result.critic1, result.critic2)
|
||||
# ponytail: adam momentum not carried in the flat snapshot — optimizer restarts
|
||||
# fresh each process; switch to saveCheckpoint/loadCheckpoint end-to-end when
|
||||
# resume quality matters (#49 harness owns checkpoint management).
|
||||
|
||||
proc initTrainState*(initial: FullSnap): TrainState =
|
||||
result.trainer = trainerFromFull(initial)
|
||||
result.buf = newReplayBuffer(getBufferCapacity(), STATE_DIM, ACTION_DIM)
|
||||
result.lastEnemyId = -1
|
||||
result.nextSave = getSaveInterval()
|
||||
|
||||
proc handleTrainingMsg*(st: var TrainState; msg: TrainingMsg): bool =
|
||||
## Process one message. Returns false for Shutdown (caller stops).
|
||||
## Tensors are born HERE from the message's plain arrays — training thread only.
|
||||
case msg.kind
|
||||
of tmkTransition:
|
||||
st.buf.add(Transition(
|
||||
state: msg.state.arrToTensor,
|
||||
action: msg.action.arrToTensor,
|
||||
reward: msg.reward,
|
||||
nextState: msg.nextState.arrToTensor,
|
||||
done: msg.done))
|
||||
of tmkNewBattle:
|
||||
# Q12a: keep the buffer if the opponent is unchanged, clear otherwise.
|
||||
if msg.enemyId != st.lastEnemyId:
|
||||
st.buf.clear()
|
||||
st.lastEnemyId = msg.enemyId
|
||||
of tmkShutdown:
|
||||
return false
|
||||
true
|
||||
|
||||
proc trainPass*(st: var TrainState; drained: int) =
|
||||
## UTD gradient steps for the transitions drained this pass, then publish the
|
||||
## latest actor to the shared snapshot and request periodic disk saves.
|
||||
if drained <= 0 or not st.buf.canSample:
|
||||
return
|
||||
let steps = drained * getUtdRatio() # Q2
|
||||
for i in 1 .. steps:
|
||||
let seqs = st.buf.sampleSequences(getBatchSize())
|
||||
if seqs.len == 0:
|
||||
break
|
||||
discard sacUpdate(st.trainer, seqs)
|
||||
inc st.stepCount
|
||||
# Publish latest actor (Q7): in-place write under the lock, bump version.
|
||||
withLock(gWeightLock):
|
||||
assert gSharedSnap.hiddenDim == st.trainer.actor.hiddenDim,
|
||||
"snapshot/trainer hidden size mismatch"
|
||||
var c = 0
|
||||
packActor(st.trainer.actor, gSharedSnap.data, c)
|
||||
inc gSharedSnap.version
|
||||
if st.stepCount >= st.nextSave:
|
||||
st.nextSave += getSaveInterval()
|
||||
var full = packFull(st.trainer)
|
||||
discard gSaveChan.trySend(move(full)) # cap-1: drop if I/O thread is busy (Q5)
|
||||
|
||||
# ── Threads ───────────────────────────────────────────────────────────────────
|
||||
|
||||
proc trainingThreadEntry() {.thread.} =
|
||||
{.cast(gcsafe).}:
|
||||
randomize()
|
||||
var st = initTrainState(gInitialFull)
|
||||
var running = true
|
||||
while running:
|
||||
let first = gTrainChan.recv() # block until traffic (no busy spin)
|
||||
var drained = 0
|
||||
var msg = first
|
||||
while running:
|
||||
if msg.kind == tmkTransition:
|
||||
inc drained
|
||||
if not handleTrainingMsg(st, msg):
|
||||
running = false # Shutdown
|
||||
break
|
||||
let (more, nxt) = gTrainChan.tryRecv()
|
||||
if not more:
|
||||
break # drained — now train (Q10)
|
||||
msg = nxt
|
||||
if not running:
|
||||
var full = packFull(st.trainer) # final save request, then exit
|
||||
discard gSaveChan.trySend(move(full))
|
||||
break
|
||||
trainPass(st, drained)
|
||||
|
||||
proc ioThreadEntry() {.thread.} =
|
||||
{.cast(gcsafe).}:
|
||||
while true:
|
||||
let fs = gSaveChan.recv() # blocks; exits via empty-data sentinel
|
||||
if fs.data.len == 0:
|
||||
break
|
||||
let (a, c1, c2, t1, t2, alpha) = unpackFull(fs)
|
||||
saveWeights(gWeightsPath, a, c1, c2, t1, t2, alpha) # atomic zip (weights.nim)
|
||||
|
||||
# ── Lifecycle ─────────────────────────────────────────────────────────────────
|
||||
|
||||
proc randomFull(): FullSnap =
|
||||
let t = initSACTrainer(STATE_DIM, ACTION_DIM) # random nets, born + freed here
|
||||
packFull(t)
|
||||
|
||||
proc loadOrInitFull(): FullSnap =
|
||||
let path = getWeightsPath()
|
||||
if fileExists(path):
|
||||
try:
|
||||
let cp = loadCheckpoint(path)
|
||||
var t = initSACTrainer(STATE_DIM, ACTION_DIM) # env hyperparams
|
||||
t.actor = cp.actor
|
||||
t.critic1 = cp.critic1
|
||||
t.critic2 = cp.critic2
|
||||
t.targetCritic1 = cp.targetCritic1
|
||||
t.targetCritic2 = cp.targetCritic2
|
||||
assert cp.alpha > 0.0'f32, "checkpoint alpha must be positive"
|
||||
t.logAlpha = ln(cp.alpha)
|
||||
result = packFull(t)
|
||||
except Exception as e:
|
||||
stderr.writeLine "[sac] checkpoint load failed (" & e.msg & ") — random init"
|
||||
result = randomFull()
|
||||
else:
|
||||
result = randomFull()
|
||||
|
||||
proc initIntegration*() =
|
||||
## Open channels, build the initial weight snapshot, spawn both threads.
|
||||
## Call once from the main module before start().
|
||||
gWeightsPath = getWeightsPath()
|
||||
gInitialFull = loadOrInitFull()
|
||||
gSharedSnap.hiddenDim = gInitialFull.hiddenDim
|
||||
gSharedSnap.data = newSeq[float32](actorSize(gInitialFull.hiddenDim))
|
||||
for i in 0 ..< gSharedSnap.data.len: # element-wise: no refcount traffic
|
||||
gSharedSnap.data[i] = gInitialFull.data[i]
|
||||
gSharedSnap.version = 1
|
||||
initLock(gWeightLock)
|
||||
gTrainChan.open(256) # Q10 cap-256
|
||||
gSaveChan.open(1) # Q5 cap-1
|
||||
createThread(gTrainingThread, trainingThreadEntry)
|
||||
createThread(gIoThread, ioThreadEntry)
|
||||
|
||||
proc shutdownIntegration*() =
|
||||
## Stop both threads cleanly. Called after the bot disconnects (start returned).
|
||||
while gTrainChan.tryRecv().dataAvailable:
|
||||
discard # drop pending transitions — process exiting
|
||||
discard gTrainChan.trySend(TrainingMsg(kind: tmkShutdown))
|
||||
joinThread(gTrainingThread) # training requests one final save
|
||||
# Nim's recv() blocks even on closed channels, so the I/O thread exits via an
|
||||
# empty-data sentinel. Retry while it is still busy saving: every failed
|
||||
# trySend means a save is in flight and will be consumed, so this terminates.
|
||||
var stop = FullSnap(hiddenDim: -1)
|
||||
while not gSaveChan.trySend(move(stop)):
|
||||
sleep(50)
|
||||
stop = FullSnap(hiddenDim: -1)
|
||||
joinThread(gIoThread)
|
||||
gTrainChan.close()
|
||||
@@ -71,6 +71,12 @@ proc add*(buf: var ReplayBuffer; t: Transition) =
|
||||
|
||||
proc len*(buf: ReplayBuffer): int = buf.count
|
||||
|
||||
proc clear*(buf: var ReplayBuffer) =
|
||||
## Drop all transitions (#48: opponent changed across battles).
|
||||
## Old slots keep stale tensors until the ring overwrites them.
|
||||
buf.head = 0
|
||||
buf.count = 0
|
||||
|
||||
proc canSample*(buf: ReplayBuffer): bool =
|
||||
buf.count >= buf.burnIn + buf.trainWindow
|
||||
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
## Tests for integration.nim (#48) — assert-based, no framework.
|
||||
## Covers: TrainingMsg channel round-trip (plain arrays through a channel),
|
||||
## drain-then-train NewBattle/Shutdown handling, trainPass safety below canSample.
|
||||
|
||||
import arraymancer except Linear
|
||||
import std/locks
|
||||
import SAC_LSTM_Bot/integration
|
||||
import SAC_LSTM_Bot/state # STATE_DIM
|
||||
import SAC_LSTM_Bot/actions # ACTION_DIM
|
||||
import SAC_LSTM_Bot/training # initSACTrainer
|
||||
import SAC_LSTM_Bot/replay_buffer
|
||||
|
||||
# ── 1. TrainingMsg round-trips through a channel with arrays intact ──────────
|
||||
block:
|
||||
var ch: Channel[TrainingMsg]
|
||||
ch.open(4)
|
||||
var msg = TrainingMsg(kind: tmkTransition)
|
||||
for i in 0 ..< STATE_DIM:
|
||||
msg.state[i] = float32(i) * 0.5'f32
|
||||
msg.nextState[i] = float32(i) * 2.0'f32
|
||||
for i in 0 ..< ACTION_DIM:
|
||||
msg.action[i] = float32(i) - 2.0'f32
|
||||
msg.reward = -1.25'f32
|
||||
msg.done = true
|
||||
assert ch.trySend(msg)
|
||||
assert ch.trySend(TrainingMsg(kind: tmkNewBattle, enemyId: 4242))
|
||||
assert ch.trySend(TrainingMsg(kind: tmkShutdown))
|
||||
|
||||
let r1 = ch.recv()
|
||||
assert r1.kind == tmkTransition, "first msg is a transition"
|
||||
for i in 0 ..< STATE_DIM:
|
||||
assert r1.state[i] == float32(i) * 0.5'f32, "state round-trip at " & $i
|
||||
assert r1.nextState[i] == float32(i) * 2.0'f32, "nextState round-trip at " & $i
|
||||
for i in 0 ..< ACTION_DIM:
|
||||
assert r1.action[i] == float32(i) - 2.0'f32, "action round-trip at " & $i
|
||||
assert r1.reward == -1.25'f32 and r1.done
|
||||
|
||||
let r2 = ch.recv()
|
||||
assert r2.kind == tmkNewBattle and r2.enemyId == 4242
|
||||
let r3 = ch.recv()
|
||||
assert r3.kind == tmkShutdown
|
||||
|
||||
ch.close()
|
||||
# closed + empty -> tryRecv reports no data (Nim 2.2: recv would block forever)
|
||||
let (ok4, _) = ch.tryRecv()
|
||||
assert not ok4, "closed channel must report dataAvailable=false"
|
||||
echo "PASS TrainingMsg channel round-trip"
|
||||
|
||||
# ── 2. NewBattle clears only when the opponent changes; Shutdown stops ────────
|
||||
block:
|
||||
var st: TrainState
|
||||
st.trainer = initSACTrainer(STATE_DIM, ACTION_DIM)
|
||||
st.buf = newReplayBuffer(64, STATE_DIM, ACTION_DIM, burnIn = 2, trainWindow = 3)
|
||||
|
||||
assert handleTrainingMsg(st, TrainingMsg(kind: tmkNewBattle, enemyId: 7))
|
||||
assert st.lastEnemyId == 7
|
||||
|
||||
for i in 0 ..< 5:
|
||||
var m = TrainingMsg(kind: tmkTransition)
|
||||
m.reward = float32(i)
|
||||
assert handleTrainingMsg(st, m)
|
||||
assert st.buf.len == 5, "transitions stored"
|
||||
|
||||
# Same opponent -> buffer kept (Q12a).
|
||||
assert handleTrainingMsg(st, TrainingMsg(kind: tmkNewBattle, enemyId: 7))
|
||||
assert st.buf.len == 5, "same opponent must NOT clear"
|
||||
|
||||
# Opponent changed -> clear.
|
||||
assert handleTrainingMsg(st, TrainingMsg(kind: tmkNewBattle, enemyId: 8))
|
||||
assert st.buf.len == 0, "opponent change must clear"
|
||||
assert st.lastEnemyId == 8
|
||||
|
||||
# Shutdown stops the caller's loop.
|
||||
assert not handleTrainingMsg(st, TrainingMsg(kind: tmkShutdown))
|
||||
echo "PASS drain-then-train NewBattle/Shutdown"
|
||||
|
||||
# ── 3. trainPass is a safe no-op below canSample (no steps, no publish) ───────
|
||||
block:
|
||||
var st: TrainState
|
||||
st.trainer = initSACTrainer(STATE_DIM, ACTION_DIM)
|
||||
st.buf = newReplayBuffer(64, STATE_DIM, ACTION_DIM, burnIn = 2, trainWindow = 3)
|
||||
for i in 0 ..< 4:
|
||||
assert handleTrainingMsg(st, TrainingMsg(kind: tmkTransition))
|
||||
trainPass(st, 4) # 4 < burnIn+trainWindow = 5
|
||||
assert st.stepCount == 0, "no gradient steps below canSample"
|
||||
echo "PASS trainPass no-op below canSample"
|
||||
|
||||
# ── 4. Flat snapshot layout: pack/unpack round-trips weights exactly ──────────
|
||||
block:
|
||||
let t0 = initSACTrainer(STATE_DIM, ACTION_DIM)
|
||||
let fs = packFull(t0)
|
||||
assert fs.data.len == actorSize(t0.actor.hiddenDim) + 4 * criticSize(t0.actor.hiddenDim) + 1
|
||||
let (a, c1, c2, tc1, tc2, alpha) = unpackFull(fs)
|
||||
assert a.hiddenDim == t0.actor.hiddenDim
|
||||
assert c1.fc3.b.shape[0] == 1
|
||||
let fw = a.muHead.w.flatten()
|
||||
let fw0 = t0.actor.muHead.w.flatten()
|
||||
for i in 0 ..< fw.size:
|
||||
assert fw[i] == fw0[i], "actor mu weights round-trip"
|
||||
for i in 0 ..< c2.lstm.bCombined.size:
|
||||
assert c2.lstm.bCombined[i] == t0.critic2.lstm.bCombined[i], "critic lstm bias round-trip"
|
||||
assert alpha == t0.alpha()
|
||||
discard tc1
|
||||
echo "PASS flat snapshot pack/unpack round-trip"
|
||||
Reference in New Issue
Block a user