submission 738645
gxtzhuxi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 163 lines, June 9 Researcher Reciprocity License v1.0.
mla_decode.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-738645?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:488f575bcd7230a62336409dca8ed845a93de474990838df74bf037f4b4ac9d8
license declaredunknown
license concludedunknown
authorsgxtzhuxi
imported2026-08-26
Kernel source
mla_decode.py163 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
MLA Decode — DeepSeek-R1 MLA decode attention on MI355X (CDNA4).
Config: num_heads=16, num_kv_heads=1 (MQA), qk_head_dim=576, v_head_dim=512.
Decode only (q_seq_len=1), variable-length batching via indptr.
== Optimizations ==
1. Cached metadata — eliminate metadata compute per call (~50 us)
2. Per-tensor FP8 quantization for Q
3. Optimal num_kv_splits — fill 256 CUs without over-splitting
4. Cached static tensors — kv_indices, kv_last_page_len, output buffer
5. No GPU→CPU sync — avoid .item() calls entirely
"""
import torch
from task import input_t, output_t
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
# ===== MLA constants =====
NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM # 576
V_HEAD_DIM = KV_LORA_RANK # 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
NUM_CUS = 256
FP8_MAX = torch.finfo(FP8_DTYPE).max
FP8_MIN = torch.finfo(FP8_DTYPE).min
# ===== Global caches =====
_SHAPE_CACHE = {}
_OUT_CACHE = {}
def _optimal_splits(batch_size):
"""Fill 256 CUs with minimal splitting to avoid reduction overhead."""
base = batch_size * NUM_HEADS
if base >= NUM_CUS:
return 1
splits = (NUM_CUS + base - 1) // base
return min(16, max(1, splits))
def _build_shape_cache(batch_size, kv_seq_len, device):
"""One-time setup per (batch_size, kv_seq_len) shape."""
total_kv = batch_size * kv_seq_len
num_splits = _optimal_splits(batch_size)
kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)
kv_last_page_len = torch.full(
(batch_size,), kv_seq_len, dtype=torch.int32, device=device,
)
qo_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device)
kv_indptr = (
torch.arange(batch_size + 1, dtype=torch.int32, device=device) * kv_seq_len
)
info = get_mla_metadata_info_v1(
batch_size, 1, 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]
(wm, wi, wis, ri, rfm, rpm) = work
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_last_page_len,
NUM_HEADS // NUM_KV_HEADS,
NUM_KV_HEADS,
True,
wm, wis, wi, ri, rfm, rpm,
page_size=PAGE_SIZE,
kv_granularity=max(PAGE_SIZE, 16),
max_seqlen_qo=1,
uni_seqlen_qo=1,
fast_mode=False,
max_split_per_batch=num_splits,
intra_batch_mode=True,
dtype_q=FP8_DTYPE,
dtype_kv=FP8_DTYPE,
)
return {
"kv_indices": kv_indices,
"kv_last_page_len": kv_last_page_len,
"num_splits": num_splits,
"work_meta_data": wm,
"work_indptr": wi,
"work_info_set": wis,
"reduce_indptr": ri,
"reduce_final_map": rfm,
"reduce_partial_map": rpm,
}
def _get_output(batch_size, device):
if batch_size not in _OUT_CACHE:
_OUT_CACHE[batch_size] = torch.empty(
batch_size, NUM_HEADS, V_HEAD_DIM,
dtype=torch.bfloat16, device=device,
)
return _OUT_CACHE[batch_size]
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
kv_seq_len = config["kv_seq_len"]
device = q.device
kv_fp8, kv_scale = kv_data["fp8"]
# FP8 quantize Q (per-tensor)
amax = q.abs().amax().clamp(min=1e-12)
q_scale = amax / FP8_MAX
q_fp8 = (q / q_scale).clamp(min=FP8_MIN, max=FP8_MAX).to(FP8_DTYPE)
key = (batch_size, kv_seq_len)
if key not in _SHAPE_CACHE:
_SHAPE_CACHE[key] = _build_shape_cache(batch_size, kv_seq_len, device)
c = _SHAPE_CACHE[key]
kv_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, NUM_KV_HEADS, -1)
o = _get_output(batch_size, device)
mla_decode_fwd(
q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM),
kv_4d,
o,
qo_indptr, kv_indptr,
c["kv_indices"],
c["kv_last_page_len"],
1,
page_size=PAGE_SIZE,
nhead_kv=NUM_KV_HEADS,
sm_scale=SM_SCALE,
logit_cap=0.0,
num_kv_splits=c["num_splits"],
q_scale=q_scale.to(torch.float32).reshape(1),
kv_scale=kv_scale,
intra_batch_mode=True,
work_meta_data=c["work_meta_data"],
work_indptr=c["work_indptr"],
work_info_set=c["work_info_set"],
reduce_indptr=c["reduce_indptr"],
reduce_final_map=c["reduce_final_map"],
reduce_partial_map=c["reduce_partial_map"],
)
return o
scrolls · 163 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