submission 754240
zhuang000123 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 317 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-754240?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:fcd9e64cb585ac9c842e65edcbceb818e007df413b0e127ff96d2134faee156f
license declaredunknown
license concludedunknown
authorszhuang000123
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 Triton for bs<=4 kv<=1024, FP8 ASM for rest.fp8
q_tile1, q_descale1, "e4m3", acc=qk_all, fast_math=True)num-warps = 8
num_warps=8, num_stages=2, waves_per_eu=2,online-softmax
m_new = tl.maximum(m_i, m_ij)stages = 2
num_warps=8, num_stages=2, waves_per_eu=2,tile-n = 64
BLOCK_N=64, V_DIM=V_HEAD_DIM,Kernel source
submission.py317 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
# gpumode leaderboard — v65 cost-model splits
"""
v65: Use aiter's cost model for per-config NUM_KV_SPLITS.
bs=256: 1 split (no reduce!), bs=64: 4, bs=32: 8, bs=4: 8-16.
MXFP4 Triton for bs<=4 kv<=1024, FP8 ASM for rest.
"""
import torch
import triton
import triton.language as tl
import aiter as _aiter_mod
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.ops.quant import static_per_tensor_quant
NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (576 ** 0.5)
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
# MXFP4 constants
QK_PACKED = QK_HEAD_DIM // 2
K_PART1_PACKED = 256
K_PART1_SCALES = 16
Q_PART1_DIM = 512
MXFP4_SPLITS = 16
# FP8 ASM constants
ASM_SPLITS = 16
# ====== MXFP4 Triton Kernel (v30) ======
@triton.jit
def _mxfp4_stage1(
Q_fp8, KV_packed, KV_Scale, V_fp8, Q_scale_ptr, V_scale_ptr, kv_indptr,
Out, Lse,
stride_q_batch, stride_q_head, stride_q_dim,
stride_kv_n, stride_ks_n, stride_vn,
stride_os, stride_oq, stride_oh,
stride_ls, stride_lq, stride_lh,
SM_SCALE_VAL,
BLOCK_N: tl.constexpr, V_DIM: tl.constexpr,
K_P1_PACKED: tl.constexpr, K_P1_SCALES: tl.constexpr,
Q_P1_DIM: tl.constexpr, NUM_HEADS_: tl.constexpr, NUM_SPLITS: tl.constexpr,
):
split_id = tl.program_id(0)
batch_id = tl.program_id(1)
LOG2E: tl.constexpr = 1.44269504
kv_start = tl.load(kv_indptr + batch_id)
kv_end = tl.load(kv_indptr + batch_id + 1)
kv_len = kv_end - kv_start
split_size = tl.cdiv(kv_len, NUM_SPLITS)
my_start = kv_start + split_id * split_size
my_end = tl.minimum(my_start + split_size, kv_end)
q_idx = batch_id
if my_start >= my_end:
offs_h = tl.arange(0, NUM_HEADS_)
offs_dv = tl.arange(0, V_DIM)
tl.store(Out + split_id * stride_os + q_idx * stride_oq + offs_h[:, None] * stride_oh + offs_dv[None, :],
tl.zeros([NUM_HEADS_, V_DIM], dtype=tl.float32))
tl.store(Lse + split_id * stride_ls + q_idx * stride_lq + offs_h * stride_lh,
tl.full([NUM_HEADS_], float("-inf"), dtype=tl.float32))
return
q_base = Q_fp8 + q_idx * stride_q_batch
offs_d1 = tl.arange(0, Q_P1_DIM)
offs_h = tl.arange(0, NUM_HEADS_)
q_tile1 = tl.load(q_base + offs_d1[:, None] * stride_q_dim + offs_h[None, :] * stride_q_head)
q_descale1 = tl.full([NUM_HEADS_, Q_P1_DIM // 32], 127, dtype=tl.uint8)
qk_factor = tl.load(Q_scale_ptr) * SM_SCALE_VAL
v_descale = tl.load(V_scale_ptr)
p_descale_pv = tl.full([NUM_HEADS_, BLOCK_N // 32], 127, dtype=tl.uint8)
v_descale_pv = tl.full([V_DIM, BLOCK_N // 32], 127, dtype=tl.uint8)
m_i = tl.full([NUM_HEADS_], float("-inf"), dtype=tl.float32)
l_i = tl.zeros([NUM_HEADS_], dtype=tl.float32)
acc = tl.zeros([NUM_HEADS_, V_DIM], dtype=tl.float32)
for _iter in range(0, tl.cdiv(my_end - my_start, BLOCK_N)):
kv_pos = my_start + _iter * BLOCK_N
offs_n = tl.arange(0, BLOCK_N)
valid = (kv_pos + offs_n) < my_end
k1_ptrs = KV_packed + (kv_pos + offs_n[:, None]) * stride_kv_n + tl.arange(0, K_P1_PACKED)[None, :]
k1_packed = tl.load(k1_ptrs, mask=valid[:, None], other=0)
ks1_ptrs = KV_Scale + (kv_pos + offs_n[:, None]) * stride_ks_n + tl.arange(0, K_P1_SCALES)[None, :]
k1_scales = tl.load(ks1_ptrs, mask=valid[:, None], other=127)
qk_all = tl.zeros([BLOCK_N, NUM_HEADS_], dtype=tl.float32)
qk_all = tl.dot_scaled(k1_packed, k1_scales, "e2m1",
q_tile1, q_descale1, "e4m3", acc=qk_all, fast_math=True)
qk_all = qk_all * qk_factor
qk_all = tl.where(valid[:, None], qk_all, float("-inf"))
qk = tl.trans(qk_all)
m_ij = tl.max(qk, 1)
m_new = tl.maximum(m_i, m_ij)
alpha = tl.math.exp2((m_i - m_new) * LOG2E)
p = tl.math.exp2((qk - m_new[:, None]) * LOG2E)
l_i = l_i * alpha + tl.sum(p, 1)
acc = acc * alpha[:, None]
m_i = m_new
v_ptrs = V_fp8 + (kv_pos + offs_n[:, None]) * stride_vn + tl.arange(0, V_DIM)[None, :]
v_block = tl.load(v_ptrs, mask=valid[:, None], other=0.0)
p_fp8 = p.to(tl.float8e4nv)
pv = tl.dot_scaled(p_fp8, p_descale_pv, "e4m3",
v_block, v_descale_pv, "e4m3", fast_math=True)
acc += pv
safe_l = tl.where(l_i == 0.0, 1.0, l_i)
result = acc * v_descale / safe_l[:, None]
lse_vals = m_i + tl.math.log2(tl.where(l_i == 0.0, 1.0, l_i)) / LOG2E
offs_h2 = tl.arange(0, NUM_HEADS_)
offs_dv = tl.arange(0, V_DIM)
tl.store(Out + split_id * stride_os + q_idx * stride_oq + offs_h2[:, None] * stride_oh + offs_dv[None, :], result)
tl.store(Lse + split_id * stride_ls + q_idx * stride_lq + offs_h2 * stride_lh, lse_vals)
@triton.jit
def _reduce_splits(
Split_out, Split_lse, Out,
stride_ss, stride_sq, stride_sh, stride_sd,
stride_ls, stride_lq, stride_lh,
stride_oq, stride_oh, stride_od,
V_DIM: tl.constexpr, NUM_SPLITS: tl.constexpr,
):
q_idx = tl.program_id(0)
head_id = tl.program_id(1)
LOG2E: tl.constexpr = 1.44269504
max_lse = float("-inf")
for s in range(NUM_SPLITS):
lse_s = tl.load(Split_lse + s * stride_ls + q_idx * stride_lq + head_id * stride_lh)
max_lse = tl.maximum(max_lse, lse_s)
acc = tl.zeros([V_DIM], dtype=tl.float32)
sum_exp = 0.0
for s in range(NUM_SPLITS):
lse_s = tl.load(Split_lse + s * stride_ls + q_idx * stride_lq + head_id * stride_lh)
w = tl.math.exp2((lse_s - max_lse) * LOG2E)
sum_exp += w
offs_dv = tl.arange(0, V_DIM)
partial = tl.load(Split_out + s * stride_ss + q_idx * stride_sq + head_id * stride_sh + offs_dv)
acc += partial * w
acc = acc / tl.where(sum_exp == 0.0, 1.0, sum_exp)
tl.store(Out + q_idx * stride_oq + head_id * stride_oh + tl.arange(0, V_DIM), acc.to(tl.bfloat16))
# ====== Caches ======
_mxfp4_cache = {}
_asm_cache = {}
def _build_mxfp4_cache(total_q, total_kv, device):
return {
"o": torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device),
"split_out": torch.empty((MXFP4_SPLITS, total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device=device),
"split_lse": torch.empty((MXFP4_SPLITS, total_q, NUM_HEADS), dtype=torch.float32, device=device),
"q_fp8": torch.empty((total_q, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device=device),
"q_scale": torch.tensor([5.5 / 448.0], dtype=torch.float32, device=device),
}
def _compute_optimal_splits(batch_size, kv_seq_len):
"""Compute optimal NUM_KV_SPLITS using aiter's cost model."""
cu_num = 256
overhead = 84.1
avg_kv = kv_seq_len
best_eff, best_splits = -1, 1
for i in range(1, 17):
eff = batch_size * i / (((batch_size * i + cu_num - 1) // cu_num) * cu_num) * avg_kv / (avg_kv + overhead * i)
if eff > best_eff:
best_eff, best_splits = eff, i
# FP8 constraint: min_block_n = 128 for nhead*max_seqlen_q=16
min_block_n = 128
fp8_max = int((kv_seq_len + min_block_n - 1) // min_block_n)
if best_splits > 1:
fp8_max2 = int((abs(kv_seq_len - 1) // min_block_n + 1))
fp8_max = min(fp8_max, fp8_max2)
return min(best_splits, fp8_max)
def _build_asm_cache(batch_size, q_seq_len, total_q, total_kv, device, qo_indptr, kv_indptr, kv_seq_len):
q_fp8 = torch.empty((total_q, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device=device)
q_scale = torch.tensor([5.5 / 448.0], dtype=torch.float32, device=device)
o = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device)
kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)
kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
# Use cost-model optimized splits
num_splits = _compute_optimal_splits(batch_size, kv_seq_len)
info = get_mla_metadata_info_v1(
batch_size, q_seq_len, NUM_HEADS, FP8_DTYPE, FP8_DTYPE,
is_sparse=False, fast_mode=False, num_kv_splits=num_splits, intra_batch_mode=True,
)
work = [torch.empty(s, dtype=t, device=device) for s, t in info]
(work_metadata, work_indptr_t, work_info_set, reduce_indptr, reduce_final_map, reduce_partial_map) = work
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_last_page_len,
NUM_HEADS // NUM_KV_HEADS, NUM_KV_HEADS, False,
work_metadata, work_info_set, work_indptr_t,
reduce_indptr, reduce_final_map, reduce_partial_map,
page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE, 16),
max_seqlen_qo=q_seq_len, uni_seqlen_qo=q_seq_len,
fast_mode=False, max_split_per_batch=num_splits, intra_batch_mode=True,
dtype_q=FP8_DTYPE, dtype_kv=FP8_DTYPE,
)
num_partial = reduce_partial_map.size(0)
logits = torch.empty((num_partial * q_seq_len, 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device=device)
attn_lse = torch.empty((num_partial * q_seq_len, 1, NUM_HEADS, 1), dtype=torch.float32, device=device)
return {
"q_fp8": q_fp8, "q_fp8_flat": q_fp8.view(-1, QK_HEAD_DIM), "q_scale": q_scale,
"o": o, "kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,
"work_meta_data": work_metadata, "work_indptr": work_indptr_t,
"work_info_set": work_info_set, "reduce_indptr": reduce_indptr,
"reduce_final_map": reduce_final_map, "reduce_partial_map": reduce_partial_map,
"logits": logits, "attn_lse": attn_lse,
"q_seq_len": q_seq_len, "total_kv": total_kv,
}
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
total_q = q.shape[0]
total_kv = batch_size * config["kv_seq_len"]
if batch_size <= 4 and config["kv_seq_len"] <= 1024:
# MXFP4 Triton path (better for small batch)
key = (batch_size, total_q, total_kv)
c = _mxfp4_cache.get(key)
if c is None:
c = _build_mxfp4_cache(total_q, total_kv, q.device)
_mxfp4_cache[key] = c
static_per_tensor_quant(c["q_fp8"].view(-1, QK_HEAD_DIM), q.view(-1, QK_HEAD_DIM), c["q_scale"])
kv_buffer_mxfp4, kv_scale_mxfp4 = kv_data["mxfp4"]
kv_raw = kv_buffer_mxfp4.view(torch.uint8).reshape(total_kv, QK_PACKED)
ks_raw = kv_scale_mxfp4.view(torch.uint8)
if ks_raw.dim() > 2:
ks_raw = ks_raw.reshape(total_kv, -1)
kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
v_fp8 = kv_buffer_fp8.view(total_kv, QK_HEAD_DIM)
q_fp8 = c["q_fp8"]
_mxfp4_stage1[(MXFP4_SPLITS, batch_size)](
q_fp8, kv_raw, ks_raw, v_fp8, c["q_scale"], kv_scale_fp8, kv_indptr,
c["split_out"], c["split_lse"],
q_fp8.stride(0), q_fp8.stride(1), q_fp8.stride(2),
kv_raw.stride(0), ks_raw.stride(0), v_fp8.stride(0),
c["split_out"].stride(0), c["split_out"].stride(1), c["split_out"].stride(2),
c["split_lse"].stride(0), c["split_lse"].stride(1), c["split_lse"].stride(2),
SM_SCALE,
BLOCK_N=64, V_DIM=V_HEAD_DIM,
K_P1_PACKED=K_PART1_PACKED, K_P1_SCALES=K_PART1_SCALES,
Q_P1_DIM=Q_PART1_DIM, NUM_HEADS_=NUM_HEADS, NUM_SPLITS=MXFP4_SPLITS,
num_warps=8, num_stages=2, waves_per_eu=2,
)
_reduce_splits[(total_q, NUM_HEADS)](
c["split_out"], c["split_lse"], c["o"],
c["split_out"].stride(0), c["split_out"].stride(1),
c["split_out"].stride(2), c["split_out"].stride(3),
c["split_lse"].stride(0), c["split_lse"].stride(1), c["split_lse"].stride(2),
c["o"].stride(0), c["o"].stride(1), c["o"].stride(2),
V_DIM=V_HEAD_DIM, NUM_SPLITS=MXFP4_SPLITS,
)
return c["o"]
else:
# FP8 ASM path (better for large batch)
q_seq_len = config["q_seq_len"]
key = (batch_size, q_seq_len, total_q, total_kv)
c = _asm_cache.get(key)
if c is None:
c = _build_asm_cache(batch_size, q_seq_len, total_q, total_kv,
q.device, qo_indptr, kv_indptr, config["kv_seq_len"])
_asm_cache[key] = c
static_per_tensor_quant(c["q_fp8_flat"], q.view(-1, QK_HEAD_DIM), c["q_scale"])
kv_buffer_fp8, kv_scale = kv_data["fp8"]
kv_4d = kv_buffer_fp8.view(c["total_kv"], PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)
_aiter_mod.mla_decode_stage1_asm_fwd(
c["q_fp8"], kv_4d, qo_indptr, kv_indptr,
c["kv_indices"], c["kv_last_page_len"],
None, c["work_meta_data"], c["work_indptr"], c["work_info_set"],
c["q_seq_len"], PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
c["logits"], c["attn_lse"], c["o"], c["q_scale"], kv_scale,
)
_aiter_mod.mla_reduce_v1(
c["logits"], c["attn_lse"], c["reduce_indptr"],
c["reduce_final_map"], c["reduce_partial_map"],
c["q_seq_len"], c["o"], None,
)
return c["o"]
scrolls · 317 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