Skip to content
KernelIndex
Search⌘K

submission 738645

gxtzhuxi · 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.

mla_decode.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-738645?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
112.5µs
#479 of 766
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:488f575bcd7230a62336409dca8ed845a93de474990838df74bf037f4b4ac9d8
license declaredunknown
license concludedunknown
authorsgxtzhuxi
imported2026-08-26

Kernel source

mla_decode.py163 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""
MLA Decode — DeepSeek-R1 MLA decode attention on MI355X (CDNA4).

Config: num_heads=16, num_kv_heads=1 (MQA), qk_head_dim=576, v_head_dim=512.
Decode only (q_seq_len=1), variable-length batching via indptr.

== Optimizations ==
  1. Cached metadata — eliminate metadata compute per call (~50 us)
  2. Per-tensor FP8 quantization for Q
  3. Optimal num_kv_splits — fill 256 CUs without over-splitting
  4. Cached static tensors — kv_indices, kv_last_page_len, output buffer
  5. No GPU→CPU sync — avoid .item() calls entirely
"""

import torch
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


# ===== 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)

PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
NUM_CUS = 256
FP8_MAX = torch.finfo(FP8_DTYPE).max
FP8_MIN = torch.finfo(FP8_DTYPE).min

# ===== Global caches =====
_SHAPE_CACHE = {}
_OUT_CACHE = {}


def _optimal_splits(batch_size):
    """Fill 256 CUs with minimal splitting to avoid reduction overhead."""
    base = batch_size * NUM_HEADS
    if base >= NUM_CUS:
        return 1
    splits = (NUM_CUS + base - 1) // base
    return min(16, max(1, splits))


def _build_shape_cache(batch_size, kv_seq_len, device):
    """One-time setup per (batch_size, kv_seq_len) shape."""
    total_kv = batch_size * kv_seq_len
    num_splits = _optimal_splits(batch_size)

    kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)
    kv_last_page_len = torch.full(
        (batch_size,), kv_seq_len, dtype=torch.int32, device=device,
    )

    qo_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device)
    kv_indptr = (
        torch.arange(batch_size + 1, dtype=torch.int32, device=device) * kv_seq_len
    )

    info = get_mla_metadata_info_v1(
        batch_size, 1, NUM_HEADS, FP8_DTYPE, FP8_DTYPE,
        is_sparse=False, fast_mode=False,
        num_kv_splits=num_splits, intra_batch_mode=True,
    )
    work = [torch.empty(s, dtype=t, device=device) for s, t in info]
    (wm, wi, wis, ri, rfm, rpm) = work

    get_mla_metadata_v1(
        qo_indptr, kv_indptr, kv_last_page_len,
        NUM_HEADS // NUM_KV_HEADS,
        NUM_KV_HEADS,
        True,
        wm, wis, wi, ri, rfm, rpm,
        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_splits,
        intra_batch_mode=True,
        dtype_q=FP8_DTYPE,
        dtype_kv=FP8_DTYPE,
    )

    return {
        "kv_indices": kv_indices,
        "kv_last_page_len": kv_last_page_len,
        "num_splits": num_splits,
        "work_meta_data": wm,
        "work_indptr": wi,
        "work_info_set": wis,
        "reduce_indptr": ri,
        "reduce_final_map": rfm,
        "reduce_partial_map": rpm,
    }


def _get_output(batch_size, device):
    if batch_size not in _OUT_CACHE:
        _OUT_CACHE[batch_size] = torch.empty(
            batch_size, NUM_HEADS, V_HEAD_DIM,
            dtype=torch.bfloat16, device=device,
        )
    return _OUT_CACHE[batch_size]


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data

    batch_size = config["batch_size"]
    kv_seq_len = config["kv_seq_len"]
    device = q.device

    kv_fp8, kv_scale = kv_data["fp8"]

    # FP8 quantize Q (per-tensor)
    amax = q.abs().amax().clamp(min=1e-12)
    q_scale = amax / FP8_MAX
    q_fp8 = (q / q_scale).clamp(min=FP8_MIN, max=FP8_MAX).to(FP8_DTYPE)

    key = (batch_size, kv_seq_len)
    if key not in _SHAPE_CACHE:
        _SHAPE_CACHE[key] = _build_shape_cache(batch_size, kv_seq_len, device)
    c = _SHAPE_CACHE[key]

    kv_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, NUM_KV_HEADS, -1)
    o = _get_output(batch_size, device)

    mla_decode_fwd(
        q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM),
        kv_4d,
        o,
        qo_indptr, kv_indptr,
        c["kv_indices"],
        c["kv_last_page_len"],
        1,
        page_size=PAGE_SIZE,
        nhead_kv=NUM_KV_HEADS,
        sm_scale=SM_SCALE,
        logit_cap=0.0,
        num_kv_splits=c["num_splits"],
        q_scale=q_scale.to(torch.float32).reshape(1),
        kv_scale=kv_scale,
        intra_batch_mode=True,
        work_meta_data=c["work_meta_data"],
        work_indptr=c["work_indptr"],
        work_info_set=c["work_info_set"],
        reduce_indptr=c["reduce_indptr"],
        reduce_final_map=c["reduce_final_map"],
        reduce_partial_map=c["reduce_partial_map"],
    )
    return o
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