Skip to content
KernelIndex
Search⌘K

submission 719655

RexHuang · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-719655?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
102.2µs
#463 of 766
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:89e7afebb4bc48a7d86edab0cad3ec486d649d228d0c497a1809b639e47ad10b
license declaredunknown
license concludedunknown
authorsRexHuang
imported2026-08-26

Techniques

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

persistent-kernel2. Persistent mode — metadata precomputed once, reducing kernel launch overhead

Kernel source

submission.py238 lines
"""
MLA decode kernel with aiter acceleration for AMD GPUs.

Performance strategy on AMD (aiter available):
  1. FP8 a8w8 kernel — Q and KV quantized to FP8 (~2-3x faster than bf16 on MI355X)
  2. Persistent mode — metadata precomputed once, reducing kernel launch overhead
  3. intra_batch_mode + num_kv_splits=32 — optimal KV chunking for cache utilisation

Fallback (no aiter / non-AMD):
  PyTorch SDPA FlashAttention with native MQA broadcast (HK=1 → H=16).
"""

import torch
import torch.nn.functional as F

# ---------------------------------------------------------------------------
# Aiter availability — graceful fallback when not installed
# ---------------------------------------------------------------------------
_HAS_AITER = False
try:
    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
    _HAS_AITER = True
except ImportError:
    pass

# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
PAGE_SIZE = 1
NUM_KV_SPLITS = 32

if _HAS_AITER:
    _FP8_DTYPE = aiter_dtypes.fp8


# ---------------------------------------------------------------------------
# FP8 quantization (dynamic per-tensor, sglang-style)
# ---------------------------------------------------------------------------

def _quantize_fp8(tensor: torch.Tensor) -> tuple:
    """Dynamic per-tensor FP8 quantization."""
    finfo = torch.finfo(_FP8_DTYPE)
    amax = tensor.abs().amax().clamp(min=1e-12)
    scale = amax / finfo.max
    fp8 = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(_FP8_DTYPE)
    return fp8, scale.to(torch.float32).reshape(1)


# ---------------------------------------------------------------------------
# Aiter path: persistent-mode FP8 a8w8 MLA decode
# ---------------------------------------------------------------------------

def _build_persistent_metadata(
    batch_size, max_q_len, nhead, nhead_kv,
    q_dtype, kv_dtype, qo_indptr, kv_indptr, kv_last_page_len,
):
    """Allocate and populate work buffers for persistent mla_decode_fwd."""
    info = get_mla_metadata_info_v1(
        batch_size, max_q_len, nhead, q_dtype, kv_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]
    (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,
        nhead // nhead_kv, nhead_kv, True,   # is_causal
        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=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 {
        "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,
    }


def _aiter_decode(q, kv_data, qo_indptr, kv_indptr, config):
    """MLA decode via aiter persistent-mode FP8 a8w8 kernel."""
    B = config["batch_size"]
    H = config["num_heads"]       # 16
    HK = config["num_kv_heads"]   # 1
    D = config["qk_head_dim"]     # 576
    Dv = config["v_head_dim"]     # 512
    scale = config["sm_scale"]
    max_q = config["q_seq_len"]

    # Quantize Q to FP8 on-the-fly (negligible cost vs kernel speedup)
    q_fp8, q_scale = _quantize_fp8(q)

    # Use pre-quantised FP8 KV from input; fall back to bf16 + on-the-fly quant
    if "fp8" in kv_data:
        kv_fp8, kv_scale = kv_data["fp8"]
    else:
        kv_fp8, kv_scale = _quantize_fp8(kv_data["bf16"])

    total_kv = kv_fp8.shape[0]
    kv_4d = kv_fp8.view(total_kv, PAGE_SIZE, HK, D)
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
    kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)

    # Build persistent-mode metadata
    meta = _build_persistent_metadata(
        B, max_q, H, HK,
        q_fp8.dtype, kv_fp8.dtype,
        qo_indptr, kv_indptr, kv_last_page_len,
    )

    # Output buffer
    o = torch.empty((q.shape[0], H, Dv), dtype=torch.bfloat16, device="cuda")

    mla_decode_fwd(
        q_fp8.view(-1, H, D),
        kv_4d,
        o,
        qo_indptr,
        kv_indptr,
        kv_indices,
        kv_last_page_len,
        max_q,
        page_size=PAGE_SIZE,
        nhead_kv=HK,
        sm_scale=scale,
        logit_cap=0.0,
        num_kv_splits=NUM_KV_SPLITS,
        q_scale=q_scale,
        kv_scale=kv_scale,
        intra_batch_mode=True,
        **meta,
    )
    return o


# ---------------------------------------------------------------------------
# SDPA fallback (no aiter / non-AMD)
# ---------------------------------------------------------------------------

def _sdpa_decode(q, kv_data, qo_indptr, kv_indptr, config):
    """MLA decode via PyTorch SDPA FlashAttention with native MQA broadcast."""
    B = config["batch_size"]
    H = config["num_heads"]       # 16
    HK = config["num_kv_heads"]   # 1
    D = config["qk_head_dim"]     # 576
    Dv = config["v_head_dim"]     # 512
    scale = config["sm_scale"]

    kv_buf = kv_data["bf16"]                         # (total_kv, 1, 576)
    K = kv_buf                                        # (total_kv, 1, 576)
    V = kv_buf[..., :Dv].contiguous()                # (total_kv, 1, 512)

    kv_lens = kv_indptr[1:] - kv_indptr[:-1]
    q_lens = qo_indptr[1:] - qo_indptr[:-1]

    all_equal_kv = (kv_lens == kv_lens[0]).all().item()
    all_equal_q = (q_lens == q_lens[0]).all().item()

    if all_equal_kv and all_equal_q:
        return _sdpa_equal(q, K, V, B, H, HK, D, Dv, scale,
                           q_lens[0].item(), kv_lens[0].item())
    return _sdpa_variable(q, K, V, qo_indptr, kv_indptr,
                          B, H, HK, D, Dv, scale, q_lens, kv_lens)


def _sdpa_equal(q, K, V, B, H, HK, D, Dv, scale, Lq, Lk):
    """Fast path: all sequences have equal length (common decode case)."""
    Qs = q.view(B, Lq, H, D).permute(0, 2, 1, 3).contiguous()    # (B, H, Lq, D)
    Ks = K.view(B, Lk, HK, D).permute(0, 2, 1, 3).contiguous()   # (B, HK, Lk, D)
    Vs = V.view(B, Lk, HK, Dv).permute(0, 2, 1, 3).contiguous()  # (B, HK, Lk, Dv)

    # SDPA handles MQA broadcast internally (HK=1 → H=16)
    O = F.scaled_dot_product_attention(Qs, Ks, Vs, scale=scale)    # (B, H, Lq, Dv)
    return O.permute(0, 2, 1, 3).reshape(-1, H, Dv).contiguous()


def _sdpa_variable(q, K, V, qo_indptr, kv_indptr,
                   B, H, HK, D, Dv, scale, q_lens, kv_lens):
    """Fallback: variable-length sequences via padding + attention mask."""
    max_kv = kv_lens.max().item()
    max_q = q_lens.max().item()

    Q_pad = q.new_zeros(B, max_q, H, D)
    K_pad = K.new_zeros(B, max_kv, HK, D)
    V_pad = K.new_zeros(B, max_kv, HK, Dv)
    kv_mask = torch.ones(B, max_kv, dtype=torch.bool, device=q.device)

    for i in range(B):
        qs, qe = int(qo_indptr[i]), int(qo_indptr[i + 1])
        ks, ke = int(kv_indptr[i]), int(kv_indptr[i + 1])
        ql, kl = qe - qs, ke - ks
        Q_pad[i, :ql] = q[qs:qe].view(ql, H, D)
        K_pad[i, :kl] = K[ks:ke].view(kl, HK, D)
        V_pad[i, :kl] = V[ks:ke].view(kl, HK, Dv)
        kv_mask[i, kl:] = False

    Qs = Q_pad.permute(0, 2, 1, 3)     # (B, H, max_q, D)
    Ks = K_pad.permute(0, 2, 1, 3)     # (B, HK, max_kv, D)
    Vs = V_pad.permute(0, 2, 1, 3)     # (B, HK, max_kv, Dv)

    attn_mask = kv_mask.unsqueeze(1).unsqueeze(2)   # (B, 1, 1, max_kv)

    O = F.scaled_dot_product_attention(Qs, Ks, Vs, attn_mask=attn_mask, scale=scale)

    results = []
    for i in range(B):
        ql = q_lens[i].item()
        results.append(O[i, :, :ql].permute(1, 0, 2))   # (ql, H, Dv)
    return torch.cat(results, dim=0).contiguous()


# ---------------------------------------------------------------------------
# Entry point
# ---------------------------------------------------------------------------

def custom_kernel(data):
    """MLA decode — aiter FP8 persistent kernel (AMD) or SDPA fallback."""
    q, kv_data, qo_indptr, kv_indptr, config = data

    if _HAS_AITER and torch.cuda.is_available():
        return _aiter_decode(q, kv_data, qo_indptr, kv_indptr, config)
    return _sdpa_decode(q, kv_data, qo_indptr, kv_indptr, config)
scrolls · 238 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