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:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user