From 415d4e373880cdc5f7e5fb55efbb13a451197b96 Mon Sep 17 00:00:00 2001 From: Davide Cappellini Date: Fri, 21 Aug 2026 00:11:35 +0200 Subject: [PATCH] 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 --- SAC_LSTM_Bot/src/SAC_LSTM_Bot/training.nim | 605 +++++++++++++++++++++ SAC_LSTM_Bot/tests/test_training.nim | 159 ++++++ 2 files changed, 764 insertions(+) create mode 100644 SAC_LSTM_Bot/src/SAC_LSTM_Bot/training.nim create mode 100644 SAC_LSTM_Bot/tests/test_training.nim diff --git a/SAC_LSTM_Bot/src/SAC_LSTM_Bot/training.nim b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/training.nim new file mode 100644 index 0000000..9b8e791 --- /dev/null +++ b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/training.nim @@ -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() diff --git a/SAC_LSTM_Bot/tests/test_training.nim b/SAC_LSTM_Bot/tests/test_training.nim new file mode 100644 index 0000000..fb83725 --- /dev/null +++ b/SAC_LSTM_Bot/tests/test_training.nim @@ -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