Skip to content
KernelIndex
Search⌘K

submission 730703

c2flowDS · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:cc6a77b4992b05ca2ca61380b78f8f7f0082bbe3caf3736961217e357566ffba
license declaredunknown
license concludedunknown
authorsc2flowDS
imported2026-08-15

Techniques

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

fp4Hybrid: MXFP4 Triton for bs<=4 kv<=1024 (lower overhead), FP8 ASM for rest (persistent + lower BW).
fp8q_tile1, q_descale1, "e4m3", acc=qk_all, fast_math=True)
num-warps = 8num_warps=8, num_stages=2, waves_per_eu=2,
online-softmaxm_new = tl.maximum(m_i, m_ij)
persistent-kernelHybrid: MXFP4 Triton for bs<=4 kv<=1024 (lower overhead), FP8 ASM for rest (persistent + lower BW).
stages = 2num_warps=8, num_stages=2, waves_per_eu=2,
tile-n = 64BLOCK_N=64, V_DIM=V_HEAD_DIM,

Kernel source

submission.py288 lines
# gpumode leaderboard — MXFP4/FP8 hybrid v43-retest
"""
Hybrid: MXFP4 Triton for bs<=4 kv<=1024 (lower overhead), FP8 ASM for rest (persistent + lower BW).
"""

import torch
import triton
import triton.language as tl
import aiter as _aiter_mod
from task import input_t, output_t

from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.ops.quant import static_per_tensor_quant

NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (576 ** 0.5)
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8

# MXFP4 constants
QK_PACKED = QK_HEAD_DIM // 2
K_PART1_PACKED = 256
K_PART1_SCALES = 16
Q_PART1_DIM = 512
MXFP4_SPLITS = 16

# FP8 ASM constants
ASM_SPLITS = 16


# ====== MXFP4 Triton Kernel (v30) ======

@triton.jit
def _mxfp4_stage1(
    Q_fp8, KV_packed, KV_Scale, V_fp8, Q_scale_ptr, V_scale_ptr, kv_indptr,
    Out, Lse,
    stride_q_batch, stride_q_head, stride_q_dim,
    stride_kv_n, stride_ks_n, stride_vn,
    stride_os, stride_oq, stride_oh,
    stride_ls, stride_lq, stride_lh,
    SM_SCALE_VAL,
    BLOCK_N: tl.constexpr, V_DIM: tl.constexpr,
    K_P1_PACKED: tl.constexpr, K_P1_SCALES: tl.constexpr,
    Q_P1_DIM: tl.constexpr, NUM_HEADS_: tl.constexpr, NUM_SPLITS: tl.constexpr,
):
    split_id = tl.program_id(0)
    batch_id = tl.program_id(1)
    LOG2E: tl.constexpr = 1.44269504

    kv_start = tl.load(kv_indptr + batch_id)
    kv_end = tl.load(kv_indptr + batch_id + 1)
    kv_len = kv_end - kv_start
    split_size = tl.cdiv(kv_len, NUM_SPLITS)
    my_start = kv_start + split_id * split_size
    my_end = tl.minimum(my_start + split_size, kv_end)
    q_idx = batch_id

    if my_start >= my_end:
        offs_h = tl.arange(0, NUM_HEADS_)
        offs_dv = tl.arange(0, V_DIM)
        tl.store(Out + split_id * stride_os + q_idx * stride_oq + offs_h[:, None] * stride_oh + offs_dv[None, :],
                 tl.zeros([NUM_HEADS_, V_DIM], dtype=tl.float32))
        tl.store(Lse + split_id * stride_ls + q_idx * stride_lq + offs_h * stride_lh,
                 tl.full([NUM_HEADS_], float("-inf"), dtype=tl.float32))
        return

    q_base = Q_fp8 + q_idx * stride_q_batch
    offs_d1 = tl.arange(0, Q_P1_DIM)
    offs_h = tl.arange(0, NUM_HEADS_)
    q_tile1 = tl.load(q_base + offs_d1[:, None] * stride_q_dim + offs_h[None, :] * stride_q_head)
    q_descale1 = tl.full([NUM_HEADS_, Q_P1_DIM // 32], 127, dtype=tl.uint8)
    qk_factor = tl.load(Q_scale_ptr) * SM_SCALE_VAL
    v_descale = tl.load(V_scale_ptr)
    p_descale_pv = tl.full([NUM_HEADS_, BLOCK_N // 32], 127, dtype=tl.uint8)
    v_descale_pv = tl.full([V_DIM, BLOCK_N // 32], 127, dtype=tl.uint8)

    m_i = tl.full([NUM_HEADS_], float("-inf"), dtype=tl.float32)
    l_i = tl.zeros([NUM_HEADS_], dtype=tl.float32)
    acc = tl.zeros([NUM_HEADS_, V_DIM], dtype=tl.float32)

    for _iter in range(0, tl.cdiv(my_end - my_start, BLOCK_N)):
        kv_pos = my_start + _iter * BLOCK_N
        offs_n = tl.arange(0, BLOCK_N)
        valid = (kv_pos + offs_n) < my_end

        k1_ptrs = KV_packed + (kv_pos + offs_n[:, None]) * stride_kv_n + tl.arange(0, K_P1_PACKED)[None, :]
        k1_packed = tl.load(k1_ptrs, mask=valid[:, None], other=0)
        ks1_ptrs = KV_Scale + (kv_pos + offs_n[:, None]) * stride_ks_n + tl.arange(0, K_P1_SCALES)[None, :]
        k1_scales = tl.load(ks1_ptrs, mask=valid[:, None], other=127)

        qk_all = tl.zeros([BLOCK_N, NUM_HEADS_], dtype=tl.float32)
        qk_all = tl.dot_scaled(k1_packed, k1_scales, "e2m1",
                               q_tile1, q_descale1, "e4m3", acc=qk_all, fast_math=True)
        qk_all = qk_all * qk_factor
        qk_all = tl.where(valid[:, None], qk_all, float("-inf"))
        qk = tl.trans(qk_all)

        m_ij = tl.max(qk, 1)
        m_new = tl.maximum(m_i, m_ij)
        alpha = tl.math.exp2((m_i - m_new) * LOG2E)
        p = tl.math.exp2((qk - m_new[:, None]) * LOG2E)
        l_i = l_i * alpha + tl.sum(p, 1)
        acc = acc * alpha[:, None]
        m_i = m_new

        v_ptrs = V_fp8 + (kv_pos + offs_n[:, None]) * stride_vn + tl.arange(0, V_DIM)[None, :]
        v_block = tl.load(v_ptrs, mask=valid[:, None], other=0.0)
        p_fp8 = p.to(tl.float8e4nv)
        pv = tl.dot_scaled(p_fp8, p_descale_pv, "e4m3",
                           v_block, v_descale_pv, "e4m3", fast_math=True)
        acc += pv

    safe_l = tl.where(l_i == 0.0, 1.0, l_i)
    result = acc * v_descale / safe_l[:, None]
    lse_vals = m_i + tl.math.log2(tl.where(l_i == 0.0, 1.0, l_i)) / LOG2E

    offs_h2 = tl.arange(0, NUM_HEADS_)
    offs_dv = tl.arange(0, V_DIM)
    tl.store(Out + split_id * stride_os + q_idx * stride_oq + offs_h2[:, None] * stride_oh + offs_dv[None, :], result)
    tl.store(Lse + split_id * stride_ls + q_idx * stride_lq + offs_h2 * stride_lh, lse_vals)


@triton.jit
def _reduce_splits(
    Split_out, Split_lse, Out,
    stride_ss, stride_sq, stride_sh, stride_sd,
    stride_ls, stride_lq, stride_lh,
    stride_oq, stride_oh, stride_od,
    V_DIM: tl.constexpr, NUM_SPLITS: tl.constexpr,
):
    q_idx = tl.program_id(0)
    head_id = tl.program_id(1)
    LOG2E: tl.constexpr = 1.44269504
    max_lse = float("-inf")
    for s in range(NUM_SPLITS):
        lse_s = tl.load(Split_lse + s * stride_ls + q_idx * stride_lq + head_id * stride_lh)
        max_lse = tl.maximum(max_lse, lse_s)
    acc = tl.zeros([V_DIM], dtype=tl.float32)
    sum_exp = 0.0
    for s in range(NUM_SPLITS):
        lse_s = tl.load(Split_lse + s * stride_ls + q_idx * stride_lq + head_id * stride_lh)
        w = tl.math.exp2((lse_s - max_lse) * LOG2E)
        sum_exp += w
        offs_dv = tl.arange(0, V_DIM)
        partial = tl.load(Split_out + s * stride_ss + q_idx * stride_sq + head_id * stride_sh + offs_dv)
        acc += partial * w
    acc = acc / tl.where(sum_exp == 0.0, 1.0, sum_exp)
    tl.store(Out + q_idx * stride_oq + head_id * stride_oh + tl.arange(0, V_DIM), acc.to(tl.bfloat16))


# ====== Caches ======

_mxfp4_cache = {}
_asm_cache = {}

def _build_mxfp4_cache(total_q, total_kv, device):
    return {
        "o": torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device),
        "split_out": torch.empty((MXFP4_SPLITS, total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device=device),
        "split_lse": torch.empty((MXFP4_SPLITS, total_q, NUM_HEADS), dtype=torch.float32, device=device),
        "q_fp8": torch.empty((total_q, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device=device),
        "q_scale": torch.tensor([5.5 / 448.0], dtype=torch.float32, device=device),
    }

def _build_asm_cache(batch_size, q_seq_len, total_q, total_kv, device, qo_indptr, kv_indptr):
    q_fp8 = torch.empty((total_q, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device=device)
    q_scale = torch.tensor([5.5 / 448.0], dtype=torch.float32, device=device)
    o = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device)
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)
    kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)

    info = get_mla_metadata_info_v1(
        batch_size, q_seq_len, NUM_HEADS, FP8_DTYPE, FP8_DTYPE,
        is_sparse=False, fast_mode=False, num_kv_splits=ASM_SPLITS, intra_batch_mode=True,
    )
    work = [torch.empty(s, dtype=t, device=device) for s, t in info]
    (work_metadata, work_indptr_t, work_info_set, reduce_indptr, reduce_final_map, reduce_partial_map) = work

    get_mla_metadata_v1(
        qo_indptr, kv_indptr, kv_last_page_len,
        NUM_HEADS // NUM_KV_HEADS, NUM_KV_HEADS, False,
        work_metadata, work_info_set, work_indptr_t,
        reduce_indptr, reduce_final_map, reduce_partial_map,
        page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE, 16),
        max_seqlen_qo=q_seq_len, uni_seqlen_qo=q_seq_len,
        fast_mode=False, max_split_per_batch=ASM_SPLITS, intra_batch_mode=True,
        dtype_q=FP8_DTYPE, dtype_kv=FP8_DTYPE,
    )

    num_partial = reduce_partial_map.size(0)
    logits = torch.empty((num_partial * q_seq_len, 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device=device)
    attn_lse = torch.empty((num_partial * q_seq_len, 1, NUM_HEADS, 1), dtype=torch.float32, device=device)

    return {
        "q_fp8": q_fp8, "q_fp8_flat": q_fp8.view(-1, QK_HEAD_DIM), "q_scale": q_scale,
        "o": o, "kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,
        "work_meta_data": work_metadata, "work_indptr": work_indptr_t,
        "work_info_set": work_info_set, "reduce_indptr": reduce_indptr,
        "reduce_final_map": reduce_final_map, "reduce_partial_map": reduce_partial_map,
        "logits": logits, "attn_lse": attn_lse,
        "q_seq_len": q_seq_len, "total_kv": total_kv,
    }


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    batch_size = config["batch_size"]
    total_q = q.shape[0]
    total_kv = batch_size * config["kv_seq_len"]

    if batch_size <= 4 and config["kv_seq_len"] <= 1024:
        # MXFP4 Triton path (better for small batch)
        key = (batch_size, total_q, total_kv)
        c = _mxfp4_cache.get(key)
        if c is None:
            c = _build_mxfp4_cache(total_q, total_kv, q.device)
            _mxfp4_cache[key] = c

        static_per_tensor_quant(c["q_fp8"].view(-1, QK_HEAD_DIM), q.view(-1, QK_HEAD_DIM), c["q_scale"])

        kv_buffer_mxfp4, kv_scale_mxfp4 = kv_data["mxfp4"]
        kv_raw = kv_buffer_mxfp4.view(torch.uint8).reshape(total_kv, QK_PACKED)
        ks_raw = kv_scale_mxfp4.view(torch.uint8)
        if ks_raw.dim() > 2:
            ks_raw = ks_raw.reshape(total_kv, -1)

        kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
        v_fp8 = kv_buffer_fp8.view(total_kv, QK_HEAD_DIM)

        q_fp8 = c["q_fp8"]
        _mxfp4_stage1[(MXFP4_SPLITS, batch_size)](
            q_fp8, kv_raw, ks_raw, v_fp8, c["q_scale"], kv_scale_fp8, kv_indptr,
            c["split_out"], c["split_lse"],
            q_fp8.stride(0), q_fp8.stride(1), q_fp8.stride(2),
            kv_raw.stride(0), ks_raw.stride(0), v_fp8.stride(0),
            c["split_out"].stride(0), c["split_out"].stride(1), c["split_out"].stride(2),
            c["split_lse"].stride(0), c["split_lse"].stride(1), c["split_lse"].stride(2),
            SM_SCALE,
            BLOCK_N=64, V_DIM=V_HEAD_DIM,
            K_P1_PACKED=K_PART1_PACKED, K_P1_SCALES=K_PART1_SCALES,
            Q_P1_DIM=Q_PART1_DIM, NUM_HEADS_=NUM_HEADS, NUM_SPLITS=MXFP4_SPLITS,
            num_warps=8, num_stages=2, waves_per_eu=2,
        )

        _reduce_splits[(total_q, NUM_HEADS)](
            c["split_out"], c["split_lse"], c["o"],
            c["split_out"].stride(0), c["split_out"].stride(1),
            c["split_out"].stride(2), c["split_out"].stride(3),
            c["split_lse"].stride(0), c["split_lse"].stride(1), c["split_lse"].stride(2),
            c["o"].stride(0), c["o"].stride(1), c["o"].stride(2),
            V_DIM=V_HEAD_DIM, NUM_SPLITS=MXFP4_SPLITS,
        )
        return c["o"]

    else:
        # FP8 ASM path (better for large batch)
        q_seq_len = config["q_seq_len"]
        key = (batch_size, q_seq_len, total_q, total_kv)
        c = _asm_cache.get(key)
        if c is None:
            c = _build_asm_cache(batch_size, q_seq_len, total_q, total_kv,
                                 q.device, qo_indptr, kv_indptr)
            _asm_cache[key] = c

        static_per_tensor_quant(c["q_fp8_flat"], q.view(-1, QK_HEAD_DIM), c["q_scale"])

        kv_buffer_fp8, kv_scale = kv_data["fp8"]
        kv_4d = kv_buffer_fp8.view(c["total_kv"], PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)

        _aiter_mod.mla_decode_stage1_asm_fwd(
            c["q_fp8"], kv_4d, qo_indptr, kv_indptr,
            c["kv_indices"], c["kv_last_page_len"],
            None, c["work_meta_data"], c["work_indptr"], c["work_info_set"],
            c["q_seq_len"], PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
            c["logits"], c["attn_lse"], c["o"], c["q_scale"], kv_scale,
        )

        _aiter_mod.mla_reduce_v1(
            c["logits"], c["attn_lse"], c["reduce_indptr"],
            c["reduce_final_map"], c["reduce_partial_map"],
            c["q_seq_len"], c["o"], None,
        )
        return c["o"]
scrolls · 288 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