From 96065597c0d7d7e15efd6e97bbd2681e1edb3df9 Mon Sep 17 00:00:00 2001 From: Davide Cappellini Date: Mon, 24 Aug 2026 01:03:04 +0200 Subject: [PATCH] fix(divergence): guard reward-normalizer cold start; drop arctanh recovery MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Welford variance-collapse divided by 1e-8 producing bit-exact +/−5e8 / +−1.25e8 poisoned rewards into TD targets; guard skips normalization until stats meaningful; tanh-inversion removal bounds log-prob path. Fixes #60. --- SAC_LSTM_Bot/src/SAC_LSTM_Bot/rewards.nim | 16 ++++++++++++---- SAC_LSTM_Bot/src/SAC_LSTM_Bot/training.nim | 5 +---- 2 files changed, 13 insertions(+), 8 deletions(-) diff --git a/SAC_LSTM_Bot/src/SAC_LSTM_Bot/rewards.nim b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/rewards.nim index a66ba05..83218af 100644 --- a/SAC_LSTM_Bot/src/SAC_LSTM_Bot/rewards.nim +++ b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/rewards.nim @@ -69,8 +69,16 @@ proc update*(rn: var RewardNormalizer; r: float64) = rn.m2 += delta * delta2 proc normalize*(rn: RewardNormalizer; r: float64): float64 = - ## Returns (r - mean) / (std + eps). - ## Cold start (n < 2): returns 0.0 to avoid NaN/inf. - if rn.n < 2: return 0.0 + ## Returns (r - mean) / (std + eps) once statistics are meaningful + ## (n >= 4 and spread well above zero). Before that, returns the RAW + ## reward unchanged — Welford M2 collapses to exactly 0 when early raw + ## rewards are identical, and dividing by the 1e-8 floor then z-scores + ## the first differing reward to ~1e8, poisoning TD targets. + if rn.n < 4: return r let variance = rn.m2 / rn.n.float64 # ponytail: population var; switch to n-1 if bias matters - result = (r - rn.mean) / (sqrt(variance) + NormEps) + let stddev = sqrt(variance) + if stddev <= 1e-3 * (abs(rn.mean) + 1.0): return r + # ponytail: warm-up pass-through ceiling — raw rewards bypass normalization + # until stats are meaningful; upgrade = persist Welford state in checkpoint + # if warm-up noise ever hurts learning. + result = (r - rn.mean) / (stddev + NormEps) diff --git a/SAC_LSTM_Bot/src/SAC_LSTM_Bot/training.nim b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/training.nim index ef281e4..10e7056 100644 --- a/SAC_LSTM_Bot/src/SAC_LSTM_Bot/training.nim +++ b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/training.nim @@ -200,10 +200,7 @@ proc squashedLogProb(mu, logStd, action: Tensor[float32]): tuple[logProb: float32; dLogProbDMu, dLogProbDLogStd: Tensor[float32]] = let std = logStd.map(proc(v: float32): float32 = exp(v)) - # Recover pre-tanh z ≈ arctanh(action) - let z = action.map(proc(a: float32): float32 = - let ac = clamp(a, -1.0'f32 + 1e-6'f32, 1.0'f32 - 1e-6'f32) - 0.5'f32 * ln((1.0'f32 + ac) / (1.0'f32 - ac))) + let z = mu # deterministic reparam: action = tanh(mu), so z ≡ mu, diff ≡ 0 result.dLogProbDMu = newTensor[float32](4) result.dLogProbDLogStd = newTensor[float32](4) let twoPiLog = 0.5'f32 * ln(2.0'f32 * PI.float32)