submission 589553
Ananda Sai A · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 106 lines, June 9 Researcher Reciprocity License v1.0.
submission_v4_hybrid.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-589553?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:7c11b77e85e1bfd038dda95e95ceb1d773944cd90361c11c44e508b84ec300c1
license declaredunknown
license concludedunknown
authorsAnanda Sai A
imported2026-08-15
Kernel source
submission_v4_hybrid.py106 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
import torch
import torch.nn.functional as F
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
NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
_warm = set()
def _quantize_fp8(tensor):
finfo = torch.finfo(FP8_DTYPE)
amax = tensor.abs().amax().clamp(min=1e-12)
scale = amax / finfo.max
return (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE), scale.float().reshape(1)
def _aiter_fp8_decode(q, kv_data, qo_indptr, kv_indptr, config, num_splits=32):
"""Standard AITER fp8 path for large shapes."""
bs = config["batch_size"]
q_len = config["q_seq_len"]
total_kv = int(kv_indptr[-1].item())
q_fp8, q_scale = _quantize_fp8(q)
kv_fp8, kv_scale = kv_data["fp8"]
kv_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, NUM_KV_HEADS, kv_fp8.shape[-1])
info = get_mla_metadata_info_v1(bs, q_len, NUM_HEADS, q_fp8.dtype, kv_fp8.dtype,
is_sparse=False, fast_mode=False, num_kv_splits=num_splits, intra_batch_mode=True)
work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
kv_last = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
get_mla_metadata_v1(qo_indptr, kv_indptr, kv_last,
NUM_HEADS, NUM_KV_HEADS, True,
work[0], work[2], work[1], work[3], work[4], work[5],
page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE, 16),
max_seqlen_qo=q_len, uni_seqlen_qo=q_len,
fast_mode=False, max_split_per_batch=num_splits,
intra_batch_mode=True, dtype_q=q_fp8.dtype, dtype_kv=kv_fp8.dtype)
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, q_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,
work_meta_data=work[0], work_indptr=work[1], work_info_set=work[2],
reduce_indptr=work[3], reduce_final_map=work[4], reduce_partial_map=work[5])
return o
def _sdpa_decode(q, kv_data, config):
"""Direct SDPA for small batches — bypasses flash attention overhead."""
bs = config["batch_size"]
kv_len = config["kv_seq_len"]
# Use bf16 KV for SDPA (variable-length not supported, assume uniform kv_len)
kv_bf16 = kv_data["bf16"] # (total_kv, 1, 576)
# Reshape for batched attention
Q = q.view(bs, 1, NUM_HEADS, QK_HEAD_DIM).transpose(1, 2) # (bs, 16, 1, 576)
K = kv_bf16.view(bs, kv_len, 1, QK_HEAD_DIM).permute(0, 2, 1, 3).expand(bs, NUM_HEADS, kv_len, QK_HEAD_DIM) # (bs, 16, kv_len, 576)
V = kv_bf16[:, :, :V_HEAD_DIM].view(bs, kv_len, 1, V_HEAD_DIM).permute(0, 2, 1, 3).expand(bs, NUM_HEADS, kv_len, V_HEAD_DIM) # (bs, 16, kv_len, 512)
out = F.scaled_dot_product_attention(Q, K, V, scale=SM_SCALE, is_causal=False)
return out.transpose(1, 2).reshape(bs, NUM_HEADS, V_HEAD_DIM) # (bs, 16, 512)
def _bmm_decode(q, kv_data, config):
"""Direct bmm for smallest shapes — minimum overhead."""
bs = config["batch_size"]
kv_len = config["kv_seq_len"]
kv_bf16 = kv_data["bf16"].view(bs, kv_len, QK_HEAD_DIM) # (bs, kv_len, 576)
Q = q.view(bs, NUM_HEADS, QK_HEAD_DIM) # (bs, 16, 576)
K = kv_bf16 # (bs, kv_len, 576)
V = kv_bf16[:, :, :V_HEAD_DIM] # (bs, kv_len, 512)
# scores: (bs, 16, kv_len)
scores = torch.bmm(Q, K.transpose(1, 2)) * SM_SCALE
weights = F.softmax(scores, dim=-1).to(torch.bfloat16)
# output: (bs, 16, 512)
out = torch.bmm(weights, V)
return out
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
kv_len = config["kv_seq_len"]
# Shape dispatch — use fastest path per regime
if bs <= 4:
# Small batch: direct bmm is fastest (minimal overhead)
return _bmm_decode(q, kv_data, config)
else:
# All other shapes: AITER fp8 (best throughput)
return _aiter_fp8_decode(q, kv_data, qo_indptr, kv_indptr, config, num_splits=32)
scrolls · 106 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