submission 646049
JohnHe · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 206 lines, June 9 Researcher Reciprocity License v1.0.
solution.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-646049?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:76cf2f8d3141574a5bfcf4ae131c1714db7b3ba113b926ade596b1cdffd9f9be
license declaredunknown
license concludedunknown
authorsJohnHe
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
persistent-kernel
Optimizations over the reference (aiter a8w8 persistent-mode kernel):Kernel source
solution.py206 lines
"""
Optimized MLA decode kernel for AMD MI355X (CDNA4).
Optimizations over the reference (aiter a8w8 persistent-mode kernel):
1. Adaptive NUM_KV_SPLITS per workload shape
MI355X has 256 CUs across 8 XCDs. The reference uses a fixed 32 splits.
For small batches (4×16=64 head-batch pairs), 32 splits gives 2048 blocks
which is fine, but for batch=256 the reduction overhead of 32 splits is
wasteful since batch parallelism alone (4096 pairs) saturates the CUs.
We adaptively select 16-128 splits based on batch×heads and kv_seq_len.
2. Metadata buffer caching
Persistent-mode requires 6 work buffers allocated via cudaMalloc + a
metadata population kernel. For repeated calls with identical geometry
(common in continuous batching), we cache the allocated buffers and only
re-run the cheap population kernel, saving cudaMalloc latency.
3. kv_indices caching
Simple contiguous range reused across calls with same total_kv_len.
4. Minimized Python-side overhead
Fewer intermediate variables, direct dict lookups, avoid unnecessary
tensor operations.
"""
import torch
from task import input_t, output_t
from utils import make_match_reference
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
# ---------------------------------------------------------------------------
# DeepSeek R1 MLA constants (forward_absorb path)
# ---------------------------------------------------------------------------
NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = 576 # kv_lora_rank + qk_rope_head_dim
V_HEAD_DIM = 512 # = kv_lora_rank
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
_FP8_FINFO = torch.finfo(FP8_DTYPE)
# ---------------------------------------------------------------------------
# FP8 quantization (hot path — kept minimal)
# ---------------------------------------------------------------------------
def _quantize_q_fp8(q: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
amax = q.abs().amax().clamp(min=1e-12)
scale = amax / _FP8_FINFO.max
q_fp8 = (q / scale).clamp(min=_FP8_FINFO.min, max=_FP8_FINFO.max).to(FP8_DTYPE)
return q_fp8, scale.to(torch.float32).reshape(1)
# ---------------------------------------------------------------------------
# Adaptive NUM_KV_SPLITS
# ---------------------------------------------------------------------------
def _select_splits(batch_size: int, kv_seq_len: int) -> int:
"""
MI355X: 256 CUs, 8 XCDs of 32 CUs each.
Total blocks = batch_size * NUM_HEADS * num_kv_splits.
Want >= 256 blocks (1/CU), ideally 512-2048 for latency hiding.
But more splits = more reduction overhead.
Also: each split processes kv_seq_len/num_kv_splits tokens.
Too few tokens per split -> underutilized compute.
"""
head_batches = batch_size * NUM_HEADS
if head_batches >= 256:
# batch>=16: 256+ head-batch pairs, CUs saturated from batch alone
return 16 if kv_seq_len <= 2048 else 32
elif head_batches >= 64:
# batch=4-15: moderate parallelism
return 32 if kv_seq_len <= 2048 else 64
else:
# batch=1-3: need splits for parallelism
return 64 if kv_seq_len <= 2048 else 128
# ---------------------------------------------------------------------------
# Caches
# ---------------------------------------------------------------------------
_meta_buf_cache: dict = {} # keyed on geometry -> pre-allocated buffers
_kv_idx_cache: dict = {} # keyed on total_kv_len -> int32 range tensor
def _cached_kv_indices(n: int) -> torch.Tensor:
if n not in _kv_idx_cache:
_kv_idx_cache[n] = torch.arange(n, dtype=torch.int32, device="cuda")
return _kv_idx_cache[n]
def _get_metadata(
batch_size, max_q_len, nq, nkv,
q_dtype, kv_dtype,
qo_indptr, kv_indptr, kv_last_page_len,
num_kv_splits,
):
total_kv = int(kv_indptr[-1].item())
key = (batch_size, max_q_len, nq, nkv, q_dtype, kv_dtype, num_kv_splits, total_kv)
if key not in _meta_buf_cache:
# Allocate work buffers (expensive: cudaMalloc)
info = get_mla_metadata_info_v1(
batch_size, max_q_len, nq, q_dtype, kv_dtype,
is_sparse=False, fast_mode=False,
num_kv_splits=num_kv_splits, intra_batch_mode=True,
)
bufs = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
_meta_buf_cache[key] = {
"work_meta_data": bufs[0],
"work_indptr": bufs[1],
"work_info_set": bufs[2],
"reduce_indptr": bufs[3],
"reduce_final_map": bufs[4],
"reduce_partial_map": bufs[5],
}
m = _meta_buf_cache[key]
# Populate (cheap kernel - must run every call as indptrs may differ)
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_last_page_len,
nq // nkv, nkv, True,
m["work_meta_data"], m["work_info_set"], m["work_indptr"],
m["reduce_indptr"], m["reduce_final_map"], m["reduce_partial_map"],
page_size=PAGE_SIZE,
kv_granularity=max(PAGE_SIZE, 16),
max_seqlen_qo=max_q_len,
uni_seqlen_qo=max_q_len,
fast_mode=False,
max_split_per_batch=num_kv_splits,
intra_batch_mode=True,
dtype_q=q_dtype,
dtype_kv=kv_dtype,
)
return m
# ---------------------------------------------------------------------------
# Main entry point
# ---------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
"""
Optimized MLA decode: FP8 Q + FP8 KV via aiter persistent-mode kernel,
with MI355X-tuned adaptive splitting and metadata caching.
"""
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
q_seq_len = config["q_seq_len"]
kv_seq_len = config["kv_seq_len"]
# Adaptive split count for MI355X
num_splits = _select_splits(batch_size, kv_seq_len)
# FP8 quantize Q on-the-fly
q_fp8, q_scale = _quantize_q_fp8(q)
# Pre-quantized FP8 KV
kv_fp8, kv_scale = kv_data["fp8"]
# 4D view for aiter: (total_kv, page_size=1, nkv=1, 576) - zero-copy
kv_4d = kv_fp8.view(kv_fp8.shape[0], 1, 1, kv_fp8.shape[-1])
total_kv = int(kv_indptr[-1].item())
kv_indices = _cached_kv_indices(total_kv)
kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
meta = _get_metadata(
batch_size, q_seq_len, NUM_HEADS, NUM_KV_HEADS,
q_fp8.dtype, kv_fp8.dtype,
qo_indptr, kv_indptr, kv_last_page_len,
num_splits,
)
# Fresh output buffer (must not reuse - caller may retain reference)
o = torch.empty((q.shape[0], NUM_HEADS, V_HEAD_DIM),
dtype=torch.bfloat16, device="cuda")
mla_decode_fwd(
q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM),
kv_4d,
o,
qo_indptr,
kv_indptr,
kv_indices,
kv_last_page_len,
q_seq_len,
page_size=PAGE_SIZE,
nhead_kv=NUM_KV_HEADS,
sm_scale=SM_SCALE,
logit_cap=0.0,
num_kv_splits=num_splits,
q_scale=q_scale,
kv_scale=kv_scale,
intra_batch_mode=True,
**meta,
)
return oscrolls · 206 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