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
|
rn.m2 += delta * delta2
|
||||||
|
|
||||||
proc normalize*(rn: RewardNormalizer; r: float64): float64 =
|
proc normalize*(rn: RewardNormalizer; r: float64): float64 =
|
||||||
## Returns (r - mean) / (std + eps).
|
## Returns (r - mean) / (std + eps) once statistics are meaningful
|
||||||
## Cold start (n < 2): returns 0.0 to avoid NaN/inf.
|
## (n >= 4 and spread well above zero). Before that, returns the RAW
|
||||||
if rn.n < 2: return 0.0
|
## 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
|
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;
|
tuple[logProb: float32;
|
||||||
dLogProbDMu, dLogProbDLogStd: Tensor[float32]] =
|
dLogProbDMu, dLogProbDLogStd: Tensor[float32]] =
|
||||||
let std = logStd.map(proc(v: float32): float32 = exp(v))
|
let std = logStd.map(proc(v: float32): float32 = exp(v))
|
||||||
# Recover pre-tanh z ≈ arctanh(action)
|
let z = mu # deterministic reparam: action = tanh(mu), so z ≡ mu, diff ≡ 0
|
||||||
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)))
|
|
||||||
result.dLogProbDMu = newTensor[float32](4)
|
result.dLogProbDMu = newTensor[float32](4)
|
||||||
result.dLogProbDLogStd = newTensor[float32](4)
|
result.dLogProbDLogStd = newTensor[float32](4)
|
||||||
let twoPiLog = 0.5'f32 * ln(2.0'f32 * PI.float32)
|
let twoPiLog = 0.5'f32 * ln(2.0'f32 * PI.float32)
|
||||||
|
|||||||
Reference in New Issue
Block a user