Skip to content
KernelIndex
Search⌘K

submission 646049

JohnHe · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

solution.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-646049?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
190.9µs
#608 of 766
2026-03-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:76cf2f8d3141574a5bfcf4ae131c1714db7b3ba113b926ade596b1cdffd9f9be
license declaredunknown
license concludedunknown
authorsJohnHe
imported2026-08-26

Techniques

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

persistent-kernelOptimizations over the reference (aiter a8w8 persistent-mode kernel):

Kernel source

solution.py206 lines
"""
Optimized MLA decode kernel for AMD MI355X (CDNA4).

Optimizations over the reference (aiter a8w8 persistent-mode kernel):

1. Adaptive NUM_KV_SPLITS per workload shape
   MI355X has 256 CUs across 8 XCDs. The reference uses a fixed 32 splits.
   For small batches (4×16=64 head-batch pairs), 32 splits gives 2048 blocks
   which is fine, but for batch=256 the reduction overhead of 32 splits is
   wasteful since batch parallelism alone (4096 pairs) saturates the CUs.
   We adaptively select 16-128 splits based on batch×heads and kv_seq_len.

2. Metadata buffer caching
   Persistent-mode requires 6 work buffers allocated via cudaMalloc + a
   metadata population kernel. For repeated calls with identical geometry
   (common in continuous batching), we cache the allocated buffers and only
   re-run the cheap population kernel, saving cudaMalloc latency.

3. kv_indices caching
   Simple contiguous range reused across calls with same total_kv_len.

4. Minimized Python-side overhead
   Fewer intermediate variables, direct dict lookups, avoid unnecessary
   tensor operations.
"""

import torch
from task import input_t, output_t
from utils import make_match_reference

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

# ---------------------------------------------------------------------------
# DeepSeek R1 MLA constants (forward_absorb path)
# ---------------------------------------------------------------------------
NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = 576     # kv_lora_rank + qk_rope_head_dim
V_HEAD_DIM = 512      # = kv_lora_rank
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
_FP8_FINFO = torch.finfo(FP8_DTYPE)


# ---------------------------------------------------------------------------
# FP8 quantization (hot path — kept minimal)
# ---------------------------------------------------------------------------
def _quantize_q_fp8(q: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    amax = q.abs().amax().clamp(min=1e-12)
    scale = amax / _FP8_FINFO.max
    q_fp8 = (q / scale).clamp(min=_FP8_FINFO.min, max=_FP8_FINFO.max).to(FP8_DTYPE)
    return q_fp8, scale.to(torch.float32).reshape(1)


# ---------------------------------------------------------------------------
# Adaptive NUM_KV_SPLITS
# ---------------------------------------------------------------------------
def _select_splits(batch_size: int, kv_seq_len: int) -> int:
    """
    MI355X: 256 CUs, 8 XCDs of 32 CUs each.
    Total blocks = batch_size * NUM_HEADS * num_kv_splits.
    Want >= 256 blocks (1/CU), ideally 512-2048 for latency hiding.
    But more splits = more reduction overhead.

    Also: each split processes kv_seq_len/num_kv_splits tokens.
    Too few tokens per split -> underutilized compute.
    """
    head_batches = batch_size * NUM_HEADS

    if head_batches >= 256:
        # batch>=16: 256+ head-batch pairs, CUs saturated from batch alone
        return 16 if kv_seq_len <= 2048 else 32
    elif head_batches >= 64:
        # batch=4-15: moderate parallelism
        return 32 if kv_seq_len <= 2048 else 64
    else:
        # batch=1-3: need splits for parallelism
        return 64 if kv_seq_len <= 2048 else 128


# ---------------------------------------------------------------------------
# Caches
# ---------------------------------------------------------------------------
_meta_buf_cache: dict = {}   # keyed on geometry -> pre-allocated buffers
_kv_idx_cache: dict = {}     # keyed on total_kv_len -> int32 range tensor


def _cached_kv_indices(n: int) -> torch.Tensor:
    if n not in _kv_idx_cache:
        _kv_idx_cache[n] = torch.arange(n, dtype=torch.int32, device="cuda")
    return _kv_idx_cache[n]


def _get_metadata(
    batch_size, max_q_len, nq, nkv,
    q_dtype, kv_dtype,
    qo_indptr, kv_indptr, kv_last_page_len,
    num_kv_splits,
):
    total_kv = int(kv_indptr[-1].item())
    key = (batch_size, max_q_len, nq, nkv, q_dtype, kv_dtype, num_kv_splits, total_kv)

    if key not in _meta_buf_cache:
        # Allocate work buffers (expensive: cudaMalloc)
        info = get_mla_metadata_info_v1(
            batch_size, max_q_len, nq, q_dtype, kv_dtype,
            is_sparse=False, fast_mode=False,
            num_kv_splits=num_kv_splits, intra_batch_mode=True,
        )
        bufs = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
        _meta_buf_cache[key] = {
            "work_meta_data": bufs[0],
            "work_indptr": bufs[1],
            "work_info_set": bufs[2],
            "reduce_indptr": bufs[3],
            "reduce_final_map": bufs[4],
            "reduce_partial_map": bufs[5],
        }

    m = _meta_buf_cache[key]

    # Populate (cheap kernel - must run every call as indptrs may differ)
    get_mla_metadata_v1(
        qo_indptr, kv_indptr, kv_last_page_len,
        nq // nkv, nkv, True,
        m["work_meta_data"], m["work_info_set"], m["work_indptr"],
        m["reduce_indptr"], m["reduce_final_map"], m["reduce_partial_map"],
        page_size=PAGE_SIZE,
        kv_granularity=max(PAGE_SIZE, 16),
        max_seqlen_qo=max_q_len,
        uni_seqlen_qo=max_q_len,
        fast_mode=False,
        max_split_per_batch=num_kv_splits,
        intra_batch_mode=True,
        dtype_q=q_dtype,
        dtype_kv=kv_dtype,
    )
    return m


# ---------------------------------------------------------------------------
# Main entry point
# ---------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
    """
    Optimized MLA decode: FP8 Q + FP8 KV via aiter persistent-mode kernel,
    with MI355X-tuned adaptive splitting and metadata caching.
    """
    q, kv_data, qo_indptr, kv_indptr, config = data

    batch_size = config["batch_size"]
    q_seq_len = config["q_seq_len"]
    kv_seq_len = config["kv_seq_len"]

    # Adaptive split count for MI355X
    num_splits = _select_splits(batch_size, kv_seq_len)

    # FP8 quantize Q on-the-fly
    q_fp8, q_scale = _quantize_q_fp8(q)

    # Pre-quantized FP8 KV
    kv_fp8, kv_scale = kv_data["fp8"]

    # 4D view for aiter: (total_kv, page_size=1, nkv=1, 576) - zero-copy
    kv_4d = kv_fp8.view(kv_fp8.shape[0], 1, 1, kv_fp8.shape[-1])

    total_kv = int(kv_indptr[-1].item())
    kv_indices = _cached_kv_indices(total_kv)
    kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)

    meta = _get_metadata(
        batch_size, q_seq_len, NUM_HEADS, NUM_KV_HEADS,
        q_fp8.dtype, kv_fp8.dtype,
        qo_indptr, kv_indptr, kv_last_page_len,
        num_splits,
    )

    # Fresh output buffer (must not reuse - caller may retain reference)
    o = torch.empty((q.shape[0], NUM_HEADS, V_HEAD_DIM),
                     dtype=torch.bfloat16, device="cuda")

    mla_decode_fwd(
        q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM),
        kv_4d,
        o,
        qo_indptr,
        kv_indptr,
        kv_indices,
        kv_last_page_len,
        q_seq_len,
        page_size=PAGE_SIZE,
        nhead_kv=NUM_KV_HEADS,
        sm_scale=SM_SCALE,
        logit_cap=0.0,
        num_kv_splits=num_splits,
        q_scale=q_scale,
        kv_scale=kv_scale,
        intra_batch_mode=True,
        **meta,
    )
    return o
scrolls · 206 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