chore: rename libs→common_libs, all bot dirs to _garage suffix, fix all path refs

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-08-27 18:18:41 +02:00
parent f8c0c871c6
commit b509195ee9
832 changed files with 4967 additions and 368 deletions
+11
View File
@@ -0,0 +1,11 @@
{
"name": "SAC_LSTM_Bot",
"version": "0.1.0",
"authors": ["Davide Cappellini"],
"description": "SAC+LSTM Tank Royale bot — self-reported identity; MUST match the name in ../SAC_LSTM_Bot.json (booter identity) or the training runner never sees this bot join",
"homepage": "",
"countryCodes": ["IT"],
"gameTypes": ["classic", "melee", "1v1"],
"platform": "Nim",
"programmingLang": "Nim"
}
+349
View File
@@ -0,0 +1,349 @@
## SAC_LSTM_Bot — Recurrent SAC-v2 bot (Gitea #48).
##
## Thread layout (decisions Q1–Q14, see integration.nim for plumbing):
## bot thread — this file's run(): inference only, <2ms/tick.
## training thread — permanent background SAC updates (integration.nim).
## I/O thread — atomic weight saves (integration.nim).
## No Arraymancer tensor ever crosses a thread boundary: the bot object and
## channels carry plain scalars / fixed arrays / plain seqs only.
import std/[os, math, random, algorithm]
import arraymancer except Linear
import tankroyale_botapi
import radar_lock
import SAC_LSTM_Bot/state
import SAC_LSTM_Bot/network
import SAC_LSTM_Bot/actions
import SAC_LSTM_Bot/rewards
import SAC_LSTM_Bot/integration
# Identity json: baked-in src json by default; SACLSTM_BOT_JSON lets a mirror
# twin (#54) boot the same binary under its own name (loadBotInfo gives the
# json total precedence over env, so the twin must point at its own file).
let botJsonPath = getEnv("SACLSTM_BOT_JSON",
currentSourcePath().parentDir / "SAC_LSTM_Bot.json")
# ── Colors (Recurrent Royalty palette) ───────────────────────────────────────
const
ColBody = fromHex("#7B2FBE")
ColTurret = fromHex("#FFD700")
ColGun = fromHex("#4A0E6B")
ColRadar = fromHex("#FFD700")
ColScan = fromHex("#FFB000")
ColBullet = fromHex("#FFC125")
ColTracks = fromHex("#2C2C34")
proc applyColors() =
setBodyColor(ColBody)
setTurretColor(ColTurret)
setGunColor(ColGun)
setRadarColor(ColRadar)
setScanColor(ColScan)
setBulletColor(ColBullet)
setTracksColor(ColTracks)
# ── Enemy bullet tracking ────────────────────────────────────────────────────
# PPO_Bot-proven pattern: dead-reckoned fixed buffer. getBulletStates() is NOT
# used from the bot thread — its seq refcount is shared with the main thread.
type InFlightBullet = object
x, y, vx, vy, power: float64
const MaxBotBullets = 4
const MaxHidden = 512 # ponytail: cap for the plain-array LSTM persistence; raise if SACLSTM_HIDDEN_SIZE > 512
# ── Bot type — PLAIN DATA ONLY on the shared object (no tensors/heap seqs:
# each round runs a fresh bot thread; heap blocks owned by the previous
# round's thread must not be freed from another thread) ─────────────────────
type SacBot = ref object of Bot
enemyBearing: float # last known absolute bearing to enemy
battleId: int # main thread bumps in onGameStarted; bot thread compares
seenBattle: int # bot-thread copy for battle-change detection
newBattleSent: bool # first scan of THIS battle emits NewBattle
hasContact: bool
enemy: EnemyData
ticksSinceScan: int
# per-step reward accumulators (consumed by the next tick's transition)
dmgDealt, dmgTaken, wastedPower: float64
wallHits, hits, ramTaken: int
# pending transition (episode spans the whole battle; round end is NOT a boundary)
hasLastTrans: bool
lastState: array[STATE_DIM, float32]
lastAction: array[ACTION_DIM, float32]
rn: RewardNormalizer # Welford running stats, persists across battles
bullets: array[MaxBotBullets, InFlightBullet]
bulletCount: int
hArr, cArr: array[MaxHidden, float32] # LSTM state across rounds; zeros at battle start
# ── Plain-array <-> tensor helpers (bot thread only) ──────────────────────────
proc stateToArr(t: Tensor[float32]): array[STATE_DIM, float32] =
for i in 0 ..< STATE_DIM: result[i] = t[i]
proc actionToArr(t: Tensor[float32]): array[ACTION_DIM, float32] =
for i in 0 ..< ACTION_DIM: result[i] = t[i]
proc hiddenToTensor(arr: array[MaxHidden, float32]; n: int): Tensor[float32] =
result = newTensor[float32](n)
for i in 0 ..< n: result[i] = arr[i]
proc tensorToHidden(t: Tensor[float32]; arr: var array[MaxHidden, float32]) =
for i in 0 ..< t.shape[0]: arr[i] = t[i]
# ── Reward ────────────────────────────────────────────────────────────────────
proc takeReward(bot: SacBot; win = false, loss = false): float32 =
## Consume accumulated step events -> Welford-normalized reward (#44).
## Lever 2 (#59): pass enemy distance (frac of arena diagonal) for the
## anti-charge term; sentinel 2.0 (> ChargeDistFrac) when no contact.
var distFrac = 2.0
if bot.hasContact:
let diag = hypot(getArenaWidth().float64, getArenaHeight().float64)
distFrac = hypot(bot.enemy.x - getX(), bot.enemy.y - getY()) / diag
let raw = computeReward(
damageInflicted = bot.dmgDealt,
damageReceived = bot.dmgTaken,
wallHitTicks = bot.wallHits,
wastedShotPower = bot.wastedPower,
hitCount = bot.hits,
ramTakenCount = bot.ramTaken,
enemyDistFrac = distFrac,
win = win, loss = loss)
# Lever-2 (#59) observability: env-gated one-liner for smoke/calibration
# greps — proves hit/ram/charge terms fire and shows raw magnitudes. File
# (not stderr): the battle runner swallows bot process streams.
# ponytail: grows unbounded if left on; keep off outside smokes.
if getEnv("SACLSTM_REWARD_DEBUG") == "1" and
(bot.hits > 0 or bot.ramTaken > 0 or (distFrac < ChargeDistFrac and bot.dmgDealt <= 0.0)):
try:
let f = open(getWeightsPath().parentDir.parentDir / "reward_debug.log", fmAppend)
f.writeLine("raw=" & $raw & " hits=" & $bot.hits & " ram=" & $bot.ramTaken &
" distFrac=" & $distFrac)
f.close()
except CatchableError:
discard
bot.dmgDealt = 0; bot.dmgTaken = 0; bot.wastedPower = 0; bot.wallHits = 0
bot.hits = 0; bot.ramTaken = 0
let norm = rewards.normalize(bot.rn, raw)
rewards.update(bot.rn, raw)
norm.float32
# ── Event handlers ────────────────────────────────────────────────────────────
# onGameStarted/onRoundStarted fire on the MAIN thread (bot thread not yet
# started or already joined) — plain-field writes only, no tensors here.
method onGameStarted*(bot: SacBot, e: GameStartedEventForBot) =
inc bot.battleId # bot thread zeroes LSTM + per-battle flags at next tick (Q4/Q12)
method onRoundStarted*(bot: SacBot, e: RoundStartedEvent) =
setAdjustRadarForBodyTurn(true)
setAdjustRadarForGunTurn(true)
radar_lock.init()
applyColors()
# Per-round reset ONLY. NOT the LSTM hidden state (persists across rounds, Q4);
# NOT hasLastTrans (the pending transition spans the round boundary — episode
# ends at battle end only).
bot.hasContact = false
bot.enemy = EnemyData()
bot.ticksSinceScan = 0
bot.bulletCount = 0
method onScannedBot*(bot: SacBot, e: ScannedBotEvent) =
bot.enemyBearing = directionTo(getX(), getY(), e.x, e.y)
# Same-tick radar lock: apply turn rate immediately so it takes effect this tick.
setRadarTurnRate(radar_lock.doRadar(getRadarDirection(), bot.enemyBearing))
# Fire detection: energy drop in [0.1, 3.0] between scans (PPO_Bot heuristic).
let prevE = if bot.hasContact: bot.enemy.energy else: e.energy
let drop = prevE - e.energy
bot.enemy.hasFired = bot.hasContact and drop >= 0.1 and drop <= 3.0
if bot.enemy.hasFired:
bot.enemy.lastFirePower = drop
# Keep previous-scan deltas before overwriting (state.nim derives accel/turn rate).
bot.enemy.prevSpeed = bot.enemy.speed
bot.enemy.prevDirection = bot.enemy.direction
bot.enemy.hasPrevScan = bot.hasContact
bot.enemy.x = e.x
bot.enemy.y = e.y
bot.enemy.direction = e.direction
bot.enemy.speed = e.speed
bot.enemy.energy = e.energy
bot.hasContact = true
bot.ticksSinceScan = 0
# Q12b/Q14+#49: one NewBattle per battle, numeric scannedBotId. Identity is
# keyed on getBotName(id) training-side (integration.opponentKey), numeric
# fallback in the pre-BotListUpdate window.
if not bot.newBattleSent:
bot.newBattleSent = true
discard sendTrainingMsg(TrainingMsg(kind: tmkNewBattle, enemyId: e.scannedBotId))
method onBulletHit*(bot: SacBot, e: BulletHitBotEvent) =
bot.dmgDealt += e.damage
inc bot.hits
method onHitByBullet*(bot: SacBot, e: HitByBulletEvent) =
bot.dmgTaken += e.damage
# Lever 2 (#59): anti-ram — every BotHitBotEvent receipt means a bot-bot
# collision happened and we ate RAM_DAMAGE (server deals 0.6 to both parties;
# only the hitter gets notified). Flat per-event penalty; being rammed without
# hitting back stays event-invisible.
# ponytail: enemy-initiated rams undetected — add energy-residual detection if
# v2 battle data shows ram-heavy losses.
method onHitBot*(bot: SacBot, e: BotHitBotEvent) =
inc bot.ramTaken
method onHitWall*(bot: SacBot, e: BotHitWallEvent) =
inc bot.wallHits
method onBulletHitWall*(bot: SacBot, e: BulletHitWallEvent) =
if e.bullet.ownerId == getMyId():
bot.wastedPower += e.bullet.power
method onGameAborted*(bot: SacBot) =
# Mid-round abort: drop the pending transition rather than leak it into the
# next battle's data.
bot.hasLastTrans = false
bot.dmgDealt = 0; bot.dmgTaken = 0; bot.wastedPower = 0; bot.wallHits = 0
bot.hits = 0; bot.ramTaken = 0
method onRoundEnded*(bot: SacBot, e: RoundEndedEventForBot) =
## Harness liveness signal (#49): RunTraining.java watches round_counter.txt
## and aborts the battle if it freezes (dead bot process). Main thread.
bumpRoundCounter()
method onGameEnded*(bot: SacBot, e: GameEndedEventForBot) =
## Battle end -> terminal transition with done=true. Main thread; the API has
## joined the bot thread before this fires, so these plain fields are quiescent.
if bot.hasLastTrans:
let r = bot.takeReward(win = e.results.rank == 1, loss = e.results.rank != 1)
discard sendTrainingMsg(TrainingMsg(kind: tmkTransition,
state: bot.lastState, action: bot.lastAction,
reward: r, nextState: bot.lastState, done: true))
bot.hasLastTrans = false
# ── Run loop (bot thread) ─────────────────────────────────────────────────────
method run(bot: SacBot) =
randomize()
# Per-thread locals: born and freed on THIS thread, every round. Nothing
# heap-owned survives the round boundary except the bot object's plain fields.
var actor: ActorNet
var actorReady = false
var myVersion = 0
var myHidden = 0
var flat: seq[float32]
while isRunning():
# Battle boundary (Q4/Q12): zero LSTM persistence + per-battle flags.
if bot.battleId != bot.seenBattle:
bot.seenBattle = bot.battleId
bot.newBattleSent = false
zeroMem(addr bot.hArr, sizeof(bot.hArr))
zeroMem(addr bot.cArr, sizeof(bot.cArr))
# Weight sync (Q7/Q11): always-latest; rebuild this thread's tensors on change.
if pullWeights(myVersion, myHidden, flat):
if myHidden > MaxHidden:
# hArr/cArr are fixed-capacity; a bigger SACLSTM_HIDDEN_SIZE would
# heap-overflow them in tensorToHidden. Loud misconfig beats corruption.
raise newException(ValueError, "SACLSTM_HIDDEN_SIZE=" & $myHidden &
" exceeds MaxHidden=" & $MaxHidden & " (bot-side LSTM persistence cap)")
var cur = 0
actor = actorFromFlat(flat, cur, myHidden)
actorReady = true
# Spawn an enemy bullet when a fresh scan shows they fired; then advance and
# prune the tracked bullets (positions feed state slots 22–33).
if bot.hasContact and bot.enemy.hasFired and bot.bulletCount < MaxBotBullets:
let p = bot.enemy.lastFirePower
let spd = 20.0 - 3.0 * p
let ang = arctan2(getY() - bot.enemy.y, getX() - bot.enemy.x)
bot.bullets[bot.bulletCount] = InFlightBullet(x: bot.enemy.x, y: bot.enemy.y,
vx: spd * cos(ang), vy: spd * sin(ang), power: p)
inc bot.bulletCount
let aW = float64(getArenaWidth())
let aH = float64(getArenaHeight())
var alive = 0
for i in 0 ..< bot.bulletCount:
let b = bot.bullets[i]
let nx = b.x + b.vx
let ny = b.y + b.vy
if nx >= 0.0 and nx <= aW and ny >= 0.0 and ny <= aH:
bot.bullets[alive] = InFlightBullet(x: nx, y: ny, vx: b.vx, vy: b.vy, power: b.power)
inc alive
bot.bulletCount = alive
# State build (35-dim, this thread's tensor).
inc bot.ticksSinceScan
var gs: GameState
gs.x = getX()
gs.y = getY()
gs.direction = getDirection()
gs.speed = getSpeed()
gs.energy = getEnergy()
gs.gunDirection = getGunDirection()
gs.gunHeat = getGunHeat()
gs.arenaWidth = aW
gs.arenaHeight = aH
gs.hasContact = bot.hasContact
gs.enemy = bot.enemy
gs.ticksSinceLastScan = bot.ticksSinceScan
var bd: array[MaxBotBullets, BulletData]
for i in 0 ..< bot.bulletCount:
bd[i] = BulletData(x: bot.bullets[i].x, y: bot.bullets[i].y, power: bot.bullets[i].power)
if bot.bulletCount > 1: # closest threats fill slots 0-2
bd.toOpenArray(0, bot.bulletCount - 1).sort(proc(a, b: BulletData): int =
cmp(hypot(a.x - gs.x, a.y - gs.y), hypot(b.x - gs.x, b.y - gs.y)))
gs.bulletCount = min(bot.bulletCount, 3)
for i in 0 ..< gs.bulletCount:
gs.bullets[i] = bd[i]
# Consume the fired pulse AFTER the state saw it (one shot -> one bullet).
bot.enemy.hasFired = false
let stateT = buildState(gs)
# Finalize the PREVIOUS transition: reward from events since the last tick,
# nextState is this tick's observation (PPO_Bot alignment).
if actorReady and bot.hasLastTrans:
discard sendTrainingMsg(TrainingMsg(kind: tmkTransition,
state: bot.lastState, action: bot.lastAction,
reward: bot.takeReward(), nextState: stateToArr(stateT), done: false))
bot.lastState = stateToArr(stateT)
if not actorReady:
setRadarTurnRate(45.0) # no weights yet (defensive; main pre-inits) — just sweep
go()
continue
# Inference: hidden state lives as plain arrays on the bot object (persists
# across rounds); tensors are rebuilt per tick on this thread.
let h = hiddenToTensor(bot.hArr, myHidden)
let c = hiddenToTensor(bot.cArr, myHidden)
let fwd = actor.actorForward(stateT, (h: h, c: c))
tensorToHidden(fwd.lstm.h, bot.hArr)
tensorToHidden(fwd.lstm.c, bot.cArr)
bot.lastAction = actionToArr(fwd.actions)
bot.hasLastTrans = true
# Actions -> intents (go() snapshots them at send time).
let mapped = mapActions(fwd.actions, getSpeed(), getGunHeat())
setTurnRate(mapped.turnRate)
setTargetSpeed(getSpeed() + mapped.acceleration) # actions.nim contract
setGunTurnRate(mapped.gunTurnRate)
if mapped.firePower > 0.0:
discard setFire(mapped.firePower)
if not bot.hasContact:
setRadarTurnRate(45.0) # sweep until first lock (onScannedBot overrides same-tick)
go()
# ── Entry point ───────────────────────────────────────────────────────────────
when isMainModule:
initIntegration() # spawn training + I/O threads, seed weight snapshot
var bot = SacBot()
start(bot, botJsonPath) # blocks until server disconnect
shutdownIntegration() # Shutdown msg -> final save -> joins
Binary file not shown.
@@ -0,0 +1,36 @@
## actions.nim — map raw network output (4 tanh values) to bot intent fields.
##
## Note on acceleration vs targetSpeed:
## TankRoyale uses setTargetSpeed(), not setAcceleration().
## The mapped `acceleration` field is a delta; callers must compute:
## newTargetSpeed = clamp(currentSpeed + acceleration, -8.0, 8.0)
## and call setTargetSpeed(newTargetSpeed).
import arraymancer
const ACTION_DIM* = 4
type
MappedActions* = object
turnRate*: float ## degrees/tick, speed-aware; [-10, 10] at speed 0
acceleration*: float ## delta speed in [-2, +1]; caller adds to currentSpeed
gunTurnRate*: float ## degrees/tick in [-20, 20]
firePower*: float ## 0 = don't fire; (0.1, 3.0] = fire with this power
proc mapActions*(networkOutput: Tensor[float32],
currentSpeed: float,
gunHeat: float): MappedActions =
## networkOutput: [4] tensor of tanh values in [-1, 1].
let a0 = networkOutput[0].float
let a1 = networkOutput[1].float
let a2 = networkOutput[2].float
let a3 = networkOutput[3].float
result.turnRate = a0 * (10.0 - 0.75 * abs(currentSpeed))
# asymmetric accel: [-1,1] -> [-2, +1] via (value * 1.5 - 0.5)
result.acceleration = a1 * 1.5 - 0.5
result.gunTurnRate = a2 * 20.0
if a3 > 0.0 and gunHeat <= 0.0:
result.firePower = a3 * 2.9 + 0.1
else:
result.firePower = 0.0
@@ -0,0 +1,463 @@
## 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()
Binary file not shown.
@@ -0,0 +1,161 @@
## network.nim — LSTM-based Actor and dual Critic for SAC-v2.
## No autograd; inference only. Manual LSTM cell from scratch.
import arraymancer
import std/[math, random, os, strutils]
# ── Configuration ─────────────────────────────────────────────────────────────
proc getHiddenSize*(): int =
let s = getEnv("SACLSTM_HIDDEN_SIZE", "256")
result = parseInt(s)
proc isEvalMode*(): bool =
getEnv("SACLSTM_EVAL_MODE", "0") == "1"
# ── Types ─────────────────────────────────────────────────────────────────────
type
Linear* = object
w*, b*: Tensor[float32] # w: [out, in], b: [out]
LSTMCell* = object
## Combined weight matrix Wi|Wf|Wg|Wo stacked: [4*hidden, input+hidden]
## Combined bias stacked: [4*hidden]
wCombined*: Tensor[float32]
bCombined*: Tensor[float32]
hiddenDim*: int
ActorNet* = object
fc1*: Linear
lstm*: LSTMCell
fc2*: Linear
muHead*: Linear
logStdHead*: Linear
hiddenDim*: int
CriticNet* = object
fc1*: Linear
lstm*: LSTMCell
fc2*: Linear
fc3*: Linear
hiddenDim*: int
LSTMState* = tuple[h, c: Tensor[float32]] # each [hiddenDim]
# ── Init helpers ──────────────────────────────────────────────────────────────
proc initLinear*(inDim, outDim: int; scale: float32): Linear =
result.w = randomNormalTensor[float32]([outDim, inDim]) *. scale
result.b = zeros[float32](outDim)
proc initLinearHe*(inDim, outDim: int): Linear =
initLinear(inDim, outDim, sqrt(2.0'f32 / inDim.float32))
proc initLinearOut*(inDim, outDim: int): Linear =
initLinear(inDim, outDim, sqrt(1.0'f32 / inDim.float32))
proc initLSTMCell*(inputDim, hiddenDim: int): LSTMCell =
result.hiddenDim = hiddenDim
let fanIn = (inputDim + hiddenDim).float32
let scale = sqrt(1.0'f32 / fanIn)
result.wCombined = randomNormalTensor[float32]([4 * hiddenDim, inputDim + hiddenDim]) *. scale
result.bCombined = zeros[float32](4 * hiddenDim)
proc zeroState*(hiddenDim: int): LSTMState =
result = (h: zeros[float32](hiddenDim), c: zeros[float32](hiddenDim))
proc initActorNet*(stateDim: int): ActorNet =
let h = getHiddenSize()
result.hiddenDim = h
result.fc1 = initLinearHe(stateDim, h)
result.lstm = initLSTMCell(h, h)
result.fc2 = initLinearHe(h, 128)
result.muHead = initLinearOut(128, 4)
result.logStdHead = initLinearOut(128, 4)
proc initCriticNet*(stateDim, actionDim: int): CriticNet =
let h = getHiddenSize()
result.hiddenDim = h
result.fc1 = initLinearHe(stateDim + actionDim, h)
result.lstm = initLSTMCell(h, h)
result.fc2 = initLinearHe(h, 128)
result.fc3 = initLinearOut(128, 1)
# ── Forward helpers ───────────────────────────────────────────────────────────
proc linear*(l: Linear; x: Tensor[float32]): Tensor[float32] =
l.w * x + l.b
proc relu*(x: Tensor[float32]): Tensor[float32] =
x.map(proc(v: float32): float32 = max(0.0'f32, v))
proc sigmoid*(x: Tensor[float32]): Tensor[float32] =
x.map(proc(v: float32): float32 = 1.0'f32 / (1.0'f32 + exp(-v)))
proc tanhT*(x: Tensor[float32]): Tensor[float32] =
x.map(proc(v: float32): float32 = tanh(v))
proc lstmStep*(cell: LSTMCell; x, h, c: Tensor[float32]): LSTMState =
## x: [inputDim], h/c: [hiddenDim] → h', c': [hiddenDim]
let xh = concat(x, h, axis = 0) # [inputDim + hiddenDim]
let gates = cell.wCombined * xh + cell.bCombined # [4*hidden]
let hd = cell.hiddenDim
let iGate = sigmoid(gates[0 ..< hd])
let fGate = sigmoid(gates[hd ..< 2*hd])
let gGate = tanhT(gates[2*hd ..< 3*hd])
let oGate = sigmoid(gates[3*hd ..< 4*hd])
let cPrime = fGate *. c + iGate *. gGate
let hPrime = oGate *. tanhT(cPrime)
result = (h: hPrime, c: cPrime)
# ── Actor forward ─────────────────────────────────────────────────────────────
const
LOG_STD_MIN = -5.0'f32
LOG_STD_MAX = 2.0'f32
LOG_PROB_EPS = 1e-6'f32
proc actorForward*(net: ActorNet; state: Tensor[float32]; lstm: LSTMState;
deterministic = false):
tuple[actions: Tensor[float32]; logProb: float32; lstm: LSTMState] =
## state: [stateDim], lstm: (h,c) each [hiddenDim]
## Returns actions [4], scalar logProb, updated (h',c').
let h1 = relu(net.fc1.linear(state))
let lstmOut = lstmStep(net.lstm, h1, lstm.h, lstm.c)
let h2 = relu(net.fc2.linear(lstmOut.h))
let mu = net.muHead.linear(h2)
let logStdRaw = net.logStdHead.linear(h2)
let logStd = logStdRaw.map(proc(v: float32): float32 = clamp(v, LOG_STD_MIN, LOG_STD_MAX))
if deterministic or isEvalMode():
let actions = tanhT(mu)
return (actions: actions, logProb: 0.0'f32, lstm: lstmOut)
# Reparameterization: z = mu + std * eps, action = tanh(z)
let std = logStd.map(proc(v: float32): float32 = exp(v))
var actions = newTensor[float32](4)
var logProb = 0.0'f32
let twoPiLog = 0.5'f32 * ln(2.0'f32 * PI.float32)
for i in 0 ..< 4:
let eps = gauss(0.0'f64, 1.0'f64).float32
let z = mu[i] + std[i] * eps
actions[i] = tanh(z)
# log N(z | mu, std) - log(1 - tanh²(z) + eps)
let diff = (z - mu[i]) / std[i]
let logNorm = -0.5'f32 * diff * diff - ln(std[i]) - twoPiLog
let tanhCorr = ln(1.0'f32 - actions[i] * actions[i] + LOG_PROB_EPS)
logProb += logNorm - tanhCorr
result = (actions: actions, logProb: logProb, lstm: lstmOut)
# ── Critic forward ────────────────────────────────────────────────────────────
proc criticForward*(net: CriticNet; stateAction: Tensor[float32]; lstm: LSTMState):
tuple[q: float32; lstm: LSTMState] =
## stateAction: [stateDim + actionDim], lstm: (h,c) each [hiddenDim]
let h1 = relu(net.fc1.linear(stateAction))
let lstmOut = lstmStep(net.lstm, h1, lstm.h, lstm.c)
let h2 = relu(net.fc2.linear(lstmOut.h))
let q = net.fc3.linear(h2)
result = (q: q[0], lstm: lstmOut)
@@ -0,0 +1,136 @@
## replay_buffer.nim — sequential ring buffer for off-policy SAC+LSTM training.
##
## Stores transitions and samples contiguous sequences for recurrent training.
## Sequences NEVER cross battle boundaries (done=true).
##
## Config env vars:
## SACLSTM_BUFFER_CAPACITY (default: 500_000)
## SACLSTM_BURN_IN (default: 8)
## SACLSTM_TRAIN_WINDOW (default: 16)
import arraymancer
import std/[os, strutils, random]
# ── Config ────────────────────────────────────────────────────────────────────
proc getBufferCapacity*(): int =
parseInt(getEnv("SACLSTM_BUFFER_CAPACITY", "500000"))
proc getBurnIn*(): int =
parseInt(getEnv("SACLSTM_BURN_IN", "8"))
proc getTrainWindow*(): int =
parseInt(getEnv("SACLSTM_TRAIN_WINDOW", "16"))
# ── Types ─────────────────────────────────────────────────────────────────────
type
Transition* = object
state*: Tensor[float32] # [stateDim]
action*: Tensor[float32] # [actionDim]
reward*: float32
nextState*: Tensor[float32] # [stateDim]
done*: bool # true = battle end
Sequence* = object
burnIn*: seq[Transition] # first burnIn steps (for LSTM warm-up)
train*: seq[Transition] # next trainWindow steps (for gradient computation)
ReplayBuffer* = object
## Ring buffer. `head` is the next write position. `count` tracks fill level.
transitions: seq[Transition]
capacity: int
stateDim: int
actionDim: int
head: int # next write index
count: int # number of valid transitions stored
burnIn: int
trainWindow: int
# ── Construction ──────────────────────────────────────────────────────────────
proc newReplayBuffer*(capacity, stateDim, actionDim: int;
burnIn = getBurnIn();
trainWindow = getTrainWindow()): ReplayBuffer =
result.capacity = capacity
result.stateDim = stateDim
result.actionDim = actionDim
result.burnIn = burnIn
result.trainWindow = trainWindow
result.head = 0
result.count = 0
result.transitions = newSeq[Transition](capacity)
# ── Core operations ───────────────────────────────────────────────────────────
proc add*(buf: var ReplayBuffer; t: Transition) =
buf.transitions[buf.head] = t
buf.head = (buf.head + 1) mod buf.capacity
if buf.count < buf.capacity:
inc buf.count
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
# ── Sampling ──────────────────────────────────────────────────────────────────
proc sampleSequences*(buf: ReplayBuffer; batchSize: int): seq[Sequence] =
## Sample `batchSize` contiguous sequences of length burnIn+trainWindow.
## Sequences never cross a done=true boundary and never wrap the ring buffer.
##
## Returns fewer than batchSize sequences if not enough valid starts exist.
## Returns empty seq if canSample is false.
if not buf.canSample: return @[]
let seqLen = buf.burnIn + buf.trainWindow
let oldest = if buf.count < buf.capacity: 0
else: buf.head # oldest valid index when full
# Build valid starting indices.
# ponytail: O(count) scan per sample call; upgrade to an indexed set of
# boundary positions if count reaches hundreds of thousands and profiling shows
# this is a bottleneck.
var validStarts: seq[int]
for i in 0 ..< buf.count - seqLen + 1:
# Absolute ring-buffer index for the i-th oldest transition
let startIdx = (oldest + i) mod buf.capacity
# Check: the sequence [startIdx .. startIdx+seqLen-2] must not contain done=true
# (a done at position k means the battle ended there; the next transition is
# from a new battle, so the sequence would cross a boundary).
# Also, the sequence must not wrap around the ring buffer.
let endIdx = startIdx + seqLen - 1 # exclusive of wrap check
if endIdx >= buf.capacity:
# Sequence wraps the ring buffer — invalid starting point.
continue
var crosses = false
for j in 0 ..< seqLen - 1:
if buf.transitions[startIdx + j].done:
crosses = true
break
if not crosses:
validStarts.add(startIdx)
if validStarts.len == 0: return @[]
result = newSeq[Sequence](min(batchSize, validStarts.len))
# Sample with replacement if batchSize > validStarts.len, else sample without.
# ponytail: sampling with replacement for simplicity; shuffle+take for
# without-replacement if the caller needs it.
for i in 0 ..< result.len:
let startIdx = validStarts[rand(validStarts.len - 1)]
var s: Sequence
s.burnIn = newSeq[Transition](buf.burnIn)
s.train = newSeq[Transition](buf.trainWindow)
for j in 0 ..< buf.burnIn:
s.burnIn[j] = buf.transitions[startIdx + j]
for j in 0 ..< buf.trainWindow:
s.train[j] = buf.transitions[startIdx + buf.burnIn + j]
result[i] = s
@@ -0,0 +1,84 @@
## rewards.nim — Raw reward computation + running mean/variance normalizer.
## Welford online algorithm; safe cold-start (0 or 1 samples).
import std/math
# ── Lever-2 shaping constants (#59, campaign v2) — TUNABLE ────────────────────
# Scale discipline: commensurate with existing magnitudes (dealt p=1 was +4,
# wall tick -5/tick, win +20). Death/loss and win terms stay dominant; these
# only re-rank mid-band behaviors (fight vs outlive vs get-rammed).
const
# Multiplier on the bullet-damage-dealt term: p=1 hit +4 -> +5. Low-power
# spam stays unprofitable (6*0.1-2 = -1.4 < 0 even after x1.25).
AggressionMult* = 1.25 # ponytail: TUNABLE — raise toward 1.5 if v2 bot still passivity-leaning
# Flat per landed shot on top of damage: discrete accuracy signal.
HitBonus* = 0.5 # ponytail: TUNABLE — keep < 6p-2 at min viable power (~0.34)
# Per bot-bot collision (BotHitBotEvent): server deals RAM_DAMAGE=0.6 to
# both parties but only notifies the hitter — each receipt = damage taken.
RamTakenPenalty* = 3.0 # ponytail: TUNABLE — vs p=0.8 bullet received (-6.8)
# Enemy-charging deterrent: distance/diagonal below this => escalating
# negative (max at zero distance), suppressed while we deal damage that step.
ChargeDistFrac* = 0.12 # ponytail: TUNABLE — ~120u of 800x600 diag (1000)
ChargePenalty* = 2.0 # ponytail: TUNABLE — per-tick ceiling, milder than wall (-5/tick)
# ── Raw reward ────────────────────────────────────────────────────────────────
proc computeReward*(
damageInflicted: float64 = 0.0, # fire power p of own shot that hit
damageReceived: float64 = 0.0, # fire power p_e of enemy shot that hit
wallHitTicks: int = 0, # ticks in wall contact this step
wastedShotPower: float64 = 0.0, # fire power of shot that missed/hit wall
hitCount: int = 0, # own bullets that hit the enemy this step (#59)
ramTakenCount: int = 0, # collisions where we were the victim (#59)
enemyDistFrac: float64 = 2.0, # enemy dist / arena diag; >ChargeDistFrac when no contact (#59)
win: bool = false,
loss: bool = false
): float64 =
## Returns the raw (un-normalized) reward for one decision step.
## Damage formula: 4p + 2(p-1) = 6p - 2 (matches Tank Royale bullet rules).
let p = damageInflicted
let pe = damageReceived
if p > 0.0:
result += AggressionMult * (6.0 * p - 2.0)
if hitCount > 0: result += HitBonus * hitCount.float64
if pe > 0.0: result -= 6.0 * pe - 2.0
result -= RamTakenPenalty * ramTakenCount.float64
if enemyDistFrac < ChargeDistFrac and p <= 0.0:
result -= ChargePenalty * (1.0 - enemyDistFrac / ChargeDistFrac)
result -= 5.0 * wallHitTicks.float64
if wastedShotPower > 0.0: result -= 0.1 * wastedShotPower
if win: result += 20.0
if loss: result -= 10.0
# ── Running normalizer (Welford) ──────────────────────────────────────────────
const NormEps = 1e-8
type
RewardNormalizer* = object
n*: int # samples seen
mean*: float64
m2*: float64 # sum of squared deviations (Welford M2)
proc update*(rn: var RewardNormalizer; r: float64) =
rn.n += 1
let delta = r - rn.mean
rn.mean += delta / rn.n.float64
let delta2 = r - rn.mean
rn.m2 += delta * delta2
proc normalize*(rn: RewardNormalizer; r: float64): float64 =
## Returns (r - mean) / (std + eps) once statistics are meaningful
## (n >= 4 and spread well above zero). Before that, returns the RAW
## reward unchanged — Welford M2 collapses to exactly 0 when early raw
## rewards are identical, and dividing by the 1e-8 floor then z-scores
## the first differing reward to ~1e8, poisoning TD targets.
if rn.n < 4: return r
let variance = rn.m2 / rn.n.float64 # ponytail: population var; switch to n-1 if bias matters
let stddev = sqrt(variance)
if stddev <= 1e-3 * (abs(rn.mean) + 1.0): return r
# ponytail: warm-up pass-through ceiling — raw rewards bypass normalization
# until stats are meaningful; upgrade = persist Welford state in checkpoint
# if warm-up noise ever hurts learning.
result = (r - rn.mean) / (stddev + NormEps)
Binary file not shown.
@@ -0,0 +1,117 @@
## State vector module — produces a 35-dimensional normalized tensor for SAC+LSTM policy.
## No bot API imports; takes plain data structs populated from game events.
## The LSTM handles temporal context, so no explicit history window here.
import std/math
import arraymancer
const STATE_DIM* = 35
type
BulletData* = object
## Enemy bullet in flight (absolute arena coords + fire power).
x*, y*: float64
power*: float64 # fire power in [0.1, 3.0]; speed = 20 - 3*power
EnemyData* = object
## Current enemy state, from the most recent onScannedBot event.
x*, y*: float64
direction*: float64
speed*: float64
energy*: float64
hasFired*: bool
lastFirePower*: float64
prevSpeed*: float64 # speed from the previous scan (for acceleration)
prevDirection*: float64 # direction from the previous scan (for turn rate)
hasPrevScan*: bool # true once we have at least two scans
GameState* = object
## Accumulates data from bot events. Populate fields before calling buildState.
# Own bot
x*, y*: float64
direction*: float64
speed*: float64
energy*: float64
gunDirection*: float64
gunHeat*: float64
arenaWidth*, arenaHeight*: float64
# Enemy
hasContact*: bool
enemy*: EnemyData
ticksSinceLastScan*: int
# Bullets in flight (up to 3 tracked)
bullets*: array[3, BulletData]
bulletCount*: int
proc buildState*(gs: GameState): Tensor[float32] =
## Build the 35-float normalized state tensor.
##
## Layout:
## [0-6] own bot: x/aW, y/aH, dir/360, speed/8, energy/100, gunDir/360, gunHeat/1.8
## [7-13] enemy: x/aW, y/aH, dir/360, speed/8, energy/100, hasFired, lastFirePower/3
## [14-17] derived: enemyAccel/8, enemyTurnRate/180, relBearing/180, distance/diag
## [18-21] walls: top, bottom, left, right — each / max(aW,aH)
## [22-33] bullets: up to 3 × (relX/aW, relY/aH, speed/20, ticksToImpact clamped to 1)
## [34] scan staleness: ticksSinceLastScan/30 clamped to 1
result = zeros[float32](STATE_DIM)
let aW = gs.arenaWidth
let aH = gs.arenaHeight
let diag = sqrt(aW * aW + aH * aH)
let wMax = max(aW, aH)
# --- Own bot (0-6) ---
result[0] = float32(gs.x / aW)
result[1] = float32(gs.y / aH)
result[2] = float32(gs.direction / 360.0)
result[3] = float32(gs.speed / 8.0)
result[4] = float32(gs.energy / 100.0)
result[5] = float32(gs.gunDirection / 360.0)
result[6] = float32(gs.gunHeat / 1.8)
# --- Enemy current (7-13) ---
if gs.hasContact:
result[7] = float32(gs.enemy.x / aW)
result[8] = float32(gs.enemy.y / aH)
result[9] = float32(gs.enemy.direction / 360.0)
result[10] = float32(gs.enemy.speed / 8.0)
result[11] = float32(gs.enemy.energy / 100.0)
result[12] = float32(if gs.enemy.hasFired: 1.0 else: 0.0)
result[13] = float32(gs.enemy.lastFirePower / 3.0)
# --- Derived (14-17) ---
if gs.hasContact:
if gs.enemy.hasPrevScan:
result[14] = float32((gs.enemy.speed - gs.enemy.prevSpeed) / 8.0)
let dDir = ((gs.enemy.direction - gs.enemy.prevDirection) + 540.0) mod 360.0 - 180.0
result[15] = float32(dDir / 180.0)
let dx = gs.enemy.x - gs.x
let dy = gs.enemy.y - gs.y
let absDir = (180.0 * arctan2(dx, dy) / PI + 360.0) mod 360.0
let relBearing = ((absDir - gs.direction) + 540.0) mod 360.0 - 180.0
result[16] = float32(relBearing / 180.0)
result[17] = float32(sqrt(dx * dx + dy * dy) / diag)
# --- Wall distances (18-21): top, bottom, left, right ---
result[18] = float32((aH - gs.y) / wMax)
result[19] = float32(gs.y / wMax)
result[20] = float32(gs.x / wMax)
result[21] = float32((aW - gs.x) / wMax)
# --- Bullet tracking (22-33): up to 3 bullets × 4 floats ---
# Per slot: relX/aW, relY/aH, speed/20, ticksToImpact/diag (clamped to 1)
for i in 0 ..< min(gs.bulletCount, 3):
let b = gs.bullets[i]
let bSpd = 20.0 - 3.0 * b.power
let bdx = b.x - gs.x
let bdy = b.y - gs.y
let bdist = sqrt(bdx * bdx + bdy * bdy)
let ticks = if bSpd > 0.0: min(bdist / bSpd / diag, 1.0) else: 0.0
let base = 22 + i * 4
result[base + 0] = float32(bdx / aW)
result[base + 1] = float32(bdy / aH)
result[base + 2] = float32(bSpd / 20.0)
result[base + 3] = float32(ticks)
# --- Scan staleness (34) ---
result[34] = float32(min(gs.ticksSinceLastScan.float64 / 30.0, 1.0))
Binary file not shown.
@@ -0,0 +1,611 @@
## training.nim — SAC-v2 update for the LSTM Actor + twin Critic.
## Manual backprop; no autograd. Uses Arraymancer tensors throughout.
##
## Config env vars:
## SACLSTM_LR_ACTOR (default: 3e-4)
## SACLSTM_LR_CRITIC (default: 3e-4)
## SACLSTM_LR_ALPHA (default: 3e-4)
## SACLSTM_GAMMA (default: 0.99)
## SACLSTM_TAU (default: 0.005)
## SACLSTM_TARGET_ENTROPY (default: -4.0)
import arraymancer except Linear
import std/[math, os, strutils]
import SAC_LSTM_Bot/network
import SAC_LSTM_Bot/weights
import SAC_LSTM_Bot/replay_buffer
# ── Config ────────────────────────────────────────────────────────────────────
proc getLrActor*(): float32 = parseFloat(getEnv("SACLSTM_LR_ACTOR", "3e-4")).float32
proc getLrCritic*(): float32 = parseFloat(getEnv("SACLSTM_LR_CRITIC", "3e-4")).float32
proc getLrAlpha*(): float32 = parseFloat(getEnv("SACLSTM_LR_ALPHA", "3e-4")).float32
proc getGamma*(): float32 = parseFloat(getEnv("SACLSTM_GAMMA", "0.99")).float32
proc getTau*(): float32 = parseFloat(getEnv("SACLSTM_TAU", "0.005")).float32
proc getTargetEntropy*(): float32 =
parseFloat(getEnv("SACLSTM_TARGET_ENTROPY", "-4.0")).float32
# ── SACTrainer ────────────────────────────────────────────────────────────────
type
SACTrainer* = object
actor*: ActorNet
critic1*: CriticNet
critic2*: CriticNet
targetCritic1*: CriticNet
targetCritic2*: CriticNet
logAlpha*: float32 ## log of entropy temperature; alpha = exp(logAlpha)
targetEntropy*: float32
tau*: float32
lrActor*: float32
lrCritic*: float32
lrAlpha*: float32
gamma*: float32
adam*: SACAdamStates
SACMetrics* = object
criticLoss*: float32
actorLoss*: float32
alphaLoss*: float32
alpha*: float32
proc initSACTrainer*(stateDim, actionDim: int): SACTrainer =
result.actor = initActorNet(stateDim)
result.critic1 = initCriticNet(stateDim, actionDim)
result.critic2 = initCriticNet(stateDim, actionDim)
result.targetCritic1 = result.critic1
result.targetCritic2 = result.critic2
result.logAlpha = 0.0'f32
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)
proc alpha*(t: SACTrainer): float32 = exp(t.logAlpha)
# ── Adam steps ────────────────────────────────────────────────────────────────
proc adamStepScalar(param: var float32; grad: float32;
state: var AdamVar; lr: float32) =
## Scalar Adam for logAlpha (state.m/v are shape [1] tensors).
inc state.t
let b1 = 0.9'f32; let b2 = 0.999'f32; let eps = 1e-8'f32
state.m[0] = b1 * state.m[0] + (1.0'f32 - b1) * grad
state.v[0] = b2 * state.v[0] + (1.0'f32 - b2) * grad * grad
let mHat = state.m[0] / (1.0'f32 - b1 ^ state.t.float32)
let vHat = state.v[0] / (1.0'f32 - b2 ^ state.t.float32)
param -= lr * mHat / (sqrt(vHat) + eps)
proc adamStepTensor(param: var Tensor[float32]; grad: Tensor[float32];
state: var AdamVar; lr: float32) =
## Tensor Adam (same pattern as PPO_Bot/training.nim adamStep).
inc state.t
let b1 = 0.9'f32; let b2 = 0.999'f32; let eps = 1e-8'f32
state.m = b1 *. state.m + (1.0'f32 - b1) *. grad
state.v = b2 *. state.v + (1.0'f32 - b2) *. (grad *. grad)
let mHat = state.m /. (1.0'f32 - b1 ^ state.t.float32)
let vHat = state.v /. (1.0'f32 - b2 ^ state.t.float32)
param -= lr *. mHat /. vHat.map(proc(x: float32): float32 = sqrt(x) + eps)
# ── Gradient clipping ─────────────────────────────────────────────────────────
proc globalNorm(grads: seq[Tensor[float32]]): float32 =
var sumSq = 0.0'f32
for g in grads:
for v in g: sumSq += v * v
sqrt(sumSq)
proc clipGrads(grads: var seq[Tensor[float32]]; maxNorm: float32) =
let norm = globalNorm(grads)
if norm > maxNorm and norm == norm:
let scale = maxNorm / norm
for g in grads.mitems: g = g *. scale
# ── Forward caches (for backprop) ─────────────────────────────────────────────
type
LinearFwd = object
inp, pre, act: Tensor[float32] # input, pre-relu, post-relu (or linear)
proc linearReluFwd(l: Linear; x: Tensor[float32]): LinearFwd =
result.inp = x
result.pre = l.w * x + l.b
result.act = relu(result.pre)
proc linearFwd(l: Linear; x: Tensor[float32]): LinearFwd =
result.inp = x
result.pre = l.w * x + l.b
result.act = result.pre # no nonlinearity
type
LSTMFwdCache = object
xh, gatesPre: Tensor[float32] # [inputDim+hd], [4*hd]
iGate, fGate, gGate, oGate: Tensor[float32] # [hd] each
cPrev, cPrime, hPrime: Tensor[float32] # [hd] each
proc lstmStepCached(cell: LSTMCell; x, h, c: Tensor[float32]): LSTMFwdCache =
result.cPrev = c
result.xh = concat(x, h, axis = 0)
result.gatesPre = cell.wCombined * result.xh + cell.bCombined
let hd = cell.hiddenDim
result.iGate = sigmoid(result.gatesPre[0 ..< hd])
result.fGate = sigmoid(result.gatesPre[hd ..< 2*hd])
result.gGate = tanhT(result.gatesPre[2*hd ..< 3*hd])
result.oGate = sigmoid(result.gatesPre[3*hd ..< 4*hd])
result.cPrime = result.fGate *. c + result.iGate *. result.gGate
result.hPrime = result.oGate *. tanhT(result.cPrime)
# ── Backward helpers ──────────────────────────────────────────────────────────
proc reluGrad(pre, dAct: Tensor[float32]): Tensor[float32] =
result = newTensor[float32](dAct.shape)
for i in 0 ..< dAct.shape[0]:
result[i] = if pre[i] > 0.0'f32: dAct[i] else: 0.0'f32
## Linear layer backward: returns (dx, dw, db) given upstream grad dAct.
## If hasRelu, applies relu' gate before computing gradients.
proc linearBack(w: Tensor[float32]; fwd: LinearFwd;
dAct: Tensor[float32]; hasRelu: bool):
tuple[dx, dw, db: Tensor[float32]] =
let dPre = if hasRelu: reluGrad(fwd.pre, dAct) else: dAct
result.dw = dPre.unsqueeze(1) * fwd.inp.unsqueeze(0) # [out, in]
result.db = dPre
result.dx = w.transpose * dPre # [in]
## LSTM single-step backward. dHPrime: [hd], dCPrime: [hd] (use zeros for truncated BPTT).
## Returns (dwCombined, dbCombined, dxh).
proc lstmBack(cell: LSTMCell; cache: LSTMFwdCache;
dHPrime, dCPrime: Tensor[float32]):
tuple[dwCombined, dbCombined, dxh: Tensor[float32]] =
let hd = cell.hiddenDim
let tanhCPrime = tanhT(cache.cPrime)
# Output gate
let dOGate_post = dHPrime *. tanhCPrime
# Cell state: gradient from h' and from downstream dCPrime
let dCPrimeTotal = dHPrime *. cache.oGate *.
(ones[float32](hd) - tanhCPrime *. tanhCPrime) + dCPrime
# Gate post-activation gradients
let dFGate_post = dCPrimeTotal *. cache.cPrev
let dIGate_post = dCPrimeTotal *. cache.gGate
let dGGate_post = dCPrimeTotal *. cache.iGate
# Gate pre-activation gradients (sigmoid', tanh')
let dIPre = dIGate_post *. cache.iGate *. (ones[float32](hd) - cache.iGate)
let dFPre = dFGate_post *. cache.fGate *. (ones[float32](hd) - cache.fGate)
let dGPre = dGGate_post *. (ones[float32](hd) - cache.gGate *. cache.gGate)
let dOPre = dOGate_post *. cache.oGate *. (ones[float32](hd) - cache.oGate)
# Concatenated gate gradient [4*hd]
let dGatesPre = concat(dIPre, dFPre, dGPre, dOPre, axis = 0)
result.dwCombined = dGatesPre.unsqueeze(1) * cache.xh.unsqueeze(0)
result.dbCombined = dGatesPre
result.dxh = cell.wCombined.transpose * dGatesPre
# ── Squashed-Gaussian log-prob and its gradients ──────────────────────────────
const
LOG_PROB_EPS = 1e-6'f32
LOG_STD_MIN = -5.0'f32
LOG_STD_MAX = 2.0'f32
## Given stored mu, clamped logStd, and sampled action = tanh(z), recover
## log π(a|s) and gradients w.r.t. mu and logStd.
proc squashedLogProb(mu, logStd, action: Tensor[float32]):
tuple[logProb: float32;
dLogProbDMu, dLogProbDLogStd: Tensor[float32]] =
let std = logStd.map(proc(v: float32): float32 = exp(v))
let z = mu # deterministic reparam: action = tanh(mu), so z ≡ mu, diff ≡ 0
result.dLogProbDMu = newTensor[float32](4)
result.dLogProbDLogStd = newTensor[float32](4)
let twoPiLog = 0.5'f32 * ln(2.0'f32 * PI.float32)
var lp = 0.0'f32
for i in 0 ..< 4:
let diff = (z[i] - mu[i]) / std[i]
let logNorm = -0.5'f32 * diff * diff - ln(std[i]) - twoPiLog
let tanhCorr = ln(1.0'f32 - action[i] * action[i] + LOG_PROB_EPS)
lp += logNorm - tanhCorr
result.dLogProbDMu[i] = diff / std[i] # (z-mu)/std²
result.dLogProbDLogStd[i] = diff * diff - 1.0'f32 # d logN / d logStd
result.logProb = lp
# ── Critic forward with activation cache ──────────────────────────────────────
type
CriticFwdCache = object
fc1: LinearFwd
lstm: LSTMFwdCache
fc2: LinearFwd
fc3: LinearFwd
q: float32
proc criticFwdCached(net: CriticNet; stateAction: Tensor[float32];
h, c: Tensor[float32]): CriticFwdCache =
result.fc1 = linearReluFwd(net.fc1, stateAction)
result.lstm = lstmStepCached(net.lstm, result.fc1.act, h, c)
result.fc2 = linearReluFwd(net.fc2, result.lstm.hPrime)
result.fc3 = linearFwd(net.fc3, result.fc2.act)
result.q = result.fc3.act[0]
# ── Critic backward ───────────────────────────────────────────────────────────
type
CriticGrads = object
dFc1W, dFc1B: Tensor[float32]
dFc2W, dFc2B: Tensor[float32]
dFc3W, dFc3B: Tensor[float32]
dLstmW, dLstmB: Tensor[float32]
dInput: Tensor[float32] ## grad w.r.t. stateAction input
proc criticBack(net: CriticNet; cache: CriticFwdCache; dQ: float32): CriticGrads =
let dFc3Act = [dQ].toTensor()
let fc3b = linearBack(net.fc3.w, cache.fc3, dFc3Act, hasRelu = false)
result.dFc3W = fc3b.dw; result.dFc3B = fc3b.db
let fc2b = linearBack(net.fc2.w, cache.fc2, fc3b.dx, hasRelu = true)
result.dFc2W = fc2b.dw; result.dFc2B = fc2b.db
let zeros_hd = zeros[float32](net.hiddenDim)
let lstmb = lstmBack(net.lstm, cache.lstm, fc2b.dx, zeros_hd)
result.dLstmW = lstmb.dwCombined; result.dLstmB = lstmb.dbCombined
# xh = [fc1.act | h_prev], dx is the x-part (fc1 output dim = hiddenDim)
let dLstmX = lstmb.dxh[0 ..< cache.fc1.act.shape[0]]
let fc1b = linearBack(net.fc1.w, cache.fc1, dLstmX, hasRelu = true)
result.dFc1W = fc1b.dw; result.dFc1B = fc1b.db
result.dInput = fc1b.dx # [stateDim + actionDim]
# ── Actor forward with activation cache ───────────────────────────────────────
type
ActorFwdCache = object
fc1: LinearFwd
lstm: LSTMFwdCache
fc2: LinearFwd
muHead: LinearFwd
lsHead: LinearFwd ## logStd head
mu: Tensor[float32] ## [4]
logStd: Tensor[float32] ## [4] clamped
action: Tensor[float32] ## [4] tanh(mu) — deterministic for gradient
proc actorFwdCached(net: ActorNet; state: Tensor[float32];
h, c: Tensor[float32]): ActorFwdCache =
result.fc1 = linearReluFwd(net.fc1, state)
result.lstm = lstmStepCached(net.lstm, result.fc1.act, h, c)
result.fc2 = linearReluFwd(net.fc2, result.lstm.hPrime)
result.muHead = linearFwd(net.muHead, result.fc2.act)
result.lsHead = linearFwd(net.logStdHead, result.fc2.act)
result.mu = result.muHead.act
result.logStd = result.lsHead.act.map(
proc(v: float32): float32 = clamp(v, LOG_STD_MIN, LOG_STD_MAX))
# Use tanh(mu) as the action for gradient computation (reparameterization).
# ponytail: deterministic here; add stochastic sample if off-policy bias matters.
result.action = tanhT(result.mu)
# ── Actor backward ────────────────────────────────────────────────────────────
type
ActorGrads = object
dFc1W, dFc1B: Tensor[float32]
dFc2W, dFc2B: Tensor[float32]
dMuW, dMuB: Tensor[float32]
dLogStdW, dLogStdB: Tensor[float32]
dLstmW, dLstmB: Tensor[float32]
proc actorBack(net: ActorNet; cache: ActorFwdCache;
dMu, dLogStd: Tensor[float32]): ActorGrads =
let muBack = linearBack(net.muHead.w, cache.muHead, dMu, hasRelu = false)
result.dMuW = muBack.dw; result.dMuB = muBack.db
let lsBack = linearBack(net.logStdHead.w, cache.lsHead, dLogStd, hasRelu = false)
result.dLogStdW = lsBack.dw; result.dLogStdB = lsBack.db
# fc2 gets grads from both output heads
let fc2b = linearBack(net.fc2.w, cache.fc2, muBack.dx + lsBack.dx, hasRelu = true)
result.dFc2W = fc2b.dw; result.dFc2B = fc2b.db
let zeros_hd = zeros[float32](net.hiddenDim)
let lstmb = lstmBack(net.lstm, cache.lstm, fc2b.dx, zeros_hd)
result.dLstmW = lstmb.dwCombined; result.dLstmB = lstmb.dbCombined
let dLstmX = lstmb.dxh[0 ..< cache.fc1.act.shape[0]]
let fc1b = linearBack(net.fc1.w, cache.fc1, dLstmX, hasRelu = true)
result.dFc1W = fc1b.dw; result.dFc1B = fc1b.db
# ── Adam application ──────────────────────────────────────────────────────────
proc applyActorAdam(net: var ActorNet; g: ActorGrads;
adam: var ActorAdam; lr: float32) =
adamStepTensor(net.fc1.w, g.dFc1W, adam.fc1.w, lr)
adamStepTensor(net.fc1.b, g.dFc1B, adam.fc1.b, lr)
adamStepTensor(net.fc2.w, g.dFc2W, adam.fc2.w, lr)
adamStepTensor(net.fc2.b, g.dFc2B, adam.fc2.b, lr)
adamStepTensor(net.muHead.w, g.dMuW, adam.muHead.w, lr)
adamStepTensor(net.muHead.b, g.dMuB, adam.muHead.b, lr)
adamStepTensor(net.logStdHead.w, g.dLogStdW, adam.logStdHead.w, lr)
adamStepTensor(net.logStdHead.b, g.dLogStdB, adam.logStdHead.b, lr)
adamStepTensor(net.lstm.wCombined, g.dLstmW, adam.lstm.wCombined, lr)
adamStepTensor(net.lstm.bCombined, g.dLstmB, adam.lstm.bCombined, lr)
proc applyCriticAdam(net: var CriticNet; g: CriticGrads;
adam: var CriticAdam; lr: float32) =
adamStepTensor(net.fc1.w, g.dFc1W, adam.fc1.w, lr)
adamStepTensor(net.fc1.b, g.dFc1B, adam.fc1.b, lr)
adamStepTensor(net.fc2.w, g.dFc2W, adam.fc2.w, lr)
adamStepTensor(net.fc2.b, g.dFc2B, adam.fc2.b, lr)
adamStepTensor(net.fc3.w, g.dFc3W, adam.fc3.w, lr)
adamStepTensor(net.fc3.b, g.dFc3B, adam.fc3.b, lr)
adamStepTensor(net.lstm.wCombined, g.dLstmW, adam.lstm.wCombined, lr)
adamStepTensor(net.lstm.bCombined, g.dLstmB, adam.lstm.bCombined, lr)
# ── Soft target update ────────────────────────────────────────────────────────
proc softUpdateLinear(target: var Linear; src: Linear; tau: float32) =
target.w = tau *. src.w + (1.0'f32 - tau) *. target.w
target.b = tau *. src.b + (1.0'f32 - tau) *. target.b
proc softUpdateLSTM(target: var LSTMCell; src: LSTMCell; tau: float32) =
target.wCombined = tau *. src.wCombined + (1.0'f32 - tau) *. target.wCombined
target.bCombined = tau *. src.bCombined + (1.0'f32 - tau) *. target.bCombined
proc softUpdateCritic(target: var CriticNet; src: CriticNet; tau: float32) =
softUpdateLinear(target.fc1, src.fc1, tau)
softUpdateLSTM(target.lstm, src.lstm, tau)
softUpdateLinear(target.fc2, src.fc2, tau)
softUpdateLinear(target.fc3, src.fc3, tau)
# ── Gradient accumulators ─────────────────────────────────────────────────────
proc zeroCriticGrads(net: CriticNet): CriticGrads =
result.dFc1W = zeros[float32](net.fc1.w.shape)
result.dFc1B = zeros[float32](net.fc1.b.shape)
result.dFc2W = zeros[float32](net.fc2.w.shape)
result.dFc2B = zeros[float32](net.fc2.b.shape)
result.dFc3W = zeros[float32](net.fc3.w.shape)
result.dFc3B = zeros[float32](net.fc3.b.shape)
result.dLstmW = zeros[float32](net.lstm.wCombined.shape)
result.dLstmB = zeros[float32](net.lstm.bCombined.shape)
result.dInput = zeros[float32](net.fc1.w.shape[1]) # [stateDim+actionDim]
proc zeroActorGrads(net: ActorNet): ActorGrads =
result.dFc1W = zeros[float32](net.fc1.w.shape)
result.dFc1B = zeros[float32](net.fc1.b.shape)
result.dFc2W = zeros[float32](net.fc2.w.shape)
result.dFc2B = zeros[float32](net.fc2.b.shape)
result.dMuW = zeros[float32](net.muHead.w.shape)
result.dMuB = zeros[float32](net.muHead.b.shape)
result.dLogStdW = zeros[float32](net.logStdHead.w.shape)
result.dLogStdB = zeros[float32](net.logStdHead.b.shape)
result.dLstmW = zeros[float32](net.lstm.wCombined.shape)
result.dLstmB = zeros[float32](net.lstm.bCombined.shape)
proc addCriticGrads(a: var CriticGrads; b: CriticGrads) =
a.dFc1W += b.dFc1W; a.dFc1B += b.dFc1B
a.dFc2W += b.dFc2W; a.dFc2B += b.dFc2B
a.dFc3W += b.dFc3W; a.dFc3B += b.dFc3B
a.dLstmW += b.dLstmW; a.dLstmB += b.dLstmB
# dInput not accumulated (not used for parameter update)
proc addActorGrads(a: var ActorGrads; b: ActorGrads) =
a.dFc1W += b.dFc1W; a.dFc1B += b.dFc1B
a.dFc2W += b.dFc2W; a.dFc2B += b.dFc2B
a.dMuW += b.dMuW; a.dMuB += b.dMuB
a.dLogStdW += b.dLogStdW; a.dLogStdB += b.dLogStdB
a.dLstmW += b.dLstmW; a.dLstmB += b.dLstmB
proc scaleCriticGrads(g: var CriticGrads; s: float32) =
g.dFc1W = g.dFc1W *. s; g.dFc1B = g.dFc1B *. s
g.dFc2W = g.dFc2W *. s; g.dFc2B = g.dFc2B *. s
g.dFc3W = g.dFc3W *. s; g.dFc3B = g.dFc3B *. s
g.dLstmW = g.dLstmW *. s; g.dLstmB = g.dLstmB *. s
proc scaleActorGrads(g: var ActorGrads; s: float32) =
g.dFc1W = g.dFc1W *. s; g.dFc1B = g.dFc1B *. s
g.dFc2W = g.dFc2W *. s; g.dFc2B = g.dFc2B *. s
g.dMuW = g.dMuW *. s; g.dMuB = g.dMuB *. s
g.dLogStdW = g.dLogStdW *. s; g.dLogStdB = g.dLogStdB *. s
g.dLstmW = g.dLstmW *. s; g.dLstmB = g.dLstmB *. s
proc criticGradsAsSeq(g: CriticGrads): seq[Tensor[float32]] =
@[g.dFc1W, g.dFc1B, g.dFc2W, g.dFc2B, g.dFc3W, g.dFc3B, g.dLstmW, g.dLstmB]
proc applyClipToCritic(g: var CriticGrads; maxNorm: float32) =
var gs = criticGradsAsSeq(g)
clipGrads(gs, maxNorm)
g.dFc1W = gs[0]; g.dFc1B = gs[1]
g.dFc2W = gs[2]; g.dFc2B = gs[3]
g.dFc3W = gs[4]; g.dFc3B = gs[5]
g.dLstmW = gs[6]; g.dLstmB = gs[7]
proc applyClipToActor(g: var ActorGrads; maxNorm: float32) =
var gs = @[g.dFc1W, g.dFc1B, g.dFc2W, g.dFc2B,
g.dMuW, g.dMuB, g.dLogStdW, g.dLogStdB, g.dLstmW, g.dLstmB]
clipGrads(gs, maxNorm)
g.dFc1W = gs[0]; g.dFc1B = gs[1]
g.dFc2W = gs[2]; g.dFc2B = gs[3]
g.dMuW = gs[4]; g.dMuB = gs[5]
g.dLogStdW = gs[6]; g.dLogStdB = gs[7]
g.dLstmW = gs[8]; g.dLstmB = gs[9]
# ── SAC update ────────────────────────────────────────────────────────────────
proc sacUpdate*(trainer: var SACTrainer; sequences: seq[Sequence]): SACMetrics =
## One SAC-v2 update given a batch of sequences. No-op if empty.
if sequences.len == 0: return
let N = sequences.len.float32
let alph = trainer.alpha()
let gamma = trainer.gamma
var totalCriticLoss = 0.0'f32
var totalActorLoss = 0.0'f32
var totalAlphaLoss = 0.0'f32
var accC1Grads = zeroCriticGrads(trainer.critic1)
var accC2Grads = zeroCriticGrads(trainer.critic2)
var accAGrads = zeroActorGrads(trainer.actor)
var dLogAlpha = 0.0'f32
for sq in sequences:
# ── 1. Burn-in: warm up hidden states, no gradient ──────────────────────
var actorH = zeros[float32](trainer.actor.hiddenDim)
var actorC = zeros[float32](trainer.actor.hiddenDim)
var c1H = zeros[float32](trainer.critic1.hiddenDim)
var c1C = zeros[float32](trainer.critic1.hiddenDim)
var c2H = zeros[float32](trainer.critic2.hiddenDim)
var c2C = zeros[float32](trainer.critic2.hiddenDim)
var tc1H = zeros[float32](trainer.targetCritic1.hiddenDim)
var tc1C = zeros[float32](trainer.targetCritic1.hiddenDim)
var tc2H = zeros[float32](trainer.targetCritic2.hiddenDim)
var tc2C = zeros[float32](trainer.targetCritic2.hiddenDim)
for tr in sq.burnIn:
let sa = concat(tr.state, tr.action, axis = 0)
let af = lstmStepCached(trainer.actor.lstm,
relu(trainer.actor.fc1.linear(tr.state)), actorH, actorC)
actorH = af.hPrime; actorC = af.cPrime
let c1f = lstmStepCached(trainer.critic1.lstm,
relu(trainer.critic1.fc1.linear(sa)), c1H, c1C)
c1H = c1f.hPrime; c1C = c1f.cPrime
let c2f = lstmStepCached(trainer.critic2.lstm,
relu(trainer.critic2.fc1.linear(sa)), c2H, c2C)
c2H = c2f.hPrime; c2C = c2f.cPrime
let tc1f = lstmStepCached(trainer.targetCritic1.lstm,
relu(trainer.targetCritic1.fc1.linear(sa)), tc1H, tc1C)
tc1H = tc1f.hPrime; tc1C = tc1f.cPrime
let tc2f = lstmStepCached(trainer.targetCritic2.lstm,
relu(trainer.targetCritic2.fc1.linear(sa)), tc2H, tc2C)
tc2H = tc2f.hPrime; tc2C = tc2f.cPrime
# ── 2–4. Training window ─────────────────────────────────────────────────
let T = sq.train.len.float32
var seqC1Grads = zeroCriticGrads(trainer.critic1)
var seqC2Grads = zeroCriticGrads(trainer.critic2)
var seqAGrads = zeroActorGrads(trainer.actor)
var seqDLogAlpha = 0.0'f32
for tr in sq.train:
let s = tr.state
let a = tr.action
let r = tr.reward
let sn = tr.nextState
let d = if tr.done: 0.0'f32 else: 1.0'f32
let sa = concat(s, a, axis = 0)
# ── 2. Critic update ─────────────────────────────────────────────────
let c1Cache = criticFwdCached(trainer.critic1, sa, c1H, c1C)
let c2Cache = criticFwdCached(trainer.critic2, sa, c2H, c2C)
# ── 3. Actor update (run first to get actorFwd on s before advancing h/c) ─
let actorFwd = actorFwdCached(trainer.actor, s, actorH, actorC)
let aCurr = actorFwd.action
let lpResult = squashedLogProb(actorFwd.mu, actorFwd.logStd, aCurr)
let logProbA = lpResult.logProb
# Advance actor hidden state from s → sn before computing actorNxt
actorH = actorFwd.lstm.hPrime; actorC = actorFwd.lstm.cPrime
# Next-state action from current actor (uses h/c advanced through s)
let actorNxt = actorFwdCached(trainer.actor, sn, actorH, actorC)
let aN = actorNxt.action
let lpN = squashedLogProb(actorNxt.mu, actorNxt.logStd, aN).logProb
let saN = concat(sn, aN, axis = 0)
# Target Q
let tc1Cache = criticFwdCached(trainer.targetCritic1, saN, tc1H, tc1C)
let tc2Cache = criticFwdCached(trainer.targetCritic2, saN, tc2H, tc2C)
let minQTarg = min(tc1Cache.q, tc2Cache.q)
# Bellman target
let y = r + gamma * d * (minQTarg - alph * lpN)
let errQ1 = c1Cache.q - y
let errQ2 = c2Cache.q - y
totalCriticLoss += 0.5'f32 * (errQ1 * errQ1 + errQ2 * errQ2)
# MSE gradient: d_loss/d_q = (q - y) [scaling applied at accumulation]
addCriticGrads(seqC1Grads, criticBack(trainer.critic1, c1Cache, errQ1))
addCriticGrads(seqC2Grads, criticBack(trainer.critic2, c2Cache, errQ2))
# Advance critic hidden states
c1H = c1Cache.lstm.hPrime; c1C = c1Cache.lstm.cPrime
c2H = c2Cache.lstm.hPrime; c2C = c2Cache.lstm.cPrime
tc1H = tc1Cache.lstm.hPrime; tc1C = tc1Cache.lstm.cPrime
tc2H = tc2Cache.lstm.hPrime; tc2C = tc2Cache.lstm.cPrime
# ── 3 (cont). Actor gradient via critic ─────────────────────────────────
# Q-values for current policy action (critics used as frozen estimators)
let saCurr = concat(s, aCurr, axis = 0)
let qA1Cache = criticFwdCached(trainer.critic1, saCurr, c1H, c1C)
let qA2Cache = criticFwdCached(trainer.critic2, saCurr, c2H, c2C)
let q1Val = qA1Cache.q
let q2Val = qA2Cache.q
let qA1Back = criticBack(trainer.critic1, qA1Cache, -1.0'f32)
let qA2Back = criticBack(trainer.critic2, qA2Cache, -1.0'f32)
totalActorLoss += alph * logProbA - min(q1Val, q2Val)
# Gradient of -minQ w.r.t. action = dInput[stateDim ..< stateDim+actionDim]
# from the critic whose Q was smaller.
let minQBack = if q1Val <= q2Val: qA1Back else: qA2Back
let stateDim = s.shape[0]
let actionDim = aCurr.shape[0]
let dQdA = minQBack.dInput[stateDim ..< stateDim + actionDim]
# Chain through tanh: d(tanh(mu))/d(mu) = 1 - action²
let dTanh = aCurr.map(proc(a: float32): float32 = 1.0'f32 - a * a)
# Total gradient w.r.t. mu: (alpha * dLogP/dMu + dQ/dA) * dTanh/dMu
let dMu = (alph *. lpResult.dLogProbDMu + dQdA) *. dTanh
let dLogStd = alph *. lpResult.dLogProbDLogStd
addActorGrads(seqAGrads, actorBack(trainer.actor, actorFwd, dMu, dLogStd))
# ── 4. Alpha update ──────────────────────────────────────────────────
# Loss = -log_alpha * stop_grad(logProb + targetEntropy)
# d_loss/d_log_alpha = -(logProb + targetEntropy)
totalAlphaLoss += -trainer.logAlpha * (logProbA + trainer.targetEntropy)
seqDLogAlpha += -(logProbA + trainer.targetEntropy)
# Average sequence grads over T steps, accumulate over batch
scaleCriticGrads(seqC1Grads, 1.0'f32 / T)
scaleCriticGrads(seqC2Grads, 1.0'f32 / T)
scaleActorGrads(seqAGrads, 1.0'f32 / T)
addCriticGrads(accC1Grads, seqC1Grads)
addCriticGrads(accC2Grads, seqC2Grads)
addActorGrads(accAGrads, seqAGrads)
dLogAlpha += seqDLogAlpha / T
# Average over batch
scaleCriticGrads(accC1Grads, 1.0'f32 / N)
scaleCriticGrads(accC2Grads, 1.0'f32 / N)
scaleActorGrads(accAGrads, 1.0'f32 / N)
dLogAlpha /= N
# Gradient clipping (max_norm = 1.0)
applyClipToCritic(accC1Grads, 1.0'f32)
applyClipToCritic(accC2Grads, 1.0'f32)
applyClipToActor(accAGrads, 1.0'f32)
# Apply Adam updates
applyCriticAdam(trainer.critic1, accC1Grads, trainer.adam.critic1, trainer.lrCritic)
applyCriticAdam(trainer.critic2, accC2Grads, trainer.adam.critic2, trainer.lrCritic)
applyActorAdam(trainer.actor, accAGrads, trainer.adam.actor, trainer.lrActor)
adamStepScalar(trainer.logAlpha, dLogAlpha, trainer.adam.alpha, trainer.lrAlpha)
# ── 5. Soft target update ──────────────────────────────────────────────────
softUpdateCritic(trainer.targetCritic1, trainer.critic1, trainer.tau)
softUpdateCritic(trainer.targetCritic2, trainer.critic2, trainer.tau)
let totalSteps = N * sequences[0].train.len.float32
result.criticLoss = totalCriticLoss / totalSteps
result.actorLoss = totalActorLoss / totalSteps
result.alphaLoss = totalAlphaLoss / totalSteps
result.alpha = trainer.alpha()
@@ -0,0 +1,352 @@
## weights.nim — save/load all SAC-LSTM network tensors as .npy inside a .zip.
##
## Strategy: write_npy writes to paths; zip/zipfiles.addFile reads from paths.
## So we write each tensor to a temp .npy, add it to the zip, then delete temps.
## Load reverses: extract each entry to a temp .npy, read_npy, delete.
## Atomic save: build the zip in a temp path, then rename over the target.
import arraymancer except Linear
import zip/zipfiles
import std/[os, times, strutils]
import SAC_LSTM_Bot/network
# ── Adam state types (used by training.nim) ───────────────────────────────────
type
AdamVar* = object
m*, v*: Tensor[float32]
t*: int
## Adam states for one Linear layer (w and b).
LinearAdam* = object
w*, b*: AdamVar
## Adam states for one LSTMCell (wCombined and bCombined).
LSTMCellAdam* = object
wCombined*, bCombined*: AdamVar
## Adam states for one ActorNet.
ActorAdam* = object
fc1*, fc2*, muHead*, logStdHead*: LinearAdam
lstm*: LSTMCellAdam
## Adam states for one CriticNet.
CriticAdam* = object
fc1*, fc2*, fc3*: LinearAdam
lstm*: LSTMCellAdam
SACAdamStates* = object
actor*: ActorAdam
critic1*: CriticAdam
critic2*: CriticAdam
alpha*: AdamVar # scalar, shape [1]
initialized*: bool
# ── Init helpers ──────────────────────────────────────────────────────────────
proc initAdamVar(t: Tensor[float32]): AdamVar =
AdamVar(m: zeros[float32](t.shape), v: zeros[float32](t.shape), t: 0)
proc initLinearAdam*(l: Linear): LinearAdam =
LinearAdam(w: initAdamVar(l.w), b: initAdamVar(l.b))
proc initLSTMCellAdam*(c: LSTMCell): LSTMCellAdam =
LSTMCellAdam(
wCombined: initAdamVar(c.wCombined),
bCombined: initAdamVar(c.bCombined))
proc initActorAdam*(a: ActorNet): ActorAdam =
ActorAdam(
fc1: initLinearAdam(a.fc1),
fc2: initLinearAdam(a.fc2),
muHead: initLinearAdam(a.muHead),
logStdHead: initLinearAdam(a.logStdHead),
lstm: initLSTMCellAdam(a.lstm))
proc initCriticAdam*(c: CriticNet): CriticAdam =
CriticAdam(
fc1: initLinearAdam(c.fc1),
fc2: initLinearAdam(c.fc2),
fc3: initLinearAdam(c.fc3),
lstm: initLSTMCellAdam(c.lstm))
proc initSACAdamStates*(actor: ActorNet; critic1, critic2: CriticNet): SACAdamStates =
result.actor = initActorAdam(actor)
result.critic1 = initCriticAdam(critic1)
result.critic2 = initCriticAdam(critic2)
result.alpha = initAdamVar(ones[float32](1))
result.initialized = true
# ── Internal: temp dir per save ───────────────────────────────────────────────
proc tmpDir(): string =
getTempDir() / ("sacw_" & $int(epochTime() * 1000))
# ── Save helpers ──────────────────────────────────────────────────────────────
template addT(z: var ZipArchive; name: string; t: Tensor[float32]; tmp: string) =
## Write tensor to a temp file, add to zip, delete temp file.
let p = tmp / name
t.write_npy(p)
z.addFile(name, p)
proc addLinear(z: var ZipArchive; prefix: string; l: Linear; tmp: string) =
addT(z, prefix & "_w.npy", l.w, tmp)
addT(z, prefix & "_b.npy", l.b, tmp)
proc addLSTMCell(z: var ZipArchive; prefix: string; c: LSTMCell; tmp: string) =
addT(z, prefix & "_wc.npy", c.wCombined, tmp)
addT(z, prefix & "_bc.npy", c.bCombined, tmp)
proc addActorNet(z: var ZipArchive; prefix: string; a: ActorNet; tmp: string) =
addLinear(z, prefix & "_fc1", a.fc1, tmp)
addLSTMCell(z, prefix & "_lstm", a.lstm, tmp)
addLinear(z, prefix & "_fc2", a.fc2, tmp)
addLinear(z, prefix & "_mu", a.muHead, tmp)
addLinear(z, prefix & "_logstd", a.logStdHead, tmp)
proc addCriticNet(z: var ZipArchive; prefix: string; c: CriticNet; tmp: string) =
addLinear(z, prefix & "_fc1", c.fc1, tmp)
addLSTMCell(z, prefix & "_lstm", c.lstm, tmp)
addLinear(z, prefix & "_fc2", c.fc2, tmp)
addLinear(z, prefix & "_fc3", c.fc3, tmp)
proc addAdamVar(z: var ZipArchive; prefix: string; v: AdamVar; tmp: string) =
addT(z, prefix & "_m.npy", v.m, tmp)
addT(z, prefix & "_v.npy", v.v, tmp)
proc addLinearAdam(z: var ZipArchive; prefix: string; la: LinearAdam; tmp: string) =
addAdamVar(z, prefix & "_w", la.w, tmp)
addAdamVar(z, prefix & "_b", la.b, tmp)
proc addLSTMCellAdam(z: var ZipArchive; prefix: string; la: LSTMCellAdam; tmp: string) =
addAdamVar(z, prefix & "_wc", la.wCombined, tmp)
addAdamVar(z, prefix & "_bc", la.bCombined, tmp)
proc addActorAdam(z: var ZipArchive; prefix: string; a: ActorAdam; tmp: string) =
addLinearAdam(z, prefix & "_fc1", a.fc1, tmp)
addLSTMCellAdam(z, prefix & "_lstm", a.lstm, tmp)
addLinearAdam(z, prefix & "_fc2", a.fc2, tmp)
addLinearAdam(z, prefix & "_mu", a.muHead, tmp)
addLinearAdam(z, prefix & "_logstd", a.logStdHead, tmp)
proc addCriticAdam(z: var ZipArchive; prefix: string; c: CriticAdam; tmp: string) =
addLinearAdam(z, prefix & "_fc1", c.fc1, tmp)
addLSTMCellAdam(z, prefix & "_lstm", c.lstm, tmp)
addLinearAdam(z, prefix & "_fc2", c.fc2, tmp)
addLinearAdam(z, prefix & "_fc3", c.fc3, tmp)
# ── Public API ────────────────────────────────────────────────────────────────
proc saveWeights*(path: string;
actor: ActorNet;
critic1, critic2: CriticNet;
targetCritic1, targetCritic2: CriticNet;
alpha: float32) =
## Save network tensors (no Adam states) to `path` (.zip).
## Atomic: writes to a temp path first, then renames.
let tmp = tmpDir()
createDir(tmp)
let tmpZip = path & ".tmp"
try:
var z: ZipArchive
if not z.open(tmpZip, fmWrite):
raise newException(IOError, "cannot create zip: " & tmpZip)
addActorNet(z, "actor", actor, tmp)
addCriticNet(z, "c1", critic1, tmp)
addCriticNet(z, "c2", critic2, tmp)
addCriticNet(z, "tc1", targetCritic1,tmp)
addCriticNet(z, "tc2", targetCritic2,tmp)
# alpha: store as a 1-element tensor
addT(z, "alpha.npy", [alpha].toTensor.asType(float32), tmp)
z.close()
createDir(path.parentDir)
moveFile(tmpZip, path)
finally:
removeDir(tmp)
if fileExists(tmpZip): removeFile(tmpZip)
proc saveCheckpoint*(path: string;
actor: ActorNet;
critic1, critic2: CriticNet;
targetCritic1, targetCritic2: CriticNet;
alpha: float32;
adam: SACAdamStates) =
## Save networks + Adam states to `path` (.zip). Atomic.
let tmp = tmpDir()
createDir(tmp)
let tmpZip = path & ".tmp"
try:
var z: ZipArchive
if not z.open(tmpZip, fmWrite):
raise newException(IOError, "cannot create zip: " & tmpZip)
addActorNet(z, "actor", actor, tmp)
addCriticNet(z, "c1", critic1, tmp)
addCriticNet(z, "c2", critic2, tmp)
addCriticNet(z, "tc1", targetCritic1,tmp)
addCriticNet(z, "tc2", targetCritic2,tmp)
addT(z, "alpha.npy", [alpha].toTensor.asType(float32), tmp)
if adam.initialized:
addActorAdam(z, "adam_actor", adam.actor, tmp)
addCriticAdam(z, "adam_c1", adam.critic1, tmp)
addCriticAdam(z, "adam_c2", adam.critic2, tmp)
addAdamVar(z, "adam_alpha", adam.alpha, tmp)
# t counters (all in lockstep; store as text)
writeFile(tmp / "adam_t.txt",
$adam.actor.fc1.w.t & "\n" &
$adam.critic1.fc1.w.t & "\n" &
$adam.critic2.fc1.w.t & "\n" &
$adam.alpha.t)
z.addFile("adam_t.txt", tmp / "adam_t.txt")
z.close()
createDir(path.parentDir)
moveFile(tmpZip, path)
finally:
removeDir(tmp)
if fileExists(tmpZip): removeFile(tmpZip)
# ── Load helpers ──────────────────────────────────────────────────────────────
template loadT(name: string; tmp: string): Tensor[float32] =
read_npy[float32](tmp / name)
proc loadLinear(z: var ZipArchive; prefix, tmp: string): Linear =
z.extractFile(prefix & "_w.npy", tmp / (prefix & "_w.npy"))
z.extractFile(prefix & "_b.npy", tmp / (prefix & "_b.npy"))
result.w = read_npy[float32](tmp / (prefix & "_w.npy"))
result.b = read_npy[float32](tmp / (prefix & "_b.npy"))
proc loadLSTMCell(z: var ZipArchive; prefix, tmp: string): LSTMCell =
z.extractFile(prefix & "_wc.npy", tmp / (prefix & "_wc.npy"))
z.extractFile(prefix & "_bc.npy", tmp / (prefix & "_bc.npy"))
result.wCombined = read_npy[float32](tmp / (prefix & "_wc.npy"))
result.bCombined = read_npy[float32](tmp / (prefix & "_bc.npy"))
result.hiddenDim = result.bCombined.shape[0] div 4
proc loadActorNet(z: var ZipArchive; prefix, tmp: string): ActorNet =
result.fc1 = loadLinear(z, prefix & "_fc1", tmp)
result.lstm = loadLSTMCell(z, prefix & "_lstm", tmp)
result.fc2 = loadLinear(z, prefix & "_fc2", tmp)
result.muHead = loadLinear(z, prefix & "_mu", tmp)
result.logStdHead = loadLinear(z, prefix & "_logstd", tmp)
result.hiddenDim = result.lstm.hiddenDim
proc loadCriticNet(z: var ZipArchive; prefix, tmp: string): CriticNet =
result.fc1 = loadLinear(z, prefix & "_fc1", tmp)
result.lstm = loadLSTMCell(z, prefix & "_lstm", tmp)
result.fc2 = loadLinear(z, prefix & "_fc2", tmp)
result.fc3 = loadLinear(z, prefix & "_fc3", tmp)
result.hiddenDim = result.lstm.hiddenDim
proc loadAdamVarFromZip(z: var ZipArchive; prefix, tmp: string): AdamVar =
z.extractFile(prefix & "_m.npy", tmp / (prefix & "_m.npy"))
z.extractFile(prefix & "_v.npy", tmp / (prefix & "_v.npy"))
result.m = read_npy[float32](tmp / (prefix & "_m.npy"))
result.v = read_npy[float32](tmp / (prefix & "_v.npy"))
proc loadLinearAdam(z: var ZipArchive; prefix, tmp: string): LinearAdam =
result.w = loadAdamVarFromZip(z, prefix & "_w", tmp)
result.b = loadAdamVarFromZip(z, prefix & "_b", tmp)
proc loadLSTMCellAdam(z: var ZipArchive; prefix, tmp: string): LSTMCellAdam =
result.wCombined = loadAdamVarFromZip(z, prefix & "_wc", tmp)
result.bCombined = loadAdamVarFromZip(z, prefix & "_bc", tmp)
proc loadActorAdam(z: var ZipArchive; prefix, tmp: string): ActorAdam =
result.fc1 = loadLinearAdam(z, prefix & "_fc1", tmp)
result.lstm = loadLSTMCellAdam(z, prefix & "_lstm", tmp)
result.fc2 = loadLinearAdam(z, prefix & "_fc2", tmp)
result.muHead = loadLinearAdam(z, prefix & "_mu", tmp)
result.logStdHead = loadLinearAdam(z, prefix & "_logstd", tmp)
proc loadCriticAdam(z: var ZipArchive; prefix, tmp: string): CriticAdam =
result.fc1 = loadLinearAdam(z, prefix & "_fc1", tmp)
result.lstm = loadLSTMCellAdam(z, prefix & "_lstm", tmp)
result.fc2 = loadLinearAdam(z, prefix & "_fc2", tmp)
result.fc3 = loadLinearAdam(z, prefix & "_fc3", tmp)
type
WeightCheckpoint* = object
actor*: ActorNet
critic1*: CriticNet
critic2*: CriticNet
targetCritic1*: CriticNet
targetCritic2*: CriticNet
alpha*: float32
adam*: SACAdamStates ## initialized=false if not present in zip
proc loadCheckpoint*(path: string): WeightCheckpoint =
## Load all tensors from `path` (.zip). Raises IOError if file not found.
## Adam states loaded only if present; result.adam.initialized reflects this.
if not fileExists(path):
raise newException(IOError, "checkpoint not found: " & path)
let tmp = tmpDir()
createDir(tmp)
try:
var z: ZipArchive
if not z.open(path, fmRead):
raise newException(IOError, "cannot open zip: " & path)
result.actor = loadActorNet(z, "actor", tmp)
result.critic1 = loadCriticNet(z, "c1", tmp)
result.critic2 = loadCriticNet(z, "c2", tmp)
result.targetCritic1 = loadCriticNet(z, "tc1", tmp)
result.targetCritic2 = loadCriticNet(z, "tc2", tmp)
z.extractFile("alpha.npy", tmp / "alpha.npy")
let alphaTensor = read_npy[float32](tmp / "alpha.npy")
result.alpha = alphaTensor[0]
# Adam states — optional
var hasAdam = false
for f in z.walkFiles:
if f.startsWith("adam_"):
hasAdam = true
break
if hasAdam:
result.adam.actor = loadActorAdam(z, "adam_actor", tmp)
result.adam.critic1 = loadCriticAdam(z, "adam_c1", tmp)
result.adam.critic2 = loadCriticAdam(z, "adam_c2", tmp)
result.adam.alpha = loadAdamVarFromZip(z, "adam_alpha", tmp)
# t counters
z.extractFile("adam_t.txt", tmp / "adam_t.txt")
let ts = readFile(tmp / "adam_t.txt").strip().splitLines()
if ts.len >= 4:
let tActor = parseInt(ts[0])
let tCritic1 = parseInt(ts[1])
let tCritic2 = parseInt(ts[2])
let tAlpha = parseInt(ts[3])
# propagate t to all Adam vars
template setT(v: var AdamVar; tval: int) = v.t = tval
setT(result.adam.actor.fc1.w, tActor)
setT(result.adam.actor.fc1.b, tActor)
setT(result.adam.actor.lstm.wCombined,tActor)
setT(result.adam.actor.lstm.bCombined,tActor)
setT(result.adam.actor.fc2.w, tActor)
setT(result.adam.actor.fc2.b, tActor)
setT(result.adam.actor.muHead.w, tActor)
setT(result.adam.actor.muHead.b, tActor)
setT(result.adam.actor.logStdHead.w, tActor)
setT(result.adam.actor.logStdHead.b, tActor)
setT(result.adam.critic1.fc1.w, tCritic1)
setT(result.adam.critic1.fc1.b, tCritic1)
setT(result.adam.critic1.lstm.wCombined,tCritic1)
setT(result.adam.critic1.lstm.bCombined,tCritic1)
setT(result.adam.critic1.fc2.w, tCritic1)
setT(result.adam.critic1.fc2.b, tCritic1)
setT(result.adam.critic1.fc3.w, tCritic1)
setT(result.adam.critic1.fc3.b, tCritic1)
setT(result.adam.critic2.fc1.w, tCritic2)
setT(result.adam.critic2.fc1.b, tCritic2)
setT(result.adam.critic2.lstm.wCombined,tCritic2)
setT(result.adam.critic2.lstm.bCombined,tCritic2)
setT(result.adam.critic2.fc2.w, tCritic2)
setT(result.adam.critic2.fc2.b, tCritic2)
setT(result.adam.critic2.fc3.w, tCritic2)
setT(result.adam.critic2.fc3.b, tCritic2)
setT(result.adam.alpha, tAlpha)
result.adam.initialized = true
z.close()
finally:
removeDir(tmp)