arch(scan): dispatch_scan.cuh Phase 1/Phase 2 separation in GDN chunk loop
Direct translation of CCCL dispatch_scan.cuh (1469 lines) architecture:
CCCL dispatch_scan has two kernels:
1. DeviceScanInitKernel — initializes tile_state (parallelizable)
2. DeviceScanKernel — sequential scan using tile_state propagation
Our _torch_chunk_gated_delta_rule now separates:
Phase 1 (init, parallelizable): pre-compute ALL chunk-local attn matrices
attn_i[c] = q[c] @ k[c].T * decay[c] — does NOT depend on state
Also pre-compute g.exp() and clamped g once, outside loop
Phase 2 (scan, sequential): only state-dependent ops in the loop
v_prime, v_new, attn_inter, core_out, state update
This matches CCCL's insight: everything that doesn't need tile_state
should be computed before the scan kernel, not interleaved with it.
This commit is contained in:
@@ -218,17 +218,34 @@ 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=1)
|
diagonal=1)
|
||||||
|
|
||||||
for i in range(total_len // chunk_size):
|
# dispatch_scan.cuh Phase 1: pre-compute ALL chunk-local attention matrices
|
||||||
q_i, k_i, v_i = query[:, :, i], key[:, :, i], value[:, :, i]
|
# outside the state loop. attn_i[c] only depends on q, k, decay_mask — NOT state.
|
||||||
attn_i = (_ix_matmul(q_i, k_i.transpose(-1, -2)) * decay_mask[:, :, i]).masked_fill_(mask_upper2, 0)
|
# This is the CCCL "init kernel" pattern: compute everything possible
|
||||||
|
# before the sequential scan kernel that needs tile_state propagation.
|
||||||
|
num_chunks = total_len // chunk_size
|
||||||
|
attn_i_all = torch.empty(
|
||||||
|
batch, num_heads, num_chunks, chunk_size, chunk_size,
|
||||||
|
dtype=value.dtype, device=value.device)
|
||||||
|
for i in range(num_chunks):
|
||||||
|
attn_i_all[:, :, i] = (
|
||||||
|
_ix_matmul(query[:, :, i], key[:, :, i].transpose(-1, -2))
|
||||||
|
* decay_mask[:, :, i]
|
||||||
|
).masked_fill_(mask_upper2, 0)
|
||||||
|
|
||||||
|
# dispatch_scan.cuh Phase 2: sequential state propagation (scan kernel).
|
||||||
|
# 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):
|
||||||
v_prime = _ix_matmul(k_cumdecay[:, :, i], last_state)
|
v_prime = _ix_matmul(k_cumdecay[:, :, i], last_state)
|
||||||
v_new = v_i - v_prime
|
v_new = value[:, :, i] - v_prime
|
||||||
attn_inter = _ix_matmul(q_i * g[:, :, i, :, None].clamp(-20, 20).exp(), last_state)
|
attn_inter = _ix_matmul(query[:, :, i] * g_exp_cache[:, :, i, :, None], last_state)
|
||||||
core_out[:, :, i] = attn_inter + _ix_matmul(attn_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
|
||||||
last_state = (
|
last_state = (
|
||||||
last_state * g[:, :, i, -1, None, None].clamp(-20, 20).exp()
|
last_state * g_exp_cache[:, :, i, -1, None, None]
|
||||||
+ _ix_matmul(
|
+ _ix_matmul(
|
||||||
(k_i * (g[:, :, i, -1, None] - g[:, :, i]).clamp(-20, 20).exp()[..., None])
|
(key[:, :, i] * (g_clamped[:, :, i, -1, None] - g_clamped[:, :, i]).exp()[..., None])
|
||||||
.transpose(-1, -2), v_new)
|
.transpose(-1, -2), v_new)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user