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:
@@ -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"
|
||||
}
|
||||
@@ -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
|
||||
Executable
BIN
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()
|
||||
Executable
BIN
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)
|
||||
Executable
BIN
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))
|
||||
BIN
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)
|
||||
Reference in New Issue
Block a user