Skip to content
KernelIndex
Search⌘K

submission 745751

farhan-navas · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e4118cf31163c5401986acf571a05065b9019d9a8e752213e9f79b883e70eb1b
license declaredunknown
license concludedunknown
authorsfarhan-navas
imported2026-08-15

Techniques

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

fp8q_nope = q_nope_raw.to(tl.float8e4nv) # [H, 512] fp8
mmas = tl.dot(q_nope, tl.trans(k_nope))
num-warps = 4num_warps=4,
online-softmaxm_new = tl.maximum(m_i, tl.max(s_scaled, axis=1))
stages = 2num_stages=2,

Kernel source

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

"""
MLA Decode Submission — AMD MI355X (gfx950)

Hybrid: AITER (a16w8/a8w8) + Triton MLA decode kernel.
Based on SGLang's production Triton MLA decode patterns.
"""

from task import input_t, output_t
import os
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")  # -2-3μs per kernel launch
import torch
import triton
import triton.language as tl


# Eval server: Triton 3.6.0, GPUTarget(backend='hip', arch='gfx950', warp_size=64)

PAGE_SIZE = 1
NUM_KV_SPLITS = 32
SM_SCALE = 1.0 / (576 ** 0.5)

_cache = {}


# ============================================================
# AITER path (existing optimized baseline)
# ============================================================

def _quantize_fp8(tensor):
    from aiter import dtypes as aiter_dtypes
    from aiter.ops.quant import per_tensor_quant_hip
    return per_tensor_quant_hip(tensor, scale=None, quant_dtype=aiter_dtypes.fp8)


_metadata_done = set()

def _get_cached_buffers(batch_size, nq, nkv, dv, total_q, total_kv_len, q_dtype, kv_dtype, kv_splits=32):
    from aiter import get_mla_metadata_info_v1
    key = (batch_size, nq, total_q, total_kv_len, q_dtype, kv_splits)
    if key not in _cache:
        max_q_len = 1
        info = get_mla_metadata_info_v1(
            batch_size, max_q_len, nq, q_dtype, kv_dtype,
            is_sparse=False, fast_mode=False,
            num_kv_splits=kv_splits, intra_batch_mode=True,
        )
        work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
        o = torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda")
        kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
        _cache[key] = (work, o, kv_indices)
    return _cache[key]


_aiter_mla = None
_aiter_meta = None

def _aiter_path(data, use_bf16_q=False, num_splits=None):
    global _aiter_mla, _aiter_meta
    if _aiter_mla is None:
        from aiter.mla import mla_decode_fwd
        from aiter import get_mla_metadata_v1
        _aiter_mla = mla_decode_fwd
        _aiter_meta = get_mla_metadata_v1

    q, kv_data, qo_indptr, kv_indptr, config = data

    batch_size = config["batch_size"]
    nq = config["num_heads"]
    nkv = config["num_kv_heads"]
    dq = config["qk_head_dim"]
    dv = config["v_head_dim"]
    max_q_len = 1

    if use_bf16_q:
        q_input = q.view(-1, nq, dq)
        q_scale = None
    else:
        q_fp8, q_scale = _quantize_fp8(q)
        q_input = q_fp8.view(-1, nq, dq)

    kv_buffer_fp8, kv_scale = kv_data["fp8"]

    # Avoid .item() CPU-GPU sync — compute from config
    total_kv_len = batch_size * config["kv_seq_len"]
    total_q = q.shape[0]

    # Cache kv_last_page_len (same for same shape)
    kv_lpl_key = ("kv_lpl", batch_size, config["kv_seq_len"])
    if kv_lpl_key not in _cache:
        _cache[kv_lpl_key] = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
    kv_last_page_len = _cache[kv_lpl_key]
    kv_splits = num_splits if num_splits is not None else NUM_KV_SPLITS
    work, o, kv_indices = _get_cached_buffers(
        batch_size, nq, nkv, dv, total_q, total_kv_len,
        q_input.dtype, kv_buffer_fp8.dtype, kv_splits=kv_splits,
    )
    (work_metadata, work_indptr, work_info_set,
     reduce_indptr, reduce_final_map, reduce_partial_map) = work

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

    meta_key = (batch_size, nq, total_q, total_kv_len, q_input.dtype, kv_splits)
    if meta_key not in _metadata_done:
        _aiter_meta(
            qo_indptr, kv_indptr, kv_last_page_len,
            nq // nkv, nkv, True,
            work_metadata, work_info_set, work_indptr,
            reduce_indptr, reduce_final_map, reduce_partial_map,
            page_size=PAGE_SIZE,
            kv_granularity=max(PAGE_SIZE, 16),
            max_seqlen_qo=max_q_len,
            uni_seqlen_qo=max_q_len,
            fast_mode=False,
            max_split_per_batch=kv_splits,
            intra_batch_mode=True,
            dtype_q=q_input.dtype,
            dtype_kv=kv_buffer_fp8.dtype,
        )
        _metadata_done.add(meta_key)

    _aiter_mla(
        q_input,
        kv_buffer_4d,
        o,
        qo_indptr,
        kv_indptr,
        kv_indices,
        kv_last_page_len,
        max_q_len,
        page_size=PAGE_SIZE,
        nhead_kv=nkv,
        sm_scale=SM_SCALE,
        logit_cap=0.0,
        num_kv_splits=kv_splits,
        q_scale=q_scale,
        kv_scale=kv_scale,
        intra_batch_mode=True,
        work_meta_data=work_metadata,
        work_indptr=work_indptr,
        work_info_set=work_info_set,
        reduce_indptr=reduce_indptr,
        reduce_final_map=reduce_final_map,
        reduce_partial_map=reduce_partial_map,
    )
    return o


# ============================================================
# Triton MLA Decode — Stage 1 (split-K attention)
# ============================================================
# Approach 1: Native fp8×fp8 tl.dot — no cast overhead
# Q cast to fp8 once before tile loop, KV loaded as raw fp8
# tl.dot(fp8, fp8) → native MFMA fp8 on gfx950
# bf16 fallback when USE_FP8_DOT=False

@triton.jit
def _mla_stage1(
    Q,              # [batch_size, num_heads, 576] bf16 — contiguous
    KV,             # [total_kv, 576] fp8 or bf16
    kv_scale,       # scalar float (1.0 for bf16, actual scale for fp8)
    kv_indptr,      # [batch_size + 1] int32
    partial_out,    # [batch_size * num_kv_splits * num_heads * DV] f32
    partial_lse,    # [batch_size * num_kv_splits * num_heads] f32
    stride_qb,      # Q stride: batch
    stride_qh,      # Q stride: head
    stride_kvt,     # KV stride: token
    NUM_KV_SPLITS: tl.constexpr,
    NUM_HEADS: tl.constexpr,
    SM_SCALE: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_DMODEL: tl.constexpr,   # 512
    BLOCK_DPE: tl.constexpr,      # 64
    BLOCK_DV: tl.constexpr,       # 512
    USE_FP8_DOT: tl.constexpr,    # True = native fp8 dot, False = bf16 dot
):
    batch_id = tl.program_id(0)
    split_id = tl.program_id(1)

    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_KV_SPLITS)
    my_start = kv_start + split_id * split_size
    my_end = tl.minimum(my_start + split_size, kv_end)
    my_len = my_end - my_start

    partial_base = batch_id * NUM_KV_SPLITS * NUM_HEADS * BLOCK_DV + split_id * NUM_HEADS * BLOCK_DV
    lse_base = batch_id * NUM_KV_SPLITS * NUM_HEADS + split_id * NUM_HEADS

    offs_h = tl.arange(0, NUM_HEADS)

    if my_len <= 0:
        tl.store(partial_lse + lse_base + offs_h, tl.full([NUM_HEADS], float("-inf"), dtype=tl.float32))
        return

    # Load Q and cast once — fp8 for native dot, bf16 for fallback
    offs_d = tl.arange(0, BLOCK_DMODEL)
    offs_pe = tl.arange(0, BLOCK_DPE)

    q_base = batch_id * stride_qb
    q_nope_raw = tl.load(Q + q_base + offs_h[:, None] * stride_qh + offs_d[None, :])
    q_pe_raw = tl.load(Q + q_base + offs_h[:, None] * stride_qh + (BLOCK_DMODEL + offs_pe[None, :]))

    if USE_FP8_DOT:
        # Cast Q bf16 → fp8 ONCE (amortized over all KV tiles)
        q_nope = q_nope_raw.to(tl.float8e4nv)   # [H, 512] fp8
        q_pe = q_pe_raw.to(tl.float8e4nv)        # [H, 64] fp8
    else:
        q_nope = q_nope_raw                       # [H, 512] bf16
        q_pe = q_pe_raw                           # [H, 64] bf16

    # Online softmax state
    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, BLOCK_DV], dtype=tl.float32)

    # Absorb kv_scale into QK prescale
    scale = SM_SCALE * kv_scale * 1.4426950408889634

    offs_n = tl.arange(0, BLOCK_N)

    for start in range(0, my_len, BLOCK_N):
        n_valid = tl.minimum(BLOCK_N, my_len - start)
        kv_idx = my_start + start + offs_n
        mask_n = offs_n < n_valid

        # Load K_nope and K_pe — no cast for fp8 dot, bf16 cast for fallback
        k_ptrs_nope = KV + kv_idx[:, None] * stride_kvt + offs_d[None, :]
        k_ptrs_pe = KV + kv_idx[:, None] * stride_kvt + (BLOCK_DMODEL + offs_pe[None, :])

        if USE_FP8_DOT:
            k_nope = tl.load(k_ptrs_nope, mask=mask_n[:, None], other=0.0)  # raw fp8
            k_pe = tl.load(k_ptrs_pe, mask=mask_n[:, None], other=0.0)      # raw fp8
        else:
            k_nope = tl.load(k_ptrs_nope, mask=mask_n[:, None], other=0.0).to(tl.bfloat16)
            k_pe = tl.load(k_ptrs_pe, mask=mask_n[:, None], other=0.0).to(tl.bfloat16)

        # QK: native fp8×fp8 dot or bf16 dot → FP32 accumulator
        s = tl.dot(q_nope, tl.trans(k_nope))
        s += tl.dot(q_pe, tl.trans(k_pe))

        # Mask + scale
        s = tl.where(mask_n[None, :], s, float("-inf"))
        s_scaled = s * scale

        # Online softmax
        m_new = tl.maximum(m_i, tl.max(s_scaled, axis=1))
        alpha = tl.math.exp2(m_i - m_new)
        p = tl.math.exp2(s_scaled - m_new[:, None])
        l_i = l_i * alpha + tl.sum(p, axis=1)
        acc = acc * alpha[:, None]

        # Load V — no cast for fp8 dot, bf16 cast for fallback
        offs_v = tl.arange(0, BLOCK_DV)
        v_ptrs = KV + kv_idx[:, None] * stride_kvt + offs_v[None, :]

        if USE_FP8_DOT:
            v = tl.load(v_ptrs, mask=mask_n[:, None], other=0.0)  # raw fp8
            # Cast P to fp8 for native fp8×fp8 PV dot
            acc = tl.dot(p.to(tl.float8e4nv), v, acc)
        else:
            v = tl.load(v_ptrs, mask=mask_n[:, None], other=0.0).to(tl.bfloat16)
            acc += tl.dot(p.to(tl.bfloat16), v)

        m_i = m_new

    # Normalize and apply deferred V scale
    acc = acc * (kv_scale / l_i[:, None])

    # Store LSE in log2 domain: lse = log2(l_i) + m_i
    lse = tl.math.log2(l_i) + m_i
    tl.store(partial_lse + lse_base + offs_h, lse)

    # Store normalized partial output [NUM_HEADS, BLOCK_DV] as 2D block
    offs_v = tl.arange(0, BLOCK_DV)
    out_offs = offs_h[:, None] * BLOCK_DV + offs_v[None, :]
    tl.store(partial_out + partial_base + out_offs, acc)


# ============================================================
# Triton MLA Decode — Stage 2 (reduction)
# ============================================================

@triton.jit
def _mla_stage2(
    partial_out,    # [batch_size * num_kv_splits * num_heads * DV] f32
    partial_lse,    # [batch_size * num_kv_splits * num_heads] f32
    output,         # [batch_size, num_heads, DV] bf16
    stride_ob,      # output stride batch
    stride_oh,      # output stride head
    NUM_KV_SPLITS: tl.constexpr,
    NUM_HEADS: tl.constexpr,
    BLOCK_DV: tl.constexpr,
):
    batch_id = tl.program_id(0)
    head_id = tl.program_id(1)

    # Load all LSEs for this batch x head, find global max
    lse_base = batch_id * NUM_KV_SPLITS * NUM_HEADS + head_id
    out_base = batch_id * NUM_KV_SPLITS * NUM_HEADS * BLOCK_DV + head_id * BLOCK_DV

    offs_v = tl.arange(0, BLOCK_DV)
    acc = tl.zeros([BLOCK_DV], dtype=tl.float32)
    weight_sum = tl.zeros([1], dtype=tl.float32)

    # First pass: find global max LSE
    m_global = tl.full([1], float("-inf"), dtype=tl.float32)
    for s in range(NUM_KV_SPLITS):
        lse_s = tl.load(partial_lse + lse_base + s * NUM_HEADS)
        m_global = tl.maximum(m_global, lse_s)

    # Second pass: weighted sum with rescaling
    for s in range(NUM_KV_SPLITS):
        lse_s = tl.load(partial_lse + lse_base + s * NUM_HEADS)
        w = tl.math.exp2(lse_s - m_global)
        partial = tl.load(partial_out + out_base + s * NUM_HEADS * BLOCK_DV + offs_v)
        acc += w * partial
        weight_sum += w

    acc = acc / weight_sum

    # Store
    out_ptr = output + batch_id * stride_ob + head_id * stride_oh + offs_v
    tl.store(out_ptr, acc.to(tl.bfloat16))


# ============================================================
# Triton path wrapper
# ============================================================

_triton_cache = {}


def _triton_path(data, use_fp8=True, block_n=None, splits=None, stages=None, s1_waves=None, s2_waves=None, warps=4):
    q, kv_data, qo_indptr, kv_indptr, config = data

    batch_size = config["batch_size"]
    num_heads = config["num_heads"]
    dv = config["v_head_dim"]
    total_q = q.shape[0]

    # fp8 KV: 2x less bandwidth than bf16 (576 vs 1152 bytes/token)
    if use_fp8:
        kv_tensor, kv_scale_val = kv_data["fp8"]  # [total_kv, 1, 576] fp8 + scalar
        kv_scale_f = float(kv_scale_val)
    else:
        kv_tensor = kv_data["bf16"]  # [total_kv, 1, 576] bf16
        kv_scale_f = 1.0

    # Reshape KV: [total_kv, 1, 576] → [total_kv, 576]
    kv_flat = kv_tensor.view(-1, 576)

    # Q is [total_q, num_heads, 576] bf16, already contiguous
    q_3d = q.view(total_q, num_heads, 576)

    BLOCK_N = block_n if block_n is not None else (32 if use_fp8 else 16)
    BLOCK_DMODEL = 512
    BLOCK_DPE = 64
    BLOCK_DV = 512
    num_kv_splits = splits if splits is not None else NUM_KV_SPLITS

    # Allocate/reuse partial buffers
    cache_key = ("triton", batch_size, num_heads, num_kv_splits, BLOCK_N)
    if cache_key not in _triton_cache:
        partial_out = torch.empty(
            batch_size * num_kv_splits * num_heads * BLOCK_DV,
            dtype=torch.float32, device="cuda"
        )
        partial_lse = torch.empty(
            batch_size * num_kv_splits * num_heads,
            dtype=torch.float32, device="cuda"
        )
        out = torch.empty(
            (total_q, num_heads, BLOCK_DV),
            dtype=torch.bfloat16, device="cuda"
        )
        _triton_cache[cache_key] = (partial_out, partial_lse, out)

    partial_out, partial_lse, out = _triton_cache[cache_key]

    if out.shape[0] != total_q:
        out = torch.empty((total_q, num_heads, BLOCK_DV), dtype=torch.bfloat16, device="cuda")
        _triton_cache[cache_key] = (partial_out, partial_lse, out)

    # Stage 1: split-K attention
    grid1 = (batch_size, num_kv_splits)
    s1_kwargs = dict(
        NUM_KV_SPLITS=num_kv_splits,
        NUM_HEADS=num_heads,
        SM_SCALE=SM_SCALE,
        BLOCK_N=BLOCK_N,
        BLOCK_DMODEL=BLOCK_DMODEL,
        BLOCK_DPE=BLOCK_DPE,
        BLOCK_DV=BLOCK_DV,
        USE_FP8_DOT=use_fp8,
        num_warps=warps,
        num_stages=stages if stages is not None else 1,
    )
    if s1_waves is not None:
        s1_kwargs["waves_per_eu"] = s1_waves
    _mla_stage1[grid1](
        q_3d, kv_flat, kv_scale_f, kv_indptr,
        partial_out, partial_lse,
        q_3d.stride(0), q_3d.stride(1),
        kv_flat.stride(0),
        **s1_kwargs,
    )

    # Stage 2: reduction
    grid2 = (batch_size, num_heads)
    s2_kwargs = dict(
        NUM_KV_SPLITS=num_kv_splits,
        NUM_HEADS=num_heads,
        BLOCK_DV=BLOCK_DV,
        num_warps=4,
        num_stages=2,
    )
    if s2_waves is not None:
        s2_kwargs["waves_per_eu"] = s2_waves
    _mla_stage2[grid2](
        partial_out, partial_lse, out,
        out.stride(0), out.stride(1),
        **s2_kwargs,
    )

    return out


# ============================================================
# Entry point
# ============================================================

# EVOLVE-BLOCK-START mla_kernel
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"]
    total_kv = batch_size * kv_seq_len

    if batch_size >= 256:
        # AITER for BS=256 — per-shape splits
        if kv_seq_len >= 8192:
            return _aiter_path(data, use_bf16_q=False, num_splits=64)
        else:
            return _aiter_path(data, use_bf16_q=True, num_splits=16)
    elif batch_size >= 64:
        # AITER for BS≥64 — per-shape splits
        if kv_seq_len >= 8192:
            return _aiter_path(data, use_bf16_q=False, num_splits=32)
        else:
            return _aiter_path(data, use_bf16_q=True, num_splits=16)
    elif batch_size >= 32 and kv_seq_len >= 8192:
        # AITER a8w8 for BS=32 KV=8192
        return _aiter_path(data, use_bf16_q=False, num_splits=16)
    elif batch_size <= 4:
        # Triton bf16 N=32, warps=8, stages=2, waves_per_eu=1
        return _triton_path(data, use_fp8=False, block_n=32,
                            splits=16 if kv_seq_len <= 1024 else 32,
                            stages=2, s1_waves=1, s2_waves=4, warps=8)
    else:
        # BS=32 KV=1024: Triton bf16 N=32, warps=8, splits=8
        return _triton_path(data, use_fp8=False, block_n=32,
                            splits=8, stages=2, s2_waves=4, warps=8)
# EVOLVE-BLOCK-END mla_kernel
scrolls · 470 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