fix(GDN): remove pre-cumsum clamp — match xllm reference, fix 99.98% NaN
ROOT CAUSE: g.clamp(-5,2) before cumsum corrupted gate values. The GDN algorithm computes decay_mask = exp(g_i - g_j) which is numerically stable via subtraction cancelling cumsum growth. Pre-clamping g distorts these differences → wrong decay rates → NaN. xllm reference: qwen3_gated_delta_net_base.cpp lines 170-238 - cumsum first (no pre-clamp) - difference form: (g_i_last - g[:, i]).exp() for state update - k_cumdecay uses g.exp() directly (not clamped) Removed: g.clamp(-5,2), g.clamp(-20,20), g_exp_cache, g_clamped Added: xllm-style g_i_last/g_exp_term/k_g_exp state update
This commit is contained in:
@@ -264,13 +264,12 @@ def _torch_chunk_gated_delta_rule(
|
|||||||
torch.ones(chunk_size, chunk_size, dtype=torch.bool, device=query.device),
|
torch.ones(chunk_size, chunk_size, dtype=torch.bool, device=query.device),
|
||||||
diagonal=0)
|
diagonal=0)
|
||||||
|
|
||||||
# CCCL accumulator_t pattern: clamp BEFORE cumsum to prevent overflow
|
# Match xllm qwen3_gated_delta_net_base.cpp line 170-175:
|
||||||
# at the source. Without this, individual g values of ±10 accumulate
|
# cumsum first, then difference form (g_i - g_j) which is numerically
|
||||||
# over 64 positions to ±640 — far beyond float32 exp() safe range (~88).
|
# stable — the subtraction cancels cumsum growth so exp() stays bounded.
|
||||||
g = g.clamp(-5.0, 2.0)
|
# Do NOT clamp g before cumsum — that corrupts gate values and causes NaN.
|
||||||
g = g.cumsum(dim=-1)
|
g = g.cumsum(dim=-1)
|
||||||
g = g.clamp(-20.0, 20.0)
|
decay_mask = (g.unsqueeze(-1) - g.unsqueeze(-2)).tril().exp().to(torch.float32).tril()
|
||||||
decay_mask = ((g.unsqueeze(-1) - g.unsqueeze(-2)).tril().exp().float()).tril()
|
|
||||||
attn = -((_ix_matmul(k_beta, key.transpose(-1, -2))) * decay_mask).masked_fill(mask_upper, 0)
|
attn = -((_ix_matmul(k_beta, key.transpose(-1, -2))) * decay_mask).masked_fill(mask_upper, 0)
|
||||||
for i in range(1, chunk_size):
|
for i in range(1, chunk_size):
|
||||||
row = attn[..., i, :i].clone()
|
row = attn[..., i, :i].clone()
|
||||||
@@ -278,7 +277,7 @@ def _torch_chunk_gated_delta_rule(
|
|||||||
attn[..., i, :i] = row + (row.unsqueeze(-1) * sub).sum(-2)
|
attn[..., i, :i] = row + (row.unsqueeze(-1) * sub).sum(-2)
|
||||||
attn = attn + torch.eye(chunk_size, dtype=attn.dtype, device=attn.device)
|
attn = attn + torch.eye(chunk_size, dtype=attn.dtype, device=attn.device)
|
||||||
value = _ix_matmul(attn, v_beta)
|
value = _ix_matmul(attn, v_beta)
|
||||||
k_cumdecay = _ix_matmul(attn, k_beta * g.clamp(-20, 20).exp().unsqueeze(-1))
|
k_cumdecay = _ix_matmul(attn, k_beta * g.exp().unsqueeze(-1))
|
||||||
|
|
||||||
last_state = (
|
last_state = (
|
||||||
torch.zeros(batch, num_heads, k_dim, v_dim, dtype=value.dtype, device=value.device)
|
torch.zeros(batch, num_heads, k_dim, v_dim, dtype=value.dtype, device=value.device)
|
||||||
@@ -304,22 +303,22 @@ def _torch_chunk_gated_delta_rule(
|
|||||||
* decay_mask[:, :, i]
|
* decay_mask[:, :, i]
|
||||||
).masked_fill_(mask_upper2, 0)
|
).masked_fill_(mask_upper2, 0)
|
||||||
|
|
||||||
# dispatch_scan.cuh Phase 2: sequential state propagation (scan kernel).
|
# State propagation — match xllm qwen3_gated_delta_net_base.cpp line 218-238
|
||||||
# Only state-dependent ops remain in this loop.
|
|
||||||
g_exp_cache = g.clamp(-20, 20).exp() # pre-compute once
|
|
||||||
g_clamped = g.clamp(-20, 20) # keep raw clamped g for difference computation
|
|
||||||
for i in range(num_chunks):
|
for i in range(num_chunks):
|
||||||
|
q_i = query[:, :, i]
|
||||||
|
k_i = key[:, :, i]
|
||||||
|
v_i = value[:, :, i]
|
||||||
v_prime = _ix_matmul(k_cumdecay[:, :, i], last_state)
|
v_prime = _ix_matmul(k_cumdecay[:, :, i], last_state)
|
||||||
v_new = value[:, :, i] - v_prime
|
v_new = v_i - v_prime
|
||||||
attn_inter = _ix_matmul(query[:, :, i] * g_exp_cache[:, :, i, :, None], last_state)
|
# attn_inter: q * exp(g) @ state — xllm line 228
|
||||||
|
attn_inter = _ix_matmul(q_i * g[:, :, i].unsqueeze(-1).exp(), last_state)
|
||||||
core_out[:, :, i] = attn_inter + _ix_matmul(attn_i_all[:, :, i], v_new)
|
core_out[:, :, i] = attn_inter + _ix_matmul(attn_i_all[:, :, i], v_new)
|
||||||
# State update uses difference form: exp(g[-1] - g[:]) to avoid division
|
# State update — xllm line 230-237: difference form for numerical stability
|
||||||
last_state = (
|
g_i_last = g[:, :, i, -1].unsqueeze(-1) # (B, H, 1)
|
||||||
last_state * g_exp_cache[:, :, i, -1, None, None]
|
g_exp_term = (g_i_last - g[:, :, i]).exp().unsqueeze(-1) # (B, H, C, 1)
|
||||||
+ _ix_matmul(
|
k_g_exp = (k_i * g_exp_term).transpose(-1, -2).contiguous()
|
||||||
(key[:, :, i] * (g_clamped[:, :, i, -1, None] - g_clamped[:, :, i]).exp()[..., None])
|
last_state = (last_state * g_i_last.unsqueeze(-1).exp()
|
||||||
.transpose(-1, -2), v_new)
|
+ _ix_matmul(k_g_exp, v_new))
|
||||||
)
|
|
||||||
|
|
||||||
if not output_final_state:
|
if not output_final_state:
|
||||||
last_state = None
|
last_state = None
|
||||||
|
|||||||
Reference in New Issue
Block a user