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:
@@ -135,7 +135,7 @@ __global__ void gdn_forward_kernel(const scalar_t* __restrict__ q,
|
|||||||
float beta_value = 0.0F;
|
float beta_value = 0.0F;
|
||||||
if (threadIdx.x == 0) {
|
if (threadIdx.x == 0) {
|
||||||
const float gate_raw = load_as_float(gate, gate_index);
|
const float gate_raw = load_as_float(gate, gate_index);
|
||||||
gate_value = GateIsExp ? gate_raw : __expf(gate_raw);
|
{ const float gc = fminf(fmaxf(gate_raw, -5.0F), 2.0F); gate_value = GateIsExp ? fminf(gate_raw, 7.389F) : __expf(gc); }
|
||||||
beta_value = load_as_float(beta, gate_index);
|
beta_value = load_as_float(beta, gate_index);
|
||||||
}
|
}
|
||||||
gate_value = __shfl_sync(0xffffffffU, gate_value, 0);
|
gate_value = __shfl_sync(0xffffffffU, gate_value, 0);
|
||||||
@@ -191,7 +191,7 @@ __global__ void gdn_forward_kernel(const scalar_t* __restrict__ q,
|
|||||||
for (int c = 0; c < COLS; ++c) {
|
for (int c = 0; c < COLS; ++c) {
|
||||||
const float new_state =
|
const float new_state =
|
||||||
fmaf(k_reg[r], delta[c], gate_value * state_shard[c][r]);
|
fmaf(k_reg[r], delta[c], gate_value * state_shard[c][r]);
|
||||||
state_shard[c][r] = new_state;
|
state_shard[c][r] = fminf(fmaxf(new_state, -65504.0F), 65504.0F);
|
||||||
attn_partial[c] += new_state * q_reg[r];
|
attn_partial[c] += new_state * q_reg[r];
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -294,7 +294,7 @@ __global__ void gdn_forward_vlk_varlen_kernel(
|
|||||||
float beta_value = 0.0F;
|
float beta_value = 0.0F;
|
||||||
if (threadIdx.x == 0) {
|
if (threadIdx.x == 0) {
|
||||||
const float gate_raw = load_as_float(gate, gate_index);
|
const float gate_raw = load_as_float(gate, gate_index);
|
||||||
gate_value = GateIsExp ? gate_raw : __expf(gate_raw);
|
{ const float gc = fminf(fmaxf(gate_raw, -5.0F), 2.0F); gate_value = GateIsExp ? fminf(gate_raw, 7.389F) : __expf(gc); }
|
||||||
beta_value = load_as_float(beta, gate_index);
|
beta_value = load_as_float(beta, gate_index);
|
||||||
}
|
}
|
||||||
gate_value = __shfl_sync(0xffffffffU, gate_value, 0);
|
gate_value = __shfl_sync(0xffffffffU, gate_value, 0);
|
||||||
@@ -350,7 +350,7 @@ __global__ void gdn_forward_vlk_varlen_kernel(
|
|||||||
for (int c = 0; c < COLS; ++c) {
|
for (int c = 0; c < COLS; ++c) {
|
||||||
const float new_state =
|
const float new_state =
|
||||||
fmaf(k_reg[r], delta[c], gate_value * state_shard[c][r]);
|
fmaf(k_reg[r], delta[c], gate_value * state_shard[c][r]);
|
||||||
state_shard[c][r] = new_state;
|
state_shard[c][r] = fminf(fmaxf(new_state, -65504.0F), 65504.0F);
|
||||||
attn_partial[c] += new_state * q_reg[r];
|
attn_partial[c] += new_state * q_reg[r];
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -543,7 +543,7 @@ __global__ void gdn_decode_mixed_qkv_global_state_kernel(
|
|||||||
for (int c = 0; c < COLS; ++c) {
|
for (int c = 0; c < COLS; ++c) {
|
||||||
const float new_state =
|
const float new_state =
|
||||||
fmaf(k_reg[r], delta[c], gate_value * state_shard[c][r]);
|
fmaf(k_reg[r], delta[c], gate_value * state_shard[c][r]);
|
||||||
state_shard[c][r] = new_state;
|
state_shard[c][r] = fminf(fmaxf(new_state, -65504.0F), 65504.0F);
|
||||||
attn_partial[c] += new_state * q_reg[r];
|
attn_partial[c] += new_state * q_reg[r];
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -765,7 +765,7 @@ __global__ void gdn_decode_mixed_qkv_ddtree_state_kernel(
|
|||||||
for (int c = 0; c < COLS; ++c) {
|
for (int c = 0; c < COLS; ++c) {
|
||||||
const float new_state =
|
const float new_state =
|
||||||
fmaf(k_reg[r], delta[c], gate_value * state_shard[c][r]);
|
fmaf(k_reg[r], delta[c], gate_value * state_shard[c][r]);
|
||||||
state_shard[c][r] = new_state;
|
state_shard[c][r] = fminf(fmaxf(new_state, -65504.0F), 65504.0F);
|
||||||
attn_partial[c] += new_state * q_reg[r];
|
attn_partial[c] += new_state * q_reg[r];
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -581,6 +581,13 @@ class GatedDeltaNet(nn.Module):
|
|||||||
k_4d = k.unsqueeze(0) # (1, L, Hk, K)
|
k_4d = k.unsqueeze(0) # (1, L, Hk, K)
|
||||||
v_4d = v_raw.unsqueeze(0) # (1, L, Hv, V)
|
v_4d = v_raw.unsqueeze(0) # (1, L, Hv, V)
|
||||||
g_3d = gate.unsqueeze(0) # (1, L, Hv)
|
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)
|
beta_3d = b_seq.unsqueeze(0) # (1, L, Hv)
|
||||||
|
|
||||||
# Initial state from temporal_state
|
# Initial state from temporal_state
|
||||||
@@ -800,6 +807,8 @@ class GatedDeltaNet(nn.Module):
|
|||||||
k_t.view(BH, self.head_k_dim, 1),
|
k_t.view(BH, self.head_k_dim, 1),
|
||||||
delta.view(BH, 1, self.head_v_dim),
|
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
|
# Output: core_out = q_t @ updated temporal_state
|
||||||
core_out = _ix_bmm(
|
core_out = _ix_bmm(
|
||||||
|
|||||||
Reference in New Issue
Block a user