submission 753817
sean_nobricks · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 473 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-753817?include=source"interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, int32
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:cc23736b2700f4b273bb18e9aafd822bb60526b8302904ef4aafd430b725da97
license declaredunknown
license concludedunknown
authorssean_nobricks
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
qk = tl.dot(q_nope, tl.trans(k_nope))num-warps = 4
num_warps=4, num_stages=2, waves_per_eu=1,online-softmax
m_i_new = tl.maximum(m_i, tl.max(qk, 1))persistent-kernel
per-tile mask. An AITER non-persistent fallback handles very large workloadssplit-k
"""MLA decode — Triton fp8 flash-decoding with split-K and LSE reduction.stages = 2
num_warps=4, num_stages=2, waves_per_eu=1,tile-n = 16
BLOCK_N = 16Kernel source
submission.py473 lines
"""MLA decode — Triton fp8 flash-decoding with split-K and LSE reduction.
Two-stage flash-decoding for DeepSeek R1 forward_absorb MLA on MI355X:
- Stage 1 splits the KV sequence across CTAs, runs Q@K^T + online softmax + P@V
with fp8 KV loads and 16-head MQA packing, and writes per-split partial
outputs and LSEs.
- Stage 2 reduces the per-split partials with an LSE-weighted sum.
An exact-no-mask stage 1 specialization is dispatched when the KV length
divides evenly into the chosen split count and BLOCK_N, eliminating the
per-tile mask. An AITER non-persistent fallback handles very large workloads
where library ASM bandwidth exceeds the custom path.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
# DeepSeek R1 forward_absorb MLA constants
NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576 # kv_lora_rank (512) + qk_rope_head_dim (64)
V_HEAD_DIM = 512 # = kv_lora_rank
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
PAGE_SIZE = 1
LOG2E = 1.44269504
from aiter import dtypes as aiter_dtypes
FP8_DTYPE = aiter_dtypes.fp8
# Pre-allocated workspace buffers
_workspace_partial_out = None
_workspace_partial_lse = None
_workspace_final_out = None
# ---------------------------------------------------------------------------
# Stage 1: flash-decoding with fp8 KV loads and MQA head packing
# ---------------------------------------------------------------------------
@triton.jit
def mla_flash_decode_stage1(
Q, KV_fp8, Out_partial, LSE_partial,
qo_indptr, kv_indptr,
KV_scale_ptr,
sm_scale_log2e,
stride_qt, stride_qh, stride_qd,
stride_kvt, stride_kvd,
stride_opt_b, stride_opt_s, stride_opt_h, stride_opt_d,
stride_lse_b, stride_lse_s, stride_lse_h,
BLOCK_H: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_DK: tl.constexpr,
BLOCK_DPE: tl.constexpr,
BLOCK_DV: tl.constexpr,
NUM_KV_SPLITS: tl.constexpr,
):
batch_id = tl.program_id(0)
split_id = tl.program_id(1)
tl.assume(stride_qt > 0)
tl.assume(stride_qh > 0)
tl.assume(stride_qd > 0)
tl.assume(stride_kvt > 0)
tl.assume(stride_kvd > 0)
tl.assume(stride_opt_b > 0)
tl.assume(stride_opt_s > 0)
tl.assume(stride_opt_h > 0)
tl.assume(stride_opt_d > 0)
tl.assume(stride_lse_b > 0)
tl.assume(stride_lse_s > 0)
tl.assume(stride_lse_h > 0)
q_start = tl.load(qo_indptr + batch_id)
kv_start = tl.load(kv_indptr + batch_id)
kv_end = tl.load(kv_indptr + batch_id + 1)
kv_len = kv_end - kv_start
kv_per_split = tl.cdiv(kv_len, NUM_KV_SPLITS)
split_kv_start = split_id * kv_per_split
split_kv_end = tl.minimum(split_kv_start + kv_per_split, kv_len)
# Early exit for unused splits
if split_kv_start >= kv_len:
offs_h = tl.arange(0, BLOCK_H)
offs_dv = tl.arange(0, BLOCK_DV)
out_ptrs = (Out_partial + batch_id * stride_opt_b + split_id * stride_opt_s
+ offs_h[:, None] * stride_opt_h + offs_dv[None, :] * stride_opt_d)
tl.store(out_ptrs, tl.zeros([BLOCK_H, BLOCK_DV], dtype=tl.bfloat16))
lse_ptrs = (LSE_partial + batch_id * stride_lse_b + split_id * stride_lse_s
+ offs_h * stride_lse_h)
tl.store(lse_ptrs, tl.full([BLOCK_H], value=float('-inf'), dtype=tl.float32))
return
kv_scale = tl.load(KV_scale_ptr).to(tl.float32)
qk_scale = kv_scale * sm_scale_log2e
offs_h = tl.arange(0, BLOCK_H)
offs_dk = tl.max_contiguous(tl.multiple_of(tl.arange(0, BLOCK_DK), BLOCK_DK), BLOCK_DK)
offs_dpe = tl.max_contiguous(tl.multiple_of(tl.arange(0, BLOCK_DPE), BLOCK_DPE), BLOCK_DPE)
q_base = Q + q_start * stride_qt
q_nope = tl.load(
q_base + offs_h[:, None] * stride_qh + offs_dk[None, :] * stride_qd,
cache_modifier=".cg",
)
q_pe = tl.load(
q_base + offs_h[:, None] * stride_qh + (BLOCK_DK + offs_dpe[None, :]) * stride_qd,
cache_modifier=".cg",
)
q_nope = (q_nope.to(tl.float32) * qk_scale).to(tl.bfloat16)
q_pe = (q_pe.to(tl.float32) * qk_scale).to(tl.bfloat16)
m_i = tl.full([BLOCK_H], value=float('-inf'), dtype=tl.float32)
l_i = tl.zeros([BLOCK_H], dtype=tl.float32)
acc = tl.zeros([BLOCK_H, BLOCK_DV], dtype=tl.float32)
offs_n = tl.arange(0, BLOCK_N)
kv_base = KV_fp8 + kv_start * stride_kvt
for n_start in range(split_kv_start, split_kv_end, BLOCK_N):
n_offs = n_start + offs_n
kv_mask = n_offs < split_kv_end
kv_ptrs_base = kv_base + n_offs[:, None] * stride_kvt
k_nope_fp8 = tl.load(
kv_ptrs_base + offs_dk[None, :] * stride_kvd,
mask=kv_mask[:, None], other=0.0,
cache_modifier=".cg",
)
k_pe_fp8 = tl.load(
kv_ptrs_base + (BLOCK_DK + offs_dpe[None, :]) * stride_kvd,
mask=kv_mask[:, None], other=0.0,
cache_modifier=".cg",
)
k_nope = k_nope_fp8.to(tl.bfloat16)
k_pe = k_pe_fp8.to(tl.bfloat16)
qk = tl.dot(q_nope, tl.trans(k_nope))
qk += tl.dot(q_pe, tl.trans(k_pe))
qk = tl.where(kv_mask[None, :], qk, float('-inf'))
m_i_new = tl.maximum(m_i, tl.max(qk, 1))
alpha = tl.math.exp2(m_i - m_i_new)
p = tl.math.exp2(qk - m_i_new[:, None])
l_i = l_i * alpha + tl.sum(p, 1)
acc = acc * alpha[:, None]
acc += tl.dot(p.to(tl.bfloat16), k_nope)
m_i = m_i_new
acc = acc * kv_scale / l_i[:, None]
offs_dv = tl.arange(0, BLOCK_DV)
out_ptrs = (Out_partial + batch_id * stride_opt_b + split_id * stride_opt_s
+ offs_h[:, None] * stride_opt_h + offs_dv[None, :] * stride_opt_d)
tl.store(out_ptrs, acc.to(tl.bfloat16))
lse = (tl.math.log2(l_i) + m_i) * 0.6931471805599453
lse_ptrs = (LSE_partial + batch_id * stride_lse_b + split_id * stride_lse_s
+ offs_h * stride_lse_h)
tl.store(lse_ptrs, lse)
@triton.jit
def mla_flash_decode_stage1_exact_nomask(
Q, KV_fp8, Out_partial, LSE_partial,
qo_indptr, kv_indptr,
KV_scale_ptr,
sm_scale_log2e,
stride_qt, stride_qh, stride_qd,
stride_kvt, stride_kvd,
stride_opt_b, stride_opt_s, stride_opt_h, stride_opt_d,
stride_lse_b, stride_lse_s, stride_lse_h,
BLOCK_H: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_DK: tl.constexpr,
BLOCK_DPE: tl.constexpr,
BLOCK_DV: tl.constexpr,
NUM_KV_SPLITS: tl.constexpr,
):
batch_id = tl.program_id(0)
split_id = tl.program_id(1)
tl.assume(stride_qt > 0)
tl.assume(stride_qh > 0)
tl.assume(stride_qd > 0)
tl.assume(stride_kvt > 0)
tl.assume(stride_kvd > 0)
tl.assume(stride_opt_b > 0)
tl.assume(stride_opt_s > 0)
tl.assume(stride_opt_h > 0)
tl.assume(stride_opt_d > 0)
tl.assume(stride_lse_b > 0)
tl.assume(stride_lse_s > 0)
tl.assume(stride_lse_h > 0)
q_start = tl.load(qo_indptr + batch_id)
kv_start = tl.load(kv_indptr + batch_id)
kv_end = tl.load(kv_indptr + batch_id + 1)
kv_len = kv_end - kv_start
split_kv_tokens = kv_len // NUM_KV_SPLITS
split_kv_tokens = tl.multiple_of(split_kv_tokens, BLOCK_N)
tl.assume(split_kv_tokens >= BLOCK_N)
split_kv_start = split_id * split_kv_tokens
split_kv_start = tl.multiple_of(split_kv_start, BLOCK_N)
split_kv_end = split_kv_start + split_kv_tokens
kv_scale = tl.load(KV_scale_ptr).to(tl.float32)
qk_scale = kv_scale * sm_scale_log2e
offs_h = tl.arange(0, BLOCK_H)
offs_dk = tl.max_contiguous(tl.multiple_of(tl.arange(0, BLOCK_DK), BLOCK_DK), BLOCK_DK)
offs_dpe = tl.max_contiguous(tl.multiple_of(tl.arange(0, BLOCK_DPE), BLOCK_DPE), BLOCK_DPE)
q_base = Q + q_start * stride_qt
q_nope = tl.load(
q_base + offs_h[:, None] * stride_qh + offs_dk[None, :] * stride_qd,
cache_modifier=".cg",
)
q_pe = tl.load(
q_base + offs_h[:, None] * stride_qh + (BLOCK_DK + offs_dpe[None, :]) * stride_qd,
cache_modifier=".cg",
)
q_nope = (q_nope.to(tl.float32) * qk_scale).to(tl.bfloat16)
q_pe = (q_pe.to(tl.float32) * qk_scale).to(tl.bfloat16)
m_i = tl.full([BLOCK_H], value=float("-inf"), dtype=tl.float32)
l_i = tl.zeros([BLOCK_H], dtype=tl.float32)
acc = tl.zeros([BLOCK_H, BLOCK_DV], dtype=tl.float32)
offs_n = tl.arange(0, BLOCK_N)
kv_base = KV_fp8 + kv_start * stride_kvt
for n_start in range(split_kv_start, split_kv_end, BLOCK_N):
n_offs = n_start + offs_n
kv_ptrs_base = kv_base + n_offs[:, None] * stride_kvt
k_nope_fp8 = tl.load(
kv_ptrs_base + offs_dk[None, :] * stride_kvd,
cache_modifier=".cg",
)
k_pe_fp8 = tl.load(
kv_ptrs_base + (BLOCK_DK + offs_dpe[None, :]) * stride_kvd,
cache_modifier=".cg",
)
k_nope = k_nope_fp8.to(tl.bfloat16)
k_pe = k_pe_fp8.to(tl.bfloat16)
qk = tl.dot(q_nope, tl.trans(k_nope))
qk += tl.dot(q_pe, tl.trans(k_pe))
m_i_new = tl.maximum(m_i, tl.max(qk, 1))
alpha = tl.math.exp2(m_i - m_i_new)
p = tl.math.exp2(qk - m_i_new[:, None])
l_i = l_i * alpha + tl.sum(p, 1)
acc = acc * alpha[:, None]
acc += tl.dot(p.to(tl.bfloat16), k_nope)
m_i = m_i_new
acc = acc * kv_scale / l_i[:, None]
offs_dv = tl.arange(0, BLOCK_DV)
out_ptrs = (Out_partial + batch_id * stride_opt_b + split_id * stride_opt_s
+ offs_h[:, None] * stride_opt_h + offs_dv[None, :] * stride_opt_d)
tl.store(out_ptrs, acc.to(tl.bfloat16))
lse = (tl.math.log2(l_i) + m_i) * 0.6931471805599453
lse_ptrs = (LSE_partial + batch_id * stride_lse_b + split_id * stride_lse_s
+ offs_h * stride_lse_h)
tl.store(lse_ptrs, lse)
# ---------------------------------------------------------------------------
# Stage 2: LSE-weighted reduction of split partial outputs
# ---------------------------------------------------------------------------
@triton.jit
def mla_flash_decode_stage2(
Out_partial, LSE_partial, Out_final,
stride_opt_b, stride_opt_s, stride_opt_h, stride_opt_d,
stride_lse_b, stride_lse_s, stride_lse_h,
stride_of_t, stride_of_h, stride_of_d,
BLOCK_DV: tl.constexpr,
NUM_KV_SPLITS: tl.constexpr,
):
batch_id = tl.program_id(0)
head_id = tl.program_id(1)
lse_base = LSE_partial + batch_id * stride_lse_b + head_id * stride_lse_h
out_base = Out_partial + batch_id * stride_opt_b + head_id * stride_opt_h
lse_max = tl.load(lse_base)
for s in tl.static_range(1, NUM_KV_SPLITS):
lse_s = tl.load(lse_base + s * stride_lse_s)
lse_max = tl.where(lse_s > lse_max, lse_s, lse_max)
offs_dv = tl.arange(0, BLOCK_DV)
acc = tl.zeros([BLOCK_DV], dtype=tl.float32)
weight_sum = 0.0
for s in tl.static_range(NUM_KV_SPLITS):
lse_s = tl.load(lse_base + s * stride_lse_s)
w = tl.exp(lse_s - lse_max)
partial = tl.load(out_base + s * stride_opt_s + offs_dv * stride_opt_d).to(tl.float32)
acc += w * partial
weight_sum += w
acc = acc / weight_sum
out_ptrs = (Out_final + batch_id * stride_of_t + head_id * stride_of_h
+ offs_dv * stride_of_d)
tl.store(out_ptrs, acc.to(tl.bfloat16))
def _quantize_fp8(tensor):
finfo = torch.finfo(FP8_DTYPE)
amax = tensor.abs().amax().clamp(min=1e-12)
scale = amax / finfo.max
fp8_tensor = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
return fp8_tensor, scale.to(torch.float32).reshape(1)
_cached_aiter_data = {}
def _library_mla_decode_nonpersistent(q, kv_data, qo_indptr, kv_indptr, config):
global _cached_aiter_data
from aiter.mla import mla_decode_fwd
q_fp8, q_scale = _quantize_fp8(q)
kv_buffer_fp8, kv_scale = kv_data["fp8"]
batch_size = config["batch_size"]
kv_seq_len = config["kv_seq_len"]
total_kv_len = batch_size * kv_seq_len
kv_buffer_4d = kv_buffer_fp8.view(
kv_buffer_fp8.shape[0], PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM
)
cache_key = (batch_size, kv_seq_len)
if cache_key not in _cached_aiter_data:
kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
o = torch.empty((q.shape[0], NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
_cached_aiter_data[cache_key] = (kv_indices, kv_last_page_len, o)
kv_indices, kv_last_page_len, o = _cached_aiter_data[cache_key]
mla_decode_fwd(
q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM),
kv_buffer_4d, o,
qo_indptr, kv_indptr, kv_indices, kv_last_page_len,
config["q_seq_len"],
page_size=PAGE_SIZE, nhead_kv=NUM_KV_HEADS,
sm_scale=SM_SCALE, logit_cap=0.0,
q_scale=q_scale, kv_scale=kv_scale,
)
return o
def _choose_kv_splits(batch_size, kv_len, block_n=64, num_cus=256):
"""Choose splits to balance CU utilization vs stage2 overhead."""
total_ctas = batch_size
if total_ctas >= num_cus:
return 1
target_cu = max(1, num_cus // batch_size)
max_useful = max(1, kv_len // (4 * block_n))
splits = min(target_cu, max_useful)
for po2 in [1, 2, 4, 8, 16, 32, 64]:
if po2 >= splits:
return po2
return 64
# ---------------------------------------------------------------------------
# Entry point
# ---------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
global _workspace_partial_out, _workspace_partial_lse, _workspace_final_out
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
kv_seq_len = config["kv_seq_len"]
# AITER fallback for very large workloads where library ASM bandwidth
# exceeds the custom path; bs=256/kv=8192 stays on the custom kernel.
total_kv_tokens = batch_size * kv_seq_len
use_aiter_fallback = total_kv_tokens >= 1048576 and not (batch_size == 256 and kv_seq_len == 8192)
if use_aiter_fallback:
return _library_mla_decode_nonpersistent(q, kv_data, qo_indptr, kv_indptr, config)
kv_buffer_fp8, kv_scale_tensor = kv_data["fp8"]
BLOCK_H = 16
if batch_size == 4 and kv_seq_len == 1024:
BLOCK_N = 16
elif batch_size <= 32 and kv_seq_len <= 1024:
BLOCK_N = 32
else:
BLOCK_N = 64
BLOCK_DK = 512
BLOCK_DPE = 64
BLOCK_DV = 512
NUM_KV_SPLITS = _choose_kv_splits(batch_size, kv_seq_len, BLOCK_N)
uses_exact_split_stage1 = (kv_seq_len % (NUM_KV_SPLITS * BLOCK_N) == 0)
po_shape = (batch_size, NUM_KV_SPLITS, NUM_HEADS, V_HEAD_DIM)
pl_shape = (batch_size, NUM_KV_SPLITS, NUM_HEADS)
if _workspace_partial_out is None or _workspace_partial_out.shape != po_shape:
_workspace_partial_out = torch.empty(po_shape, dtype=torch.bfloat16, device=q.device)
_workspace_partial_lse = torch.empty(pl_shape, dtype=torch.float32, device=q.device)
sm_scale_log2e = SM_SCALE * LOG2E
stride_kvt = kv_buffer_fp8.stride(0)
stride_kvd = kv_buffer_fp8.stride(2)
grid_stage1 = (batch_size, NUM_KV_SPLITS)
stage1_args = (
q, kv_buffer_fp8, _workspace_partial_out, _workspace_partial_lse,
qo_indptr, kv_indptr,
kv_scale_tensor,
sm_scale_log2e,
q.stride(0), q.stride(1), q.stride(2),
stride_kvt, stride_kvd,
_workspace_partial_out.stride(0), _workspace_partial_out.stride(1),
_workspace_partial_out.stride(2), _workspace_partial_out.stride(3),
_workspace_partial_lse.stride(0), _workspace_partial_lse.stride(1),
_workspace_partial_lse.stride(2),
)
stage1_meta = dict(
BLOCK_H=BLOCK_H, BLOCK_N=BLOCK_N,
BLOCK_DK=BLOCK_DK, BLOCK_DPE=BLOCK_DPE, BLOCK_DV=BLOCK_DV,
NUM_KV_SPLITS=NUM_KV_SPLITS,
num_warps=4, num_stages=2, waves_per_eu=1,
schedule_hint="memory-bound-attention",
)
if uses_exact_split_stage1:
mla_flash_decode_stage1_exact_nomask[grid_stage1](
*stage1_args,
**stage1_meta,
)
else:
mla_flash_decode_stage1[grid_stage1](
*stage1_args,
**stage1_meta,
)
if NUM_KV_SPLITS == 1:
return _workspace_partial_out.squeeze(1)
fo_shape = (q.shape[0], NUM_HEADS, V_HEAD_DIM)
if _workspace_final_out is None or _workspace_final_out.shape != fo_shape:
_workspace_final_out = torch.empty(fo_shape, dtype=torch.bfloat16, device=q.device)
grid_stage2 = (batch_size, NUM_HEADS)
mla_flash_decode_stage2[grid_stage2](
_workspace_partial_out, _workspace_partial_lse, _workspace_final_out,
_workspace_partial_out.stride(0), _workspace_partial_out.stride(1),
_workspace_partial_out.stride(2), _workspace_partial_out.stride(3),
_workspace_partial_lse.stride(0), _workspace_partial_lse.stride(1),
_workspace_partial_lse.stride(2),
_workspace_final_out.stride(0), _workspace_final_out.stride(1),
_workspace_final_out.stride(2),
BLOCK_DV=V_HEAD_DIM,
NUM_KV_SPLITS=NUM_KV_SPLITS,
num_warps=4, num_stages=1,
)
return _workspace_final_out
scrolls · 473 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON