diff --git a/SAC_LSTM_Bot/config.nims b/SAC_LSTM_Bot/config.nims index 29e2072..7bcea66 100644 --- a/SAC_LSTM_Bot/config.nims +++ b/SAC_LSTM_Bot/config.nims @@ -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/; 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" diff --git a/SAC_LSTM_Bot/src/SAC_LSTM_Bot.nim b/SAC_LSTM_Bot/src/SAC_LSTM_Bot.nim index 69f3bb8..bb16eac 100644 --- a/SAC_LSTM_Bot/src/SAC_LSTM_Bot.nim +++ b/SAC_LSTM_Bot/src/SAC_LSTM_Bot.nim @@ -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 diff --git a/SAC_LSTM_Bot/src/SAC_LSTM_Bot/integration.nim b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/integration.nim new file mode 100644 index 0000000..e8628ac --- /dev/null +++ b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/integration.nim @@ -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() diff --git a/SAC_LSTM_Bot/src/SAC_LSTM_Bot/replay_buffer.nim b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/replay_buffer.nim index 7672deb..85b68ec 100644 --- a/SAC_LSTM_Bot/src/SAC_LSTM_Bot/replay_buffer.nim +++ b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/replay_buffer.nim @@ -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 diff --git a/SAC_LSTM_Bot/tests/test_integration.nim b/SAC_LSTM_Bot/tests/test_integration.nim new file mode 100644 index 0000000..fc11276 --- /dev/null +++ b/SAC_LSTM_Bot/tests/test_integration.nim @@ -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"