Skip to content
KernelIndex
Search⌘K

submission 668813

jiajia931 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v0092a.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-668813?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
40.4µs
#116 of 766
2026-03-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3677aa4e83aa0c9793dca7c69d1cd55f4ad80de05dbffcd73ec5f8ecf45aab0a
license declaredunknown
license concludedunknown
authorsjiajia931
imported2026-08-15

Techniques

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

fp4v0092a: three-way hybrid a8w8 + a16w8 + mxfp4.
mmaacc_e0 += tl.dot(p_f16, vlo0_s, out_dtype=tl.float32)
num-warps = 4num_warps=4,
online-softmaxm_new = tl.maximum(m_i, qk_max)
persistent-kernelNone, # num_kv_splits_indptr = None -> persistent mode
stages = 1num_stages=1,
tile-n = 32BLOCK_N = 32

Kernel source

submission_v0092a.py838 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
# submission_version: v0092a

"""
v0092a: three-way hybrid a8w8 + a16w8 + mxfp4.

Based on v0092 mxfp4 kernel (KV reuse, log2 domain, scale register gather).
Strategy per shape:
  - bs=4: a8w8 (tiny Q quant cost, fp8 Q saves read bandwidth)
  - bs>=32 kv=1024: a16w8 (skip Q quant, dominates on short sequences)
  - bs>=32 kv=8192: mxfp4 (hardware dot_scaled, 2x KV bandwidth savings)
"""

import os
os.environ.setdefault("TRITON_CACHE_DIR", "/tmp/triton_cache_v0092a")

from typing import Any

import torch
import triton
import triton.language as tl

import aiter as aiter_module
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.mla import mla_decode_fwd

try:
    from aiter.utility.fp4_utils import dynamic_mxfp4_quant
    HAS_MXFP4 = True
except ImportError:
    HAS_MXFP4 = False

try:
    from aiter.ops.quant import dynamic_per_tensor_quant
    HAS_FUSED_QUANT = True
except ImportError:
    HAS_FUSED_QUANT = False

input_t = Any
output_t = Any

NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM
V_HEAD_DIM = KV_LORA_RANK
HALF_V_DIM = V_HEAD_DIM // 2
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
LOG2E = 1.4426950408889634
SM_SCALE_LOG2E = SM_SCALE * LOG2E
FP8_DTYPE = aiter_dtypes.fp8

PAGE_SIZE_TABLE = {
    (4, 1024): 1,
    (32, 1024): 1,
    (64, 1024): 2,
    (256, 1024): 2,
    (4, 8192): 8,
    (32, 8192): 8,
    (64, 8192): 8,
    (256, 8192): 8,
}
NUM_KV_SPLITS_TABLE = {
    (4, 1024): 16,
    (4, 8192): 32,
    (32, 1024): 8,
    (32, 8192): 8,
    (64, 1024): 4,
    (64, 8192): 8,
    (256, 1024): 1,
    (256, 8192): 16,
}
MXFP4_SPLITS_TABLE = {
    (32, 8192): 8,   # 32 * 8 = 256 stage1 programs
    (64, 8192): 4,   # 64 * 4 = 256 stage1 programs
    (256, 8192): 2,  # 256 * 2 = 512 stage1 programs; lower stage2 traffic than 4-way split
}
SHAPE_STRATEGY = {
    (4, 1024): "a8w8",   (4, 8192): "a8w8",    # small batch: fp8 Q saves bandwidth
    (32, 1024): "a16w8", (32, 8192): "mxfp4",   # large batch short kv: skip Q quant
    (64, 1024): "a16w8", (64, 8192): "mxfp4",
    (256, 1024): "a16w8", (256, 8192): "mxfp4",
}
DEFAULT_STRATEGY = "a16w8"


def _get_page_size(batch_size: int, kv_seq_len: int) -> int:
    return PAGE_SIZE_TABLE.get((batch_size, kv_seq_len), 2)


def _get_num_kv_splits(batch_size: int, kv_seq_len: int) -> int:
    return NUM_KV_SPLITS_TABLE.get((batch_size, kv_seq_len), 16)


def _get_mxfp4_splits(batch_size: int, kv_seq_len: int) -> int:
    return MXFP4_SPLITS_TABLE.get((batch_size, kv_seq_len), 1)


_E2M1_LUT = None


def _get_e2m1_lut() -> torch.Tensor:
    global _E2M1_LUT
    if _E2M1_LUT is None:
        _E2M1_LUT = torch.tensor(
            [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0,
             0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0],
            dtype=torch.float16,
            device="cuda",
        )
    return _E2M1_LUT


# =====================================================================================
# MXFP4 qh16 kernel: one program handles all 16 query heads for one (request, split)
# =====================================================================================

@triton.jit
def _mla_mxfp4_qh16_stage1_reuse(
    Q_packed,      # [batch, H, 288] uint8
    Q_scale,       # [batch*H, 18] uint8 (viewed as byte rows; actual row stride >= 18)
    KV_packed,     # [total_kv, 288] uint8
    KV_scale,      # [total_kv, 18] uint8
    E2M1_LUT,      # [16] fp16
    Mid_O,         # [batch*H, S, 512] fp32, layout = [even256 | odd256]
    Mid_LSE,       # [batch*H, S] fp32, log2-domain lse
    kv_indptr,     # [batch+1] int32
    sm_scale_log2e,
    stride_qp_b: tl.constexpr,
    stride_qp_h: tl.constexpr,
    stride_qs,
    stride_kv,
    stride_ks,
    stride_mid_bh,
    stride_mid_s,
    stride_lse_bh,
    stride_lse_s,
    NUM_HEADS: tl.constexpr,
    BLOCK_N: tl.constexpr,
    NUM_SPLITS: tl.constexpr,
    HALF_C: tl.constexpr,
    BLOCK_V: tl.constexpr,
):
    pid = tl.program_id(0)
    batch_id = pid // NUM_SPLITS
    split_id = pid % NUM_SPLITS

    kv_begin = tl.load(kv_indptr + batch_id)
    seq_len = tl.load(kv_indptr + batch_id + 1) - kv_begin

    kv_per_split = tl.cdiv(seq_len, NUM_SPLITS)
    split_start = kv_per_split * split_id
    split_end = tl.minimum(split_start + kv_per_split, seq_len)

    offs_h = tl.arange(0, NUM_HEADS)
    offs_v = tl.arange(0, BLOCK_V)

    if split_start >= split_end:
        zptrs = Mid_O + (batch_id * NUM_HEADS + offs_h[:, None]) * stride_mid_bh + split_id * stride_mid_s + offs_v[None, :]
        tl.store(zptrs, tl.zeros([NUM_HEADS, BLOCK_V], dtype=tl.float32))
        tl.store(zptrs + BLOCK_V, tl.zeros([NUM_HEADS, BLOCK_V], dtype=tl.float32))
        tl.store(zptrs + HALF_C, tl.zeros([NUM_HEADS, BLOCK_V], dtype=tl.float32))
        tl.store(zptrs + HALF_C + BLOCK_V, tl.zeros([NUM_HEADS, BLOCK_V], dtype=tl.float32))
        lptrs = Mid_LSE + (batch_id * NUM_HEADS + offs_h) * stride_lse_bh + split_id * stride_lse_s
        tl.store(lptrs, tl.full([NUM_HEADS], float("-inf"), dtype=tl.float32))
        return

    # Q layout: first 512 dims => 2 x 128 packed-byte halves; last 64 dims => 32 packed bytes.
    offs_c0 = tl.arange(0, BLOCK_V)
    offs_c1 = BLOCK_V + offs_c0
    offs_r = tl.arange(0, 32)

    q_row = batch_id * NUM_HEADS
    q_nope0 = tl.load(
        Q_packed + batch_id * stride_qp_b + offs_h[:, None] * stride_qp_h + offs_c0[None, :]
    )
    q_sc0 = tl.load(
        Q_scale + (q_row + offs_h[:, None]) * stride_qs + tl.arange(0, 8)[None, :]
    )
    q_nope1 = tl.load(
        Q_packed + batch_id * stride_qp_b + offs_h[:, None] * stride_qp_h + offs_c1[None, :]
    )
    q_sc1 = tl.load(
        Q_scale + (q_row + offs_h[:, None]) * stride_qs + 8 + tl.arange(0, 8)[None, :]
    )
    q_rope = tl.load(
        Q_packed + batch_id * stride_qp_b + offs_h[:, None] * stride_qp_h + HALF_C + offs_r[None, :]
    )
    q_sc_rope = tl.load(
        Q_scale + (q_row + offs_h[:, None]) * stride_qs + 16 + tl.arange(0, 2)[None, :]
    )

    acc_e0 = tl.zeros([NUM_HEADS, BLOCK_V], dtype=tl.float32)
    acc_e1 = tl.zeros([NUM_HEADS, BLOCK_V], dtype=tl.float32)
    acc_o0 = tl.zeros([NUM_HEADS, BLOCK_V], dtype=tl.float32)
    acc_o1 = tl.zeros([NUM_HEADS, BLOCK_V], dtype=tl.float32)

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

    offs_n = tl.arange(0, BLOCK_N)
    sc_v_idx = offs_v // 16  # 128 bytes -> 8 scale groups

    for start_n in range(split_start, split_end, BLOCK_N):
        cur_n = start_n + offs_n
        mask_n = cur_n < split_end
        kv_locs = kv_begin + cur_n  # contiguous within request; no separate kv_indices indirection

        # Load first 512 dims once as two 128-byte halves. Each half is reused for both QK(nope) and PV.
        k_nope0 = tl.load(
            KV_packed + kv_locs[:, None] * stride_kv + offs_c0[None, :],
            mask=mask_n[:, None],
            other=0,
        )
        k_sc0 = tl.load(
            KV_scale + kv_locs[:, None] * stride_ks + tl.arange(0, 8)[None, :],
            mask=mask_n[:, None],
            other=0,
        )
        k_nope1 = tl.load(
            KV_packed + kv_locs[:, None] * stride_kv + offs_c1[None, :],
            mask=mask_n[:, None],
            other=0,
        )
        k_sc1 = tl.load(
            KV_scale + kv_locs[:, None] * stride_ks + 8 + tl.arange(0, 8)[None, :],
            mask=mask_n[:, None],
            other=0,
        )

        qk = tl.zeros([NUM_HEADS, BLOCK_N], dtype=tl.float32)
        qk = tl.dot_scaled(
            q_nope0, q_sc0, "e2m1",
            tl.trans(k_nope0), k_sc0, "e2m1",
            acc=qk,
            fast_math=True,
        )
        qk = tl.dot_scaled(
            q_nope1, q_sc1, "e2m1",
            tl.trans(k_nope1), k_sc1, "e2m1",
            acc=qk,
            fast_math=True,
        )

        k_rope = tl.load(
            KV_packed + kv_locs[:, None] * stride_kv + HALF_C + offs_r[None, :],
            mask=mask_n[:, None],
            other=0,
        )
        k_sc_rope = tl.load(
            KV_scale + kv_locs[:, None] * stride_ks + 16 + tl.arange(0, 2)[None, :],
            mask=mask_n[:, None],
            other=0,
        )
        qk = tl.dot_scaled(
            q_rope, q_sc_rope, "e2m1",
            tl.trans(k_rope), k_sc_rope, "e2m1",
            acc=qk,
            fast_math=True,
        )

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

        qk_max = tl.max(qk, axis=1)
        m_new = tl.maximum(m_i, qk_max)
        alpha = tl.exp2(m_i - m_new)
        p = tl.exp2(qk - tl.reshape(m_new, [NUM_HEADS, 1]))
        l_i = l_i * alpha + tl.sum(p, axis=1)

        alpha_2d = tl.reshape(alpha, [NUM_HEADS, 1])
        acc_e0 = acc_e0 * alpha_2d
        acc_e1 = acc_e1 * alpha_2d
        acc_o0 = acc_o0 * alpha_2d
        acc_o1 = acc_o1 * alpha_2d
        m_i = m_new

        p_f16 = p.to(tl.float16)

        # Reuse the already-loaded first 512 dims for PV.
        v0_i = k_nope0.to(tl.int32)
        sc0 = k_sc0[:, sc_v_idx]
        sc0_f = tl.exp2(sc0.to(tl.float32) - 127.0)

        lo0 = v0_i & 0xF
        hi0 = (v0_i >> 4) & 0xF
        vlo0 = tl.load(E2M1_LUT + lo0)
        vhi0 = tl.load(E2M1_LUT + hi0)
        vlo0_s = (vlo0.to(tl.float32) * sc0_f).to(tl.float16)
        vhi0_s = (vhi0.to(tl.float32) * sc0_f).to(tl.float16)
        acc_e0 += tl.dot(p_f16, vlo0_s, out_dtype=tl.float32)
        acc_o0 += tl.dot(p_f16, vhi0_s, out_dtype=tl.float32)

        v1_i = k_nope1.to(tl.int32)
        sc1 = k_sc1[:, sc_v_idx]
        sc1_f = tl.exp2(sc1.to(tl.float32) - 127.0)

        lo1 = v1_i & 0xF
        hi1 = (v1_i >> 4) & 0xF
        vlo1 = tl.load(E2M1_LUT + lo1)
        vhi1 = tl.load(E2M1_LUT + hi1)
        vlo1_s = (vlo1.to(tl.float32) * sc1_f).to(tl.float16)
        vhi1_s = (vhi1.to(tl.float32) * sc1_f).to(tl.float16)
        acc_e1 += tl.dot(p_f16, vlo1_s, out_dtype=tl.float32)
        acc_o1 += tl.dot(p_f16, vhi1_s, out_dtype=tl.float32)

    safe_l = tl.where(l_i > 0.0, l_i, 1.0)
    inv_l = 1.0 / safe_l
    inv_l_2d = tl.reshape(inv_l, [NUM_HEADS, 1])
    acc_e0 = acc_e0 * inv_l_2d
    acc_e1 = acc_e1 * inv_l_2d
    acc_o0 = acc_o0 * inv_l_2d
    acc_o1 = acc_o1 * inv_l_2d
    lse = m_i + tl.log(safe_l) * LOG2E

    base_ptrs = Mid_O + (batch_id * NUM_HEADS + offs_h[:, None]) * stride_mid_bh + split_id * stride_mid_s
    tl.store(base_ptrs + offs_v[None, :], acc_e0)
    tl.store(base_ptrs + BLOCK_V + offs_v[None, :], acc_e1)
    tl.store(base_ptrs + HALF_C + offs_v[None, :], acc_o0)
    tl.store(base_ptrs + HALF_C + BLOCK_V + offs_v[None, :], acc_o1)

    lptrs = Mid_LSE + (batch_id * NUM_HEADS + offs_h) * stride_lse_bh + split_id * stride_lse_s
    tl.store(lptrs, lse)


@triton.jit
def _mla_mxfp4_stage2_log2(
    Mid_O,
    Mid_LSE,
    O,
    stride_mid_bh,
    stride_mid_s,
    stride_lse_bh,
    stride_lse_s,
    stride_o_bh,
    HALF_C: tl.constexpr,
    NUM_SPLITS: tl.constexpr,
):
    pid = tl.program_id(0)
    offs = tl.arange(0, HALF_C)

    e_max = tl.full([], float("-inf"), dtype=tl.float32)
    e_sum = tl.full([], 0.0, dtype=tl.float32)
    r_even = tl.zeros([HALF_C], dtype=tl.float32)
    r_odd = tl.zeros([HALF_C], dtype=tl.float32)

    mid_base = Mid_O + pid * stride_mid_bh
    lse_base = Mid_LSE + pid * stride_lse_bh

    for s in tl.static_range(0, NUM_SPLITS):
        sb = mid_base + s * stride_mid_s
        sv_e = tl.load(sb + offs)
        sv_o = tl.load(sb + HALF_C + offs)
        lse = tl.load(lse_base + s * stride_lse_s)

        n_max = tl.maximum(lse, e_max)
        old_sc = tl.exp2(e_max - n_max)
        el = tl.exp2(lse - n_max)
        r_even = r_even * old_sc + el * sv_e
        r_odd = r_odd * old_sc + el * sv_o
        e_sum = e_sum * old_sc + el
        e_max = n_max

    inv_s = 1.0 / tl.where(e_sum > 0.0, e_sum, 1.0)
    r_even = r_even * inv_s
    r_odd = r_odd * inv_s

    o_base = O + pid * stride_o_bh
    tl.store(o_base + offs * 2, r_even.to(tl.bfloat16))
    tl.store(o_base + offs * 2 + 1, r_odd.to(tl.bfloat16))


# =====================================================================================
# AITER direct stage1/reduce fallback
# =====================================================================================

def _mla_stage1_direct(
    q_ready: torch.Tensor,
    kv_buffer_4d: torch.Tensor,
    cached: dict,
    output: torch.Tensor,
    q_scale,
    kv_scale: torch.Tensor,
):
    aiter_module.mla_decode_stage1_asm_fwd(
        q_ready,
        kv_buffer_4d,
        cached["qo_indptr"],
        cached["kv_indptr"],
        cached["kv_indices"],
        cached["kv_last_page_len"],
        None,  # num_kv_splits_indptr = None -> persistent mode
        cached["work_meta_data"],
        cached["work_indptr"],
        cached["work_info_set"],
        1,
        cached["page_size"],
        NUM_KV_HEADS,
        SM_SCALE,
        cached["logits"],
        cached["attn_lse"],
        output,
        q_scale,
        kv_scale,
    )


def _mla_reduce_direct(cached: dict, output: torch.Tensor):
    aiter_module.mla_reduce_v1(
        cached["logits"],
        cached["attn_lse"],
        cached["reduce_indptr"],
        cached["reduce_final_map"],
        cached["reduce_partial_map"],
        1,
        output,
    )


# =====================================================================================
# Cache builders
# =====================================================================================

_cache = {}
_decided_strategy = {}


def _build_mxfp4_cache(batch_size: int, kv_seq_len: int) -> dict:
    key = ("mxfp4", batch_size, kv_seq_len)
    if key in _cache:
        return _cache[key]

    num_splits = _get_mxfp4_splits(batch_size, kv_seq_len)
    entry = {
        "num_splits": num_splits,
        "mid_o": torch.empty(
            (batch_size * NUM_HEADS, num_splits, V_HEAD_DIM),
            dtype=torch.float32,
            device="cuda",
        ),
        "mid_lse": torch.empty(
            (batch_size * NUM_HEADS, num_splits),
            dtype=torch.float32,
            device="cuda",
        ),
        "output": torch.empty(
            (batch_size, NUM_HEADS, V_HEAD_DIM),
            dtype=torch.bfloat16,
            device="cuda",
        ),
    }
    _cache[key] = entry
    return entry


def _build_a8w8_cache(batch_size: int, kv_seq_len: int) -> dict:
    key = ("a8w8", batch_size, kv_seq_len)
    if key in _cache:
        return _cache[key]

    total_q = batch_size
    num_kv_splits = _get_num_kv_splits(batch_size, kv_seq_len)
    page_size = _get_page_size(batch_size, kv_seq_len)

    num_pages_per_batch = kv_seq_len // page_size
    total_pages = batch_size * num_pages_per_batch

    qo_indptr = torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda")
    kv_indptr = torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda") * num_pages_per_batch
    kv_last_page_len = torch.full((batch_size,), page_size, dtype=torch.int32, device="cuda")
    kv_indices = torch.arange(total_pages, dtype=torch.int32, device="cuda")

    info = get_mla_metadata_info_v1(
        batch_size,
        1,
        NUM_HEADS,
        FP8_DTYPE,
        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]
    (work_metadata, work_indptr, 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,
        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=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=FP8_DTYPE,
    )

    num_partials = reduce_partial_map.size(0)
    entry = {
        "page_size": page_size,
        "num_kv_splits": num_kv_splits,
        "qo_indptr": qo_indptr,
        "kv_indptr": kv_indptr,
        "kv_indices": kv_indices,
        "kv_last_page_len": kv_last_page_len,
        "output": torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
        "logits": torch.empty((num_partials, 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device="cuda"),
        "attn_lse": torch.empty((num_partials, 1, NUM_HEADS, 1), dtype=torch.float32, device="cuda"),
        "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,
        "q_fp8_buf": torch.empty((total_q, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda"),
        "q_scale_buf": torch.empty(1, dtype=torch.float32, device="cuda"),
    }
    _cache[key] = entry
    return entry


def _build_a16w8_cache(batch_size: int, kv_seq_len: int) -> dict:
    key = ("a16w8", batch_size, kv_seq_len)
    if key in _cache:
        return _cache[key]

    total_q = batch_size
    num_kv_splits = _get_num_kv_splits(batch_size, kv_seq_len)
    page_size = _get_page_size(batch_size, kv_seq_len)

    num_pages_per_batch = kv_seq_len // page_size
    total_pages = batch_size * num_pages_per_batch

    qo_indptr = torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda")
    kv_indptr = torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda") * num_pages_per_batch
    kv_last_page_len = torch.full((batch_size,), page_size, dtype=torch.int32, device="cuda")
    kv_indices = torch.arange(total_pages, dtype=torch.int32, device="cuda")

    q_dtype = torch.bfloat16
    kv_dtype = FP8_DTYPE

    info = get_mla_metadata_info_v1(
        batch_size, 1, NUM_HEADS, q_dtype, kv_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]
    (work_metadata, work_indptr, 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, 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=1, uni_seqlen_qo=1, fast_mode=False,
        max_split_per_batch=num_kv_splits, intra_batch_mode=True,
        dtype_q=q_dtype, dtype_kv=kv_dtype,
    )

    num_partials = reduce_partial_map.size(0)
    entry = {
        "page_size": page_size,
        "num_kv_splits": num_kv_splits,
        "qo_indptr": qo_indptr,
        "kv_indptr": kv_indptr,
        "kv_indices": kv_indices,
        "kv_last_page_len": kv_last_page_len,
        "output": torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
        "logits": torch.empty((num_partials, 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device="cuda"),
        "attn_lse": torch.empty((num_partials, 1, NUM_HEADS, 1), dtype=torch.float32, device="cuda"),
        "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,
    }
    _cache[key] = entry
    return entry


# =====================================================================================
# Runtime helpers
# =====================================================================================


def _quantize_fp8_naive(tensor: torch.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)



def _run_a8w8_direct(q_bf16: torch.Tensor, kv_buffer_4d: torch.Tensor, kv_scale: torch.Tensor, cached: dict) -> torch.Tensor:
    if HAS_FUSED_QUANT:
        try:
            dynamic_per_tensor_quant(cached["q_fp8_buf"], q_bf16.view_as(cached["q_fp8_buf"]), cached["q_scale_buf"])
            q_fp8 = cached["q_fp8_buf"]
            q_scale = cached["q_scale_buf"]
        except Exception:
            q_fp8, q_scale = _quantize_fp8_naive(q_bf16)
    else:
        q_fp8, q_scale = _quantize_fp8_naive(q_bf16)

    output = cached["output"]
    _mla_stage1_direct(q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM), kv_buffer_4d, cached, output, q_scale, kv_scale)
    _mla_reduce_direct(cached, output)
    return output


def _run_a16w8_direct(q_bf16: torch.Tensor, kv_buffer_4d: torch.Tensor, kv_scale: torch.Tensor, cached: dict) -> torch.Tensor:
    output = cached["output"]
    _mla_stage1_direct(q_bf16.view(-1, NUM_HEADS, QK_HEAD_DIM), kv_buffer_4d, cached, output, None, kv_scale)
    _mla_reduce_direct(cached, output)
    return output


def _run_mxfp4_qh16(q_bf16: torch.Tensor, kv_data_mxfp4, kv_indptr: torch.Tensor, cached: dict) -> torch.Tensor:
    kv_packed, kv_scale = kv_data_mxfp4
    num_splits = cached["num_splits"]

    q_2d = q_bf16.view(-1, QK_HEAD_DIM)
    q_packed_2d, q_scale_2d = dynamic_mxfp4_quant(q_2d)
    q_packed_u8 = q_packed_2d.view(torch.uint8)
    q_scale_u8 = q_scale_2d.view(torch.uint8)
    q_packed_3d = q_packed_u8.reshape(q_bf16.shape[0], NUM_HEADS, QK_HEAD_DIM // 2)

    total_kv = kv_packed.shape[0]
    kv_packed_2d = kv_packed.reshape(total_kv, -1).view(torch.uint8)
    kv_scale_2d = kv_scale.reshape(total_kv, -1).view(torch.uint8)

    mid_o = cached["mid_o"]
    mid_lse = cached["mid_lse"]
    output = cached["output"]
    output_flat = output.view(-1, V_HEAD_DIM)
    lut = _get_e2m1_lut()

    BLOCK_N = 32
    BLOCK_V = 128

    grid1 = (q_bf16.shape[0] * num_splits,)
    _mla_mxfp4_qh16_stage1_reuse[grid1](
        q_packed_3d,
        q_scale_u8,
        kv_packed_2d,
        kv_scale_2d,
        lut,
        mid_o,
        mid_lse,
        kv_indptr,
        SM_SCALE_LOG2E,
        q_packed_3d.stride(0),
        q_packed_3d.stride(1),
        q_scale_u8.stride(0),
        kv_packed_2d.stride(0),
        kv_scale_2d.stride(0),
        mid_o.stride(0),
        mid_o.stride(1),
        mid_lse.stride(0),
        mid_lse.stride(1),
        NUM_HEADS=NUM_HEADS,
        BLOCK_N=BLOCK_N,
        NUM_SPLITS=num_splits,
        HALF_C=HALF_V_DIM,
        BLOCK_V=BLOCK_V,
        num_warps=4,
        num_stages=1,
    )

    grid2 = (q_bf16.shape[0] * NUM_HEADS,)
    _mla_mxfp4_stage2_log2[grid2](
        mid_o,
        mid_lse,
        output_flat,
        mid_o.stride(0),
        mid_o.stride(1),
        mid_lse.stride(0),
        mid_lse.stride(1),
        output_flat.stride(0),
        HALF_C=HALF_V_DIM,
        NUM_SPLITS=num_splits,
        num_warps=4,
        num_stages=1,
    )
    return output


# =====================================================================================
# Reference and custom entrypoints
# =====================================================================================


def ref_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    batch_size = config["batch_size"]
    kv_seq_len = config.get("kv_seq_len", 1024)

    q_input, q_scale = _quantize_fp8_naive(q)
    kv_input, kv_scale = kv_data["fp8"]

    page_size = _get_page_size(batch_size, kv_seq_len)
    total_pages = kv_input.shape[0] // page_size
    kv_buffer_4d = kv_input.view(total_pages, page_size, NUM_KV_HEADS, kv_input.shape[-1])

    num_pages_per_batch = kv_seq_len // page_size
    kv_indptr_paged = torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda") * num_pages_per_batch
    kv_indices_ref = torch.arange(total_pages, dtype=torch.int32, device="cuda")
    kv_last_page_len = torch.full((batch_size,), page_size, dtype=torch.int32, device="cuda")
    num_kv_splits = _get_num_kv_splits(batch_size, kv_seq_len)

    info = get_mla_metadata_info_v1(
        batch_size,
        1,
        NUM_HEADS,
        q_input.dtype,
        kv_input.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]
    (wm, wi, wis, ri, rfm, rpm) = work
    get_mla_metadata_v1(
        qo_indptr,
        kv_indptr_paged,
        kv_last_page_len,
        NUM_HEADS // NUM_KV_HEADS,
        NUM_KV_HEADS,
        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=q_input.dtype,
        dtype_kv=kv_input.dtype,
    )
    output = torch.empty((q.shape[0], NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
    mla_decode_fwd(
        q_input.view(-1, NUM_HEADS, QK_HEAD_DIM),
        kv_buffer_4d,
        output,
        qo_indptr,
        kv_indptr_paged,
        kv_indices_ref,
        kv_last_page_len,
        1,
        page_size=page_size,
        nhead_kv=NUM_KV_HEADS,
        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 output



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"]
    shape_key = (batch_size, kv_seq_len)

    strategy = _decided_strategy.get(shape_key)
    if strategy is None:
        strategy = SHAPE_STRATEGY.get(shape_key, DEFAULT_STRATEGY)
        if strategy == "mxfp4" and not HAS_MXFP4:
            strategy = "a16w8"

    if strategy == "mxfp4":
        try:
            cached = _build_mxfp4_cache(batch_size, kv_seq_len)
            result = _run_mxfp4_qh16(q, kv_data["mxfp4"], kv_indptr, cached)
            _decided_strategy[shape_key] = "mxfp4"
            return result
        except Exception:
            strategy = "a16w8"

    kv_input, kv_scale = kv_data["fp8"]
    page_size = _get_page_size(batch_size, kv_seq_len)
    kv_buffer_4d = kv_input.view(
        kv_input.shape[0] // page_size,
        page_size,
        NUM_KV_HEADS,
        kv_input.shape[-1],
    )

    if strategy == "a16w8":
        cached = _build_a16w8_cache(batch_size, kv_seq_len)
        result = _run_a16w8_direct(q, kv_buffer_4d, kv_scale, cached)
        _decided_strategy[shape_key] = "a16w8"
        return result

    cached = _build_a8w8_cache(batch_size, kv_seq_len)
    result = _run_a8w8_direct(q, kv_buffer_4d, kv_scale, cached)
    _decided_strategy[shape_key] = "a8w8"
    return result
scrolls · 838 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