Skip to content
KernelIndex
Search⌘K

submission 586648

Ronit Kapoor · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-586648?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
5.86ms
#759 of 766
2026-03-18

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:ba5b8c6d84eeae0d7d81e82d946809dcd0e338e06482219fa398b4eb6249b487
license declaredunknown
license concludedunknown
authorsRonit Kapoor
imported2026-08-26

Kernel source

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

"""
MLA (Multi-head Latent Attention) decode kernel.

Next experiment:
- bf16 default
- first try ROCm SDPA for the uniform decode case:
    * q_seq_len == 1 for every request
    * kv_seq_len is uniform across the batch
    * dense packed q / kv buffers
- if SDPA is unavailable or slower-path conditions are not met, fall back to:
    * batched matmul fast path for short KV (<= 1024)
    * segmented fallback for long KV
- fp8 path kept unchanged
"""

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

QKV_DTYPE = "bf16"
FASTPATH_MAX_KV = 1024


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

def custom_kernel(data: input_t) -> output_t:
    if QKV_DTYPE == "fp8":
        return custom_kernel_fp8(data)
    elif QKV_DTYPE == "bf16":
        return custom_kernel_bf16(data)
    else:
        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 _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


# ---------------------------------------------------------------------------
# Uniform decode fast path using SDPA (best next experiment)
# ---------------------------------------------------------------------------

def _try_uniform_bf16_sdpa_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:
        return None

    try:
        # q:   [B, 16, 1, 576]
        q_sdpa = q.contiguous().unsqueeze(2)

        # kv:  [B, L, 576]
        kv_bld = kv_buffer_bf16[:, 0, :].contiguous().view(batch_size, kv_len, q.shape[-1])

        # k: [B, 1, L, 576], v: [B, 1, L, 512]
        k_sdpa = kv_bld.unsqueeze(1)
        v_sdpa = kv_bld[:, :, :v_head_dim].unsqueeze(1)

        out = F.scaled_dot_product_attention(
            q_sdpa,
            k_sdpa,
            v_sdpa,
            attn_mask=None,
            dropout_p=0.0,
            is_causal=False,
            scale=sm_scale,
            enable_gqa=True,
        )  # [B, 16, 1, 512]

        return out.squeeze(2).to(torch.bfloat16).contiguous()
    except Exception:
        return None


# ---------------------------------------------------------------------------
# Existing short-KV batched fast path
# ---------------------------------------------------------------------------

def _try_uniform_bf16_matmul_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()


# ---------------------------------------------------------------------------
# Segmented bf16 fallback
# ---------------------------------------------------------------------------

def _segmented_bf16_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())

        qi = q[q_s:q_e]                        # [seq_q, 16, 576]
        kvc = kv_buffer_bf16[kv_s:kv_e, 0]    # [seq_kv, 576]

        ki = kvc
        vi = kvc[:, :v_head_dim]

        qi_t = qi.float().permute(1, 0, 2)     # [16, seq_q, 576]
        scores = torch.matmul(qi_t * sm_scale, ki.float().T)
        probs = F.softmax(scores, dim=-1)

        oi = torch.matmul(probs, vi.float())   # [16, seq_q, 512]
        oi = oi.permute(1, 0, 2)               # [seq_q, 16, 512]
        out_list.append(oi.to(torch.bfloat16))

    return torch.cat(out_list, dim=0)


# ---------------------------------------------------------------------------
# Public bf16 kernel
# ---------------------------------------------------------------------------

def custom_kernel_bf16(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"]

    # 1) Best next experiment: ROCm SDPA with GQA.
    fast = _try_uniform_bf16_sdpa_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

    # 2) Proven short-KV fast path.
    fast = _try_uniform_bf16_matmul_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

    # 3) Generic fallback.
    return _segmented_bf16_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 path
# ---------------------------------------------------------------------------

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 · 329 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