fix(divergence): guard reward-normalizer cold start; drop arctanh recovery
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.
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user