fix(SAC_LSTM_Bot): training review fixes — hidden state ordering, redundant forwards, actor grad clip (#47)

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-08-21 00:16:47 +02:00
parent 415d4e3738
commit 62a6cc8ccf
+31 -22
View File
@@ -425,6 +425,16 @@ proc applyClipToCritic(g: var CriticGrads; maxNorm: float32) =
g.dFc3W = gs[4]; g.dFc3B = gs[5] g.dFc3W = gs[4]; g.dFc3B = gs[5]
g.dLstmW = gs[6]; g.dLstmB = gs[7] 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 ──────────────────────────────────────────────────────────────── # ── SAC update ────────────────────────────────────────────────────────────────
proc sacUpdate*(trainer: var SACTrainer; sequences: seq[Sequence]): SACMetrics = 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 c1Cache = criticFwdCached(trainer.critic1, sa, c1H, c1C)
let c2Cache = criticFwdCached(trainer.critic2, sa, c2H, c2C) 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 actorNxt = actorFwdCached(trainer.actor, sn, actorH, actorC)
let aN = actorNxt.action let aN = actorNxt.action
let lpN = squashedLogProb(actorNxt.mu, actorNxt.logStd, aN).logProb 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 tc1H = tc1Cache.lstm.hPrime; tc1C = tc1Cache.lstm.cPrime
tc2H = tc2Cache.lstm.hPrime; tc2C = tc2Cache.lstm.cPrime tc2H = tc2Cache.lstm.hPrime; tc2C = tc2Cache.lstm.cPrime
# ── 3. Actor update ────────────────────────────────────────────────── # ── 3 (cont). Actor gradient via critic ─────────────────────────────────
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) # Q-values for current policy action (critics used as frozen estimators)
let saCurr = concat(s, aCurr, axis = 0) let saCurr = concat(s, aCurr, axis = 0)
let qA1Back = criticBack(trainer.critic1, let qA1Cache = criticFwdCached(trainer.critic1, saCurr, c1H, c1C)
criticFwdCached(trainer.critic1, saCurr, c1H, c1C), let qA2Cache = criticFwdCached(trainer.critic2, saCurr, c2H, c2C)
-1.0'f32) # d_loss/d_q = -1 (maximize Q) let q1Val = qA1Cache.q
let qA2Back = criticBack(trainer.critic2, let q2Val = qA2Cache.q
criticFwdCached(trainer.critic2, saCurr, c2H, c2C), let qA1Back = criticBack(trainer.critic1, qA1Cache, -1.0'f32)
-1.0'f32) let qA2Back = criticBack(trainer.critic2, qA2Cache, -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) totalActorLoss += alph * logProbA - min(q1Val, q2Val)
# Gradient of -minQ w.r.t. action = dInput[stateDim ..< stateDim+actionDim] # 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)) addActorGrads(seqAGrads, actorBack(trainer.actor, actorFwd, dMu, dLogStd))
# Advance actor hidden state
actorH = actorFwd.lstm.hPrime; actorC = actorFwd.lstm.cPrime
# ── 4. Alpha update ────────────────────────────────────────────────── # ── 4. Alpha update ──────────────────────────────────────────────────
# Loss = -log_alpha * stop_grad(logProb + targetEntropy) # Loss = -log_alpha * stop_grad(logProb + targetEntropy)
# d_loss/d_log_alpha = -(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) scaleActorGrads(accAGrads, 1.0'f32 / N)
dLogAlpha /= N dLogAlpha /= N
# Gradient clipping on critics (max_norm = 1.0) # Gradient clipping (max_norm = 1.0)
applyClipToCritic(accC1Grads, 1.0'f32) applyClipToCritic(accC1Grads, 1.0'f32)
applyClipToCritic(accC2Grads, 1.0'f32) applyClipToCritic(accC2Grads, 1.0'f32)
applyClipToActor(accAGrads, 1.0'f32)
# Apply Adam updates # Apply Adam updates
applyCriticAdam(trainer.critic1, accC1Grads, trainer.adam.critic1, trainer.lrCritic) applyCriticAdam(trainer.critic1, accC1Grads, trainer.adam.critic1, trainer.lrCritic)