fix(GDN): gate clamp [-5,0] (decay only) + state clamp ±100

Root cause: gate=3.0 → exp(2.0)=7.389 per step → state explodes even with state clamp 65504
- 65504 * 7.389 = 483900 → re-clamped to 65504 → oscillates at max → output inf

Fix: gate_raw ∈ [-5, 0] so exp(gate) ∈ [0.007, 1.0] — pure decay, never grows
     GateIsExp path: clamp ≤ 1.0 — same invariant
     state ∈ [-100, 100] — tight enough to prevent output overflow

GDN gate is -dt * A_log.exp() where dt>0, A_log>0 → always negative in normal weights.
Clamping to ≤0 enforces this invariant even for pathological inputs.
This commit is contained in:
EX Engine
2026-08-10 04:44:05 +00:00
parent 1ae398eeee
commit 41aec33955

View File

@@ -135,7 +135,7 @@ __global__ void gdn_forward_kernel(const scalar_t* __restrict__ q,
float beta_value = 0.0F;
if (threadIdx.x == 0) {
const float gate_raw = load_as_float(gate, gate_index);
{ const float gc = fminf(fmaxf(gate_raw, -5.0F), 2.0F); gate_value = GateIsExp ? fminf(gate_raw, 7.389F) : __expf(gc); }
{ const float gc = fminf(fmaxf(gate_raw, -5.0F), 0.0F); gate_value = GateIsExp ? fminf(gate_raw, 1.0F) : __expf(gc); }
beta_value = load_as_float(beta, gate_index);
}
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) {
const float new_state =
fmaf(k_reg[r], delta[c], gate_value * state_shard[c][r]);
state_shard[c][r] = fminf(fmaxf(new_state, -65504.0F), 65504.0F);
state_shard[c][r] = fminf(fmaxf(new_state, -100.0F), 100.0F);
attn_partial[c] += new_state * q_reg[r];
}
}
@@ -294,7 +294,7 @@ __global__ void gdn_forward_vlk_varlen_kernel(
float beta_value = 0.0F;
if (threadIdx.x == 0) {
const float gate_raw = load_as_float(gate, gate_index);
{ const float gc = fminf(fmaxf(gate_raw, -5.0F), 2.0F); gate_value = GateIsExp ? fminf(gate_raw, 7.389F) : __expf(gc); }
{ const float gc = fminf(fmaxf(gate_raw, -5.0F), 0.0F); gate_value = GateIsExp ? fminf(gate_raw, 1.0F) : __expf(gc); }
beta_value = load_as_float(beta, gate_index);
}
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) {
const float new_state =
fmaf(k_reg[r], delta[c], gate_value * state_shard[c][r]);
state_shard[c][r] = fminf(fmaxf(new_state, -65504.0F), 65504.0F);
state_shard[c][r] = fminf(fmaxf(new_state, -100.0F), 100.0F);
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) {
const float new_state =
fmaf(k_reg[r], delta[c], gate_value * state_shard[c][r]);
state_shard[c][r] = fminf(fmaxf(new_state, -65504.0F), 65504.0F);
state_shard[c][r] = fminf(fmaxf(new_state, -100.0F), 100.0F);
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) {
const float new_state =
fmaf(k_reg[r], delta[c], gate_value * state_shard[c][r]);
state_shard[c][r] = fminf(fmaxf(new_state, -65504.0F), 65504.0F);
state_shard[c][r] = fminf(fmaxf(new_state, -100.0F), 100.0F);
attn_partial[c] += new_state * q_reg[r];
}
}