[AMD] Tilelang sparse fwd for dsv32 mi355/mi300 (#19945)
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user