submission 722089
ihansel · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 358 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-722089?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:7f72deaeedca54c9e16803b2296c798d4bdc5ef63d2a294d48333b7b88a8afc8
license declaredunknown
license concludedunknown
authorsihansel
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
kv_mxfp4_buf, kv_mxfp4_scale = kv_data["mxfp4"]fp8
q_fp8 = q_chunk.to(tl.float8e4nv)mma
acc += tl.dot(p.to(tl.bfloat16), v_block, out_dtype=tl.float32)num-warps = 8
num_warps=8, num_stages=2,online-softmax
m_new = tl.maximum(m_i, m_ij)split-k
split_kv_len = split_end - split_startstages = 2
num_warps=8, num_stages=2,tile-k = 256
BLOCK_K = 256tile-n = 64
BLOCK_N = 64Kernel source
submission.py358 lines
"""MLA Decode v60: Multi-head KV reuse + metadata caching.
M2 experiment: load KV ONCE per program, process ALL 16 query heads.
16x reduction in KV bandwidth vs single-head kernel.
Grid: (batch * num_splits) instead of (batch * heads * num_splits).
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
from reference import ref_kernel # noqa: F401
from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
SM_SCALE = 1.0 / (576 ** 0.5)
SM_SCALE_LOG2 = SM_SCALE * 1.44269504088896
FP8_DTYPE = aiter_dtypes.fp8
PAGE_SIZE = 1
PADDED_V = 512
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)
# ==================== MULTI-HEAD MXFP4 KERNEL (M2) ====================
@triton.jit
def _mla_multihead_stage1(
Q, stride_qt, stride_qh, stride_qd,
K_fp4, stride_kf_t, stride_kf_d,
K_scales, stride_ks_t, stride_ks_d,
V_bf16, stride_vt, stride_vd,
Mid_O, stride_mo_s, stride_mo_h, stride_mo_d,
Mid_LSE, stride_ml_s, stride_ml_h,
qo_indptr, kv_indptr,
BLOCK_H: tl.constexpr,
NUM_SPLITS: tl.constexpr,
sm_scale_log2,
QK_DIM: tl.constexpr, V_DIM: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr, BLOCK_DV: tl.constexpr,
):
"""Process ALL query heads per program, sharing KV load."""
# Grid: (batch_size * NUM_SPLITS,)
pid = tl.program_id(0)
split_id = pid % NUM_SPLITS
batch_id = pid // NUM_SPLITS
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
tokens_per_split = (kv_len + NUM_SPLITS - 1) // NUM_SPLITS
split_start = split_id * tokens_per_split
split_end = tl.minimum(split_start + tokens_per_split, kv_len)
split_kv_len = split_end - split_start
offs_h = tl.arange(0, BLOCK_H)
offs_dv = tl.arange(0, BLOCK_DV)
flat_idx = q_start * BLOCK_H + offs_h
if split_kv_len <= 0:
o_ptrs = Mid_O + split_id * stride_mo_s + flat_idx[:, None] * stride_mo_h + offs_dv[None, :] * stride_mo_d
tl.store(o_ptrs, tl.zeros([BLOCK_H, BLOCK_DV], dtype=tl.bfloat16), mask=offs_dv[None, :] < V_DIM)
lse_ptrs = Mid_LSE + split_id * stride_ml_s + flat_idx * stride_ml_h
tl.store(lse_ptrs, tl.full([BLOCK_H], float('-inf'), dtype=tl.float32))
return
NUM_K_ITERS: tl.constexpr = (QK_DIM + BLOCK_K - 1) // BLOCK_K
PACKED_BK: tl.constexpr = BLOCK_K // 2
SCALE_BK: tl.constexpr = BLOCK_K // 32
q_base = Q + q_start * stride_qt
m_i = tl.full([BLOCK_H], float('-inf'), dtype=tl.float32)
l_i = tl.full([BLOCK_H], 0.0, dtype=tl.float32)
acc = tl.zeros([BLOCK_H, BLOCK_DV], dtype=tl.float32)
offs_n = tl.arange(0, BLOCK_N)
offs_bk = tl.arange(0, PACKED_BK)
offs_sk = tl.arange(0, SCALE_BK)
offs_qk = tl.arange(0, BLOCK_K)
for block_start in range(0, split_kv_len, BLOCK_N):
mask_n = (block_start + offs_n) < split_kv_len
kv_pos = kv_start + split_start + block_start
qk = tl.zeros([BLOCK_H, BLOCK_N], dtype=tl.float32)
for k_iter in range(NUM_K_ITERS):
k_offset = k_iter * BLOCK_K
k_remaining = QK_DIM - k_offset
# Q for ALL heads: [BLOCK_H, BLOCK_K]
q_ptrs = q_base + offs_h[:, None] * stride_qh + (k_offset + offs_qk[None, :]) * stride_qd
q_chunk = tl.load(q_ptrs, mask=offs_qk[None, :] < k_remaining, other=0.0)
q_fp8 = q_chunk.to(tl.float8e4nv)
q_scale = tl.full([BLOCK_H, SCALE_BK], 127, dtype=tl.uint8)
# K (SHARED across all heads): [PACKED_BK, BLOCK_N]
k_ptrs = K_fp4 + (kv_pos + offs_n[None, :]) * stride_kf_t + (k_offset // 2 + offs_bk[:, None]) * stride_kf_d
k_mask = (offs_bk[:, None] < (k_remaining + 1) // 2) & mask_n[None, :]
k_block = tl.load(k_ptrs, mask=k_mask, other=0)
# K scales (SHARED): [BLOCK_N, SCALE_BK]
ks_ptrs = K_scales + (kv_pos + offs_n[:, None]) * stride_ks_t + (k_offset // 32 + offs_sk[None, :]) * stride_ks_d
ks_mask = mask_n[:, None] & (offs_sk[None, :] < (k_remaining + 31) // 32)
k_scales_chunk = tl.load(ks_ptrs, mask=ks_mask, other=0)
# [BLOCK_H, BLOCK_N] — all heads, shared K
qk = tl.dot_scaled(q_fp8, q_scale, "e4m3", k_block, k_scales_chunk, "e2m1",
acc=qk, fast_math=True)
qk = qk * sm_scale_log2
qk = tl.where(mask_n[None, :], qk, float('-inf'))
m_ij = tl.max(qk, axis=1)
m_new = tl.maximum(m_i, m_ij)
alpha = tl.math.exp2(m_i - m_new)
p = tl.math.exp2(qk - m_new[:, None])
l_ij = tl.sum(p, axis=1)
acc = acc * alpha[:, None]
l_i = l_i * alpha + l_ij
m_i = m_new
# V (SHARED across all heads): [BLOCK_N, V_DIM]
v_ptrs = V_bf16 + (kv_pos + offs_n[:, None]) * stride_vt + offs_dv[None, :] * stride_vd
v_mask = mask_n[:, None] & (offs_dv[None, :] < V_DIM)
v_block = tl.load(v_ptrs, mask=v_mask, other=0.0)
# OV: [BLOCK_H, BLOCK_N] × [BLOCK_N, BLOCK_DV] = [BLOCK_H, BLOCK_DV]
acc += tl.dot(p.to(tl.bfloat16), v_block, out_dtype=tl.float32)
partial_out = acc / l_i[:, None]
o_ptrs = Mid_O + split_id * stride_mo_s + flat_idx[:, None] * stride_mo_h + offs_dv[None, :] * stride_mo_d
tl.store(o_ptrs, partial_out.to(tl.bfloat16), mask=offs_dv[None, :] < V_DIM)
lse_vals = m_i + tl.math.log2(l_i)
lse_ptrs = Mid_LSE + split_id * stride_ml_s + flat_idx * stride_ml_h
tl.store(lse_ptrs, lse_vals)
@triton.jit
def _mla_reduce(
Mid_O, stride_mo_s, stride_mo_h, stride_mo_d,
Mid_LSE, stride_ml_s, stride_ml_h,
Out, stride_ot, stride_oh, stride_od,
qo_indptr, num_q_heads: tl.constexpr,
NUM_SPLITS: tl.constexpr, V_DIM: tl.constexpr, BLOCK_DV: tl.constexpr,
):
pid = tl.program_id(0)
batch_id = pid // num_q_heads
head_id = pid % num_q_heads
q_start = tl.load(qo_indptr + batch_id)
flat_idx = q_start * num_q_heads + head_id
offs_dv = tl.arange(0, BLOCK_DV)
max_lse = tl.full([1], float('-inf'), dtype=tl.float32)
for s in range(NUM_SPLITS):
lse_s = tl.load(Mid_LSE + s * stride_ml_s + flat_idx * stride_ml_h)
max_lse = tl.maximum(max_lse, lse_s)
acc = tl.zeros([BLOCK_DV], dtype=tl.float32)
sum_w = tl.zeros([1], dtype=tl.float32)
for s in range(NUM_SPLITS):
lse_s = tl.load(Mid_LSE + s * stride_ml_s + flat_idx * stride_ml_h)
w = tl.math.exp2(lse_s - max_lse)
sum_w += w
o_ptrs = Mid_O + s * stride_mo_s + flat_idx * stride_mo_h + offs_dv * stride_mo_d
partial = tl.load(o_ptrs, mask=offs_dv < V_DIM, other=0.0).to(tl.float32)
acc += w * partial
acc = acc / sum_w
o_base = Out + q_start * stride_ot + head_id * stride_oh
tl.store(o_base + offs_dv * stride_od, acc.to(tl.bfloat16), mask=offs_dv < V_DIM)
# ==================== METADATA CACHE ====================
_meta_cache = {}
_aiter_buffers = {}
def _get_cached_metadata(batch_size, q_seq_len, kv_seq_len, nq, nkv, dtype_q, dtype_kv,
num_kv_splits, device):
key = (batch_size, q_seq_len, kv_seq_len, nq, nkv, dtype_q, dtype_kv, num_kv_splits)
if key not in _meta_cache:
total_kv_len = batch_size * kv_seq_len
qo_indptr_c = torch.arange(0, batch_size + 1, dtype=torch.int32, device=device) * q_seq_len
kv_indptr_c = torch.arange(0, batch_size + 1, dtype=torch.int32, device=device) * kv_seq_len
kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)
kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device=device)
info = get_mla_metadata_info_v1(
batch_size, q_seq_len, nq, dtype_q, dtype_kv,
is_sparse=False, fast_mode=True,
num_kv_splits=num_kv_splits, intra_batch_mode=True,
)
work = [torch.empty(s, dtype=t, device=device) for s, t in info]
(wm, wi, wis, ri, rfm, rpm) = work
get_mla_metadata_v1(
qo_indptr_c, kv_indptr_c, kv_last_page_len,
nq // nkv, nkv, True,
wm, wis, wi, ri, rfm, rpm,
page_size=PAGE_SIZE, kv_granularity=32,
max_seqlen_qo=q_seq_len, uni_seqlen_qo=q_seq_len,
fast_mode=True, max_split_per_batch=num_kv_splits,
intra_batch_mode=True, dtype_q=dtype_q, dtype_kv=dtype_kv,
)
meta = {"work_meta_data": wm, "work_indptr": wi, "work_info_set": wis,
"reduce_indptr": ri, "reduce_final_map": rfm, "reduce_partial_map": rpm}
_meta_cache[key] = (meta, qo_indptr_c, kv_indptr_c, kv_indices, kv_last_page_len)
return _meta_cache[key]
# ==================== AITER PATH ====================
def _aiter_path(q, kv_data, qo_indptr, kv_indptr, config):
batch_size = config["batch_size"]
nq = config["num_heads"]
nkv = config["num_kv_heads"]
dq = config["qk_head_dim"]
dv = config["v_head_dim"]
q_seq_len = config["q_seq_len"]
kv_seq_len = config["kv_seq_len"]
total_kv_tokens = batch_size * kv_seq_len
use_fp8 = total_kv_tokens > 131072
if use_fp8:
q_input, q_scale = quantize_fp8(q)
kv_buffer, kv_scale = kv_data["fp8"]
else:
q_input, q_scale = q, None
kv_buffer = kv_data["bf16"]
kv_scale = None
kv_buffer_4d = kv_buffer.view(kv_buffer.shape[0], PAGE_SIZE, nkv, kv_buffer.shape[-1])
if total_kv_tokens <= 4096:
num_kv_splits = 16
elif total_kv_tokens <= 32768:
num_kv_splits = 16
elif total_kv_tokens <= 65536:
num_kv_splits = 16
elif total_kv_tokens <= 262144:
num_kv_splits = 32
else:
num_kv_splits = 64
meta, qo_c, kv_c, kv_idx, kv_lpl = _get_cached_metadata(
batch_size, q_seq_len, kv_seq_len, nq, nkv,
q_input.dtype, kv_buffer.dtype, num_kv_splits, q.device,
)
buf_key = (batch_size, q_seq_len, nq, dv)
if buf_key not in _aiter_buffers:
_aiter_buffers[buf_key] = torch.empty(
(q.shape[0], nq, dv), dtype=torch.bfloat16, device=q.device
)
o = _aiter_buffers[buf_key]
mla_decode_fwd(
q_input.view(-1, nq, dq), kv_buffer_4d, o,
qo_c, kv_c, kv_idx, kv_lpl,
q_seq_len, page_size=PAGE_SIZE, nhead_kv=nkv, sm_scale=SM_SCALE,
logit_cap=0.0, num_kv_splits=num_kv_splits,
q_scale=q_scale, kv_scale=kv_scale, intra_batch_mode=True,
**meta,
)
return o
# ==================== CUSTOM MULTI-HEAD MXFP4 PATH ====================
def _custom_multihead_path(q, kv_data, qo_indptr, kv_indptr, config):
batch_size = config["batch_size"]
nq = config["num_heads"]
dq = config["qk_head_dim"]
dv = config["v_head_dim"]
kv_seq_len = config["kv_seq_len"]
kv_mxfp4_buf, kv_mxfp4_scale = kv_data["mxfp4"]
total_kv = kv_mxfp4_buf.shape[0]
total_q = q.shape[0]
total_qh = total_q * nq
k_fp4 = kv_mxfp4_buf.view(total_kv, -1).view(torch.uint8)
num_scale_blocks = (dq + 31) // 32
k_scales = kv_mxfp4_scale.view(torch.uint8)[:total_kv, :num_scale_blocks]
kv_bf16 = kv_data["bf16"]
kv_2d = kv_bf16.view(total_kv, -1)
# Tune splits for CU occupancy: grid = batch_size * NUM_SPLITS
if batch_size <= 4:
NUM_SPLITS = max(4, 256 // batch_size) # Target ~256 programs
elif batch_size <= 32:
NUM_SPLITS = max(4, 256 // batch_size)
elif batch_size <= 128:
NUM_SPLITS = 4
else:
NUM_SPLITS = 1
# BLOCK_N=64 gives -18% on bs4/kv1k (better occupancy for small shapes)
# vs BLOCK_N=128 which is better for large KV sequences
if kv_seq_len <= 1024:
BLOCK_N = 64
else:
BLOCK_N = 128
max_useful_splits = max(1, kv_seq_len // BLOCK_N)
NUM_SPLITS = min(NUM_SPLITS, max_useful_splits)
BLOCK_K = 256
mid_o = torch.empty((NUM_SPLITS, total_qh, dv), dtype=torch.bfloat16, device=q.device)
mid_lse = torch.full((NUM_SPLITS, total_qh), float('-inf'), dtype=torch.float32, device=q.device)
grid1 = (batch_size * NUM_SPLITS,)
_mla_multihead_stage1[grid1](
q, q.stride(0), q.stride(1), q.stride(2),
k_fp4, k_fp4.stride(0), k_fp4.stride(1),
k_scales, k_scales.stride(0), k_scales.stride(1),
kv_2d, kv_2d.stride(0), kv_2d.stride(1),
mid_o, mid_o.stride(0), mid_o.stride(1), mid_o.stride(2),
mid_lse, mid_lse.stride(0), mid_lse.stride(1),
qo_indptr, kv_indptr,
BLOCK_H=nq, NUM_SPLITS=NUM_SPLITS,
sm_scale_log2=SM_SCALE_LOG2,
QK_DIM=dq, V_DIM=dv,
BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K, BLOCK_DV=PADDED_V,
num_warps=8, num_stages=2,
)
o = torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device=q.device)
grid2 = (batch_size * nq,)
_mla_reduce[grid2](
mid_o, mid_o.stride(0), mid_o.stride(1), mid_o.stride(2),
mid_lse, mid_lse.stride(0), mid_lse.stride(1),
o, o.stride(0), o.stride(1), o.stride(2),
qo_indptr, num_q_heads=nq,
NUM_SPLITS=NUM_SPLITS, V_DIM=dv, BLOCK_DV=PADDED_V, num_warps=4,
)
return o
# ==================== DISPATCH ====================
CUSTOM_THRESHOLD = 8192 # Use custom for small shapes
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
total_kv_tokens = config["batch_size"] * config["kv_seq_len"]
if total_kv_tokens <= CUSTOM_THRESHOLD:
return _custom_multihead_path(q, kv_data, qo_indptr, kv_indptr, config)
else:
return _aiter_path(q, kv_data, qo_indptr, kv_indptr, config)
scrolls · 358 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