[GDN/KDA] Fuse SM100 CuteDSL prefill state I/O into the chunk h kernel (#30169)

Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
Yuan Luo
2026-07-16 16:40:33 +08:00
committed by GitHub
co-authored by luoyuan.luo
parent db7e6807de
commit a798a2aeea
8 changed files with 299 additions and 37 deletions
@@ -153,6 +153,7 @@ def chunk_gated_delta_rule_cutedsl(
chunk_indices: torch.Tensor, chunk_indices: torch.Tensor,
chunk_offsets: torch.Tensor, chunk_offsets: torch.Tensor,
core_attn_out: torch.Tensor | None = None, core_attn_out: torch.Tensor | None = None,
initial_state_indices: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]: ) -> tuple[torch.Tensor, torch.Tensor]:
"""Run the GDN chunk CuteDSL prefill kernels. """Run the GDN chunk CuteDSL prefill kernels.
@@ -162,11 +163,17 @@ def chunk_gated_delta_rule_cutedsl(
v: Value tensor with shape ``[1, T, Hv, V]``. v: Value tensor with shape ``[1, T, Hv, V]``.
g: Log-space decay tensor with shape ``[1, T, Hv]``. g: Log-space decay tensor with shape ``[1, T, Hv]``.
beta: Delta-rule beta tensor with shape ``[1, T, Hv]``. beta: Delta-rule beta tensor with shape ``[1, T, Hv]``.
initial_state: Recurrent state with shape ``[N, Hv, V, K]``. initial_state: Recurrent state with shape ``[N, Hv, V, K]``, or the
state POOL ``[num_slots, Hv, V, K]`` when ``initial_state_indices``
is given.
cu_seqlens: Cumulative sequence lengths with shape ``[N + 1]``. cu_seqlens: Cumulative sequence lengths with shape ``[N + 1]``.
chunk_indices: Chunk index metadata with shape ``[NT, 2]``. chunk_indices: Chunk index metadata with shape ``[NT, 2]``.
chunk_offsets: Cumulative chunk offsets with shape ``[N + 1]``. chunk_offsets: Cumulative chunk offsets with shape ``[N + 1]``.
core_attn_out: Optional output buffer with shape ``[T, Hv, V]``. core_attn_out: Optional output buffer with shape ``[T, Hv, V]``.
initial_state_indices: Optional ``[N]`` int32 pool slots. When given,
the h kernel reads AND writes the pool rows in place (fused state
gather/scatter — no ``[N, Hv, V, K]`` intermediates) and the
returned ``final_state`` is the pool tensor itself.
Returns: Returns:
A tuple ``(output, final_state)`` where ``output`` has shape A tuple ``(output, final_state)`` where ``output`` has shape
@@ -213,7 +220,16 @@ def chunk_gated_delta_rule_cutedsl(
head_k_dim, head_k_dim,
) )
v_new = q_3d.new_empty(pad_t, num_v_heads, head_v_dim) v_new = q_3d.new_empty(pad_t, num_v_heads, head_v_dim)
final_state = torch.empty_like(initial_state) if initial_state_indices is None:
# Dense mode: preserve the return-fresh-final_state contract.
final_state = torch.empty_like(initial_state)
state_indices = torch.arange(
cu_seqlens.numel() - 1, device=q_3d.device, dtype=torch.int32
)
else:
# Pool mode: read and write the pool rows in place.
final_state = initial_state
state_indices = initial_state_indices
h_cutedsl( h_cutedsl(
k_3d, k_3d,
u, u,
@@ -225,6 +241,7 @@ def chunk_gated_delta_rule_cutedsl(
final_state, final_state,
cu_seqlens, cu_seqlens,
chunk_offsets, chunk_offsets,
state_indices,
) )
output = core_attn_out if core_attn_out is not None else torch.empty_like(v_3d) output = core_attn_out if core_attn_out is not None else torch.empty_like(v_3d)
@@ -106,6 +106,7 @@ class Sm100ChunkHKernel:
ht: cute.Tensor, ht: cute.Tensor,
cu_seqlens: cute.Tensor, cu_seqlens: cute.Tensor,
chunk_offsets: cute.Tensor, chunk_offsets: cute.Tensor,
state_indices: cute.Tensor,
stream: CUstream, stream: CUstream,
): ):
tma_g2s = cpasync.CopyBulkTensorTileG2SOp() tma_g2s = cpasync.CopyBulkTensorTileG2SOp()
@@ -119,7 +120,10 @@ class Sm100ChunkHKernel:
HT_args = self._make_h_tma_args(ht, tma_s2g) HT_args = self._make_h_tma_args(ht, tma_s2g)
H_args = self._make_h_tma_args(h, tma_s2g) H_args = self._make_h_tma_args(h, tma_s2g)
grid = (self.Hv, h0.shape[0], 1) # h0/ht may be the full state pool ([num_slots, ...]) rather than a
# per-sequence gather, so the sequence count comes from cu_seqlens and
# each block resolves its state row through state_indices.
grid = (self.Hv, cu_seqlens.shape[0] - 1, 1)
block = (self.num_warps * 32, 1, 1) block = (self.num_warps * 32, 1, 1)
self.kernel( self.kernel(
K_args, K_args,
@@ -132,6 +136,7 @@ class Sm100ChunkHKernel:
g_cu, g_cu,
cu_seqlens, cu_seqlens,
chunk_offsets, chunk_offsets,
state_indices,
).launch(grid=grid, block=block, stream=stream) ).launch(grid=grid, block=block, stream=stream)
@cute.kernel @cute.kernel
@@ -147,6 +152,7 @@ class Sm100ChunkHKernel:
g_cu: cute.Tensor, g_cu: cute.Tensor,
cu_seqlens: cute.Tensor, cu_seqlens: cute.Tensor,
chunk_offsets: cute.Tensor, chunk_offsets: cute.Tensor,
state_indices: cute.Tensor,
): ):
tid, _, _ = cute.arch.thread_idx() tid, _, _ = cute.arch.thread_idx()
head_id, seq_id, _ = cute.arch.block_idx() head_id, seq_id, _ = cute.arch.block_idx()
@@ -220,6 +226,8 @@ class Sm100ChunkHKernel:
eos = cu_seqlens[seq_id + 1] eos = cu_seqlens[seq_id + 1]
seqlen = eos - bos seqlen = eos - bos
num_chunks = cute.ceil_div(seqlen, BT) num_chunks = cute.ceil_div(seqlen, BT)
# Row of h0/ht for this sequence (pool slot; fused state gather/scatter).
state_slot = state_indices[seq_id]
if warp_id == 9: if warp_id == 9:
# TMA warp # TMA warp
@@ -234,7 +242,7 @@ class Sm100ChunkHKernel:
H0_size = V_dim * K_dim * self.h_dtype.width // 8 H0_size = V_dim * K_dim * self.h_dtype.width // 8
cute.arch.mbarrier_arrive_and_expect_tx(h0_mbar, H0_size) cute.arch.mbarrier_arrive_and_expect_tx(h0_mbar, H0_size)
simple_tma_copy( simple_tma_copy(
H0_tma_atom, tmaH0[seq_id, head_id, None, None], sH0, h0_mbar H0_tma_atom, tmaH0[state_slot, head_id, None, None], sH0, h0_mbar
) )
# shape: ((BT, num_BT_tiles), (64, 2)) # shape: ((BT, num_BT_tiles), (64, 2))
@@ -531,7 +539,7 @@ class Sm100ChunkHKernel:
cute.arch.barrier(barrier_id=1, number_of_threads=128) cute.arch.barrier(barrier_id=1, number_of_threads=128)
if warp_id_ == 0: if warp_id_ == 0:
ht_dst = tmaHT[seq_id, head_id, None, None] ht_dst = tmaHT[state_slot, head_id, None, None]
simple_tma_copy(HT_tma_atom, sH0, ht_dst) simple_tma_copy(HT_tma_atom, sH0, ht_dst)
with cute.arch.elect_one(): with cute.arch.elect_one():
cute.arch.cp_async_bulk_commit_group() cute.arch.cp_async_bulk_commit_group()
@@ -676,6 +684,7 @@ class Sm100ChunkHKernel:
total_t = cute.sym_int() total_t = cute.sym_int()
pad_t = cute.sym_int() pad_t = cute.sym_int()
total_chunks_n = cute.sym_int() total_chunks_n = cute.sym_int()
num_state_slots = cute.sym_int()
num_sequences = cute.sym_int() num_sequences = cute.sym_int()
cu_entries = cute.sym_int() cu_entries = cute.sym_int()
@@ -688,13 +697,14 @@ class Sm100ChunkHKernel:
BFloat16, (total_chunks_n, Hv, V_dim, K_dim), divisibility=16 BFloat16, (total_chunks_n, Hv, V_dim, K_dim), divisibility=16
) )
h0 = make_fake_tensor( h0 = make_fake_tensor(
h_dtype, (num_sequences, Hv, V_dim, K_dim), divisibility=16 h_dtype, (num_state_slots, Hv, V_dim, K_dim), divisibility=16
) )
ht = make_fake_tensor( ht = make_fake_tensor(
h_dtype, (num_sequences, Hv, V_dim, K_dim), divisibility=16 h_dtype, (num_state_slots, Hv, V_dim, K_dim), divisibility=16
) )
cu_seqlens = make_fake_tensor(Int32, (cu_entries,), divisibility=1) cu_seqlens = make_fake_tensor(Int32, (cu_entries,), divisibility=1)
chunk_offsets = make_fake_tensor(Int32, (cu_entries,), divisibility=1) chunk_offsets = make_fake_tensor(Int32, (cu_entries,), divisibility=1)
state_indices = make_fake_tensor(Int32, (num_sequences,), divisibility=1)
kernel = Sm100ChunkHKernel(H, Hv, K_dim, V_dim, h_dtype, BT, num_stages) kernel = Sm100ChunkHKernel(H, Hv, K_dim, V_dim, h_dtype, BT, num_stages)
stream = cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True) stream = cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True)
@@ -710,6 +720,7 @@ class Sm100ChunkHKernel:
ht, ht,
cu_seqlens, cu_seqlens,
chunk_offsets, chunk_offsets,
state_indices,
stream, stream,
options="--enable-tvm-ffi", options="--enable-tvm-ffi",
) )
@@ -726,10 +737,16 @@ def h_cutedsl(
ht: torch.Tensor, ht: torch.Tensor,
cu_seqlens: torch.Tensor, cu_seqlens: torch.Tensor,
chunk_offsets: torch.Tensor, chunk_offsets: torch.Tensor,
state_indices: torch.Tensor,
BT: int = 64, BT: int = 64,
num_stages: int = 2, num_stages: int = 2,
) -> None: ) -> None:
"""Compute H/V_new with the same argument order as the CUDA wrapper.""" """Compute H/V_new with the same argument order as the CUDA wrapper.
``h0``/``ht`` may be the full state pool; ``state_indices`` [N] int32 maps
each sequence to its row, so state gather/scatter fuses into the kernel's
TMA load/store (no per-call state intermediates).
"""
_, H, K_dim = K.shape _, H, K_dim = K.shape
_, Hv, V_dim = V.shape _, Hv, V_dim = V.shape
@@ -748,6 +765,7 @@ def h_cutedsl(
ht, ht,
cu_seqlens, cu_seqlens,
chunk_offsets, chunk_offsets,
state_indices,
) )
@@ -112,15 +112,25 @@ def chunk_kda_cutedsl(
v: torch.Tensor, # [T, Hv, V] bf16 v: torch.Tensor, # [T, Hv, V] bf16
g: torch.Tensor, # [T, Hv, K] log-decay. RAW if A_log given, else pre-activated g: torch.Tensor, # [T, Hv, K] log-decay. RAW if A_log given, else pre-activated
beta: torch.Tensor, # [T, Hv] fp32, post-sigmoid beta: torch.Tensor, # [T, Hv] fp32, post-sigmoid
h0: torch.Tensor, # [N, Hv, V, K] (initial recurrent state, [V,K] layout) h0: torch.Tensor, # [N, Hv, V, K] state, or the state POOL with h0_indices
cu_seqlens: torch.Tensor, cu_seqlens: torch.Tensor,
scale: float | None = None, scale: float | None = None,
num_sms: int | None = None, num_sms: int | None = None,
A_log: torch.Tensor | None = None, # [Hv]; if set, activate g internally A_log: torch.Tensor | None = None, # [Hv]; if set, activate g internally
dt_bias: torch.Tensor | None = None, # [Hv, K] or [Hv*K] dt_bias: torch.Tensor | None = None, # [Hv, K] or [Hv*K]
lower_bound: float | None = None, lower_bound: float | None = None,
h0_indices: torch.Tensor | None = None, # [N] int32 pool slots
): ):
"""Run the KDA chunk gated-delta-rule prefill. Returns (o [T,Hv,V], ht [N,Hv,V,K]).""" """Run the KDA chunk gated-delta-rule prefill. Returns (o [T,Hv,V], ht).
Dense mode (``h0_indices is None``): ``h0`` is [N, Hv, V, K]; the final state
is returned in a fresh ``ht`` and ``h0`` is left untouched.
Pool mode: ``h0`` is the state pool [num_slots, Hv, V, K] and ``h0_indices``
maps each sequence to its slot; the h kernel reads AND writes the pool rows
in place (fused state gather/scatter — no [N, Hv, V, K] intermediates), and
the returned ``ht`` is the pool tensor itself.
"""
import torch.nn.functional as F import torch.nn.functional as F
T, Hv, K = q.shape T, Hv, K = q.shape
@@ -202,8 +212,29 @@ def chunk_kda_cutedsl(
V_new = ws["Vn"][:pad_t] V_new = ws["Vn"][:pad_t]
h_chunks = ws["hc"][:total] h_chunks = ws["hc"][:total]
ht = torch.empty_like(h0) if h0_indices is None:
kda_h_cutedsl(KR, U, W, V_new, g_cu, h_chunks, h0, ht, cu_seqlens, chunk_offsets) # Dense mode: preserve the return-fresh-ht contract.
ht = torch.empty_like(h0)
state_indices = torch.arange(
cu_seqlens.numel() - 1, device=q.device, dtype=torch.int32
)
else:
# Pool mode: read and write the pool rows in place.
ht = h0
state_indices = h0_indices
kda_h_cutedsl(
KR,
U,
W,
V_new,
g_cu,
h_chunks,
h0,
ht,
cu_seqlens,
chunk_offsets,
state_indices,
)
o = q.new_empty(T, Hv, V, dtype=torch.bfloat16) o = q.new_empty(T, Hv, V, dtype=torch.bfloat16)
kda_o_cutedsl( kda_o_cutedsl(
@@ -112,6 +112,7 @@ class Sm100KdaChunkHKernel:
ht: cute.Tensor, ht: cute.Tensor,
cu_seqlens: cute.Tensor, cu_seqlens: cute.Tensor,
chunk_offsets: cute.Tensor, chunk_offsets: cute.Tensor,
state_indices: cute.Tensor,
stream: CUstream, stream: CUstream,
): ):
tma_g2s = cpasync.CopyBulkTensorTileG2SOp() tma_g2s = cpasync.CopyBulkTensorTileG2SOp()
@@ -125,7 +126,10 @@ class Sm100KdaChunkHKernel:
HT_args = self._make_h_tma_args(ht, tma_s2g) HT_args = self._make_h_tma_args(ht, tma_s2g)
H_args = self._make_h_tma_args(h, tma_s2g) H_args = self._make_h_tma_args(h, tma_s2g)
grid = (self.Hv, h0.shape[0], 1) # h0/ht may be the full state pool ([num_slots, ...]) rather than a
# per-sequence gather, so the sequence count comes from cu_seqlens and
# each block resolves its state row through state_indices.
grid = (self.Hv, cu_seqlens.shape[0] - 1, 1)
block = (self.num_warps * 32, 1, 1) block = (self.num_warps * 32, 1, 1)
self.kernel( self.kernel(
K_args, K_args,
@@ -138,6 +142,7 @@ class Sm100KdaChunkHKernel:
g_cu, g_cu,
cu_seqlens, cu_seqlens,
chunk_offsets, chunk_offsets,
state_indices,
).launch(grid=grid, block=block, stream=stream) ).launch(grid=grid, block=block, stream=stream)
@cute.kernel @cute.kernel
@@ -153,6 +158,7 @@ class Sm100KdaChunkHKernel:
g_cu: cute.Tensor, g_cu: cute.Tensor,
cu_seqlens: cute.Tensor, cu_seqlens: cute.Tensor,
chunk_offsets: cute.Tensor, chunk_offsets: cute.Tensor,
state_indices: cute.Tensor,
): ):
tid, _, _ = cute.arch.thread_idx() tid, _, _ = cute.arch.thread_idx()
head_id, seq_id, _ = cute.arch.block_idx() head_id, seq_id, _ = cute.arch.block_idx()
@@ -226,6 +232,8 @@ class Sm100KdaChunkHKernel:
eos = cu_seqlens[seq_id + 1] eos = cu_seqlens[seq_id + 1]
seqlen = eos - bos seqlen = eos - bos
num_chunks = cute.ceil_div(seqlen, BT) num_chunks = cute.ceil_div(seqlen, BT)
# Row of h0/ht for this sequence (pool slot; fused state gather/scatter).
state_slot = state_indices[seq_id]
if warp_id == 9: if warp_id == 9:
# TMA warp # TMA warp
@@ -239,7 +247,7 @@ class Sm100KdaChunkHKernel:
H0_size = V_dim * K_dim * self.h_dtype.width // 8 H0_size = V_dim * K_dim * self.h_dtype.width // 8
cute.arch.mbarrier_arrive_and_expect_tx(h0_mbar, H0_size) cute.arch.mbarrier_arrive_and_expect_tx(h0_mbar, H0_size)
simple_tma_copy( simple_tma_copy(
H0_tma_atom, tmaH0[seq_id, head_id, None, None], sH0, h0_mbar H0_tma_atom, tmaH0[state_slot, head_id, None, None], sH0, h0_mbar
) )
gW_tiles = cute.logical_divide(tmaW[None, head_id, None], (BT, None)) gW_tiles = cute.logical_divide(tmaW[None, head_id, None], (BT, None))
@@ -514,7 +522,7 @@ class Sm100KdaChunkHKernel:
cute.arch.barrier(barrier_id=1, number_of_threads=128) cute.arch.barrier(barrier_id=1, number_of_threads=128)
if warp_id_ == 0: if warp_id_ == 0:
ht_dst = tmaHT[seq_id, head_id, None, None] ht_dst = tmaHT[state_slot, head_id, None, None]
simple_tma_copy(HT_tma_atom, sH0, ht_dst) simple_tma_copy(HT_tma_atom, sH0, ht_dst)
with cute.arch.elect_one(): with cute.arch.elect_one():
cute.arch.cp_async_bulk_commit_group() cute.arch.cp_async_bulk_commit_group()
@@ -628,6 +636,7 @@ class Sm100KdaChunkHKernel:
total_t = cute.sym_int() total_t = cute.sym_int()
pad_t = cute.sym_int() pad_t = cute.sym_int()
total_chunks_n = cute.sym_int() total_chunks_n = cute.sym_int()
num_state_slots = cute.sym_int()
num_sequences = cute.sym_int() num_sequences = cute.sym_int()
cu_entries = cute.sym_int() cu_entries = cute.sym_int()
@@ -640,13 +649,14 @@ class Sm100KdaChunkHKernel:
BFloat16, (total_chunks_n, Hv, V_dim, K_dim), divisibility=16 BFloat16, (total_chunks_n, Hv, V_dim, K_dim), divisibility=16
) )
h0 = make_fake_tensor( h0 = make_fake_tensor(
h_dtype, (num_sequences, Hv, V_dim, K_dim), divisibility=16 h_dtype, (num_state_slots, Hv, V_dim, K_dim), divisibility=16
) )
ht = make_fake_tensor( ht = make_fake_tensor(
h_dtype, (num_sequences, Hv, V_dim, K_dim), divisibility=16 h_dtype, (num_state_slots, Hv, V_dim, K_dim), divisibility=16
) )
cu_seqlens = make_fake_tensor(Int32, (cu_entries,), divisibility=1) cu_seqlens = make_fake_tensor(Int32, (cu_entries,), divisibility=1)
chunk_offsets = make_fake_tensor(Int32, (cu_entries,), divisibility=1) chunk_offsets = make_fake_tensor(Int32, (cu_entries,), divisibility=1)
state_indices = make_fake_tensor(Int32, (num_sequences,), divisibility=1)
kernel = Sm100KdaChunkHKernel(H, Hv, K_dim, V_dim, h_dtype, BT, num_stages) kernel = Sm100KdaChunkHKernel(H, Hv, K_dim, V_dim, h_dtype, BT, num_stages)
stream = cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True) stream = cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True)
@@ -662,6 +672,7 @@ class Sm100KdaChunkHKernel:
ht, ht,
cu_seqlens, cu_seqlens,
chunk_offsets, chunk_offsets,
state_indices,
stream, stream,
options="--enable-tvm-ffi", options="--enable-tvm-ffi",
) )
@@ -678,13 +689,19 @@ def kda_h_cutedsl(
ht: torch.Tensor, ht: torch.Tensor,
cu_seqlens: torch.Tensor, cu_seqlens: torch.Tensor,
chunk_offsets: torch.Tensor, chunk_offsets: torch.Tensor,
state_indices: torch.Tensor,
BT: int = 64, BT: int = 64,
num_stages: int = 2, num_stages: int = 2,
) -> None: ) -> None:
"""KDA chunk-state kernel. `kg` = per-channel pre-scaled key [T, Hv, K].""" """KDA chunk-state kernel. `kg` = per-channel pre-scaled key [T, Hv, K].
``h0``/``ht`` may be the full state pool; ``state_indices`` [N] int32 maps
each sequence to its row, so state gather/scatter fuses into the kernel's
TMA load/store (no per-call state intermediates).
"""
_, Hv, K_dim = kg.shape _, Hv, K_dim = kg.shape
_, _, V_dim = V.shape _, _, V_dim = V.shape
h_dtype = {torch.bfloat16: BFloat16, torch.float32: Float32}[h0.dtype] h_dtype = {torch.bfloat16: BFloat16, torch.float32: Float32}[h0.dtype]
Sm100KdaChunkHKernel.compile(Hv, Hv, K_dim, V_dim, h_dtype, BT, num_stages)( Sm100KdaChunkHKernel.compile(Hv, Hv, K_dim, V_dim, h_dtype, BT, num_stages)(
kg, V, W, V_new, g_cu, h, h0, ht, cu_seqlens, chunk_offsets kg, V, W, V_new, g_cu, h, h0, ht, cu_seqlens, chunk_offsets, state_indices
) )
@@ -137,38 +137,35 @@ class CuteDSLGDNKernel(LinearAttnKernelBase):
cu_seqlens = query_start_loc.to(torch.int32) cu_seqlens = query_start_loc.to(torch.int32)
# Pool gather: remap padding (-1) to the last (sentinel) slot. # Pool state I/O is fused into the h kernel's TMA load/store: pass the
# pool + per-seq slots and the kernel reads h0/writes ht in place at
# those rows (no gather/scatter kernels, no [N, Hv, V, K] intermediates).
# Remap padding (-1) to the last (sentinel) slot.
ssm_cache_indices = torch.where( ssm_cache_indices = torch.where(
cache_indices >= 0, cache_indices >= 0,
cache_indices, cache_indices,
ssm_states.shape[0] - 1, ssm_states.shape[0] - 1,
).to(torch.long) ).to(torch.int32)
initial_state = ssm_states[ssm_cache_indices].contiguous()
chunk_indices, chunk_offsets = self._prepare_meta_fn( chunk_indices, chunk_offsets = self._prepare_meta_fn(
cu_seqlens, total_seq_len, chunk_size=64 cu_seqlens, total_seq_len, chunk_size=64
) )
output, final_state = self._extend_fn( output, _ = self._extend_fn(
q=q_norm, q=q_norm,
k=k_norm, k=k_norm,
v=v_in, v=v_in,
g=g_in, g=g_in,
beta=beta_in, beta=beta_in,
initial_state=initial_state, initial_state=ssm_states,
cu_seqlens=cu_seqlens, cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices, chunk_indices=chunk_indices,
chunk_offsets=chunk_offsets, chunk_offsets=chunk_offsets,
) initial_state_indices=ssm_cache_indices,
ssm_states.index_copy_(
0,
ssm_cache_indices,
final_state.to(ssm_states.dtype),
) )
# Match Triton extend interface: (output, last_recurrent_state, h). # Match Triton extend interface: (output, last_recurrent_state, h).
# We've already written state back, so no need to return it. # The kernel already wrote state back into the pool in place.
return output, None, None return output, None, None
def target_verify(self, *args, **kwargs): def target_verify(self, *args, **kwargs):
@@ -120,27 +120,29 @@ class CuteDSLKDAKernel(LinearAttnKernelBase):
beta_in = beta[0][:num_tokens].to(torch.float32) beta_in = beta[0][:num_tokens].to(torch.float32)
cu_seqlens = query_start_loc.to(torch.int32) cu_seqlens = query_start_loc.to(torch.int32)
# Pool gather: remap padding (-1) to the last (sentinel) slot. State is # Pool state I/O is fused into the h kernel's TMA load/store: pass the
# pool + per-seq slots and the kernel reads h0/writes ht in place at
# those rows (no gather/scatter kernels, no [N, HV, V, K] intermediates).
# Remap padding (-1) to the last (sentinel) slot. State is
# [slots, HV, V, K] == cutedsl [V,K] layout, no transpose needed. # [slots, HV, V, K] == cutedsl [V,K] layout, no transpose needed.
ssm_cache_indices = torch.where( ssm_cache_indices = torch.where(
cache_indices >= 0, cache_indices, ssm_states.shape[0] - 1 cache_indices >= 0, cache_indices, ssm_states.shape[0] - 1
).to(torch.long) ).to(torch.int32)
initial_state = ssm_states[ssm_cache_indices].contiguous()
o, final_state = self._extend_fn( o, _ = self._extend_fn(
q_n, q_n,
k_n, k_n,
v_in, v_in,
g_in, g_in,
beta_in, beta_in,
initial_state, ssm_states,
cu_seqlens, cu_seqlens,
A_log=A_log, A_log=A_log,
dt_bias=dt_bias, dt_bias=dt_bias,
lower_bound=lower_bound, lower_bound=lower_bound,
h0_indices=ssm_cache_indices,
) )
ssm_states.index_copy_(0, ssm_cache_indices, final_state.to(ssm_states.dtype))
# Match chunk_kda's output layout [1, T, HV, V]. # Match chunk_kda's output layout [1, T, HV, V].
return o.unsqueeze(0) return o.unsqueeze(0)
@@ -172,6 +172,113 @@ def test_gdn_chunk_cutedsl_correctness(num_seqs: int, state_dtype: torch.dtype):
assert buffer_state_error.max().item() == 0 assert buffer_state_error.max().item() == 0
@pytest.mark.parametrize("state_dtype", [torch.bfloat16, torch.float32])
def test_gdn_chunk_cutedsl_pool_mode_matches_dense(state_dtype: torch.dtype):
"""Pool mode (initial_state_indices) must reproduce the dense gather/scatter
path bit-for-bit: same o, same final-state rows written in place at the
indexed pool slots, and every other pool row untouched."""
torch.manual_seed(11)
num_seqs = 5
seq_lens = torch.randint(1, 130, (num_seqs,), dtype=torch.int32)
cu_seqlens = torch.zeros(num_seqs + 1, device="cuda", dtype=torch.int32)
cu_seqlens[1:] = seq_lens.to(device="cuda").cumsum(0)
total_tokens = int(cu_seqlens[-1].item())
num_k_heads = 4
num_v_heads = 8
head_k_dim = 128
head_v_dim = 128
dtype = torch.bfloat16
q = torch.randn(
1, total_tokens, num_k_heads, head_k_dim, device="cuda", dtype=dtype
)
k = torch.randn_like(q)
v = torch.randn(
1, total_tokens, num_v_heads, head_v_dim, device="cuda", dtype=dtype
)
q = F.normalize(q.float(), p=2, dim=-1).to(dtype)
k = F.normalize(k.float(), p=2, dim=-1).to(dtype)
a = torch.randn(1, total_tokens, num_v_heads, device="cuda", dtype=dtype)
b = torch.randn(1, total_tokens, num_v_heads, device="cuda", dtype=dtype)
A = torch.empty(num_v_heads, device="cuda", dtype=torch.float32).uniform_(0, 16)
A_log = torch.log(A)
dt = torch.exp(
torch.rand(num_v_heads, device="cuda", dtype=torch.float32)
* (math.log(0.1) - math.log(0.001))
+ math.log(0.001)
)
dt = torch.clamp(dt, min=1e-4)
dt_bias = dt + torch.log(-torch.expm1(-dt))
g = -A_log.exp().view(1, 1, num_v_heads) * F.softplus(
a.float() + dt_bias.view(1, 1, num_v_heads)
)
beta = torch.sigmoid(b.float())
h0_dense = (
torch.randn(
num_seqs,
num_v_heads,
head_v_dim,
head_k_dim,
device="cuda",
dtype=state_dtype,
)
* 0.05
)
# Same states scattered into a larger pool at shuffled slots.
num_slots = 64
pool = (
torch.randn(
num_slots,
num_v_heads,
head_v_dim,
head_k_dim,
device="cuda",
dtype=state_dtype,
)
* 0.05
)
slots = torch.randperm(num_slots, device="cuda")[:num_seqs].to(torch.int32)
pool[slots.long()] = h0_dense
pool_before = pool.clone()
chunk_indices, chunk_offsets = prepare_metadata_cutedsl(cu_seqlens, total_tokens)
o_dense, ht_dense = chunk_gated_delta_rule_cutedsl(
q=q,
k=k,
v=v,
g=g,
beta=beta,
initial_state=h0_dense.clone(),
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
chunk_offsets=chunk_offsets,
)
o_pool, ht_pool = chunk_gated_delta_rule_cutedsl(
q=q,
k=k,
v=v,
g=g,
beta=beta,
initial_state=pool,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
chunk_offsets=chunk_offsets,
initial_state_indices=slots,
)
torch.cuda.synchronize()
# Same kernels and math; only the state addressing differs -> bit-identical.
assert ht_pool is pool
assert torch.equal(o_pool, o_dense)
assert torch.equal(pool[slots.long()], ht_dense)
untouched = torch.ones(num_slots, dtype=torch.bool, device="cuda")
untouched[slots.long()] = False
assert torch.equal(pool[untouched], pool_before[untouched])
if __name__ == "__main__": if __name__ == "__main__":
import sys import sys
@@ -220,6 +220,79 @@ def test_kda_chunk_cutedsl_realistic_gate():
assert (o.float() - ref_o[0].float()).abs().max().item() < 1e-2 assert (o.float() - ref_o[0].float()).abs().max().item() < 1e-2
@pytest.mark.parametrize("state_dtype", [torch.bfloat16, torch.float32])
def test_kda_chunk_cutedsl_pool_mode_matches_dense(state_dtype: torch.dtype):
"""Pool mode (h0_indices) must reproduce the dense gather/scatter path
bit-for-bit: same o, same final-state rows written in place at the indexed
pool slots, and every other pool row untouched."""
torch.manual_seed(3)
num_seqs = 5
seq_lens = torch.randint(1, 130, (num_seqs,), dtype=torch.int32)
cu_seqlens = torch.zeros(num_seqs + 1, device="cuda", dtype=torch.int32)
cu_seqlens[1:] = seq_lens.to("cuda").cumsum(0)
total_tokens = int(cu_seqlens[-1].item())
num_heads = 8
head_dim = 128
scale = head_dim**-0.5
q = _l2norm(torch.randn(1, total_tokens, num_heads, head_dim, device="cuda"))
k = _l2norm(torch.randn(1, total_tokens, num_heads, head_dim, device="cuda"))
v = torch.randn(1, total_tokens, num_heads, head_dim, device="cuda")
A_log = torch.randn(num_heads, device="cuda") * 0.5 - 1.5
dt_bias = torch.randn(num_heads, head_dim, device="cuda") * 0.1
g_raw = torch.randn(1, total_tokens, num_heads, head_dim, device="cuda")
g_act = -A_log.exp().view(1, 1, num_heads, 1) * F.softplus(
g_raw + dt_bias.view(1, 1, num_heads, head_dim)
)
beta = torch.sigmoid(torch.randn(1, total_tokens, num_heads, device="cuda")).float()
h0_dense = (
torch.randn(num_seqs, num_heads, head_dim, head_dim, device="cuda") * 0.05
).to(state_dtype)
# Same states scattered into a larger pool at shuffled slots.
num_slots = 64
pool = (
torch.randn(num_slots, num_heads, head_dim, head_dim, device="cuda") * 0.05
).to(state_dtype)
slots = torch.randperm(num_slots, device="cuda")[:num_seqs].to(torch.int32)
pool[slots.long()] = h0_dense
pool_before = pool.clone()
q_b, k_b, v_b = q[0].bfloat16(), k[0].bfloat16(), v[0].bfloat16()
o_dense, ht_dense = chunk_kda_cutedsl(
q_b,
k_b,
v_b,
g_act[0].float(),
beta[0].float(),
h0_dense.clone(),
cu_seqlens,
scale,
)
o_pool, ht_pool = chunk_kda_cutedsl(
q_b,
k_b,
v_b,
g_act[0].float(),
beta[0].float(),
pool,
cu_seqlens,
scale,
h0_indices=slots,
)
torch.cuda.synchronize()
# Same kernels and math; only the state addressing differs -> bit-identical.
assert ht_pool is pool
assert torch.equal(o_pool, o_dense)
assert torch.equal(pool[slots.long()], ht_dense)
untouched = torch.ones(num_slots, dtype=torch.bool, device="cuda")
untouched[slots.long()] = False
assert torch.equal(pool[untouched], pool_before[untouched])
if __name__ == "__main__": if __name__ == "__main__":
import sys import sys