submission 586016
josusanmartin · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 109 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-586016?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:f769edf86c91fcb08b8159b9cdb163db6363088df97d9da609b83798f91a1432
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15
Kernel source
submission.py109 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""v697 - Direct stage1+reduce with PAGE_SIZE=1, pre-allocated split buffers."""
import os
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
os.environ.setdefault("AMD_DIRECT_DISPATCH", "1")
import torch
from task import input_t, output_t
import aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
try:
from aiter.jit.module_quant import static_per_tensor_quant
except Exception:
try:
from aiter.ops.quant import static_per_tensor_quant
except Exception:
static_per_tensor_quant = None
NUM_HEADS = 16
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
NUM_KV_SPLITS = 32
FP8_DTYPE = aiter_dtypes.fp8
_FP8_FINFO = torch.finfo(FP8_DTYPE)
_cache = {}
def _quantize_q_fp8(q, q_fp8_buf):
amax = q.abs().amax().clamp(min=1e-12)
scale = (amax / _FP8_FINFO.max).reshape(1).to(torch.float32)
if static_per_tensor_quant is not None:
static_per_tensor_quant(q_fp8_buf, q, scale)
else:
q_fp8_buf.copy_(
(q / scale).clamp(min=_FP8_FINFO.min, max=_FP8_FINFO.max).to(FP8_DTYPE)
)
return scale
def _get_cache(dev, qo_indptr, kv_indptr, bs, kvlen):
key = (dev.index, bs, kvlen)
c = _cache.get(key)
if c is not None:
return c
total_kv = bs * kvlen
kv_indices = torch.arange(total_kv, dtype=torch.int32, device=dev)
kv_lpl = torch.ones(bs, dtype=torch.int32, device=dev)
out = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)
q_fp8 = torch.empty((bs, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device=dev)
# Metadata matching reference: PAGE_SIZE=1, is_causal=True, fast_mode=False
info = get_mla_metadata_info_v1(
bs, 1, NUM_HEADS, FP8_DTYPE, FP8_DTYPE,
is_sparse=False, fast_mode=False,
num_kv_splits=NUM_KV_SPLITS, intra_batch_mode=True,
)
bufs = [torch.empty(s, dtype=t, device=dev) for s, t in info]
wmd, wi, wis, ri, rfm, rpm = bufs
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_lpl,
16, 1, True,
wmd, wis, wi, ri, rfm, rpm,
page_size=1, kv_granularity=16,
max_seqlen_qo=1, uni_seqlen_qo=1,
fast_mode=False, max_split_per_batch=NUM_KV_SPLITS,
intra_batch_mode=True,
dtype_q=FP8_DTYPE, dtype_kv=FP8_DTYPE,
)
# Pre-allocate split buffers (size from reduce_partial_map)
pt = int(rpm.numel())
split_out = torch.empty((pt, 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device=dev)
split_lse = torch.empty((pt, 1, NUM_HEADS, 1), dtype=torch.float32, device=dev)
c = (kv_indices, kv_lpl, out, q_fp8, wmd, wi, wis, ri, rfm, rpm, split_out, split_lse)
_cache[key] = c
return c
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = int(config["batch_size"])
kvlen = int(config["kv_seq_len"])
sm_scale = float(config["sm_scale"])
cache = _get_cache(q.device, qo_indptr, kv_indptr, bs, kvlen)
kv_indices, kv_lpl, out, q_fp8, wmd, wi, wis, ri, rfm, rpm, split_out, split_lse = cache
q_scale = _quantize_q_fp8(q, q_fp8)
kv_fp8, kv_scale = kv_data["fp8"]
kv_buf = kv_fp8.view(-1, 1, 1, QK_HEAD_DIM)
aiter.mla_decode_stage1_asm_fwd(
q_fp8, kv_buf, qo_indptr, kv_indptr, kv_indices, kv_lpl, None,
wmd, wi, wis, 1, 1, 1, sm_scale, split_out, split_lse, out, q_scale, kv_scale,
)
aiter.mla_reduce_v1(split_out, split_lse, ri, rfm, rpm, 1, out, None)
return out
scrolls · 109 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