Skip to content
KernelIndex
Search⌘K

submission 722089

ihansel · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:7f72deaeedca54c9e16803b2296c798d4bdc5ef63d2a294d48333b7b88a8afc8
license declaredunknown
license concludedunknown
authorsihansel
imported2026-08-26

Techniques

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

fp4kv_mxfp4_buf, kv_mxfp4_scale = kv_data["mxfp4"]
fp8q_fp8 = q_chunk.to(tl.float8e4nv)
mmaacc += tl.dot(p.to(tl.bfloat16), v_block, out_dtype=tl.float32)
num-warps = 8num_warps=8, num_stages=2,
online-softmaxm_new = tl.maximum(m_i, m_ij)
split-ksplit_kv_len = split_end - split_start
stages = 2num_warps=8, num_stages=2,
tile-k = 256BLOCK_K = 256
tile-n = 64BLOCK_N = 64

Kernel source

submission.py358 lines
"""MLA Decode v60: Multi-head KV reuse + metadata caching.
M2 experiment: load KV ONCE per program, process ALL 16 query heads.
16x reduction in KV bandwidth vs single-head kernel.
Grid: (batch * num_splits) instead of (batch * heads * num_splits).
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
from reference import ref_kernel  # noqa: F401
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

SM_SCALE = 1.0 / (576 ** 0.5)
SM_SCALE_LOG2 = SM_SCALE * 1.44269504088896
FP8_DTYPE = aiter_dtypes.fp8
PAGE_SIZE = 1
PADDED_V = 512


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


# ==================== MULTI-HEAD MXFP4 KERNEL (M2) ====================
@triton.jit
def _mla_multihead_stage1(
    Q, stride_qt, stride_qh, stride_qd,
    K_fp4, stride_kf_t, stride_kf_d,
    K_scales, stride_ks_t, stride_ks_d,
    V_bf16, stride_vt, stride_vd,
    Mid_O, stride_mo_s, stride_mo_h, stride_mo_d,
    Mid_LSE, stride_ml_s, stride_ml_h,
    qo_indptr, kv_indptr,
    BLOCK_H: tl.constexpr,
    NUM_SPLITS: tl.constexpr,
    sm_scale_log2,
    QK_DIM: tl.constexpr, V_DIM: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr, BLOCK_DV: tl.constexpr,
):
    """Process ALL query heads per program, sharing KV load."""
    # Grid: (batch_size * NUM_SPLITS,)
    pid = tl.program_id(0)
    split_id = pid % NUM_SPLITS
    batch_id = pid // NUM_SPLITS

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

    tokens_per_split = (kv_len + NUM_SPLITS - 1) // NUM_SPLITS
    split_start = split_id * tokens_per_split
    split_end = tl.minimum(split_start + tokens_per_split, kv_len)
    split_kv_len = split_end - split_start

    offs_h = tl.arange(0, BLOCK_H)
    offs_dv = tl.arange(0, BLOCK_DV)
    flat_idx = q_start * BLOCK_H + offs_h

    if split_kv_len <= 0:
        o_ptrs = Mid_O + split_id * stride_mo_s + flat_idx[:, None] * stride_mo_h + offs_dv[None, :] * stride_mo_d
        tl.store(o_ptrs, tl.zeros([BLOCK_H, BLOCK_DV], dtype=tl.bfloat16), mask=offs_dv[None, :] < V_DIM)
        lse_ptrs = Mid_LSE + split_id * stride_ml_s + flat_idx * stride_ml_h
        tl.store(lse_ptrs, tl.full([BLOCK_H], float('-inf'), dtype=tl.float32))
        return

    NUM_K_ITERS: tl.constexpr = (QK_DIM + BLOCK_K - 1) // BLOCK_K
    PACKED_BK: tl.constexpr = BLOCK_K // 2
    SCALE_BK: tl.constexpr = BLOCK_K // 32

    q_base = Q + q_start * stride_qt

    m_i = tl.full([BLOCK_H], float('-inf'), dtype=tl.float32)
    l_i = tl.full([BLOCK_H], 0.0, dtype=tl.float32)
    acc = tl.zeros([BLOCK_H, BLOCK_DV], dtype=tl.float32)

    offs_n = tl.arange(0, BLOCK_N)
    offs_bk = tl.arange(0, PACKED_BK)
    offs_sk = tl.arange(0, SCALE_BK)
    offs_qk = tl.arange(0, BLOCK_K)

    for block_start in range(0, split_kv_len, BLOCK_N):
        mask_n = (block_start + offs_n) < split_kv_len
        kv_pos = kv_start + split_start + block_start

        qk = tl.zeros([BLOCK_H, BLOCK_N], dtype=tl.float32)
        for k_iter in range(NUM_K_ITERS):
            k_offset = k_iter * BLOCK_K
            k_remaining = QK_DIM - k_offset

            # Q for ALL heads: [BLOCK_H, BLOCK_K]
            q_ptrs = q_base + offs_h[:, None] * stride_qh + (k_offset + offs_qk[None, :]) * stride_qd
            q_chunk = tl.load(q_ptrs, mask=offs_qk[None, :] < k_remaining, other=0.0)
            q_fp8 = q_chunk.to(tl.float8e4nv)
            q_scale = tl.full([BLOCK_H, SCALE_BK], 127, dtype=tl.uint8)

            # K (SHARED across all heads): [PACKED_BK, BLOCK_N]
            k_ptrs = K_fp4 + (kv_pos + offs_n[None, :]) * stride_kf_t + (k_offset // 2 + offs_bk[:, None]) * stride_kf_d
            k_mask = (offs_bk[:, None] < (k_remaining + 1) // 2) & mask_n[None, :]
            k_block = tl.load(k_ptrs, mask=k_mask, other=0)

            # K scales (SHARED): [BLOCK_N, SCALE_BK]
            ks_ptrs = K_scales + (kv_pos + offs_n[:, None]) * stride_ks_t + (k_offset // 32 + offs_sk[None, :]) * stride_ks_d
            ks_mask = mask_n[:, None] & (offs_sk[None, :] < (k_remaining + 31) // 32)
            k_scales_chunk = tl.load(ks_ptrs, mask=ks_mask, other=0)

            # [BLOCK_H, BLOCK_N] — all heads, shared K
            qk = tl.dot_scaled(q_fp8, q_scale, "e4m3", k_block, k_scales_chunk, "e2m1",
                               acc=qk, fast_math=True)

        qk = qk * sm_scale_log2
        qk = tl.where(mask_n[None, :], qk, float('-inf'))

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

        # V (SHARED across all heads): [BLOCK_N, V_DIM]
        v_ptrs = V_bf16 + (kv_pos + offs_n[:, None]) * stride_vt + offs_dv[None, :] * stride_vd
        v_mask = mask_n[:, None] & (offs_dv[None, :] < V_DIM)
        v_block = tl.load(v_ptrs, mask=v_mask, other=0.0)

        # OV: [BLOCK_H, BLOCK_N] × [BLOCK_N, BLOCK_DV] = [BLOCK_H, BLOCK_DV]
        acc += tl.dot(p.to(tl.bfloat16), v_block, out_dtype=tl.float32)

    partial_out = acc / l_i[:, None]
    o_ptrs = Mid_O + split_id * stride_mo_s + flat_idx[:, None] * stride_mo_h + offs_dv[None, :] * stride_mo_d
    tl.store(o_ptrs, partial_out.to(tl.bfloat16), mask=offs_dv[None, :] < V_DIM)

    lse_vals = m_i + tl.math.log2(l_i)
    lse_ptrs = Mid_LSE + split_id * stride_ml_s + flat_idx * stride_ml_h
    tl.store(lse_ptrs, lse_vals)


@triton.jit
def _mla_reduce(
    Mid_O, stride_mo_s, stride_mo_h, stride_mo_d,
    Mid_LSE, stride_ml_s, stride_ml_h,
    Out, stride_ot, stride_oh, stride_od,
    qo_indptr, num_q_heads: tl.constexpr,
    NUM_SPLITS: tl.constexpr, V_DIM: tl.constexpr, BLOCK_DV: tl.constexpr,
):
    pid = tl.program_id(0)
    batch_id = pid // num_q_heads
    head_id = pid % num_q_heads
    q_start = tl.load(qo_indptr + batch_id)
    flat_idx = q_start * num_q_heads + head_id
    offs_dv = tl.arange(0, BLOCK_DV)

    max_lse = tl.full([1], float('-inf'), dtype=tl.float32)
    for s in range(NUM_SPLITS):
        lse_s = tl.load(Mid_LSE + s * stride_ml_s + flat_idx * stride_ml_h)
        max_lse = tl.maximum(max_lse, lse_s)

    acc = tl.zeros([BLOCK_DV], dtype=tl.float32)
    sum_w = tl.zeros([1], dtype=tl.float32)
    for s in range(NUM_SPLITS):
        lse_s = tl.load(Mid_LSE + s * stride_ml_s + flat_idx * stride_ml_h)
        w = tl.math.exp2(lse_s - max_lse)
        sum_w += w
        o_ptrs = Mid_O + s * stride_mo_s + flat_idx * stride_mo_h + offs_dv * stride_mo_d
        partial = tl.load(o_ptrs, mask=offs_dv < V_DIM, other=0.0).to(tl.float32)
        acc += w * partial

    acc = acc / sum_w
    o_base = Out + q_start * stride_ot + head_id * stride_oh
    tl.store(o_base + offs_dv * stride_od, acc.to(tl.bfloat16), mask=offs_dv < V_DIM)


# ==================== METADATA CACHE ====================
_meta_cache = {}
_aiter_buffers = {}


def _get_cached_metadata(batch_size, q_seq_len, kv_seq_len, nq, nkv, dtype_q, dtype_kv,
                         num_kv_splits, device):
    key = (batch_size, q_seq_len, kv_seq_len, nq, nkv, dtype_q, dtype_kv, num_kv_splits)
    if key not in _meta_cache:
        total_kv_len = batch_size * kv_seq_len
        qo_indptr_c = torch.arange(0, batch_size + 1, dtype=torch.int32, device=device) * q_seq_len
        kv_indptr_c = torch.arange(0, batch_size + 1, dtype=torch.int32, device=device) * kv_seq_len
        kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)
        kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device=device)

        info = get_mla_metadata_info_v1(
            batch_size, q_seq_len, nq, dtype_q, dtype_kv,
            is_sparse=False, fast_mode=True,
            num_kv_splits=num_kv_splits, intra_batch_mode=True,
        )
        work = [torch.empty(s, dtype=t, device=device) for s, t in info]
        (wm, wi, wis, ri, rfm, rpm) = work
        get_mla_metadata_v1(
            qo_indptr_c, kv_indptr_c, kv_last_page_len,
            nq // nkv, nkv, True,
            wm, wis, wi, ri, rfm, rpm,
            page_size=PAGE_SIZE, kv_granularity=32,
            max_seqlen_qo=q_seq_len, uni_seqlen_qo=q_seq_len,
            fast_mode=True, max_split_per_batch=num_kv_splits,
            intra_batch_mode=True, dtype_q=dtype_q, dtype_kv=dtype_kv,
        )
        meta = {"work_meta_data": wm, "work_indptr": wi, "work_info_set": wis,
                "reduce_indptr": ri, "reduce_final_map": rfm, "reduce_partial_map": rpm}
        _meta_cache[key] = (meta, qo_indptr_c, kv_indptr_c, kv_indices, kv_last_page_len)
    return _meta_cache[key]


# ==================== AITER PATH ====================
def _aiter_path(q, kv_data, qo_indptr, kv_indptr, config):
    batch_size = config["batch_size"]
    nq = config["num_heads"]
    nkv = config["num_kv_heads"]
    dq = config["qk_head_dim"]
    dv = config["v_head_dim"]
    q_seq_len = config["q_seq_len"]
    kv_seq_len = config["kv_seq_len"]
    total_kv_tokens = batch_size * kv_seq_len
    use_fp8 = total_kv_tokens > 131072

    if use_fp8:
        q_input, q_scale = quantize_fp8(q)
        kv_buffer, kv_scale = kv_data["fp8"]
    else:
        q_input, q_scale = q, None
        kv_buffer = kv_data["bf16"]
        kv_scale = None

    kv_buffer_4d = kv_buffer.view(kv_buffer.shape[0], PAGE_SIZE, nkv, kv_buffer.shape[-1])

    if total_kv_tokens <= 4096:
        num_kv_splits = 16
    elif total_kv_tokens <= 32768:
        num_kv_splits = 16
    elif total_kv_tokens <= 65536:
        num_kv_splits = 16
    elif total_kv_tokens <= 262144:
        num_kv_splits = 32
    else:
        num_kv_splits = 64

    meta, qo_c, kv_c, kv_idx, kv_lpl = _get_cached_metadata(
        batch_size, q_seq_len, kv_seq_len, nq, nkv,
        q_input.dtype, kv_buffer.dtype, num_kv_splits, q.device,
    )

    buf_key = (batch_size, q_seq_len, nq, dv)
    if buf_key not in _aiter_buffers:
        _aiter_buffers[buf_key] = torch.empty(
            (q.shape[0], nq, dv), dtype=torch.bfloat16, device=q.device
        )
    o = _aiter_buffers[buf_key]

    mla_decode_fwd(
        q_input.view(-1, nq, dq), kv_buffer_4d, o,
        qo_c, kv_c, kv_idx, kv_lpl,
        q_seq_len, page_size=PAGE_SIZE, nhead_kv=nkv, 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,
        **meta,
    )
    return o


# ==================== CUSTOM MULTI-HEAD MXFP4 PATH ====================
def _custom_multihead_path(q, kv_data, qo_indptr, kv_indptr, config):
    batch_size = config["batch_size"]
    nq = config["num_heads"]
    dq = config["qk_head_dim"]
    dv = config["v_head_dim"]
    kv_seq_len = config["kv_seq_len"]

    kv_mxfp4_buf, kv_mxfp4_scale = kv_data["mxfp4"]
    total_kv = kv_mxfp4_buf.shape[0]
    total_q = q.shape[0]
    total_qh = total_q * nq

    k_fp4 = kv_mxfp4_buf.view(total_kv, -1).view(torch.uint8)
    num_scale_blocks = (dq + 31) // 32
    k_scales = kv_mxfp4_scale.view(torch.uint8)[:total_kv, :num_scale_blocks]
    kv_bf16 = kv_data["bf16"]
    kv_2d = kv_bf16.view(total_kv, -1)

    # Tune splits for CU occupancy: grid = batch_size * NUM_SPLITS
    if batch_size <= 4:
        NUM_SPLITS = max(4, 256 // batch_size)  # Target ~256 programs
    elif batch_size <= 32:
        NUM_SPLITS = max(4, 256 // batch_size)
    elif batch_size <= 128:
        NUM_SPLITS = 4
    else:
        NUM_SPLITS = 1

    # BLOCK_N=64 gives -18% on bs4/kv1k (better occupancy for small shapes)
    # vs BLOCK_N=128 which is better for large KV sequences
    if kv_seq_len <= 1024:
        BLOCK_N = 64
    else:
        BLOCK_N = 128
    max_useful_splits = max(1, kv_seq_len // BLOCK_N)
    NUM_SPLITS = min(NUM_SPLITS, max_useful_splits)

    BLOCK_K = 256

    mid_o = torch.empty((NUM_SPLITS, total_qh, dv), dtype=torch.bfloat16, device=q.device)
    mid_lse = torch.full((NUM_SPLITS, total_qh), float('-inf'), dtype=torch.float32, device=q.device)

    grid1 = (batch_size * NUM_SPLITS,)
    _mla_multihead_stage1[grid1](
        q, q.stride(0), q.stride(1), q.stride(2),
        k_fp4, k_fp4.stride(0), k_fp4.stride(1),
        k_scales, k_scales.stride(0), k_scales.stride(1),
        kv_2d, kv_2d.stride(0), kv_2d.stride(1),
        mid_o, mid_o.stride(0), mid_o.stride(1), mid_o.stride(2),
        mid_lse, mid_lse.stride(0), mid_lse.stride(1),
        qo_indptr, kv_indptr,
        BLOCK_H=nq, NUM_SPLITS=NUM_SPLITS,
        sm_scale_log2=SM_SCALE_LOG2,
        QK_DIM=dq, V_DIM=dv,
        BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K, BLOCK_DV=PADDED_V,
        num_warps=8, num_stages=2,
    )

    o = torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device=q.device)
    grid2 = (batch_size * nq,)
    _mla_reduce[grid2](
        mid_o, mid_o.stride(0), mid_o.stride(1), mid_o.stride(2),
        mid_lse, mid_lse.stride(0), mid_lse.stride(1),
        o, o.stride(0), o.stride(1), o.stride(2),
        qo_indptr, num_q_heads=nq,
        NUM_SPLITS=NUM_SPLITS, V_DIM=dv, BLOCK_DV=PADDED_V, num_warps=4,
    )
    return o


# ==================== DISPATCH ====================
CUSTOM_THRESHOLD = 8192  # Use custom for small shapes


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

    if total_kv_tokens <= CUSTOM_THRESHOLD:
        return _custom_multihead_path(q, kv_data, qo_indptr, kv_indptr, config)
    else:
        return _aiter_path(q, kv_data, qo_indptr, kv_indptr, config)
scrolls · 358 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