Skip to content
KernelIndex
Search⌘K

submission 587387

garrick99 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v4_dotscaled.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-587387?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
108.4µs
#476 of 766
2026-03-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:9c75d95dc4fce8f370c89eaf39643ede33492d02b3dc9a1c9fb92d4d85a82fbd
license declaredunknown
license concludedunknown
authorsgarrick99
imported2026-08-26

Techniques

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

fp4K scores use native fp4 MFMA via dot_scaled (eliminates 9-tile dequant overhead).
mmaacc_even += tl.dot(p_bf, ve).to(tl.float32)
online-softmaxm_new = tl.maximum(m_i, blk_max)

Kernel source

submission_v4_dotscaled.py396 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
MLA Decode v4 — tl.dot_scaled for K scores + manual dequant for V.

K scores use native fp4 MFMA via dot_scaled (eliminates 9-tile dequant overhead).
V accumulation uses even/odd manual dequant (packed axis is N, not K for V).
Falls back to aiter fp8 if dot_scaled compilation fails.
"""

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


# ============================================================================
# FP4 dequant helpers (for V accumulation only)
# ============================================================================

@triton.jit
def _fp4_to_f32(nibble):
    sign = (nibble >> 3) & 1
    abs3 = nibble & 7
    exp2 = abs3 >> 1
    man1 = abs3 & 1
    val = tl.where(
        exp2 > 0,
        tl.math.exp2((exp2 - 1).to(tl.float32)) * (1.0 + man1.to(tl.float32) * 0.5),
        man1.to(tl.float32) * 0.5,
    )
    return tl.where(sign > 0, -val, val)


@triton.jit
def _e8m0_to_f32(byte_val):
    return tl.math.exp2(byte_val.to(tl.float32) - 127.0)


# ============================================================================
# Stage 1: dot_scaled K scores + manual V dequant
# ============================================================================

@triton.jit
def _mla_v4_stage1(
    QT_ptr,          # (B, D_QK, H) bf16 — Q transposed
    QE_ptr, QO_ptr,  # (B, H, half_dk) bf16 — Q even/odd (unused if dot_scaled works for K)
    KV_ptr,          # (total_kv, packed_dim) uint8
    KS_ptr,          # (scale_rows, scale_cols) uint8
    POE_ptr, POO_ptr, PLSE_ptr,
    KVIDX_ptr,
    sm_scale,
    # QT strides
    sqt_b, sqt_k, sqt_h,
    # QE/QO strides (for V path's Q — not needed for K with dot_scaled)
    sqe_b, sqe_h, sqe_d,
    sqo_b, sqo_h, sqo_d,
    # KV packed strides
    skv_t, skv_d,
    # KV scale strides
    sks_t, sks_d,
    # POE strides
    spe_b, spe_s, spe_h, spe_d,
    # POO strides
    spo_b, spo_s, spo_h, spo_d,
    # PLSE strides
    sl_b, sl_s, sl_h,
    # constexpr
    HALF_DV: tl.constexpr,        # 256
    N_HEADS: tl.constexpr,        # 16
    NUM_SPLITS: tl.constexpr,
    BLOCK_KV: tl.constexpr,       # 64
    K_TILE_PACKED: tl.constexpr,  # 32
    K_TILE: tl.constexpr,         # 64
    N_K_TILES: tl.constexpr,      # 9
):
    bid = tl.program_id(0)
    sid = tl.program_id(1)

    kv_start = tl.load(KVIDX_ptr + bid)
    kv_end   = tl.load(KVIDX_ptr + bid + 1)
    kv_len   = kv_end - kv_start

    tps     = tl.cdiv(kv_len, NUM_SPLITS)
    s_start = sid * tps
    s_end   = tl.minimum(s_start + tps, kv_len)

    hr  = tl.arange(0, N_HEADS)       # 16
    tp  = tl.arange(0, K_TILE_PACKED)  # 32
    tk  = tl.arange(0, K_TILE)        # 64
    bk  = tl.arange(0, BLOCK_KV)      # 64
    dv  = tl.arange(0, HALF_DV)       # 256
    sc2 = tl.arange(0, 2)             # 2 scale blocks per K tile

    if s_start >= kv_len:
        tl.store(PLSE_ptr + bid * sl_b + sid * sl_s + hr * sl_h,
                 tl.full([N_HEADS], float('-inf'), tl.float32))
        return

    # Initialize
    m_i      = tl.full([N_HEADS], float('-inf'), tl.float32)
    l_i      = tl.zeros([N_HEADS], tl.float32)
    acc_even = tl.zeros([N_HEADS, HALF_DV], tl.float32)
    acc_odd  = tl.zeros([N_HEADS, HALF_DV], tl.float32)

    for blk_off in range(s_start, s_end, BLOCK_KV):
        blk_len = tl.minimum(BLOCK_KV, s_end - blk_off)
        vmask   = bk < blk_len
        tokens  = kv_start + blk_off + bk

        # ======== K scores via tl.dot_scaled (native fp4 MFMA) ========
        scores_t = tl.zeros([BLOCK_KV, N_HEADS], tl.float32)

        for kt in tl.static_range(0, N_K_TILES):
            tile_off_p = kt * K_TILE_PACKED
            tile_off_l = kt * K_TILE

            # K tile: (BK, 32_packed) fp4
            kp_off = tokens[:, None] * skv_t + (tile_off_p + tp)[None, :] * skv_d
            k_tile = tl.load(KV_ptr + kp_off, mask=vmask[:, None], other=0)

            # K scale: (BK, 2) E8M0
            ks_off = tokens[:, None] * sks_t + (kt * 2 + sc2)[None, :] * sks_d
            k_scale = tl.load(KS_ptr + ks_off, mask=vmask[:, None], other=0)

            # Q^T tile: (K_TILE=64, H=16) bf16
            qt_off = bid * sqt_b + (tile_off_l + tk)[:, None] * sqt_k + hr[None, :] * sqt_h
            qt_tile = tl.load(QT_ptr + qt_off)

            # dot_scaled: K_fp4(BK,32packed) @ Q^T_fp16(64,16) -> (BK,16)
            scores_t = tl.dot_scaled(
                lhs=k_tile,
                rhs=qt_tile.to(tl.float16),
                lhs_scale=k_scale,
                rhs_scale=None,
                lhs_format='e2m1',
                rhs_format='fp16',
                acc=scores_t,
            )

        # Transpose: (BK, H) -> (H, BK) and apply scale
        scores = tl.trans(scores_t) * sm_scale
        scores = tl.where(vmask[None, :], scores, float('-inf'))

        # ======== Online softmax ========
        blk_max = tl.max(scores, axis=1)
        m_new   = tl.maximum(m_i, blk_max)
        alpha   = tl.exp(m_i - m_new)
        p       = tl.exp(scores - m_new[:, None])
        l_i     = l_i * alpha + tl.sum(p, axis=1)
        acc_even = acc_even * alpha[:, None]
        acc_odd  = acc_odd  * alpha[:, None]
        m_i     = m_new

        # ======== V accumulation: manual fp4 dequant (even/odd) ========
        vp_off   = tokens[:, None] * skv_t + dv[None, :] * skv_d
        packed_v = tl.load(KV_ptr + vp_off, mask=vmask[:, None], other=0).to(tl.int32)

        v_lo = _fp4_to_f32(packed_v & 0x0F)
        v_hi = _fp4_to_f32((packed_v >> 4) & 0x0F)

        vs_off  = tokens[:, None] * sks_t + (dv[None, :] // 16) * sks_d
        v_scale = _e8m0_to_f32(
            tl.load(KS_ptr + vs_off, mask=vmask[:, None], other=0).to(tl.int32))

        ve = (v_lo * v_scale).to(tl.bfloat16)
        vo = (v_hi * v_scale).to(tl.bfloat16)

        p_bf = p.to(tl.bfloat16)
        acc_even += tl.dot(p_bf, ve).to(tl.float32)
        acc_odd  += tl.dot(p_bf, vo).to(tl.float32)

    # Store
    oe = acc_even / l_i[:, None]
    oo = acc_odd  / l_i[:, None]
    lse = m_i + tl.log(l_i)

    pe_base = bid * spe_b + sid * spe_s
    tl.store(POE_ptr + pe_base + hr[:, None] * spe_h + dv[None, :] * spe_d, oe)
    po_base = bid * spo_b + sid * spo_s
    tl.store(POO_ptr + po_base + hr[:, None] * spo_h + dv[None, :] * spo_d, oo)
    tl.store(PLSE_ptr + bid * sl_b + sid * sl_s + hr * sl_h, lse)


# ============================================================================
# Stage 2: Reduce + interleave (same as v3)
# ============================================================================

@triton.jit
def _mla_reduce(
    POE_ptr, POO_ptr, PLSE_ptr, O_ptr,
    spe_b, spe_s, spe_h, spe_d,
    spo_b, spo_s, spo_h, spo_d,
    sl_b, sl_s, sl_h,
    so_b, so_h, so_d,
    NUM_SPLITS: tl.constexpr,
    HALF_V: tl.constexpr,
):
    bid = tl.program_id(0)
    hid = tl.program_id(1)
    dv  = tl.arange(0, HALF_V)

    max_lse = tl.full([], float('-inf'), tl.float32)
    for s in tl.static_range(0, NUM_SPLITS):
        lse = tl.load(PLSE_ptr + bid * sl_b + s * sl_s + hid * sl_h)
        max_lse = tl.maximum(max_lse, lse)

    acc_e = tl.zeros([HALF_V], tl.float32)
    acc_o = tl.zeros([HALF_V], tl.float32)
    sum_w = tl.full([], 0.0, tl.float32)

    for s in tl.static_range(0, NUM_SPLITS):
        lse = tl.load(PLSE_ptr + bid * sl_b + s * sl_s + hid * sl_h)
        w   = tl.exp(lse - max_lse)
        sum_w += w
        acc_e += w * tl.load(POE_ptr + bid * spe_b + s * spe_s + hid * spe_h + dv * spe_d)
        acc_o += w * tl.load(POO_ptr + bid * spo_b + s * spo_s + hid * spo_h + dv * spo_d)

    acc_e = (acc_e / sum_w).to(tl.bfloat16)
    acc_o = (acc_o / sum_w).to(tl.bfloat16)

    o_base = bid * so_b + hid * so_h
    tl.store(O_ptr + o_base + (dv * 2)     * so_d, acc_e)
    tl.store(O_ptr + o_base + (dv * 2 + 1) * so_d, acc_o)


# ============================================================================
# Aiter FP8 fallback
# ============================================================================

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

FP8_DTYPE = aiter_dtypes.fp8
_SM_SCALE = 1.0 / (576 ** 0.5)
_aiter_meta_cache = {}
_aiter_kvidx_cache = {}


def _quantize_fp8(tensor):
    finfo = torch.finfo(FP8_DTYPE)
    amax = tensor.abs().amax().clamp(min=1e-12)
    scale = amax / finfo.max
    fp8_t = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
    return fp8_t, scale.to(torch.float32).reshape(1)


def _aiter_fp8_path(q, kv_data, qo_indptr, kv_indptr, config):
    B = config["batch_size"]; nq = config["num_heads"]; nkv = config["num_kv_heads"]
    dq = config["qk_head_dim"]; dv = config["v_head_dim"]; qsl = config["q_seq_len"]
    nks = 32
    q_fp8, q_scale = _quantize_fp8(q)
    kv_fp8, kv_scale = kv_data["fp8"]
    total_kv = int(kv_indptr[-1].item())
    if total_kv not in _aiter_kvidx_cache:
        _aiter_kvidx_cache[total_kv] = torch.arange(total_kv, dtype=torch.int32, device="cuda")
    kv_indices = _aiter_kvidx_cache[total_kv]
    kv_4d = kv_fp8.view(kv_fp8.shape[0], 1, nkv, kv_fp8.shape[-1])
    kv_last = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
    meta_key = (B, config["kv_seq_len"])
    if meta_key not in _aiter_meta_cache:
        info = get_mla_metadata_info_v1(B, qsl, nq, q_fp8.dtype, kv_fp8.dtype,
            is_sparse=False, fast_mode=False, num_kv_splits=nks, intra_batch_mode=True)
        work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
        wm, wi, wis, ri, rfm, rpm = work
        get_mla_metadata_v1(qo_indptr, kv_indptr, kv_last, nq // nkv, nkv, True,
            wm, wis, wi, ri, rfm, rpm, page_size=1, kv_granularity=16,
            max_seqlen_qo=qsl, uni_seqlen_qo=qsl, fast_mode=False,
            max_split_per_batch=nks, intra_batch_mode=True,
            dtype_q=q_fp8.dtype, dtype_kv=kv_fp8.dtype)
        _aiter_meta_cache[meta_key] = (wm, wi, wis, ri, rfm, rpm)
    wm, wi, wis, ri, rfm, rpm = _aiter_meta_cache[meta_key]
    o = torch.empty((q.shape[0], nq, dv), dtype=torch.bfloat16, device="cuda")
    mla_decode_fwd(q_fp8.view(-1, nq, dq), kv_4d, o, qo_indptr, kv_indptr, kv_indices,
        kv_last, qsl, page_size=1, nhead_kv=nkv, sm_scale=_SM_SCALE, logit_cap=0.0,
        num_kv_splits=nks, 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


# ============================================================================
# dot_scaled Triton path
# ============================================================================

def _dotscaled_triton_path(q, kv_data, qo_indptr, kv_indptr, config):
    B       = config["batch_size"]
    H       = config["num_heads"]
    D_QK    = config["qk_head_dim"]
    D_V     = config["v_head_dim"]
    kv_len  = config["kv_seq_len"]
    sm_scale = config["sm_scale"]

    kv_buf, kv_scl = kv_data["mxfp4"]
    total_kv = kv_buf.shape[0]
    kv_flat = kv_buf.reshape(total_kv, -1)
    if kv_flat.dtype != torch.uint8:
        kv_flat = kv_flat.view(torch.uint8)
    kv_sc = kv_scl
    if kv_sc.dtype != torch.uint8:
        kv_sc = kv_sc.view(torch.uint8)

    # Q transposed for dot_scaled K scores
    q_t = q.transpose(-2, -1).contiguous()  # (B, D_QK, H)
    # Q even/odd for manual V dequant
    q_even = q[:, :, 0::2].contiguous()
    q_odd  = q[:, :, 1::2].contiguous()

    BLOCK_KV = 64
    K_TILE_PACKED = 32
    K_TILE = 64
    half_dk = D_QK // 2  # 288
    half_dv = D_V // 2   # 256
    n_k_tiles = half_dk // K_TILE_PACKED  # 9

    max_sp  = max(1, kv_len // BLOCK_KV)
    desired = max(1, 608 // B)
    ns      = min(max_sp, desired, 64)
    if ns > 1:
        ns = 2 ** int(math.log2(ns))

    po_e = torch.empty((B, ns, H, half_dv), dtype=torch.float32, device='cuda')
    po_o = torch.empty((B, ns, H, half_dv), dtype=torch.float32, device='cuda')
    plse = torch.empty((B, ns, H),          dtype=torch.float32, device='cuda')

    _mla_v4_stage1[(B, ns)](
        q_t, q_even, q_odd,
        kv_flat, kv_sc,
        po_e, po_o, plse,
        kv_indptr,
        sm_scale,
        q_t.stride(0), q_t.stride(1), q_t.stride(2),
        q_even.stride(0), q_even.stride(1), q_even.stride(2),
        q_odd.stride(0), q_odd.stride(1), q_odd.stride(2),
        kv_flat.stride(0), kv_flat.stride(1),
        kv_sc.stride(0), kv_sc.stride(1),
        po_e.stride(0), po_e.stride(1), po_e.stride(2), po_e.stride(3),
        po_o.stride(0), po_o.stride(1), po_o.stride(2), po_o.stride(3),
        plse.stride(0), plse.stride(1), plse.stride(2),
        HALF_DV=half_dv,
        N_HEADS=H,
        NUM_SPLITS=ns,
        BLOCK_KV=BLOCK_KV,
        K_TILE_PACKED=K_TILE_PACKED,
        K_TILE=K_TILE,
        N_K_TILES=n_k_tiles,
    )

    out = torch.empty((B, H, D_V), dtype=torch.bfloat16, device='cuda')
    _mla_reduce[(B, H)](
        po_e, po_o, plse, out,
        po_e.stride(0), po_e.stride(1), po_e.stride(2), po_e.stride(3),
        po_o.stride(0), po_o.stride(1), po_o.stride(2), po_o.stride(3),
        plse.stride(0), plse.stride(1), plse.stride(2),
        out.stride(0), out.stride(1), out.stride(2),
        NUM_SPLITS=ns,
        HALF_V=half_dv,
    )
    return out


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

_use_dotscaled = None

def custom_kernel(data: input_t) -> output_t:
    global _use_dotscaled
    q, kv_data, qo_indptr, kv_indptr, config = data
    B      = config["batch_size"]
    kv_len = config["kv_seq_len"]
    total  = B * kv_len

    if _use_dotscaled is None:
        try:
            result = _dotscaled_triton_path(q, kv_data, qo_indptr, kv_indptr, config)
            _use_dotscaled = True
            return result
        except Exception as e:
            import sys, traceback
            print(f"dot_scaled FAILED: {type(e).__name__}: {e}", file=sys.stderr)
            traceback.print_exc(file=sys.stderr)
            _use_dotscaled = False
            return _aiter_fp8_path(q, kv_data, qo_indptr, kv_indptr, config)

    if _use_dotscaled:
        if total <= 65536:
            return _dotscaled_triton_path(q, kv_data, qo_indptr, kv_indptr, config)
        else:
            return _aiter_fp8_path(q, kv_data, qo_indptr, kv_indptr, config)
    else:
        return _aiter_fp8_path(q, kv_data, qo_indptr, kv_indptr, config)
scrolls · 396 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