Skip to content
KernelIndex
Search⌘K

submission 653025

yzhou442 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 163 lines, June 9 Researcher Reciprocity License v1.0.

submission_current.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-653025?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
58.2µs
#228 of 766
2026-03-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:fcda28a3dd16a1600250eb4cdcaca8cd5c065077a027db22b9605a127b90cd36
license declaredunknown
license concludedunknown
authorsyzhou442
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

persistent-kernelNone, # num_kv_splits_indptr (not used in persistent mode)

Kernel source

submission_current.py163 lines
"""MLA decode v82 — inference_mode + micro-optimizations.

Based on v77 (bypass mla_decode_fwd dispatch, pre-allocate intermediates).
Additional optimizations:
  1. @torch.inference_mode() on custom_kernel — disables autograd tracking
     for all tensor ops, saving ~0.5-1us overhead per call.
  2. Pre-compute q_fp8 3D view in _ensure_cached() — saves 1 view() call
     per invocation since views on the same storage are cheap to cache.
  3. Unpack frequently-used cache values to local variables — dict lookups
     are ~50ns each in CPython; locals are a single LOAD_FAST opcode.
"""
import math, 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

SM_SCALE = 1.0 / math.sqrt(576)
FP8_DTYPE = aiter_dtypes.fp8
PAGE_SIZE = 1
NUM_KV_SPLITS = 32
H = 16
NKV = 1
DK = 576
DV = 512

_cache = {}


def _ensure_cached(bs, kvl, device):
    key = (bs, kvl)
    if key in _cache:
        return _cache[key]

    total_kv = bs * kvl
    qo_indptr = torch.arange(0, bs + 1, dtype=torch.int32, device=device)
    kv_indptr = torch.arange(0, bs + 1, dtype=torch.int32, device=device) * kvl
    kv_last_page_len = torch.full((bs,), kvl, dtype=torch.int32, device=device)
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)

    info = get_mla_metadata_info_v1(
        bs, 1, H, FP8_DTYPE, 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=device) for s, t in info]
    (work_metadata, work_indptr, work_info_set,
     reduce_indptr, reduce_final_map, reduce_partial_map) = work

    get_mla_metadata_v1(
        qo_indptr, kv_indptr, kv_last_page_len,
        H // NKV, NKV, True,
        work_metadata, work_info_set, work_indptr,
        reduce_indptr, reduce_final_map, reduce_partial_map,
        page_size=PAGE_SIZE,
        kv_granularity=max(PAGE_SIZE, 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,
    )

    out = torch.empty((bs, H, DV), dtype=torch.bfloat16, device=device)
    # Pre-allocate Q FP8 buffer
    q_fp8 = torch.empty((bs * H, DK), dtype=FP8_DTYPE, device=device)
    # Pre-compute 3D view (saves 1 view() call per invocation)
    q_fp8_3d = q_fp8.view(bs, H, DK)
    # q_scale = 1.0 (direct BF16->FP8 cast, no dynamic scaling)
    q_scale = torch.ones(1, dtype=torch.float32, device=device)

    # Pre-allocate intermediate buffers for stage1 ASM kernel
    # These are normally allocated every call inside mla_decode_fwd()
    rp_size = reduce_partial_map.size(0)
    logits = torch.empty((rp_size, 1, H, DV), dtype=torch.float32, device=device)
    attn_lse = torch.empty((rp_size, 1, H, 1), dtype=torch.float32, device=device)

    _cache[key] = {
        "qo_indptr": qo_indptr,
        "kv_indptr": kv_indptr,
        "kv_indices": kv_indices,
        "kv_last_page_len": kv_last_page_len,
        "work_meta_data": work_metadata,
        "work_indptr": work_indptr,
        "work_info_set": work_info_set,
        "reduce_indptr": reduce_indptr,
        "reduce_final_map": reduce_final_map,
        "reduce_partial_map": reduce_partial_map,
        "output": out,
        "q_fp8": q_fp8,
        "q_fp8_3d": q_fp8_3d,
        "q_scale": q_scale,
        "logits": logits,
        "attn_lse": attn_lse,
    }
    return _cache[key]


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr_in, kv_indptr_in, cfg = data
    bs = cfg["batch_size"]
    kvl = cfg["kv_seq_len"]

    c = _ensure_cached(bs, kvl, q.device)

    # Unpack frequently-used cache values to locals (avoid repeated dict lookups)
    q_fp8 = c["q_fp8"]
    q_fp8_3d = c["q_fp8_3d"]
    logits = c["logits"]
    attn_lse = c["attn_lse"]
    output = c["output"]
    q_scale = c["q_scale"]

    # Direct BF16->FP8 cast (1 HIP launch vs 3 for dynamic_per_tensor_quant)
    q_2d = q.view(-1, DK)
    q_fp8.copy_(q_2d)

    # FP8 KV
    kv_fp8, kv_scale = kv_data["fp8"]
    tkv = bs * kvl
    kv_4d = kv_fp8[:tkv].view(tkv, PAGE_SIZE, NKV, DK)

    # Direct ASM dispatch — bypass mla_decode_fwd() Python wrapper
    # Saves: 2 tensor allocations (logits, attn_lse) + Python dispatch overhead
    aiter.mla_decode_stage1_asm_fwd(
        q_fp8_3d,
        kv_4d,
        c["qo_indptr"],
        c["kv_indptr"],
        c["kv_indices"],
        c["kv_last_page_len"],
        None,               # num_kv_splits_indptr (not used in persistent mode)
        c["work_meta_data"],
        c["work_indptr"],
        c["work_info_set"],
        1,                   # max_seqlen_q
        PAGE_SIZE,
        NKV,
        SM_SCALE,
        logits,
        attn_lse,
        output,
        q_scale,
        kv_scale,
    )

    aiter.mla_reduce_v1(
        logits,
        attn_lse,
        c["reduce_indptr"],
        c["reduce_final_map"],
        c["reduce_partial_map"],
        1,                   # max_seqlen_q
        output,
        None,                # final_lse
    )

    return output
scrolls · 163 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