From 62a6cc8ccf42f9fc40e33e73446b4389e4020d30 Mon Sep 17 00:00:00 2001 From: Davide Cappellini Date: Fri, 21 Aug 2026 00:16:47 +0200 Subject: [PATCH] =?UTF-8?q?fix(SAC=5FLSTM=5FBot):=20training=20review=20fi?= =?UTF-8?q?xes=20=E2=80=94=20hidden=20state=20ordering,=20redundant=20forw?= =?UTF-8?q?ards,=20actor=20grad=20clip=20(#47)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-Authored-By: Claude Sonnet 4.6 --- SAC_LSTM_Bot/src/SAC_LSTM_Bot/training.nim | 53 +++++++++++++--------- 1 file changed, 31 insertions(+), 22 deletions(-) diff --git a/SAC_LSTM_Bot/src/SAC_LSTM_Bot/training.nim b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/training.nim index 9b8e791..ef281e4 100644 --- a/SAC_LSTM_Bot/src/SAC_LSTM_Bot/training.nim +++ b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/training.nim @@ -425,6 +425,16 @@ proc applyClipToCritic(g: var CriticGrads; maxNorm: float32) = 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 = @@ -496,7 +506,17 @@ proc sacUpdate*(trainer: var SACTrainer; sequences: seq[Sequence]): SACMetrics = let c1Cache = criticFwdCached(trainer.critic1, sa, c1H, c1C) let c2Cache = criticFwdCached(trainer.critic2, sa, c2H, c2C) - # Next-state action from current actor + # ── 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 @@ -523,25 +543,16 @@ proc sacUpdate*(trainer: var SACTrainer; sequences: seq[Sequence]): SACMetrics = 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 + # ── 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 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 + 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] @@ -560,9 +571,6 @@ proc sacUpdate*(trainer: var SACTrainer; sequences: seq[Sequence]): SACMetrics = 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) @@ -584,9 +592,10 @@ proc sacUpdate*(trainer: var SACTrainer; sequences: seq[Sequence]): SACMetrics = scaleActorGrads(accAGrads, 1.0'f32 / N) dLogAlpha /= N - # Gradient clipping on critics (max_norm = 1.0) + # 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)