feat(SAC_LSTM_Bot): SAC training module (#47)

Implements sacUpdate with burn-in LSTM warm-up, twin-critic TD update,
actor reparameterization gradient, auto-alpha, and soft target update.
Manual backprop (linear + LSTM single-step, truncated BPTT). 11 new tests
all green; full regression suite (57+ tests) unaffected.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-08-21 00:11:35 +02:00
parent 717ef3ead8
commit 415d4e3738
2 changed files with 764 additions and 0 deletions
+605
View File
@@ -0,0 +1,605 @@
## 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))
# Recover pre-tanh z ≈ arctanh(action)
let z = action.map(proc(a: float32): float32 =
let ac = clamp(a, -1.0'f32 + 1e-6'f32, 1.0'f32 - 1e-6'f32)
0.5'f32 * ln((1.0'f32 + ac) / (1.0'f32 - ac)))
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]
# ── 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)
# Next-state action from current actor
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. Actor update ──────────────────────────────────────────────────
let actorFwd = actorFwdCached(trainer.actor, s, actorH, actorC)
let aCurr = actorFwd.action
let lpResult = squashedLogProb(actorFwd.mu, actorFwd.logStd, aCurr)
let logProbA = lpResult.logProb
# Q-values for current policy action (critics used as frozen estimators)
let saCurr = concat(s, aCurr, axis = 0)
let qA1Back = criticBack(trainer.critic1,
criticFwdCached(trainer.critic1, saCurr, c1H, c1C),
-1.0'f32) # d_loss/d_q = -1 (maximize Q)
let qA2Back = criticBack(trainer.critic2,
criticFwdCached(trainer.critic2, saCurr, c2H, c2C),
-1.0'f32)
# Choose min-Q critic gradient (clipped double-Q actor update)
let q1Val = criticFwdCached(trainer.critic1, saCurr, c1H, c1C).q
let q2Val = criticFwdCached(trainer.critic2, saCurr, c2H, c2C).q
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))
# Advance actor hidden state
actorH = actorFwd.lstm.hPrime; actorC = actorFwd.lstm.cPrime
# ── 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 on critics (max_norm = 1.0)
applyClipToCritic(accC1Grads, 1.0'f32)
applyClipToCritic(accC2Grads, 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()
+159
View File
@@ -0,0 +1,159 @@
## test_training.nim — stdlib unittest for SAC training module.
import unittest
import arraymancer
import std/[math, random, os]
import SAC_LSTM_Bot/network
import SAC_LSTM_Bot/weights
import SAC_LSTM_Bot/replay_buffer
import SAC_LSTM_Bot/training
const
STATE_DIM = 10
ACTION_DIM = 4
HIDDEN = 16 # small for speed; override via env not needed in tests
proc makeTrainer(): SACTrainer =
## Small trainer for tests — override hidden size via env before calling.
putEnv("SACLSTM_HIDDEN_SIZE", "16")
initSACTrainer(STATE_DIM, ACTION_DIM)
proc makeBuffer(): ReplayBuffer =
newReplayBuffer(capacity = 2000, stateDim = STATE_DIM, actionDim = ACTION_DIM,
burnIn = 4, trainWindow = 8)
proc randState(): Tensor[float32] =
randomNormalTensor[float32](STATE_DIM)
proc randAction(): Tensor[float32] =
randomTensor[float32](ACTION_DIM, 1.0'f32) *. 2.0'f32 -. 1.0'f32 # uniform [-1,1]
proc fillBuffer(buf: var ReplayBuffer; n: int; donePeriod = 20) =
for i in 0 ..< n:
let t = Transition(
state: randState(),
action: randAction(),
reward: rand(-1.0'f32 .. 1.0'f32),
nextState: randState(),
done: (i mod donePeriod == donePeriod - 1))
buf.add(t)
proc isFinite(x: float32): bool =
not (x != x) and x < Inf and x > -Inf # not NaN and not Inf
suite "SACTrainer — basic update":
setup:
randomize(42)
var trainer = makeTrainer()
var buf = makeBuffer()
fillBuffer(buf, 1000)
let seqs = buf.sampleSequences(4)
test "sampleSequences returns non-empty batch":
check seqs.len > 0
test "sacUpdate returns finite losses":
let m = sacUpdate(trainer, seqs)
check isFinite(m.criticLoss)
check isFinite(m.actorLoss)
check isFinite(m.alphaLoss)
test "alpha stays positive after update":
var t2 = trainer
discard sacUpdate(t2, seqs)
check t2.alpha() > 0.0'f32
test "critic1 weights change after update":
let w_before = trainer.critic1.fc3.w.clone()
discard sacUpdate(trainer, seqs)
let w_after = trainer.critic1.fc3.w
var changed = false
for i in 0 ..< w_before.size:
if abs(w_before.unsafe_raw_offset[i] - w_after.unsafe_raw_offset[i]) > 1e-10'f32:
changed = true
break
check changed
test "critic2 weights change after update":
let w_before = trainer.critic2.fc3.w.clone()
discard sacUpdate(trainer, seqs)
let w_after = trainer.critic2.fc3.w
var changed = false
for i in 0 ..< w_before.size:
if abs(w_before.unsafe_raw_offset[i] - w_after.unsafe_raw_offset[i]) > 1e-10'f32:
changed = true
break
check changed
test "actor weights change after update":
let w_before = trainer.actor.fc1.w.clone()
discard sacUpdate(trainer, seqs)
let w_after = trainer.actor.fc1.w
var changed = false
for i in 0 ..< w_before.size:
if abs(w_before.unsafe_raw_offset[i] - w_after.unsafe_raw_offset[i]) > 1e-10'f32:
changed = true
break
check changed
test "soft target update: target moves toward critic":
## After update, target fc3.w should be closer to critic1.fc3.w than before.
let targetBefore = trainer.targetCritic1.fc3.w.clone()
let criticW = trainer.critic1.fc3.w.clone()
discard sacUpdate(trainer, seqs)
let targetAfter = trainer.targetCritic1.fc3.w
# Distance before: ||targetBefore - criticW||
var distBefore = 0.0'f32
for i in 0 ..< targetBefore.size:
let d = targetBefore.unsafe_raw_offset[i] - criticW.unsafe_raw_offset[i]
distBefore += d * d
# Distance after: ||targetAfter - criticW_new|| (critic changed too, use original for ref)
var distAfter = 0.0'f32
for i in 0 ..< targetAfter.size:
let d = targetAfter.unsafe_raw_offset[i] - criticW.unsafe_raw_offset[i]
distAfter += d * d
# Target moved toward critic (distance decreased).
# With tau=0.005, it moves a tiny bit — just check direction.
check distAfter <= distBefore + 1e-3'f32 # soft bound; critic also moves
test "empty sequences is a no-op":
let m = sacUpdate(trainer, @[])
check m.criticLoss == 0.0'f32
check m.actorLoss == 0.0'f32
check m.alpha == 0.0'f32
suite "SACTrainer — done=true terminal transitions":
test "update with done=true transitions produces finite losses":
randomize(7)
putEnv("SACLSTM_HIDDEN_SIZE", "16")
var trainer = makeTrainer()
var buf = makeBuffer()
# Fill with short episodes: done every 12 steps (burnIn=4, trainWindow=8 → seqLen=12)
fillBuffer(buf, 800, donePeriod = 12)
let seqs = buf.sampleSequences(2)
if seqs.len > 0:
let m = sacUpdate(trainer, seqs)
check isFinite(m.criticLoss)
check isFinite(m.actorLoss)
check isFinite(m.alphaLoss)
check m.alpha > 0.0'f32
suite "SACTrainer — multiple updates":
test "three sequential updates all stay finite":
randomize(99)
putEnv("SACLSTM_HIDDEN_SIZE", "16")
var trainer = makeTrainer()
var buf = makeBuffer()
fillBuffer(buf, 1000)
for _ in 1..3:
let seqs = buf.sampleSequences(4)
if seqs.len > 0:
let m = sacUpdate(trainer, seqs)
check isFinite(m.criticLoss)
check isFinite(m.actorLoss)
check m.alpha > 0.0'f32