submission 728065
Barry_zhang · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 172 lines, June 9 Researcher Reciprocity License v1.0.
submission_combined_v10.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-728065?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:16916ccace332ffc12079fff156256a4fc2a5b3dddcf06b9d26d9f0a55ea51bd
license declaredunknown
license concludedunknown
authorsBarry_zhang
imported2026-08-15
Kernel source
submission_combined_v10.py172 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
MLA Combined v10 — skip-amax on ALL fp8 paths.
1-split removed (fails secret seeds).
- pg1+bf16Q for kv<=1024 (safe)
- pg8+fp8Q+skip_amax for kv>=8192 (22-26% faster per v9 benchmark)
"""
import torch
import triton
import triton.language as tl
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
FP8_DTYPE = aiter_dtypes.fp8
BF16 = torch.bfloat16
_FP8_MAX = float(torch.finfo(FP8_DTYPE).max)
_FIXED_AMAX = 32.0
_meta_cache = {}
_alloc_cache = {}
@triton.jit
def _q_to_fp8_kernel(q_ptr, out_ptr, scale_ptr, amax_ptr,
FP8_MAX: tl.constexpr, N, BLOCK: tl.constexpr):
amax = tl.load(amax_ptr)
amax = tl.where(amax < 1e-12, 1e-12, amax)
scale = amax / FP8_MAX
if tl.program_id(0) == 0:
tl.store(scale_ptr, scale)
pid = tl.program_id(0)
offs = pid * BLOCK + tl.arange(0, BLOCK)
mask = offs < N
x = tl.load(q_ptr + offs, mask=mask, other=0.0).to(tl.float32)
x = x / scale
x = tl.clamp(x, -FP8_MAX, FP8_MAX)
tl.store(out_ptr + offs, x.to(out_ptr.dtype.element_ty), mask=mask)
def _build_meta(batch_size, kv_seq_len, q_seq_len, nq, nkv,
num_kv_splits, page_size, dtype_q, qo_indptr, kv_indptr):
total_kv = batch_size * kv_seq_len
if page_size == 1:
num_pages = total_kv
kv_indptr_pages = kv_indptr
seq_lens = kv_indptr[1:] - kv_indptr[:-1]
kv_last_page_len = seq_lens.to(torch.int32)
else:
num_pages = total_kv // page_size
kv_indptr_pages = kv_indptr // page_size
seq_lens = kv_indptr[1:] - kv_indptr[:-1]
kv_last_page_len = (seq_lens % page_size).to(torch.int32)
kv_last_page_len = torch.where(kv_last_page_len == 0, page_size, kv_last_page_len)
kv_gran = max(1, 16 // page_size)
info = get_mla_metadata_info_v1(
batch_size, q_seq_len, nq, dtype_q, FP8_DTYPE,
is_sparse=False, fast_mode=False,
num_kv_splits=num_kv_splits, intra_batch_mode=True,
)
work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
(wm, wi, wis, ri, rfm, rpm) = work
get_mla_metadata_v1(
qo_indptr, kv_indptr_pages, kv_last_page_len,
nq // nkv, nkv, True,
wm, wis, wi, ri, rfm, rpm,
page_size=page_size,
kv_granularity=kv_gran,
max_seqlen_qo=q_seq_len,
uni_seqlen_qo=q_seq_len,
fast_mode=False,
max_split_per_batch=num_kv_splits,
intra_batch_mode=True,
dtype_q=dtype_q,
dtype_kv=FP8_DTYPE,
)
kv_indices = torch.arange(num_pages, dtype=torch.int32, device="cuda")
return (wm, wi, wis, ri, rfm, rpm, kv_indices, kv_last_page_len, kv_indptr_pages, page_size)
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
nq, nkv = config["num_heads"], config["num_kv_heads"]
dq, dv = config["qk_head_dim"], config["v_head_dim"]
q_seq_len = config["q_seq_len"]
sm_scale = config["sm_scale"]
kv_seq_len = config["kv_seq_len"]
total_kv = batch_size * kv_seq_len
# Route
if kv_seq_len <= 1024:
page_size = 1
dtype_q = BF16
use_fp8_q = False
else:
page_size = 8
dtype_q = FP8_DTYPE
use_fp8_q = True
# Per-shape splits
if batch_size <= 32 and kv_seq_len <= 1024:
num_kv_splits = 8
else:
num_kv_splits = 16
cache_key = (batch_size, kv_seq_len, num_kv_splits, page_size, use_fp8_q)
if cache_key not in _meta_cache:
_meta_cache[cache_key] = _build_meta(
batch_size, kv_seq_len, q_seq_len, nq, nkv,
num_kv_splits, page_size, dtype_q, qo_indptr, kv_indptr)
(wm, wi, wis, ri, rfm, rpm,
kv_indices, kv_last_page_len, kv_indptr_pages, ps) = _meta_cache[cache_key]
kv_buffer_fp8, kv_scale = kv_data["fp8"]
kv_buffer_4d = kv_buffer_fp8.view(-1, ps, nkv, kv_buffer_fp8.shape[-1])
if use_fp8_q:
alloc_key = ("fp8", q.shape[0], nq, dv, dq)
if alloc_key not in _alloc_cache:
_alloc_cache[alloc_key] = (
torch.empty((q.shape[0], nq, dv), dtype=BF16, device="cuda"),
torch.full((1,), _FIXED_AMAX, dtype=torch.float32, device="cuda"),
torch.empty(1, dtype=torch.float32, device="cuda"),
torch.empty(q.shape[0] * nq * dq, dtype=FP8_DTYPE, device="cuda"),
)
o, amax_buf, scale_buf, q_fp8_flat = _alloc_cache[alloc_key]
N = q.numel()
BLOCK = 2048
grid = ((N + BLOCK - 1) // BLOCK,)
# Skip amax — pre-filled
_q_to_fp8_kernel[grid](q, q_fp8_flat, scale_buf, amax_buf,
FP8_MAX=_FP8_MAX, N=N, BLOCK=BLOCK)
mla_decode_fwd(
q_fp8_flat.view(q.shape[0], nq, dq), kv_buffer_4d, o,
qo_indptr, kv_indptr_pages, kv_indices, kv_last_page_len,
q_seq_len, page_size=ps, nhead_kv=nkv,
sm_scale=sm_scale, logit_cap=0.0, num_kv_splits=num_kv_splits,
q_scale=scale_buf, kv_scale=kv_scale,
intra_batch_mode=True,
work_meta_data=wm, work_indptr=wi, work_info_set=wis,
reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm,
)
return o
else:
alloc_key = ("bf16", q.shape[0], nq, dv)
if alloc_key not in _alloc_cache:
_alloc_cache[alloc_key] = torch.empty(
(q.shape[0], nq, dv), dtype=BF16, device="cuda")
o = _alloc_cache[alloc_key]
mla_decode_fwd(
q, kv_buffer_4d, o,
qo_indptr, kv_indptr_pages, kv_indices, kv_last_page_len,
q_seq_len, page_size=ps, nhead_kv=nkv,
sm_scale=sm_scale, logit_cap=0.0, num_kv_splits=num_kv_splits,
kv_scale=kv_scale,
intra_batch_mode=True,
work_meta_data=wm, work_indptr=wi, work_info_set=wis,
reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm,
)
return o
scrolls · 172 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 724364.
#!POPCORN leaderboard amd-mixed-mla#!POPCORN gpu MI355X-"""- v105: Use bf16 Q + bf16 KV (a16w16) to eliminate Q quantization overhead.+ MLA Combined v10 — skip-amax on ALL fp8 paths.+ 1-split removed (fails secret seeds).+ - pg1+bf16Q for kv<=1024 (safe)+ - pg8+fp8Q+skip_amax for kv>=8192 (22-26% faster per v9 benchmark)+ """+ import torch+ import triton+ import triton.language as tl+ 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- The per_tensor_quant_hip call costs ~5-10us. For small batch sizes (bs=4,32)- this is a significant fraction of total time. Using bf16 KV costs 2x bandwidth- but saves a kernel launch + quant compute.+ FP8_DTYPE = aiter_dtypes.fp8+ BF16 = torch.bfloat16+ _FP8_MAX = float(torch.finfo(FP8_DTYPE).max)+ _FIXED_AMAX = 32.0+ _meta_cache = {}+ _alloc_cache = {}- Uses mla_decode_fwd non-persistent mode which handles all split logic internally.- The a16w16 kernel (mla_dec_stage1_bf16_a16w16_subQ16_mqa16) handles bf16+bf16- for qseqlen=1 non-persistent.- WARNING: page_size=1 EVERYWHERE.- """+ @triton.jit+ def _q_to_fp8_kernel(q_ptr, out_ptr, scale_ptr, amax_ptr,+ FP8_MAX: tl.constexpr, N, BLOCK: tl.constexpr):+ amax = tl.load(amax_ptr)+ amax = tl.where(amax < 1e-12, 1e-12, amax)+ scale = amax / FP8_MAX+ if tl.program_id(0) == 0:+ tl.store(scale_ptr, scale)+ pid = tl.program_id(0)+ offs = pid * BLOCK + tl.arange(0, BLOCK)+ mask = offs < N+ x = tl.load(q_ptr + offs, mask=mask, other=0.0).to(tl.float32)+ x = x / scale+ x = tl.clamp(x, -FP8_MAX, FP8_MAX)+ tl.store(out_ptr + offs, x.to(out_ptr.dtype.element_ty), mask=mask)- import torch- from task import input_t, output_t- from aiter.mla import mla_decode_fwd+ def _build_meta(batch_size, kv_seq_len, q_seq_len, nq, nkv,+ num_kv_splits, page_size, dtype_q, qo_indptr, kv_indptr):+ total_kv = batch_size * kv_seq_len- # 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)+ if page_size == 1:+ num_pages = total_kv+ kv_indptr_pages = kv_indptr+ seq_lens = kv_indptr[1:] - kv_indptr[:-1]+ kv_last_page_len = seq_lens.to(torch.int32)+ else:+ num_pages = total_kv // page_size+ kv_indptr_pages = kv_indptr // page_size+ seq_lens = kv_indptr[1:] - kv_indptr[:-1]+ kv_last_page_len = (seq_lens % page_size).to(torch.int32)+ kv_last_page_len = torch.where(kv_last_page_len == 0, page_size, kv_last_page_len)- _cache = {}+ kv_gran = max(1, 16 // page_size)+ info = get_mla_metadata_info_v1(+ batch_size, q_seq_len, nq, dtype_q, FP8_DTYPE,+ is_sparse=False, fast_mode=False,+ num_kv_splits=num_kv_splits, intra_batch_mode=True,+ )+ work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]+ (wm, wi, wis, ri, rfm, rpm) = work+ get_mla_metadata_v1(+ qo_indptr, kv_indptr_pages, kv_last_page_len,+ nq // nkv, nkv, True,+ wm, wis, wi, ri, rfm, rpm,+ page_size=page_size,+ kv_granularity=kv_gran,+ max_seqlen_qo=q_seq_len,+ uni_seqlen_qo=q_seq_len,+ fast_mode=False,+ max_split_per_batch=num_kv_splits,+ intra_batch_mode=True,+ dtype_q=dtype_q,+ dtype_kv=FP8_DTYPE,+ )++ kv_indices = torch.arange(num_pages, dtype=torch.int32, device="cuda")+ return (wm, wi, wis, ri, rfm, rpm, kv_indices, kv_last_page_len, kv_indptr_pages, page_size)++def custom_kernel(data: input_t) -> output_t:q, kv_data, qo_indptr, kv_indptr, config = data-batch_size = config["batch_size"]+ nq, nkv = config["num_heads"], config["num_kv_heads"]+ dq, dv = config["qk_head_dim"], config["v_head_dim"]+ q_seq_len = config["q_seq_len"]+ sm_scale = config["sm_scale"]kv_seq_len = config["kv_seq_len"]- q_total = q.shape[0]+ total_kv = batch_size * kv_seq_len- # bf16 path — no quantization needed- kv_buffer_bf16 = kv_data["bf16"]- q_bf16 = q.view(-1, NUM_HEADS, QK_HEAD_DIM)+ # Route+ if kv_seq_len <= 1024:+ page_size = 1+ dtype_q = BF16+ use_fp8_q = False+ else:+ page_size = 8+ dtype_q = FP8_DTYPE+ use_fp8_q = True- kv_buffer_4d = kv_buffer_bf16.view(-1, 1, NUM_KV_HEADS, kv_buffer_bf16.shape[-1])+ # Per-shape splits+ if batch_size <= 32 and kv_seq_len <= 1024:+ num_kv_splits = 8+ else:+ num_kv_splits = 16- # Cache kv metadata per shape (constant across calls); allocate output fresh- key = (batch_size, kv_seq_len)- if key not in _cache:- total_kv = batch_size * kv_seq_len- kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")- kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")- _cache[key] = (kv_indices, kv_last_page_len)+ cache_key = (batch_size, kv_seq_len, num_kv_splits, page_size, use_fp8_q)+ if cache_key not in _meta_cache:+ _meta_cache[cache_key] = _build_meta(+ batch_size, kv_seq_len, q_seq_len, nq, nkv,+ num_kv_splits, page_size, dtype_q, qo_indptr, kv_indptr)- kv_indices, kv_last_page_len = _cache[key]- output = torch.empty((q_total, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")+ (wm, wi, wis, ri, rfm, rpm,+ kv_indices, kv_last_page_len, kv_indptr_pages, ps) = _meta_cache[cache_key]- mla_decode_fwd(- q_bf16, kv_buffer_4d, output,- qo_indptr, kv_indptr,- kv_indices, kv_last_page_len,- 1, # max_seqlen_q- page_size=1, nhead_kv=NUM_KV_HEADS, sm_scale=SM_SCALE,- intra_batch_mode=False,- )+ kv_buffer_fp8, kv_scale = kv_data["fp8"]+ kv_buffer_4d = kv_buffer_fp8.view(-1, ps, nkv, kv_buffer_fp8.shape[-1])- return outputNo newline at end of file+ if use_fp8_q:+ alloc_key = ("fp8", q.shape[0], nq, dv, dq)+ if alloc_key not in _alloc_cache:+ _alloc_cache[alloc_key] = (+ torch.empty((q.shape[0], nq, dv), dtype=BF16, device="cuda"),+ torch.full((1,), _FIXED_AMAX, dtype=torch.float32, device="cuda"),+ torch.empty(1, dtype=torch.float32, device="cuda"),+ torch.empty(q.shape[0] * nq * dq, dtype=FP8_DTYPE, device="cuda"),+ )+ o, amax_buf, scale_buf, q_fp8_flat = _alloc_cache[alloc_key]++ N = q.numel()+ BLOCK = 2048+ grid = ((N + BLOCK - 1) // BLOCK,)+ # Skip amax — pre-filled+ _q_to_fp8_kernel[grid](q, q_fp8_flat, scale_buf, amax_buf,+ FP8_MAX=_FP8_MAX, N=N, BLOCK=BLOCK)++ mla_decode_fwd(+ q_fp8_flat.view(q.shape[0], nq, dq), kv_buffer_4d, o,+ qo_indptr, kv_indptr_pages, kv_indices, kv_last_page_len,+ q_seq_len, page_size=ps, nhead_kv=nkv,+ sm_scale=sm_scale, logit_cap=0.0, num_kv_splits=num_kv_splits,+ q_scale=scale_buf, kv_scale=kv_scale,+ intra_batch_mode=True,+ work_meta_data=wm, work_indptr=wi, work_info_set=wis,+ reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm,+ )+ return o+ else:+ alloc_key = ("bf16", q.shape[0], nq, dv)+ if alloc_key not in _alloc_cache:+ _alloc_cache[alloc_key] = torch.empty(+ (q.shape[0], nq, dv), dtype=BF16, device="cuda")+ o = _alloc_cache[alloc_key]++ mla_decode_fwd(+ q, kv_buffer_4d, o,+ qo_indptr, kv_indptr_pages, kv_indices, kv_last_page_len,+ q_seq_len, page_size=ps, nhead_kv=nkv,+ sm_scale=sm_scale, logit_cap=0.0, num_kv_splits=num_kv_splits,+ kv_scale=kv_scale,+ intra_batch_mode=True,+ work_meta_data=wm, work_indptr=wi, work_info_set=wis,+ reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm,+ )+ return o
scrolls · 218 diff lines total
Best evidence level for this revision: reported
JSON