fix(overflow): chunk_size 64→16 — CCCL counter overflow prevention

agent_radix_sort_upsweep.cuh (517 lines) key insight:
  UNROLL_COUNT = min(64, 255/KEYS_PER_THREAD)
  — limits accumulation steps to prevent unsigned char counter overflow

Same principle applied to GatedDeltaNet cumsum:
  chunk=64 + pre_clamp_max=2.0 → worst cumsum = 128 → exp(128) = inf
  chunk=16 + pre_clamp_max=2.0 → worst cumsum = 32  → clamp(-20,20) safe

This was the remaining NaN source: clamp at [-5,2] before cumsum was
necessary but not sufficient when chunk_size=64.
This commit is contained in:
Claude
2026-08-10 00:13:49 +00:00
parent 0a697f5871
commit 83d633798f

View File

@@ -155,7 +155,12 @@ def _torch_chunk_gated_delta_rule(
value: torch.Tensor, # (batch, seq, num_heads, head_v_dim)
g: torch.Tensor, # (batch, seq, num_heads)
beta: torch.Tensor, # (batch, seq, num_heads)
chunk_size: int = 64,
# CCCL agent_radix_sort_upsweep overflow pattern: UNROLL_COUNT = min(64, 255/KEYS_PER_THREAD)
# prevents counter overflow by limiting accumulation steps.
# Same principle: chunk_size limits cumsum steps. With pre-clamp [-5,2]:
# chunk=64: worst cumsum = 64*2 = 128 → exp(128) = inf
# chunk=16: worst cumsum = 16*2 = 32 → clamp(-20,20) catches it
chunk_size: int = 16,
initial_state: Optional[torch.Tensor] = None,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,