Skip to content
KernelIndex
Search⌘K

submission 646222

sizezheng_94252 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4b29be9f65bb122cbb80b18c71a3637c364f5a13d8d4c401998eb60fad7a977f
license declaredunknown
license concludedunknown
authorssizezheng_94252
imported2026-08-26

Techniques

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

mmasc = tl.dot(q_1, tl.trans(kv1f)) + tl.dot(q_2, tl.trans(kv2f))
online-softmaxm_new = tl.maximum(m_i, tl.max(sc, axis=1))
persistent-kernelbatch_size = tl.num_programs(1) // NUM_HEAD_BLOCKS
split-kSplit-KV flash attention with GQA (16 query heads, 1 KV head).

Kernel source

submission.py265 lines
"""
Triton MLA decode kernel for MI355X.
Split-KV flash attention with GQA (16 query heads, 1 KV head).
FP8 KV, bf16 Q, online softmax, BLOCK_KV=32.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t

NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)


@triton.jit
def _mla_splitkv_kernel(
    Q, KV, kv_indptr, kv_scale_ptr,
    Out, Partial_O, Partial_LSE,
    sm_scale: tl.constexpr,
    num_splits: tl.constexpr,
    BLOCK_KV: tl.constexpr,
    V_DIM: tl.constexpr,
    QK_DIM: tl.constexpr,
    HEADS_PER_BLOCK: tl.constexpr,
    NUM_HEAD_BLOCKS: tl.constexpr,
    TOTAL_HEADS: tl.constexpr,
    SINGLE_SPLIT: tl.constexpr,
):
    """
    Grid: (num_splits, batch_size * NUM_HEAD_BLOCKS)
    Each block: HEADS_PER_BLOCK heads, one batch, one KV split.
    """
    split_id = tl.program_id(0)
    bh_id = tl.program_id(1)
    batch_size = tl.num_programs(1) // NUM_HEAD_BLOCKS

    batch_id = bh_id // NUM_HEAD_BLOCKS
    hb_id = bh_id % NUM_HEAD_BLOCKS
    h_off = hb_id * HEADS_PER_BLOCK

    kv_start = tl.load(kv_indptr + batch_id)
    kv_end = tl.load(kv_indptr + batch_id + 1)
    kv_len = kv_end - kv_start

    tps = tl.cdiv(kv_len, num_splits)
    my_start = kv_start + split_id * tps
    my_end = tl.minimum(kv_start + (split_id + 1) * tps, kv_end)
    my_len = my_end - my_start

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

    q_base = batch_id * (TOTAL_HEADS * QK_DIM) + h_off * QK_DIM
    h_ids = tl.arange(0, HEADS_PER_BLOCK)
    d1 = tl.arange(0, 512)
    d2 = 512 + tl.arange(0, 64)

    q_1 = tl.load(Q + q_base + h_ids[:, None] * QK_DIM + d1[None, :])
    q_2 = tl.load(Q + q_base + h_ids[:, None] * QK_DIM + d2[None, :])

    m_i = tl.full([HEADS_PER_BLOCK], value=float("-inf"), dtype=tl.float32)
    l_i = tl.zeros([HEADS_PER_BLOCK], dtype=tl.float32)
    acc = tl.zeros([HEADS_PER_BLOCK, V_DIM], dtype=tl.float32)
    v_ids = tl.arange(0, V_DIM)

    for bi in range(tl.cdiv(my_len, BLOCK_KV)):
        off = bi * BLOCK_KV + tl.arange(0, BLOCK_KV)
        mask = off < my_len
        g = my_start + off

        kv1 = tl.load(KV + g[:, None] * QK_DIM + d1[None, :], mask=mask[:, None], other=0.0)
        kv1f = (kv1.to(tl.float32) * kv_sc).to(tl.bfloat16)

        kv2 = tl.load(KV + g[:, None] * QK_DIM + d2[None, :], mask=mask[:, None], other=0.0)
        kv2f = (kv2.to(tl.float32) * kv_sc).to(tl.bfloat16)

        sc = tl.dot(q_1, tl.trans(kv1f)) + tl.dot(q_2, tl.trans(kv2f))
        sc = sc.to(tl.float32) * sm_scale
        sc = tl.where(mask[None, :], sc, float("-inf"))

        m_new = tl.maximum(m_i, tl.max(sc, axis=1))
        alpha = tl.exp(m_i - m_new)
        p = tl.exp(sc - m_new[:, None])
        l_i = alpha * l_i + tl.sum(p, axis=1)
        acc = acc * alpha[:, None]
        acc += tl.dot(p.to(tl.bfloat16), kv1f).to(tl.float32)
        m_i = m_new

    safe_l = tl.where(l_i > 0, l_i, 1.0)
    acc = acc / safe_l[:, None]

    if SINGLE_SPLIT == 1:
        out_base = (batch_id * TOTAL_HEADS + h_off) * V_DIM
        tl.store(Out + out_base + h_ids[:, None] * V_DIM + v_ids[None, :], acc.to(tl.bfloat16))
    else:
        po_base = (split_id * batch_size + batch_id) * (TOTAL_HEADS * V_DIM) + h_off * V_DIM
        tl.store(Partial_O + po_base + h_ids[:, None] * V_DIM + v_ids[None, :], acc)
        lse = m_i + tl.log(tl.where(l_i > 0, l_i, 1.0))
        lse = tl.where(l_i > 0, lse, float("-inf"))
        lse_base = (split_id * batch_size + batch_id) * TOTAL_HEADS + h_off
        tl.store(Partial_LSE + lse_base + h_ids, lse)


@triton.jit
def _mla_reduce_kernel(
    Partial_O, Partial_LSE, Output,
    batch_size: tl.constexpr, num_splits: tl.constexpr,
    V_DIM: tl.constexpr, TOTAL_HEADS: tl.constexpr,
):
    pid = tl.program_id(0)
    bid = pid // TOTAL_HEADS
    hid = pid % TOTAL_HEADS

    gm = float("-inf")
    for s in tl.static_range(0, 64):
        if s < num_splits:
            gm = tl.maximum(gm, tl.load(Partial_LSE + (s * batch_size + bid) * TOTAL_HEADS + hid))

    v_ids = tl.arange(0, V_DIM)
    acc = tl.zeros([V_DIM], dtype=tl.float32)
    tw = 0.0

    for s in tl.static_range(0, 64):
        if s < num_splits:
            idx = (s * batch_size + bid) * TOTAL_HEADS + hid
            w = tl.exp(tl.load(Partial_LSE + idx) - gm)
            tw += w
            po = (s * batch_size + bid) * (TOTAL_HEADS * V_DIM) + hid * V_DIM
            acc += w * tl.load(Partial_O + po + v_ids)

    tl.store(Output + (bid * TOTAL_HEADS + hid) * V_DIM + v_ids, (acc / tw).to(tl.bfloat16))


# ============================================================================
# Triton kernel caches
_triton_cache = {}

def _get_triton_cached(batch_size, num_splits):
    key = (batch_size, num_splits)
    if key in _triton_cache:
        return _triton_cache[key]
    if num_splits > 1:
        po = torch.empty(num_splits * batch_size * NUM_HEADS * V_HEAD_DIM, dtype=torch.float32, device="cuda")
        pl = torch.empty(num_splits * batch_size * NUM_HEADS, dtype=torch.float32, device="cuda")
    else:
        po = torch.empty(1, dtype=torch.float32, device="cuda")
        pl = torch.empty(1, dtype=torch.float32, device="cuda")
    out = torch.empty((batch_size, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
    _triton_cache[key] = (po, pl, out)
    return po, pl, out

# Triton split table (for small shapes where Triton beats aiter)
_TRITON_SPLIT_TABLE = {
    (4, 1024): 8, (4, 8192): 32,
    (32, 1024): 8, (32, 8192): 8,
    (64, 1024): 4,
}

# Aiter path caches
_aiter_cache = {}
PAGE_SIZE = 1
FP8_DTYPE = None

def _aiter_mla(q, kv_fp8, kv_scale, kv_indptr, batch_size, kv_seq_len):
    global FP8_DTYPE
    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
    import aiter as _aiter

    if FP8_DTYPE is None:
        FP8_DTYPE = aiter_dtypes.fp8

    total_kv_len = batch_size * kv_seq_len
    kv_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, 1, kv_fp8.shape[-1])

    key = (batch_size, kv_seq_len)
    if key not in _aiter_cache:
        kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
        if batch_size >= 256:
            num_kv_splits = 32
        elif batch_size >= 64:
            num_kv_splits = 16
        else:
            num_kv_splits = 8
        q_fp8_buf = torch.empty((batch_size, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda")
        q_scale_buf = torch.empty(1, dtype=torch.float32, device="cuda")
        o = torch.empty((batch_size, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
        info = get_mla_metadata_info_v1(
            batch_size, 1, NUM_HEADS, FP8_DTYPE, kv_fp8.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]
        _aiter_cache[key] = (kv_indices, num_kv_splits, work, o, q_fp8_buf, q_scale_buf)

    kv_indices, num_kv_splits, work, o, q_fp8, q_scale = _aiter_cache[key]
    (wm, wi, wis, ri, rfm, rpm) = work

    _aiter.dynamic_per_tensor_quant(q_fp8, q, q_scale)

    kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
    # Rebuild indptr-dependent metadata
    from aiter import get_mla_metadata_v1
    get_mla_metadata_v1(
        torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda"),  # qo_indptr
        kv_indptr, kv_last_page_len,
        NUM_HEADS, 1, True,
        wm, wis, wi, ri, rfm, rpm,
        page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE, 16),
        max_seqlen_qo=1, uni_seqlen_qo=1, fast_mode=False,
        max_split_per_batch=num_kv_splits, intra_batch_mode=True,
        dtype_q=FP8_DTYPE, dtype_kv=kv_fp8.dtype,
    )

    mla_decode_fwd(
        q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM), kv_4d, o,
        torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda"),
        kv_indptr, kv_indices, kv_last_page_len, 1,
        page_size=PAGE_SIZE, nhead_kv=1, sm_scale=SM_SCALE,
        logit_cap=0.0, num_kv_splits=num_kv_splits,
        q_scale=q_scale, kv_scale=kv_scale, intra_batch_mode=True,
        work_meta_data=wm, work_indptr=wi, work_info_set=wis,
        reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm,
    )
    return o


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    batch_size = config["batch_size"]
    kv_seq_len = config["kv_seq_len"]
    kv_fp8, kv_scale = kv_data["fp8"]

    # Hybrid: Triton for small shapes (faster due to zero overhead),
    # aiter assembly for large shapes (better bandwidth utilization)
    key = (batch_size, kv_seq_len)
    if key in _TRITON_SPLIT_TABLE:
        # Triton path
        kv_flat = kv_fp8.view(-1, QK_HEAD_DIM)
        num_splits = _TRITON_SPLIT_TABLE[key]
        po, pl, out = _get_triton_cached(batch_size, num_splits)
        single_split = 1 if num_splits == 1 else 0

        _mla_splitkv_kernel[(num_splits, batch_size)](
            q, kv_flat, kv_indptr, kv_scale,
            out, po, pl,
            sm_scale=SM_SCALE, num_splits=num_splits, BLOCK_KV=32,
            V_DIM=V_HEAD_DIM, QK_DIM=QK_HEAD_DIM,
            HEADS_PER_BLOCK=16, NUM_HEAD_BLOCKS=1,
            TOTAL_HEADS=NUM_HEADS, SINGLE_SPLIT=single_split,
        )
        if num_splits > 1:
            _mla_reduce_kernel[(batch_size * NUM_HEADS,)](
                po, pl, out,
                batch_size=batch_size, num_splits=num_splits,
                V_DIM=V_HEAD_DIM, TOTAL_HEADS=NUM_HEADS,
            )
        return out
    else:
        # Aiter assembly path for large shapes
        return _aiter_mla(q, kv_fp8, kv_scale, kv_indptr, batch_size, kv_seq_len)
scrolls · 265 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