## 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, times] import tankroyale_botapi # getBotName (#49 name-based opponent identity) 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") proc opponentKey*(enemyId: int): string = ## Q14 follow-up (#49): name-based opponent identity from the v1.0.1 ## BotListUpdate table; numeric-id fallback for the window before the first ## update arrives (getBotName still returns "" then). Called on the training ## thread — the lookup is lock-guarded in the API, no cross-thread refs. let name = getBotName(enemyId) if name.len > 0: name else: $enemyId proc bumpRoundCounter*() = ## Liveness signal for tools/training_runner/RunTraining.java (#49): one ## increment per round end; the runner aborts when it freezes (dead bot). # ponytail: non-atomic read-modify-write; single writer (main thread) and the # runner re-polls every 500ms with multi-round tolerance, torn reads self-heal. let p = getWeightsPath().parentDir / "round_counter.txt" var n = 0 try: n = parseInt(readFile(p).strip()) except CatchableError: discard # absent/garbage -> start at 1 try: createDir(p.parentDir) writeFile(p, $(n + 1)) except CatchableError: discard # counter is best-effort liveness; never kill an event handler # ── 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 evalModeActive*(): bool {.inline.} = ## Lever 4 (#59): the harness's deterministic eval battles already run the bot ## with SACLSTM_EVAL_MODE=1 (sac_train.sh eval_checkpoint, mechanism from #49). ## While set, eval ticks must NOT feed the trainer — transitions would pollute ## the replay buffer with eval-only data and trigger gradient updates. getEnv("SACLSTM_EVAL_MODE") == "1" proc sendTrainingMsg*(msg: TrainingMsg): bool {.inline.} = ## Bot-side enqueue (cap-256, drops on overflow per Q10). Thread-safe. ## Lever 4 (#59): fully suppressed in eval mode — NewBattle drops too, so an ## eval battle can neither add transitions nor clear/retarget the buffer. if evalModeActive(): return false gTrainChan.trySend(msg) # ── Training state (testable without threads) ───────────────────────────────── type TrainState* = object trainer*: SACTrainer buf*: ReplayBuffer lastEnemyKey*: string # opponent identity key (#49 name-based, Q14) 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.lastEnemyKey = "" result.nextSave = getSaveInterval() # ── Training-loss metrics (campaign v2 lever 3, #59) ────────────────────────── proc metricsFilePath*(): string = ## Sits next to the weights dir's parent: SAC_LSTM_Bot/training_metrics.jsonl ## under the #49 harness (weights live in SAC_LSTM_Bot/weights/). getWeightsPath().parentDir.parentDir / "training_metrics.jsonl" proc metricsLine*(epoch: float64; stepCount, bufferLen, drained, gradSteps: int; m: SACMetrics): string = ## One JSONL line with exactly the scalars SACTrainer.sacUpdate exposes ## (#59 lever 3 — SACMetrics was already returned, no trainer change needed): ## losses/alpha averaged over this pass's gradient steps, buffer size from ## replay_buffer.len, cumulative step count and drained transition count. "{\"epoch\":" & $epoch & ",\"steps\":" & $stepCount & ",\"buffer_size\":" & $bufferLen & ",\"drained\":" & $drained & ",\"grad_steps\":" & $gradSteps & ",\"critic_loss\":" & $m.criticLoss & ",\"actor_loss\":" & $m.actorLoss & ",\"alpha_loss\":" & $m.alphaLoss & ",\"alpha\":" & $m.alpha & "}" proc appendMetricsLine(st: TrainState; drained, gradSteps: int; m: SACMetrics) = ## Lever 3 (#59): one append per trainPass (never per gradient step). Open, ## write, close — cheap and crash-tolerant; a metrics failure never kills ## training. try: let f = open(metricsFilePath(), fmAppend) f.writeLine(metricsLine(epochTime(), st.stepCount, st.buf.len, drained, gradSteps, m)) f.close() except CatchableError: discard 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. # Identity keyed on NAME (#49); numeric-id fallback pre-BotListUpdate. let key = opponentKey(msg.enemyId) if key != st.lastEnemyKey: st.buf.clear() st.lastEnemyKey = key 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 var gradSteps = 0 var sumCritic, sumActor, sumAlphaLoss, sumAlpha = 0.0'f32 for i in 1 .. steps: let seqs = st.buf.sampleSequences(getBatchSize()) if seqs.len == 0: break let m = sacUpdate(st.trainer, seqs) sumCritic += m.criticLoss; sumActor += m.actorLoss sumAlphaLoss += m.alphaLoss; sumAlpha += m.alpha inc gradSteps inc st.stepCount # Save check INSIDE the step loop (#56 launch finding): at production sizes # (hidden 256 ⇒ ~1 s/step) a drain burst queues minutes of steps; checking # only between passes meant the process died mid-loop before stepCount ever # reached nextSave — zero checkpoints persisted for the whole campaign. # Mid-loop checks + SAVE_INTERVAL≤20 (#54) keep saves ~20 s apart. 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) if gradSteps > 0: # Lever 3 (#59): one metrics line per pass, losses averaged over its steps. appendMetricsLine(st, drained, gradSteps, SACMetrics( criticLoss: sumCritic / gradSteps.float32, actorLoss: sumActor / gradSteps.float32, alphaLoss: sumAlphaLoss / gradSteps.float32, alpha: sumAlpha / gradSteps.float32)) # 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 # ── 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 # Lever 4 (#59): one-time visibility for the suppression gate (see # sendTrainingMsg) — the eval bot trains nothing by design. if evalModeActive(): stderr.writeLine "[sac] SACLSTM_EVAL_MODE=1 — training input suppressed (lever 4, #59)" 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()