[AMD] Tilelang sparse fwd for dsv32 mi355/mi300 (#19945)

This commit is contained in:
Thomas Wang
2026-03-24 02:01:39 -07:00
committed by GitHub
parent 3dfaa47d5e
commit 855d15adf6
@@ -790,16 +790,22 @@ def sparse_mla_fwd_decode_partial(
sm_scale=None, sm_scale=None,
is_causal=True, is_causal=True,
block_I=64, block_I=64,
inner_iter=1,
num_stages=1,
threads=256, threads=256,
): ):
""" """
grid: (seq_len * REPLICATE_H, top_k_blocks). grid: (seq_len * REPLICATE_H, top_k / block_I / inner_iter)
Each block does one topk block, writes partial_o, partial_lse. Each GPU block processes `inner_iter` consecutive KV tiles and writes one (partial_o, partial_lse) entry.
""" """
assert is_causal == True, "non-causal is not supported" assert is_causal == True, "non-causal is not supported"
assert kv_group == 1 assert kv_group == 1
assert topk % block_I == 0 assert topk % block_I == 0
assert topk % (block_I * inner_iter) == 0, (
f"topk ({topk}) must be divisible by block_I * inner_iter = "
f"{block_I} * {inner_iter}"
)
# log2(e) = 1.44269504 # log2(e) = 1.44269504
if sm_scale is None: if sm_scale is None:
@@ -815,20 +821,22 @@ def sparse_mla_fwd_decode_partial(
padded_H = max(tilelang.math.next_power_of_2(head_kv), 16) padded_H = max(tilelang.math.next_power_of_2(head_kv), 16)
REPLICATE_H = (head_kv // 64) if head_kv > 64 else 1 REPLICATE_H = (head_kv // 64) if head_kv > 64 else 1
H_per_block = padded_H if REPLICATE_H == 1 else 64 H_per_block = padded_H if REPLICATE_H == 1 else 64
N_GROUPS = topk // (block_I * inner_iter)
BI = block_I BI = block_I
NI = topk // block_I
D = dim D = dim
D_tail = tail_dim D_tail = tail_dim
q_shape = [batch, seq_len, heads, dim + tail_dim] q_shape = [batch, seq_len, heads, dim + tail_dim]
kv_shape = [batch, seq_len_kv, kv_group, dim + tail_dim] kv_shape = [batch, seq_len_kv, kv_group, dim + tail_dim]
indices_shape = [batch, seq_len, kv_group, topk] indices_shape = [batch, seq_len, kv_group, topk]
partial_o_shape = [batch, seq_len, NI, heads, dim] partial_o_shape = [batch, seq_len, N_GROUPS, heads, dim]
partial_lse_shape = [batch, seq_len, NI, heads] partial_lse_shape = [batch, seq_len, N_GROUPS, heads]
indices_dtype = T.int32 indices_dtype = T.int32
dtype = T.bfloat16 dtype = T.bfloat16
accum_dtype = T.float32 accum_dtype = T.float32
_q_in_shared = inner_iter == 1
@T.prim_func @T.prim_func
def main( def main(
Q: T.Tensor(q_shape, dtype), Q: T.Tensor(q_shape, dtype),
@@ -837,88 +845,108 @@ def sparse_mla_fwd_decode_partial(
Partial_O: T.Tensor(partial_o_shape, dtype), Partial_O: T.Tensor(partial_o_shape, dtype),
Partial_Lse: T.Tensor(partial_lse_shape, accum_dtype), Partial_Lse: T.Tensor(partial_lse_shape, accum_dtype),
): ):
with T.Kernel(seq_len * REPLICATE_H, NI, threads=threads) as (bx, by): with T.Kernel(seq_len * REPLICATE_H, N_GROUPS, threads=threads) as (bx, by):
Q_shared = T.alloc_shared([H_per_block, D], dtype) if _q_in_shared:
Q_tail_shared = T.alloc_shared([H_per_block, D_tail], dtype) Q_buf = T.alloc_shared([H_per_block, D], dtype)
Q_tail_buf = T.alloc_shared([H_per_block, D_tail], dtype)
else:
Q_buf = T.alloc_fragment([H_per_block, D], dtype)
Q_tail_buf = T.alloc_fragment([H_per_block, D_tail], dtype)
KV_shared = T.alloc_shared([BI, D], dtype) KV_shared = T.alloc_shared([BI, D], dtype)
K_tail_shared = T.alloc_shared([BI, D_tail], dtype) K_tail_shared = T.alloc_shared([BI, D_tail], dtype)
S_shared = T.alloc_shared([H_per_block, BI], dtype)
mask = T.alloc_fragment([BI], T.bool) mask = T.alloc_fragment([BI], T.bool)
acc_o = T.alloc_fragment([H_per_block, D], accum_dtype) acc_o = T.alloc_fragment([H_per_block, D], accum_dtype)
acc_s = T.alloc_fragment([H_per_block, BI], accum_dtype) acc_s = T.alloc_fragment([H_per_block, BI], accum_dtype)
S_shared = T.alloc_shared([H_per_block, BI], dtype) sumexp = T.alloc_fragment([H_per_block], accum_dtype)
sumexp_i = T.alloc_fragment([H_per_block], accum_dtype) sumexp_i = T.alloc_fragment([H_per_block], accum_dtype)
alpha = T.alloc_fragment([H_per_block], accum_dtype)
m_i = T.alloc_fragment([H_per_block], accum_dtype) m_i = T.alloc_fragment([H_per_block], accum_dtype)
m_i_prev = T.alloc_fragment([H_per_block], accum_dtype)
T.fill(acc_o, 0) T.fill(acc_o, 0)
T.fill(sumexp, 0)
T.fill(m_i, -(2**30))
b_i, g_i = 0, 0 b_i, g_i = 0, 0
s_i = bx if REPLICATE_H == 1 else (bx // REPLICATE_H) s_i = bx if REPLICATE_H == 1 else (bx // REPLICATE_H)
topk_block_i = by group_i = by
q_i = s_i
H0 = 0 if REPLICATE_H == 1 else (bx % REPLICATE_H) * 64 H0 = 0 if REPLICATE_H == 1 else (bx % REPLICATE_H) * 64
H1 = H0 + H_per_block H1 = H0 + H_per_block
T.copy(Q[b_i, s_i, H0:H1, :D], Q_shared) T.copy(Q[b_i, s_i, H0:H1, :D], Q_buf)
T.copy(Q[b_i, s_i, H0:H1, D:], Q_tail_shared) T.copy(Q[b_i, s_i, H0:H1, D:], Q_tail_buf)
for k_i in T.Pipelined(inner_iter, num_stages=num_stages):
topk_block_i = group_i * inner_iter + k_i
for bi_i in T.Parallel(BI): for bi_i in T.Parallel(BI):
mask[bi_i] = Indices[b_i, s_i, g_i, topk_block_i * BI + bi_i] >= 0 mask[bi_i] = Indices[b_i, s_i, g_i, topk_block_i * BI + bi_i] >= 0
for bi_i, d_i in T.Parallel(BI, D): for bi_i, d_i in T.Parallel(BI, D):
idx = Indices[b_i, s_i, g_i, topk_block_i * BI + bi_i]
KV_shared[bi_i, d_i] = KV[ KV_shared[bi_i, d_i] = KV[
b_i, Indices[b_i, s_i, g_i, topk_block_i * BI + bi_i], g_i, d_i b_i, T.if_then_else(idx >= 0, idx, 0), g_i, d_i
] ]
for bi_i, d_i in T.Parallel(BI, D_tail): for bi_i, d_i in T.Parallel(BI, D_tail):
idx = Indices[b_i, s_i, g_i, topk_block_i * BI + bi_i]
K_tail_shared[bi_i, d_i] = KV[ K_tail_shared[bi_i, d_i] = KV[
b_i, Indices[b_i, s_i, g_i, topk_block_i * BI + bi_i], g_i, D + d_i b_i, T.if_then_else(idx >= 0, idx, 0), g_i, D + d_i
] ]
for h_i, bi_i in T.Parallel(H_per_block, BI): for h_i, bi_i in T.Parallel(H_per_block, BI):
acc_s[h_i, bi_i] = T.if_then_else( acc_s[h_i, bi_i] = T.if_then_else(
mask[bi_i], 0, -T.infinity(acc_s.dtype) mask[bi_i], 0, -T.infinity(acc_s.dtype)
) )
T.gemm( T.gemm(
Q_shared, Q_buf,
KV_shared, KV_shared,
acc_s, acc_s,
transpose_B=True, transpose_B=True,
policy=T.GemmWarpPolicy.FullCol, policy=T.GemmWarpPolicy.FullCol,
) )
T.gemm( T.gemm(
Q_tail_shared, Q_tail_buf,
K_tail_shared, K_tail_shared,
acc_s, acc_s,
transpose_B=True, transpose_B=True,
policy=T.GemmWarpPolicy.FullCol, policy=T.GemmWarpPolicy.FullCol,
) )
T.reduce_max(acc_s, m_i, dim=1, clear=True) T.copy(m_i, m_i_prev)
T.reduce_max(acc_s, m_i, dim=1, clear=False)
for h_i in T.Parallel(H_per_block): for h_i in T.Parallel(H_per_block):
m_i[h_i] = T.max(m_i[h_i], -(2**30)) alpha[h_i] = T.exp2((m_i_prev[h_i] - m_i[h_i]) * sm_scale)
for h_i, bi_i in T.Parallel(H_per_block, BI): for h_i, bi_i in T.Parallel(H_per_block, BI):
acc_s[h_i, bi_i] = T.exp2( acc_s[h_i, bi_i] = T.exp2(
acc_s[h_i, bi_i] * sm_scale - m_i[h_i] * sm_scale acc_s[h_i, bi_i] * sm_scale - m_i[h_i] * sm_scale
) )
T.reduce_sum(acc_s, sumexp_i, dim=1) T.reduce_sum(acc_s, sumexp_i, dim=1)
for h_i in T.Parallel(H_per_block):
sumexp[h_i] = sumexp[h_i] * alpha[h_i] + sumexp_i[h_i]
for h_i, d_i in T.Parallel(H_per_block, D):
acc_o[h_i, d_i] *= alpha[h_i]
T.copy(acc_s, S_shared) T.copy(acc_s, S_shared)
T.gemm(S_shared, KV_shared, acc_o, policy=T.GemmWarpPolicy.FullCol) T.gemm(S_shared, KV_shared, acc_o, policy=T.GemmWarpPolicy.FullCol)
# sumexp_i==0 (all masked), divide by 1 to get 0 and avoid nan # sumexp==0 (all masked), divide by 1 to get 0 and avoid nan
for h_i, d_i in T.Parallel(H_per_block, D): for h_i, d_i in T.Parallel(H_per_block, D):
acc_o[h_i, d_i] = acc_o[h_i, d_i] / T.if_then_else( acc_o[h_i, d_i] = acc_o[h_i, d_i] / T.if_then_else(
sumexp_i[h_i] == 0.0, 1.0, sumexp_i[h_i] sumexp[h_i] == 0.0, 1.0, sumexp[h_i]
) )
# sumexp_i==0 (all masked), use large negative so combine ignores this split # sumexp==0 (all masked), use large negative so combine ignores this split
for h_i in T.Parallel(H_per_block): for h_i in T.Parallel(H_per_block):
sumexp_i[h_i] = T.if_then_else( sumexp[h_i] = T.if_then_else(
sumexp_i[h_i] == 0.0, sumexp[h_i] == 0.0,
-(2**30), -(2**30),
T.log2(sumexp_i[h_i]) + m_i[h_i] * sm_scale, T.log2(sumexp[h_i]) + m_i[h_i] * sm_scale,
) )
T.copy(acc_o, Partial_O[b_i, s_i, topk_block_i, H0:H1, :]) T.copy(acc_o, Partial_O[b_i, s_i, group_i, H0:H1, :])
T.copy(sumexp_i, Partial_Lse[b_i, s_i, topk_block_i, H0:H1]) T.copy(sumexp, Partial_Lse[b_i, s_i, group_i, H0:H1])
return main return main
@@ -1022,45 +1050,63 @@ def tilelang_sparse_fwd(
tail_dim = dim - d_v tail_dim = dim - d_v
topk = indices.shape[-1] topk = indices.shape[-1]
assert topk == 2048 assert topk == 2048
if _is_hip: if _is_hip:
# sparse_mla_fwd_decode_partial splits topk KV blocks into N_GROUPS
# independent tiles per query, then sparse_mla_fwd_decode_combine
# reduces them via online softmax.
if _is_gfx95_supported: if _is_gfx95_supported:
# decode kernel # gfx950
if q.shape[0] <= 64: block_I, threads = 64, 256
block_per_cu = 2
else:
# gfx942
block_I, threads = 32, 128
block_per_cu = 1
NI = topk // block_I
CU = 304
def _inner_iter(seq: int) -> int:
"""Largest inner_iter ≤ NI that keeps grid/CU ≥ block_per_cu."""
max_it = int(seq * NI / (CU * block_per_cu))
it = NI
while it >= 2:
if it <= max_it and NI % it == 0:
return it
it //= 2
return 1
inner_iter = _inner_iter(q.shape[0])
n_groups = NI // inner_iter
kernel_partial = sparse_mla_fwd_decode_partial( kernel_partial = sparse_mla_fwd_decode_partial(
num_heads, num_heads,
d_v, d_v,
tail_dim, tail_dim,
topk, topk,
sm_scale=sm_scale, sm_scale=sm_scale,
block_I=64, block_I=block_I,
threads=256, inner_iter=inner_iter,
num_stages=1,
threads=threads,
) )
kernel_combine = sparse_mla_fwd_decode_combine( kernel_combine = sparse_mla_fwd_decode_combine(
num_heads, d_v, topk, head_per_block=4, block_I=64, threads=256 num_heads,
d_v,
n_groups * block_I,
head_per_block=4,
block_I=block_I,
threads=threads,
) )
partial_o, partial_lse = kernel_partial( partial_o, partial_lse = kernel_partial(
q.unsqueeze(0), kv.unsqueeze(0), indices.unsqueeze(0) q.unsqueeze(0), kv.unsqueeze(0), indices.unsqueeze(0)
) )
out = kernel_combine(partial_o, partial_lse) out = kernel_combine(partial_o, partial_lse)
return out
# prefill kernel
kernel = sparse_attention_fwd_kernel_v1(
num_heads, d_v, tail_dim, topk, sm_scale=sm_scale, num_stages=1
)
else: # reduce LDS usage on gfx942 target
kernel = sparse_attention_fwd_kernel_v1(
num_heads,
d_v,
tail_dim,
topk,
sm_scale=sm_scale,
block_I=32,
num_stages=1,
threads=128,
)
else: else:
kernel = sparse_attention_fwd_kernel_v2( kernel = sparse_attention_fwd_kernel_v2(
num_heads, d_v, tail_dim, topk, sm_scale=sm_scale num_heads, d_v, tail_dim, topk, sm_scale=sm_scale
) )
return kernel(q.unsqueeze(0), kv.unsqueeze(0), indices.unsqueeze(0)) # type: ignore out = kernel(q.unsqueeze(0), kv.unsqueeze(0), indices.unsqueeze(0)) # type: ignore
return out