Skip to content
KernelIndex
Search⌘K

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
AMD Instinct MI355X
41.6µs
#123 of 766
2026-04-04

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 output
No 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