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)