fix(GDN): clamp gate [-5,2] + state [-65504,65504] to prevent inf/NaN

Root cause from real machine test: gdn_forward.cu output abs mean = inf
- gate_raw can be positive → exp(gate) > 1 → state grows exponentially
- Over 64 tokens: exp(2.0)^64 = inf
- PyTorch ref clamps g ∈ [-5, 2] but CUDA kernel did not

Fix:
  gdn_forward.cu: clamp gate_raw ∈ [-5, 2] before exp (both kernel variants)
  gdn_forward.cu: clamp state ∈ [-65504, 65504] after update (fp16 safe range)
  qwen3_5.py: clamp g_3d before passing to SM70 kernel (belt + suspenders)
  qwen3_5.py: clamp temporal_state after decode update
This commit is contained in:
EX Engine
2026-08-10 03:38:39 +00:00
parent 388f6b2d1a
commit 8eba1750fa
2 changed files with 15 additions and 6 deletions

View File

@@ -581,6 +581,13 @@ class GatedDeltaNet(nn.Module):
k_4d = k.unsqueeze(0) # (1, L, Hk, K)
v_4d = v_raw.unsqueeze(0) # (1, L, Hv, V)
g_3d = gate.unsqueeze(0) # (1, L, Hv)
# Clamp gate to prevent exp() overflow in CUDA kernel.
# gate = -dt * A_log.exp(), typically negative (decay).
# But pathological weights can produce positive values → exp > 1
# → state grows exponentially over L tokens → inf.
# PyTorch ref clamps g ∈ [-5, 2] before cumsum.
# For recurrent kernel: clamp raw gate so exp(gate) ∈ [exp(-5), exp(2)]
g_3d = g_3d.clamp(-5.0, 2.0)
beta_3d = b_seq.unsqueeze(0) # (1, L, Hv)
# Initial state from temporal_state
@@ -800,6 +807,8 @@ class GatedDeltaNet(nn.Module):
k_t.view(BH, self.head_k_dim, 1),
delta.view(BH, 1, self.head_v_dim),
)
# Clamp state to prevent gradual drift → NaN over long sequences
temporal_state.clamp_(-65504.0, 65504.0)
# Output: core_out = q_t @ updated temporal_state
core_out = _ix_bmm(