Files
SirRoboGarage/SAC_LSTM_Bot_garage/src/SAC_LSTM_Bot/integration.nim
T

464 lines
20 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.
## 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()