Skip to content
KernelIndex
Search⌘K

submission 586972

ron1tk · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-586972?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
269.2µs
#702 of 766
2026-03-18

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f6be6ec1414e74942a7d8a5d107218a59df25823c4034fbeb3f0b9c4d7f7c62e
license declaredunknown
license concludedunknown
authorsron1tk
imported2026-08-26

Kernel source

submission.py383 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""
MLA decode submission for amd-mixed-mla.

Primary strategy:
- Use AMD AITER's public mla_decode_fwd directly on dense bf16 KV by wrapping the
  dense cache as page_size=1 paged KV.
- Fall back to the best-performing hybrid bf16 implementation if AITER is not
  available or the call fails in the harness.

Why this is the strongest public-path attempt:
- AMD's official docs expose mla_decode_fwd as the optimized MLA decode API.
- vLLM's ROCm backend notes that the assembly mla_decode_fwd kernel is where most
  decode performance gains come from.
"""

from __future__ import annotations

import torch
import torch.nn.functional as F
from task import input_t, output_t

from aiter import dtypes as aiter_dtypes

FP8_DTYPE = aiter_dtypes.fp8

# Keep bf16 as the main data path; AITER handles the actual optimized decode.
QKV_DTYPE = "bf16"
FASTPATH_MAX_KV = 1024

try:
    from aiter.mla import mla_decode_fwd as aiter_mla_decode_fwd
    AITER_MLA_AVAILABLE = True
except Exception:
    aiter_mla_decode_fwd = None
    AITER_MLA_AVAILABLE = False


# -----------------------------------------------------------------------------
# Small caches for benchmark repeats
# -----------------------------------------------------------------------------

_KV_INDICES_CACHE: dict[tuple[int, int], torch.Tensor] = {}
_KV_LAST_PAGE_LENS_CACHE: dict[tuple[int, tuple[int, ...]], torch.Tensor] = {}


# -----------------------------------------------------------------------------
# Dispatcher
# -----------------------------------------------------------------------------

def custom_kernel(data: input_t) -> output_t:
    # First try the vendor-optimized path.
    if QKV_DTYPE == "bf16":
        out = _try_aiter_bf16(data)
        if out is not None:
            return out
        return custom_kernel_bf16_hybrid(data)

    if QKV_DTYPE == "fp8":
        return custom_kernel_fp8(data)

    raise ValueError(f"Invalid QKV_DTYPE: {QKV_DTYPE}")


# -----------------------------------------------------------------------------
# Helpers
# -----------------------------------------------------------------------------

def quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    finfo = torch.finfo(FP8_DTYPE)
    amax = tensor.abs().amax().clamp(min=1e-12)
    scale = amax / finfo.max
    fp8_tensor = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
    return fp8_tensor, scale.to(torch.float32).reshape(1)


def as_scalar_scale(x, device: torch.device) -> torch.Tensor:
    if torch.is_tensor(x):
        return x.to(device=device, dtype=torch.float32).reshape(1)
    return torch.tensor([float(x)], device=device, dtype=torch.float32)


def _device_key(device: torch.device) -> int:
    return -1 if device.index is None else int(device.index)


def _get_dense_kv_indices(total_kv: int, device: torch.device) -> torch.Tensor:
    key = (_device_key(device), total_kv)
    t = _KV_INDICES_CACHE.get(key)
    if t is None:
        t = torch.arange(total_kv, device=device, dtype=torch.int32)
        _KV_INDICES_CACHE[key] = t
    return t


def _get_kv_last_page_lens(kv_indptr: torch.Tensor) -> torch.Tensor:
    kv_lens = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
    # page_size = 1, so any non-empty sequence has last_page_len = 1
    vals = torch.where(kv_lens > 0, torch.ones_like(kv_lens), torch.zeros_like(kv_lens))
    key = (_device_key(kv_indptr.device), tuple(int(x) for x in vals.tolist()))
    t = _KV_LAST_PAGE_LENS_CACHE.get(key)
    if t is None:
        t = vals.contiguous()
        _KV_LAST_PAGE_LENS_CACHE[key] = t
    return t


def _check_uniform_decode_layout(
    q: torch.Tensor,
    kv_buffer_bf16: torch.Tensor,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
) -> tuple[bool, int, int]:
    batch_size = qo_indptr.numel() - 1
    if batch_size <= 0:
        return False, 0, 0

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

    if not bool(torch.all(qo_lens == 1).item()):
        return False, 0, 0
    if not bool(torch.all(kv_lens == kv_lens[0]).item()):
        return False, 0, 0

    kv_len = int(kv_lens[0].item())
    total_q = q.shape[0]
    total_kv = kv_buffer_bf16.shape[0]

    if total_q != batch_size:
        return False, 0, 0
    if total_kv != batch_size * kv_len:
        return False, 0, 0

    return True, batch_size, kv_len


# -----------------------------------------------------------------------------
# AITER primary path
# -----------------------------------------------------------------------------

def _try_aiter_bf16(data: input_t) -> torch.Tensor | None:
    if not AITER_MLA_AVAILABLE:
        return None

    q, kv_data, qo_indptr, kv_indptr, config = data

    # Public AITER decode currently targets the DeepSeek-style 16-head case.
    if int(config["num_heads"]) not in (16, 128):
        return None

    kv_buffer_bf16 = kv_data["bf16"]
    device = q.device

    try:
        total_q = q.shape[0]
        qk_head_dim = q.shape[-1]
        v_head_dim = int(config["v_head_dim"])
        sm_scale = float(config["sm_scale"])

        # Dense KV -> paged KV with page_size=1
        # AITER expects [num_pages, page_size, num_heads_kv, qk_head_dim]
        kv_paged = kv_buffer_bf16.contiguous().view(-1, 1, 1, qk_head_dim)

        # For dense storage, the page indices are just 0..total_kv-1
        total_kv = int(kv_indptr[-1].item())
        kv_indices = _get_dense_kv_indices(total_kv, device)
        kv_last_page_lens = _get_kv_last_page_lens(kv_indptr)

        max_seqlen_q = int((qo_indptr[1:] - qo_indptr[:-1]).max().item())
        if max_seqlen_q <= 0:
            max_seqlen_q = 1

        o = torch.empty((total_q, q.shape[1], v_head_dim), device=device, dtype=torch.bfloat16)

        ret = aiter_mla_decode_fwd(
            q.contiguous(),
            kv_paged,
            o,
            qo_indptr.to(torch.int32).contiguous(),
            kv_indptr.to(torch.int32).contiguous(),
            kv_indices,
            kv_last_page_lens,
            max_seqlen_q,
            sm_scale=sm_scale,
        )

        if torch.is_tensor(ret):
            return ret.to(torch.bfloat16)
        return o
    except Exception:
        return None


# -----------------------------------------------------------------------------
# Best fallback: short-KV batched fast path + decode-specialized long-KV fallback
# -----------------------------------------------------------------------------

def _try_uniform_bf16_fastpath(
    q: torch.Tensor,
    kv_buffer_bf16: torch.Tensor,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    v_head_dim: int,
    sm_scale: float,
) -> torch.Tensor | None:
    ok, batch_size, kv_len = _check_uniform_decode_layout(
        q=q,
        kv_buffer_bf16=kv_buffer_bf16,
        qo_indptr=qo_indptr,
        kv_indptr=kv_indptr,
    )
    if not ok or kv_len > FASTPATH_MAX_KV:
        return None

    q_bhd = q.contiguous()  # [B, H, D]
    kv_bld = kv_buffer_bf16[:, 0, :].contiguous().view(batch_size, kv_len, q.shape[-1])

    k_bld = kv_bld
    v_blv = kv_bld[:, :, :v_head_dim]

    scores = torch.matmul(
        (q_bhd * sm_scale).unsqueeze(2),              # [B, H, 1, D]
        k_bld.unsqueeze(1).transpose(-1, -2),         # [B, 1, D, L]
    ).squeeze(2)                                      # [B, H, L]

    probs = F.softmax(scores.float(), dim=-1).to(q_bhd.dtype)

    out = torch.matmul(
        probs.unsqueeze(2),                           # [B, H, 1, L]
        v_blv.unsqueeze(1),                           # [B, 1, L, V]
    ).squeeze(2)                                      # [B, H, V]

    return out.to(torch.bfloat16).contiguous()


def _segmented_bf16_decode_fallback(
    q: torch.Tensor,
    kv_buffer_bf16: torch.Tensor,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    v_head_dim: int,
    sm_scale: float,
) -> torch.Tensor:
    batch_size = qo_indptr.shape[0] - 1
    out_list = []

    for i in range(batch_size):
        q_s = int(qo_indptr[i].item())
        q_e = int(qo_indptr[i + 1].item())
        kv_s = int(kv_indptr[i].item())
        kv_e = int(kv_indptr[i + 1].item())

        seq_q = q_e - q_s
        if seq_q == 0:
            continue

        if seq_q == 1:
            qh = q[q_s, :, :]                         # [H, D] bf16
            kvc = kv_buffer_bf16[kv_s:kv_e, 0]       # [L, D] bf16

            scores = torch.matmul(qh * sm_scale, kvc.transpose(0, 1))  # [H, L] bf16
            probs = F.softmax(scores.float(), dim=-1).to(qh.dtype)

            vi = kvc[:, :v_head_dim]                 # [L, V] bf16
            out_hv = torch.matmul(probs, vi)         # [H, V] bf16

            out_list.append(out_hv.unsqueeze(0).to(torch.bfloat16))
            continue

        qi = q[q_s:q_e]                              # [seq_q, H, D]
        kvc = kv_buffer_bf16[kv_s:kv_e, 0]          # [L, D]

        qi_t = qi.permute(1, 0, 2).contiguous()      # [H, seq_q, D]
        ki_t = kvc.transpose(0, 1).contiguous()      # [D, L]

        scores = torch.matmul(qi_t * sm_scale, ki_t)
        probs = F.softmax(scores.float(), dim=-1).to(qi.dtype)

        vi = kvc[:, :v_head_dim]
        oi = torch.matmul(probs, vi)
        oi = oi.permute(1, 0, 2).contiguous()

        out_list.append(oi.to(torch.bfloat16))

    if not out_list:
        return torch.empty((0, q.shape[1], v_head_dim), device=q.device, dtype=torch.bfloat16)

    return torch.cat(out_list, dim=0)


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

    v_head_dim = int(config["v_head_dim"])
    sm_scale = float(config["sm_scale"])
    kv_buffer_bf16 = kv_data["bf16"]

    fast = _try_uniform_bf16_fastpath(
        q=q,
        kv_buffer_bf16=kv_buffer_bf16,
        qo_indptr=qo_indptr,
        kv_indptr=kv_indptr,
        v_head_dim=v_head_dim,
        sm_scale=sm_scale,
    )
    if fast is not None:
        return fast

    return _segmented_bf16_decode_fallback(
        q=q,
        kv_buffer_bf16=kv_buffer_bf16,
        qo_indptr=qo_indptr,
        kv_indptr=kv_indptr,
        v_head_dim=v_head_dim,
        sm_scale=sm_scale,
    )


# -----------------------------------------------------------------------------
# FP8 fallback path (kept for experimentation)
# -----------------------------------------------------------------------------

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

    num_heads = int(config["num_heads"])
    v_head_dim = int(config["v_head_dim"])
    qk_head_dim = int(config["qk_head_dim"])
    sm_scale = float(config["sm_scale"])

    kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
    kv_scale_fp8 = as_scalar_scale(kv_scale_fp8, q.device)

    kv_fp8_2d = kv_buffer_fp8.reshape(-1, qk_head_dim)
    kv_bf16 = kv_data["bf16"]

    q_fp8, q_scale = quantize_fp8(q)
    q_scale = as_scalar_scale(q_scale, q.device)

    batch_size = qo_indptr.shape[0] - 1
    out_list = []

    for i in range(batch_size):
        q_s = int(qo_indptr[i].item())
        q_e = int(qo_indptr[i + 1].item())
        kv_s = int(kv_indptr[i].item())
        kv_e = int(kv_indptr[i + 1].item())

        seq_q = q_e - q_s
        seq_kv = kv_e - kv_s

        qi_fp8 = q_fp8[q_s:q_e].reshape(seq_q * num_heads, qk_head_dim).contiguous()
        ki_fp8 = kv_fp8_2d[kv_s:kv_e].contiguous()

        try:
            raw_scores = torch._scaled_mm(
                qi_fp8,
                ki_fp8,
                scale_a=q_scale,
                scale_b=kv_scale_fp8,
                out_dtype=torch.float32,
            )
        except RuntimeError as e:
            msg = str(e)
            if "cuBLASLt" not in msg and "_scaled_mm" not in msg:
                raise
            raw_scores = (qi_fp8.float() * q_scale).matmul(
                (ki_fp8.float() * kv_scale_fp8).transpose(0, 1)
            )

        scores = raw_scores.view(seq_q, num_heads, seq_kv).permute(1, 0, 2)
        scores = scores * sm_scale
        probs = F.softmax(scores, dim=-1)

        vi = kv_bf16[kv_s:kv_e, 0, :v_head_dim].float()
        oi = torch.matmul(probs, vi)
        oi = oi.permute(1, 0, 2).to(torch.bfloat16)
        out_list.append(oi)

    return torch.cat(out_list, dim=0)
scrolls · 383 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