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