[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,
is_causal=True,
block_I=64,
inner_iter=1,
num_stages=1,
threads=256,
):
"""
grid: (seq_len * REPLICATE_H, top_k_blocks).
Each block does one topk block, writes partial_o, partial_lse.
grid: (seq_len * REPLICATE_H, top_k / block_I / inner_iter)
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 kv_group == 1
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
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)
REPLICATE_H = (head_kv // 64) if head_kv > 64 else 1
H_per_block = padded_H if REPLICATE_H == 1 else 64
N_GROUPS = topk // (block_I * inner_iter)
BI = block_I
NI = topk // block_I
D = dim
D_tail = tail_dim
q_shape = [batch, seq_len, heads, dim + tail_dim]
kv_shape = [batch, seq_len_kv, kv_group, dim + tail_dim]
indices_shape = [batch, seq_len, kv_group, topk]
partial_o_shape = [batch, seq_len, NI, heads, dim]
partial_lse_shape = [batch, seq_len, NI, heads]
partial_o_shape = [batch, seq_len, N_GROUPS, heads, dim]
partial_lse_shape = [batch, seq_len, N_GROUPS, heads]
indices_dtype = T.int32
dtype = T.bfloat16
accum_dtype = T.float32
_q_in_shared = inner_iter == 1
@T.prim_func
def main(
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_Lse: T.Tensor(partial_lse_shape, accum_dtype),
):
with T.Kernel(seq_len * REPLICATE_H, NI, threads=threads) as (bx, by):
Q_shared = T.alloc_shared([H_per_block, D], dtype)
Q_tail_shared = T.alloc_shared([H_per_block, D_tail], dtype)
with T.Kernel(seq_len * REPLICATE_H, N_GROUPS, threads=threads) as (bx, by):
if _q_in_shared:
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)
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)
acc_o = T.alloc_fragment([H_per_block, D], 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)
alpha = 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(sumexp, 0)
T.fill(m_i, -(2**30))
b_i, g_i = 0, 0
s_i = bx if REPLICATE_H == 1 else (bx // REPLICATE_H)
topk_block_i = by
q_i = s_i
group_i = by
H0 = 0 if REPLICATE_H == 1 else (bx % REPLICATE_H) * 64
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_tail_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_buf)
for bi_i in T.Parallel(BI):
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):
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
]
for bi_i, d_i in T.Parallel(BI, D_tail):
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
]
for h_i, bi_i in T.Parallel(H_per_block, BI):
acc_s[h_i, bi_i] = T.if_then_else(
mask[bi_i], 0, -T.infinity(acc_s.dtype)
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):
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):
idx = Indices[b_i, s_i, g_i, topk_block_i * BI + bi_i]
KV_shared[bi_i, d_i] = KV[
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):
idx = Indices[b_i, s_i, g_i, topk_block_i * BI + bi_i]
K_tail_shared[bi_i, d_i] = KV[
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):
acc_s[h_i, bi_i] = T.if_then_else(
mask[bi_i], 0, -T.infinity(acc_s.dtype)
)
T.gemm(
Q_buf,
KV_shared,
acc_s,
transpose_B=True,
policy=T.GemmWarpPolicy.FullCol,
)
T.gemm(
Q_shared,
KV_shared,
acc_s,
transpose_B=True,
policy=T.GemmWarpPolicy.FullCol,
)
T.gemm(
Q_tail_shared,
K_tail_shared,
acc_s,
transpose_B=True,
policy=T.GemmWarpPolicy.FullCol,
)
T.reduce_max(acc_s, m_i, dim=1, clear=True)
for h_i in T.Parallel(H_per_block):
m_i[h_i] = T.max(m_i[h_i], -(2**30))
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] * sm_scale - m_i[h_i] * sm_scale
T.gemm(
Q_tail_buf,
K_tail_shared,
acc_s,
transpose_B=True,
policy=T.GemmWarpPolicy.FullCol,
)
T.reduce_sum(acc_s, sumexp_i, dim=1)
T.copy(acc_s, S_shared)
T.gemm(S_shared, KV_shared, acc_o, policy=T.GemmWarpPolicy.FullCol)
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):
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):
acc_s[h_i, bi_i] = T.exp2(
acc_s[h_i, bi_i] * sm_scale - m_i[h_i] * sm_scale
)
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]
# sumexp_i==0 (all masked), divide by 1 to get 0 and avoid nan
T.copy(acc_s, S_shared)
T.gemm(S_shared, KV_shared, acc_o, policy=T.GemmWarpPolicy.FullCol)
# 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):
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):
sumexp_i[h_i] = T.if_then_else(
sumexp_i[h_i] == 0.0,
sumexp[h_i] = T.if_then_else(
sumexp[h_i] == 0.0,
-(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(sumexp_i, Partial_Lse[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, Partial_Lse[b_i, s_i, group_i, H0:H1])
return main
@@ -1022,45 +1050,63 @@ def tilelang_sparse_fwd(
tail_dim = dim - d_v
topk = indices.shape[-1]
assert topk == 2048
if _is_hip:
if _is_gfx95_supported:
# decode kernel
if q.shape[0] <= 64:
kernel_partial = sparse_mla_fwd_decode_partial(
num_heads,
d_v,
tail_dim,
topk,
sm_scale=sm_scale,
block_I=64,
threads=256,
)
kernel_combine = sparse_mla_fwd_decode_combine(
num_heads, d_v, topk, head_per_block=4, block_I=64, threads=256
)
partial_o, partial_lse = kernel_partial(
q.unsqueeze(0), kv.unsqueeze(0), indices.unsqueeze(0)
)
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,
)
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:
# gfx950
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(
num_heads,
d_v,
tail_dim,
topk,
sm_scale=sm_scale,
block_I=block_I,
inner_iter=inner_iter,
num_stages=1,
threads=threads,
)
kernel_combine = sparse_mla_fwd_decode_combine(
num_heads,
d_v,
n_groups * block_I,
head_per_block=4,
block_I=block_I,
threads=threads,
)
partial_o, partial_lse = kernel_partial(
q.unsqueeze(0), kv.unsqueeze(0), indices.unsqueeze(0)
)
out = kernel_combine(partial_o, partial_lse)
else:
kernel = sparse_attention_fwd_kernel_v2(
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