Skip to content
KernelIndex
Search⌘K

submission 668995

Rakesh Jarupula · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

amd_mixed_mla.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-668995?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
877.7µs
#724 of 766
2026-03-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2afeadc4bcc73a5887f7523cf91f018abf3788dd8818dd758b7c8167f08e31f4
license declaredunknown
license concludedunknown
authorsRakesh Jarupula
imported2026-08-26

Techniques

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

online-softmaxm_new = tl.maximum(m_i, score)
split-k- Triton split-K: grid=(B, H, S), each CTA handles a KV slice
tile-n = 8BN = 8 # inner unroll factor

Kernel source

amd_mixed_mla.py364 lines
"""
Custom MLA (Multi-head Latent Attention) decode kernel optimized for MI355X (CDNA3).

Key design:
  - Triton split-K: grid=(B, H, S), each CTA handles a KV slice
  - Separate K (576-dim) and V (512-dim) loads per KV token — avoids shape mismatch
  - FP8 KV: scalar dequant inline (multiply by kv_scale)
  - Online softmax (running m, l) — Flash Attention style
  - Reduction kernel merges splits via LSE (numerically stable)
  - PyTorch bmm fallback for safety

DeepSeek R1 forward_absorb MLA:
  num_heads=16, num_kv_heads=1 (MQA), qk_head_dim=576, v_head_dim=512
  Decode: q_seq_len=1, kv_seq_len up to 8k
"""

import torch
import triton
import triton.language as tl
from task import input_t, output_t
from utils import make_match_reference

# -----------------------------------------------------------------------
# Constants
# -----------------------------------------------------------------------
NUM_HEADS    = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM  = 576   # kv_lora_rank(512) + qk_rope_head_dim(64)
V_HEAD_DIM   = 512   # = kv_lora_rank
SM_SCALE     = 1.0 / (QK_HEAD_DIM ** 0.5)

_NUM_KV_SPLITS = 32


# -----------------------------------------------------------------------
# Triton kernel 1/2: per-split attention with online softmax
#
# Grid = (batch_size, H, S)
#
# Notes on K vs V:
#   KV buffer layout: (total_kv, 1, D) where D=576
#   K uses all D=576 dims for score computation
#   V uses first Dv=512 dims for output accumulation
#   We load K with tile BD=1024 (masked to D=576)
#   We load V with tile BDv=512 (masked to Dv=512) from same ptr
# -----------------------------------------------------------------------

@triton.jit
def _mla_split_fwd(
    Q_ptr,        # (total_q, H, D) bf16
    KV_ptr,       # (total_kv, 1, D) fp8
    kv_scale_ptr, # () scalar f32
    O_part_ptr,   # (B, H, S, Dv) f32
    LSE_part_ptr, # (B, H, S) f32
    QO_Indptr,    # (B+1,) i32
    KV_Indptr,    # (B+1,) i32
    # Q strides
    sq0, sq1, sq2,
    # KV strides
    skv0, skv1, skv2,
    # O_part strides
    so0, so1, so2, so3,
    # LSE_part strides
    sl0, sl1, sl2,
    # Compile-time constants
    D:   tl.constexpr,   # 576
    Dv:  tl.constexpr,   # 512
    BD:  tl.constexpr,   # >= D, power of 2 (1024)
    BDv: tl.constexpr,   # >= Dv, power of 2 (512)
    BN:  tl.constexpr,   # KV tokens per inner unroll
    S:   tl.constexpr,   # num_kv_splits
    SM:  tl.constexpr,   # sm_scale (float literal)
):
    b = tl.program_id(0)
    h = tl.program_id(1)
    s = tl.program_id(2)

    # Decode: 1 query token per batch element
    qt = tl.load(QO_Indptr + b)

    kv0    = tl.load(KV_Indptr + b)
    kv1    = tl.load(KV_Indptr + b + 1)
    kv_len = kv1 - kv0

    per_split = tl.cdiv(kv_len, S)
    s_kv0 = kv0 + s * per_split
    s_kv1 = tl.minimum(s_kv0 + per_split, kv1)

    lse_ptr = LSE_part_ptr + b * sl0 + h * sl1 + s * sl2
    if s_kv0 >= s_kv1:
        tl.store(lse_ptr, float("-inf"))
        return

    # Load Q (D=576 dims) as f32  — tile BD=1024 with mask
    d_idx = tl.arange(0, BD)
    q = tl.load(Q_ptr + qt * sq0 + h * sq1 + d_idx * sq2,
                mask=d_idx < D, other=0.0).to(tl.float32)

    # V index for separate V load (BDv=512 tile)
    v_idx = tl.arange(0, BDv)

    kv_scale = tl.load(kv_scale_ptr).to(tl.float32)

    # Online softmax state
    m_i = float("-inf")
    l_i = 0.0
    acc = tl.zeros([BDv], dtype=tl.float32)

    kv_tok = s_kv0
    while kv_tok < s_kv1:
        for i in tl.static_range(0, BN):
            tok = kv_tok + i
            if tok < s_kv1:
                kv_base = KV_ptr + tok * skv0 + 0 * skv1

                # Load K: full D=576 dims (tile BD=1024, mask D)
                k = tl.load(kv_base + d_idx * skv2,
                             mask=d_idx < D, other=0.0).to(tl.float32) * kv_scale

                # Score: dot(q, k) * sm_scale
                score = tl.sum(q * k) * SM

                # Online softmax update
                m_new = tl.maximum(m_i, score)
                e     = tl.exp(score - m_new)
                r     = tl.exp(m_i - m_new)
                l_i   = l_i * r + e
                acc   = acc * r

                # Load V: first Dv=512 dims (tile BDv=512, mask Dv)
                # Same kv_base pointer, different index tile
                v = tl.load(kv_base + v_idx * skv2,
                             mask=v_idx < Dv, other=0.0).to(tl.float32) * kv_scale

                acc = acc + e * v
                m_i = m_new
        kv_tok += BN

    acc = acc / tl.maximum(l_i, 1e-8)

    o_base = O_part_ptr + b * so0 + h * so1 + s * so2
    tl.store(o_base + v_idx * so3, acc, mask=v_idx < Dv)
    tl.store(lse_ptr, m_i + tl.log(tl.maximum(l_i, 1e-8)))


# -----------------------------------------------------------------------
# Triton kernel 2/2: reduction across splits
# Grid = (B, H)
# -----------------------------------------------------------------------

@triton.jit
def _mla_reduce(
    O_part_ptr,   # (B, H, S, Dv) f32
    LSE_part_ptr, # (B, H, S) f32
    OUT_ptr,      # (total_q, H, Dv) bf16
    QO_Indptr,    # (B+1,) i32
    so0, so1, so2, so3,
    sl0, sl1, sl2,
    ot0, ot1, ot2,
    Dv:  tl.constexpr,
    BDv: tl.constexpr,
    S:   tl.constexpr,
):
    b = tl.program_id(0)
    h = tl.program_id(1)
    qt = tl.load(QO_Indptr + b)
    v_idx = tl.arange(0, BDv)

    # Find global max LSE
    m_g = float("-inf")
    for s in tl.static_range(0, S):
        m_g = tl.maximum(m_g, tl.load(LSE_part_ptr + b * sl0 + h * sl1 + s * sl2))

    # Weighted sum
    acc   = tl.zeros([BDv], dtype=tl.float32)
    l_tot = 0.0
    for s in tl.static_range(0, S):
        lse_s = tl.load(LSE_part_ptr + b * sl0 + h * sl1 + s * sl2)
        w     = tl.exp(lse_s - m_g)
        l_tot = l_tot + w
        o_s   = tl.load(O_part_ptr + b * so0 + h * so1 + s * so2 + v_idx * so3,
                         mask=v_idx < Dv, other=0.0)
        acc   = acc + w * o_s

    acc = acc / tl.maximum(l_tot, 1e-8)
    tl.store(OUT_ptr + qt * ot0 + h * ot1 + v_idx * ot2,
             acc.to(tl.bfloat16), mask=v_idx < Dv)


# -----------------------------------------------------------------------
# Triton dispatch
# -----------------------------------------------------------------------

def _triton_decode_fp8(
    q:          torch.Tensor,
    kv_fp8:     torch.Tensor,
    kv_scale:   torch.Tensor,
    qo_indptr:  torch.Tensor,
    kv_indptr:  torch.Tensor,
    config:     dict,
    S:          int,
) -> torch.Tensor:
    B  = config["batch_size"]
    H  = config["num_heads"]
    D  = config["qk_head_dim"]    # 576
    Dv = config["v_head_dim"]     # 512
    SM = float(config["sm_scale"])

    BD  = triton.next_power_of_2(D)    # 1024
    BDv = triton.next_power_of_2(Dv)   # 512
    BN  = 8  # inner unroll factor

    O_part   = torch.empty((B, H, S, Dv), dtype=torch.float32, device=q.device)
    LSE_part = torch.full( (B, H, S),     float("-inf"), dtype=torch.float32, device=q.device)

    _mla_split_fwd[(B, H, S)](
        q, kv_fp8, kv_scale,
        O_part, LSE_part,
        qo_indptr, kv_indptr,
        q.stride(0), q.stride(1), q.stride(2),
        kv_fp8.stride(0), kv_fp8.stride(1), kv_fp8.stride(2),
        O_part.stride(0), O_part.stride(1), O_part.stride(2), O_part.stride(3),
        LSE_part.stride(0), LSE_part.stride(1), LSE_part.stride(2),
        D=D, Dv=Dv, BD=BD, BDv=BDv, BN=BN, S=S, SM=SM,
    )

    output = torch.empty((q.shape[0], H, Dv), dtype=torch.bfloat16, device=q.device)
    _mla_reduce[(B, H)](
        O_part, LSE_part, output, qo_indptr,
        O_part.stride(0), O_part.stride(1), O_part.stride(2), O_part.stride(3),
        LSE_part.stride(0), LSE_part.stride(1), LSE_part.stride(2),
        output.stride(0), output.stride(1), output.stride(2),
        Dv=Dv, BDv=BDv, S=S,
    )
    return output


# -----------------------------------------------------------------------
# PyTorch fallback: FP8 KV + batch-parallel bmm
# -----------------------------------------------------------------------

def _torch_decode_fp8(
    q:         torch.Tensor,
    kv_fp8:    torch.Tensor,
    kv_scale:  torch.Tensor,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    config:    dict,
) -> torch.Tensor:
    B  = config["batch_size"]
    H  = config["num_heads"]
    D  = config["qk_head_dim"]
    Dv = config["v_head_dim"]
    sm = float(config["sm_scale"])
    sc = float(kv_scale.item())

    # Dequantize all KV at once
    kv = kv_fp8.to(torch.float32).mul_(sc).to(torch.bfloat16)  # (total_kv, 1, D)

    total_q = q.shape[0]
    out = torch.zeros((total_q, H, Dv), dtype=torch.bfloat16, device=q.device)

    for b in range(B):
        qs = int(qo_indptr[b]);  qe = int(qo_indptr[b + 1])
        ks = int(kv_indptr[b]);  ke = int(kv_indptr[b + 1])

        q_b = q[qs:qe]           # (Lq, H, D) bf16
        k_b = kv[ks:ke, 0, :]    # (Lkv, D)
        v_b = k_b[:, :Dv]        # (Lkv, Dv)

        qt = q_b.permute(1, 0, 2).float()                   # (H, Lq, D)
        kt = k_b.unsqueeze(0).expand(H, -1, -1).float()    # (H, Lkv, D)
        sc_mat = torch.bmm(qt, kt.transpose(1, 2)).mul_(sm) # (H, Lq, Lkv)
        attn   = torch.softmax(sc_mat, dim=-1).to(torch.bfloat16)

        vt = v_b.unsqueeze(0).expand(H, -1, -1)             # (H, Lkv, Dv)
        out[qs:qe] = torch.bmm(attn, vt).permute(1, 0, 2)

    return out


# -----------------------------------------------------------------------
# PyTorch fallback: MXFP4 KV (dequant + bmm)
# -----------------------------------------------------------------------

def _torch_decode_mxfp4(
    q:              torch.Tensor,
    kv_fp4:         torch.Tensor,
    kv_scale_e8m0:  torch.Tensor,
    qo_indptr:      torch.Tensor,
    kv_indptr:      torch.Tensor,
    config:         dict,
) -> torch.Tensor:
    from aiter.utility.fp4_utils import mxfp4_to_f32, e8m0_to_f32

    B  = config["batch_size"]
    H  = config["num_heads"]
    D  = config["qk_head_dim"]
    Dv = config["v_head_dim"]
    sm = float(config["sm_scale"])

    total_kv = kv_fp4.shape[0]
    nb = D // 32  # 576/32 = 18 scale blocks

    kv_f32 = mxfp4_to_f32(kv_fp4.reshape(total_kv, D // 2))  # (total_kv, D)
    sc     = e8m0_to_f32(kv_scale_e8m0)[:total_kv, :nb]       # (total_kv, 18)
    kv_f32 = (kv_f32.view(total_kv, nb, 32) * sc.unsqueeze(-1)).view(total_kv, D)
    kv_bf  = kv_f32.to(torch.bfloat16)

    total_q = q.shape[0]
    out = torch.zeros((total_q, H, Dv), dtype=torch.bfloat16, device=q.device)

    for b in range(B):
        qs = int(qo_indptr[b]);  qe = int(qo_indptr[b + 1])
        ks = int(kv_indptr[b]);  ke = int(kv_indptr[b + 1])

        q_b = q[qs:qe]
        k_b = kv_bf[ks:ke]
        v_b = k_b[:, :Dv]

        qt = q_b.permute(1, 0, 2).float()
        kt = k_b.unsqueeze(0).expand(H, -1, -1).float()
        scores = torch.bmm(qt, kt.transpose(1, 2)).mul_(sm)
        attn   = torch.softmax(scores, dim=-1).to(torch.bfloat16)
        vt = v_b.unsqueeze(0).expand(H, -1, -1)
        out[qs:qe] = torch.bmm(attn, vt).permute(1, 0, 2)

    return out


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

_use_triton: bool = True


def custom_kernel(data: input_t) -> output_t:
    """
    MLA decode kernel — DeepSeek R1 forward_absorb path.

    1. Triton split-K + FP8 KV (primary — maximizes GPU occupancy)
    2. PyTorch bmm + FP8 KV (fallback)
    """
    global _use_triton

    q, kv_data, qo_indptr, kv_indptr, config = data
    kv_fp8, kv_scale = kv_data["fp8"]

    if _use_triton:
        try:
            return _triton_decode_fp8(
                q, kv_fp8, kv_scale,
                qo_indptr, kv_indptr, config,
                S=_NUM_KV_SPLITS,
            )
        except Exception:
            _use_triton = False

    return _torch_decode_fp8(q, kv_fp8, kv_scale, qo_indptr, kv_indptr, config)


check_implementation = make_match_reference(custom_kernel, rtol=1e-01, atol=1e-01)
scrolls · 364 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