Skip to content
KernelIndex
Search⌘K

submission 716743

ooousay · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8b5036cf47bd1b109ae93c36278c77ea1252dc0bd3bcf0d225942892b8d02acf
license declaredunknown
license concludedunknown
authorsooousay
imported2026-08-15

Techniques

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

mmascores = tl.dot(q_rope, tl.trans(k_rope_bf16))
num-warps = 4num_warps=4, num_stages=1,
online-softmaxm_new = tl.maximum(m_i, m_j)
split-k"""bs=4, kv_len=1024 ? Split-K Triton MLA decode with hardcoded strides and indptr elimination.
stages = 1num_warps=4, num_stages=1,

Kernel source

submission.py1091 lines
#!POPCORN leaderboard amd-mixed-mla
"""
Auto-generated by build.py ? do not edit directly.
Edit the per-shape kernel.py files and re-run build.py.
"""

# ============================================================
# Constants
# ============================================================

from aiter import dtypes as aiter_dtypes

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  # 576
V_HEAD_DIM = KV_LORA_RANK                       # 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
FP8_DTYPE = aiter_dtypes.fp8


# ============================================================
# bs=4, kv_len=1024
# ============================================================

"""bs=4, kv_len=1024 ? Split-K Triton MLA decode with hardcoded strides and indptr elimination.

v33: Hardcode kv_start/q_tok from batch_id, eliminate indptr loads, make strides constexpr.
"""
import torch
import triton
import triton.language as tl

V_CHUNKS = 4  # split 512 V dims into 4 chunks of 128
BLOCK_V_REDUCE = V_HEAD_DIM // V_CHUNKS  # 128


@triton.jit
def _mla_stage1_4_1024(
    Q_ptr, KV_ptr, KV_scale_ptr,
    Partial_ptr, LSE_ptr,
    sm_scale: tl.constexpr,
    STRIDE_Q_TOK: tl.constexpr, STRIDE_Q_HEAD: tl.constexpr, STRIDE_KV_TOK: tl.constexpr,
    BLOCK_KV: tl.constexpr, BLOCK_LORA: tl.constexpr, BLOCK_ROPE: tl.constexpr,
    BLOCK_V: tl.constexpr,
    HEADS_PER_GROUP: tl.constexpr, NUM_SPLITS: tl.constexpr, NUM_HEADS: tl.constexpr,
    BS: tl.constexpr, NUM_ITERS: tl.constexpr,
    KV_LEN: tl.constexpr,
):
    split_id = tl.program_id(0)
    batch_id = tl.program_id(1)

    kv_scale = tl.load(KV_scale_ptr)

    # Hardcoded: kv_start = batch_id * 1024, no indptr load
    kv_start = batch_id * KV_LEN

    # Hardcoded: q_tok = batch_id (seqlen=1), no indptr load
    q_tok = batch_id

    # Load Q split into lora and rope parts
    # Bake both sm_scale and kv_scale into Q (loaded once, eliminates per-iteration KV scaling)
    h_offs = tl.arange(0, HEADS_PER_GROUP)
    lora_offs = tl.arange(0, BLOCK_LORA)
    rope_offs = tl.arange(0, BLOCK_ROPE)

    q_base = Q_ptr + q_tok * STRIDE_Q_TOK + h_offs[:, None] * STRIDE_Q_HEAD
    q_scale = sm_scale * kv_scale
    q_lora = (tl.load(q_base + lora_offs[None, :]) * q_scale).to(tl.bfloat16)
    q_rope = (tl.load(q_base + (BLOCK_LORA + rope_offs[None, :])) * q_scale).to(tl.bfloat16)

    # Accumulators
    m_i = tl.full([HEADS_PER_GROUP], float("-inf"), dtype=tl.float32)
    l_i = tl.zeros([HEADS_PER_GROUP], dtype=tl.float32)
    acc = tl.zeros([HEADS_PER_GROUP, BLOCK_V], dtype=tl.float32)

    # tps = KV_LEN // NUM_SPLITS (exact division, no empty splits possible)
    kv_base = kv_start + split_id * (KV_LEN // NUM_SPLITS)
    v_offs = tl.arange(0, BLOCK_V)

    for it in tl.static_range(NUM_ITERS):
        tok_offs = tl.arange(0, BLOCK_KV)
        tok_idx = kv_base + it * BLOCK_KV + tok_offs

        kv_row_ptr = KV_ptr + tok_idx[:, None] * STRIDE_KV_TOK

        # Load KV[0:512] as FP8 -> cast to bf16 WITHOUT scaling (scale baked into Q)
        kv_shared = tl.load(kv_row_ptr + lora_offs[None, :])
        kv_bf16 = kv_shared.to(tl.bfloat16)

        # Load K_rope[512:576] as FP8 -> cast to bf16 WITHOUT scaling
        k_rope = tl.load(kv_row_ptr + (BLOCK_LORA + rope_offs[None, :]))
        k_rope_bf16 = k_rope.to(tl.bfloat16)

        # QK^T: kv_scale already baked into Q, so scores are correctly scaled
        scores = tl.dot(q_rope, tl.trans(k_rope_bf16))
        scores = tl.dot(q_lora, tl.trans(kv_bf16), scores)

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

        # PV: V accumulation using unscaled kv_bf16. Scale applied after loop.
        acc = tl.dot(p.to(tl.bfloat16), kv_bf16, acc)
        m_i = m_new

    # Apply kv_scale to V accumulator (once, outside loop)
    acc = acc * kv_scale

    # Normalize and store partial as bf16 + LSE as f32
    norm_acc = acc / tl.maximum(l_i[:, None], 1e-12)
    lse = m_i + tl.log(tl.maximum(l_i, 1e-12))

    # Partial: [BS, NUM_HEADS, NUM_SPLITS, BLOCK_V] as bf16
    partial_base = (Partial_ptr
        + batch_id * (NUM_HEADS * NUM_SPLITS * BLOCK_V)
        + h_offs[:, None] * (NUM_SPLITS * BLOCK_V)
        + split_id * BLOCK_V
        + v_offs[None, :])
    tl.store(partial_base, norm_acc.to(tl.bfloat16))

    # LSE: [BS, NUM_HEADS, NUM_SPLITS] as f32
    lse_base = (LSE_ptr
        + batch_id * (NUM_HEADS * NUM_SPLITS)
        + h_offs * NUM_SPLITS
        + split_id)
    tl.store(lse_base, lse)


@triton.jit
def _mla_reduce_vsplit_4_1024(
    Partial_ptr, LSE_ptr, O_ptr,
    STRIDE_O_TOK: tl.constexpr, STRIDE_O_HEAD: tl.constexpr,
    BLOCK_V_FULL: tl.constexpr, BLOCK_V_CHUNK: tl.constexpr,
    NUM_SPLITS: tl.constexpr, NUM_HEADS: tl.constexpr,
    BS: tl.constexpr,
):
    batch_id = tl.program_id(0)
    head_id = tl.program_id(1)
    v_chunk_id = tl.program_id(2)
    v_offs = v_chunk_id * BLOCK_V_CHUNK + tl.arange(0, BLOCK_V_CHUNK)

    # Layout: [BS, NUM_HEADS, NUM_SPLITS, ...]
    lse_base = LSE_ptr + batch_id * (NUM_HEADS * NUM_SPLITS) + head_id * NUM_SPLITS
    partial_base = Partial_ptr + batch_id * (NUM_HEADS * NUM_SPLITS * BLOCK_V_FULL) + head_id * (NUM_SPLITS * BLOCK_V_FULL)

    # Pass 1: global max LSE
    m_global = tl.full([1], float("-inf"), dtype=tl.float32)
    for s in range(NUM_SPLITS):
        lse_s = tl.load(lse_base + s)
        m_global = tl.maximum(m_global, lse_s)

    # Pass 2: weighted sum of normalized partials
    acc = tl.zeros([BLOCK_V_CHUNK], dtype=tl.float32)
    w_total = tl.zeros([1], dtype=tl.float32)
    for s in range(NUM_SPLITS):
        lse_s = tl.load(lse_base + s)
        w = tl.exp(lse_s - m_global)
        partial_s = tl.load(partial_base + s * BLOCK_V_FULL + v_offs).to(tl.float32)
        acc += w * partial_s
        w_total += w

    acc = acc / tl.maximum(w_total, 1e-12)

    # Hardcoded: q_tok = batch_id (seqlen=1), no indptr load
    tl.store(O_ptr + batch_id * STRIDE_O_TOK + head_id * STRIDE_O_HEAD + v_offs, acc.to(tl.bfloat16))


# === Entry point ===

_state_4_1024 = None

def _run_4_1024(q, kv_data, qo_indptr, kv_indptr, bs, kv_len):
    global _state_4_1024
    _NUM_SPLITS = 32
    _BLOCK_KV = 32
    _NUM_HEADS_PER_GROUP = 16
    _BS = 4
    _KV_LEN = 1024
    if _state_4_1024 is None:
        _state_4_1024 = {
            "partial": torch.empty((bs, NUM_HEADS, _NUM_SPLITS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
            "lse": torch.empty((bs, NUM_HEADS, _NUM_SPLITS), dtype=torch.float32, device="cuda"),
            "o": torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
        }

    kv_buf, kv_scale = kv_data["fp8"]

    # NUM_ITERS = tokens_per_split / _BLOCK_KV = (1024/32) / 32 = 1
    num_iters = _KV_LEN // _NUM_SPLITS // _BLOCK_KV

    # Q shape: (4, 16, 576) -> stride_q_tok=16*576=9216, stride_q_head=576
    # KV stride: 576
    # O shape: (4, 16, 512) -> stride_o_tok=16*512=8192, stride_o_head=512

    _mla_stage1_4_1024[(_NUM_SPLITS, bs)](
        q, kv_buf, kv_scale,
        _state_4_1024["partial"], _state_4_1024["lse"],
        SM_SCALE,
        STRIDE_Q_TOK=9216, STRIDE_Q_HEAD=576, STRIDE_KV_TOK=576,
        BLOCK_KV=_BLOCK_KV, BLOCK_LORA=KV_LORA_RANK, BLOCK_ROPE=QK_ROPE_HEAD_DIM,
        BLOCK_V=V_HEAD_DIM,
        HEADS_PER_GROUP=_NUM_HEADS_PER_GROUP, NUM_SPLITS=_NUM_SPLITS, NUM_HEADS=NUM_HEADS,
        BS=_BS, NUM_ITERS=num_iters,
        KV_LEN=_KV_LEN,
        num_warps=4, num_stages=1,
    )

    _mla_reduce_vsplit_4_1024[(bs, NUM_HEADS, V_CHUNKS)](
        _state_4_1024["partial"], _state_4_1024["lse"], _state_4_1024["o"],
        STRIDE_O_TOK=8192, STRIDE_O_HEAD=512,
        BLOCK_V_FULL=V_HEAD_DIM, BLOCK_V_CHUNK=BLOCK_V_REDUCE,
        NUM_SPLITS=_NUM_SPLITS, NUM_HEADS=NUM_HEADS,
        BS=_BS,
        num_warps=4,
    )

    return _state_4_1024["o"]

# ============================================================
# bs=4, kv_len=8192
# ============================================================

"""bs=4, kv_len=8192 ? Custom Triton MLA decode with split-K + online softmax.

v78: Hardcode shape constants. Eliminate indptr loads, constexpr strides.
Keep 3D grid, num_warps=8, num_stages=2 from v65.
"""
import torch
import triton
import triton.language as tl


@triton.jit
def _mla_stage1_4_8192(
    Q_ptr, KV_ptr, KV_scale_ptr,
    Partial_ptr, LSE_ptr,
    sm_scale: tl.constexpr,
    STRIDE_Q_TOK: tl.constexpr, STRIDE_Q_HEAD: tl.constexpr,
    STRIDE_KV_TOK: tl.constexpr,
    BLOCK_KV: tl.constexpr, BLOCK_LORA: tl.constexpr, BLOCK_ROPE: tl.constexpr,
    BLOCK_V: tl.constexpr,
    HEADS_PER_GROUP: tl.constexpr, NUM_SPLITS: tl.constexpr, NUM_HEADS: tl.constexpr,
    BS: tl.constexpr,
    KV_LEN: tl.constexpr, TOKENS_PER_SPLIT: tl.constexpr, NUM_ITERS: tl.constexpr,
):
    split_id = tl.program_id(0)
    batch_id = tl.program_id(1)
    hg = tl.program_id(2)

    kv_scale = tl.load(KV_scale_ptr)

    # Hardcoded: kv_start = batch_id * 8192, q_tok = batch_id
    kv_base = batch_id * KV_LEN + split_id * TOKENS_PER_SPLIT

    h_offs = tl.arange(0, HEADS_PER_GROUP)
    lora_offs = tl.arange(0, BLOCK_LORA)
    rope_offs = tl.arange(0, BLOCK_ROPE)

    q_base = Q_ptr + batch_id * STRIDE_Q_TOK + (hg * HEADS_PER_GROUP + h_offs[:, None]) * STRIDE_Q_HEAD
    q_scale = sm_scale * kv_scale
    q_lora = (tl.load(q_base + lora_offs[None, :]) * q_scale).to(tl.bfloat16)
    q_rope = (tl.load(q_base + (BLOCK_LORA + rope_offs[None, :])) * q_scale).to(tl.bfloat16)

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

    v_offs = tl.arange(0, BLOCK_V)

    # 8192 / 64 splits / 32 block = 4 exact iterations
    for it in range(NUM_ITERS):
        tok_offs = tl.arange(0, BLOCK_KV)
        tok_idx = kv_base + it * BLOCK_KV + tok_offs

        kv_row_ptr = KV_ptr + tok_idx[:, None] * STRIDE_KV_TOK

        kv_shared = tl.load(kv_row_ptr + lora_offs[None, :])
        kv_bf16 = kv_shared.to(tl.bfloat16)

        k_rope = tl.load(kv_row_ptr + (BLOCK_LORA + rope_offs[None, :]))
        k_rope_bf16 = k_rope.to(tl.bfloat16)

        scores = tl.dot(q_rope, tl.trans(k_rope_bf16))
        scores = tl.dot(q_lora, tl.trans(kv_bf16), scores)

        m_j = tl.max(scores, axis=1)
        m_new = tl.maximum(m_i, m_j)
        alpha = tl.exp(m_i - m_new)
        p = tl.exp(scores - m_new[:, None])
        l_i = alpha * l_i + tl.sum(p, axis=1)
        acc = acc * alpha[:, None]

        acc = tl.dot(p.to(tl.bfloat16), kv_bf16, acc)
        m_i = m_new

    acc = acc * kv_scale

    norm_acc = acc / tl.maximum(l_i[:, None], 1e-12)
    lse = m_i + tl.log(tl.maximum(l_i, 1e-12))

    partial_base = (Partial_ptr
        + batch_id * (NUM_HEADS * NUM_SPLITS * BLOCK_V)
        + (hg * HEADS_PER_GROUP + h_offs[:, None]) * (NUM_SPLITS * BLOCK_V)
        + split_id * BLOCK_V
        + v_offs[None, :])
    tl.store(partial_base, norm_acc.to(tl.bfloat16))

    lse_base = (LSE_ptr
        + batch_id * (NUM_HEADS * NUM_SPLITS)
        + (hg * HEADS_PER_GROUP + h_offs) * NUM_SPLITS
        + split_id)
    tl.store(lse_base, lse)


@triton.jit
def _mla_reduce_4_8192(
    Partial_ptr, LSE_ptr, O_ptr,
    STRIDE_O_TOK: tl.constexpr, STRIDE_O_HEAD: tl.constexpr,
    BLOCK_V: tl.constexpr, NUM_SPLITS: tl.constexpr, NUM_HEADS: tl.constexpr,
    BS: tl.constexpr,
):
    batch_id = tl.program_id(0)
    head_id = tl.program_id(1)
    v_offs = tl.arange(0, BLOCK_V)

    lse_base = LSE_ptr + batch_id * (NUM_HEADS * NUM_SPLITS) + head_id * NUM_SPLITS
    partial_base = Partial_ptr + batch_id * (NUM_HEADS * NUM_SPLITS * BLOCK_V) + head_id * (NUM_SPLITS * BLOCK_V)

    m_global = tl.full([1], float("-inf"), dtype=tl.float32)
    for s in range(NUM_SPLITS):
        lse_s = tl.load(lse_base + s)
        m_global = tl.maximum(m_global, lse_s)

    acc = tl.zeros([BLOCK_V], dtype=tl.float32)
    w_total = tl.zeros([1], dtype=tl.float32)
    for s in range(NUM_SPLITS):
        lse_s = tl.load(lse_base + s)
        w = tl.exp(lse_s - m_global)
        partial_s = tl.load(partial_base + s * BLOCK_V + v_offs)
        acc += w * partial_s
        w_total += w

    acc = acc / tl.maximum(w_total, 1e-12)

    tl.store(O_ptr + batch_id * STRIDE_O_TOK + head_id * STRIDE_O_HEAD + v_offs, acc.to(tl.bfloat16))


_state_4_8192 = None

def _run_4_8192(q, kv_data, qo_indptr, kv_indptr, bs, kv_len):
    global _state_4_8192
    _NUM_SPLITS = 64
    _BLOCK_KV = 32
    _NUM_HEADS_PER_GROUP = 16
    _BS = 4
    _NUM_ITERS = 4  # 8192 / 64 / 32

    if _state_4_8192 is None:
        _state_4_8192 = {
            "partial": torch.empty((bs, NUM_HEADS, _NUM_SPLITS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
            "lse": torch.empty((bs, NUM_HEADS, _NUM_SPLITS), dtype=torch.float32, device="cuda"),
            "o": torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
        }

    kv_buf, kv_scale = kv_data["fp8"]
    nhg = NUM_HEADS // _NUM_HEADS_PER_GROUP  # = 1

    _mla_stage1_4_8192[(_NUM_SPLITS, bs, nhg)](
        q, kv_buf, kv_scale,
        _state_4_8192["partial"], _state_4_8192["lse"],
        SM_SCALE,
        STRIDE_Q_TOK=9216, STRIDE_Q_HEAD=576,
        STRIDE_KV_TOK=576,
        BLOCK_KV=_BLOCK_KV, BLOCK_LORA=KV_LORA_RANK, BLOCK_ROPE=QK_ROPE_HEAD_DIM,
        BLOCK_V=V_HEAD_DIM,
        HEADS_PER_GROUP=_NUM_HEADS_PER_GROUP, NUM_SPLITS=_NUM_SPLITS, NUM_HEADS=NUM_HEADS,
        BS=_BS,
        KV_LEN=8192, TOKENS_PER_SPLIT=128, NUM_ITERS=_NUM_ITERS,
        num_warps=8, num_stages=2,
    )

    _mla_reduce_4_8192[(bs, NUM_HEADS)](
        _state_4_8192["partial"], _state_4_8192["lse"], _state_4_8192["o"],
        STRIDE_O_TOK=8192, STRIDE_O_HEAD=512,
        BLOCK_V=V_HEAD_DIM, NUM_SPLITS=_NUM_SPLITS, NUM_HEADS=NUM_HEADS,
        BS=_BS,
        num_warps=4,
    )

    return _state_4_8192["o"]

# ============================================================
# bs=32, kv_len=1024
# ============================================================

"""bs=32, kv_len=1024 ? Split-K Triton MLA decode with XCD remapping.

v12: Fully hardcoded for (bs=32, kv=1024, seqlen=1). No indptr loads, constexpr
strides, no empty-split check. Keeps XCD remap + static_range from v8.
"""
import torch
import triton
import triton.language as tl


@triton.jit
def remap_xcd_32_1024(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):
    pids_per_xcd = (GRID_MN + NUM_XCDS - 1) // NUM_XCDS
    tall_xcds = GRID_MN % NUM_XCDS
    tall_xcds = NUM_XCDS if tall_xcds == 0 else tall_xcds
    xcd = pid % NUM_XCDS
    local_pid = pid // NUM_XCDS
    if xcd < tall_xcds:
        pid = xcd * pids_per_xcd + local_pid
    else:
        pid = (
            tall_xcds * pids_per_xcd
            + (xcd - tall_xcds) * (pids_per_xcd - 1)
            + local_pid
        )
    return pid


@triton.jit
def _mla_stage1_32_1024(
    Q_ptr, KV_ptr, KV_scale_ptr,
    Partial_ptr, LSE_ptr,
    sm_scale: tl.constexpr,
    STRIDE_Q_TOK: tl.constexpr, STRIDE_Q_HEAD: tl.constexpr,
    STRIDE_KV_TOK: tl.constexpr,
    BLOCK_KV: tl.constexpr, BLOCK_LORA: tl.constexpr, BLOCK_ROPE: tl.constexpr,
    BLOCK_V: tl.constexpr,
    HEADS_PER_GROUP: tl.constexpr, NUM_SPLITS: tl.constexpr, NUM_HEADS: tl.constexpr,
    NUM_ITERS: tl.constexpr,
    KV_LEN: tl.constexpr, TOKENS_PER_SPLIT: tl.constexpr,
    GRID_TOTAL: tl.constexpr,
):
    pid = tl.program_id(0)
    pid = remap_xcd_32_1024(pid, GRID_TOTAL)
    split_id = pid % NUM_SPLITS
    batch_id = pid // NUM_SPLITS

    kv_scale = tl.load(KV_scale_ptr)

    # Hardcoded: kv_start = batch_id * 1024, q_tok = batch_id (seqlen=1)
    kv_base = batch_id * KV_LEN + split_id * TOKENS_PER_SPLIT

    h_offs = tl.arange(0, HEADS_PER_GROUP)
    lora_offs = tl.arange(0, BLOCK_LORA)
    rope_offs = tl.arange(0, BLOCK_ROPE)

    q_base = Q_ptr + batch_id * STRIDE_Q_TOK + h_offs[:, None] * STRIDE_Q_HEAD
    q_scale = sm_scale * kv_scale
    q_lora = (tl.load(q_base + lora_offs[None, :]) * q_scale).to(tl.bfloat16)
    q_rope = (tl.load(q_base + (BLOCK_LORA + rope_offs[None, :])) * q_scale).to(tl.bfloat16)

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

    v_offs = tl.arange(0, BLOCK_V)

    # 1024 / 8 splits / 32 block = 4 exact iterations
    for it in tl.static_range(NUM_ITERS):
        tok_offs = tl.arange(0, BLOCK_KV)
        tok_idx = kv_base + it * BLOCK_KV + tok_offs

        kv_row_ptr = KV_ptr + tok_idx[:, None] * STRIDE_KV_TOK

        kv_shared = tl.load(kv_row_ptr + lora_offs[None, :])
        kv_bf16 = kv_shared.to(tl.bfloat16)

        k_rope = tl.load(kv_row_ptr + (BLOCK_LORA + rope_offs[None, :]))
        k_rope_bf16 = k_rope.to(tl.bfloat16)

        scores = tl.dot(q_rope, tl.trans(k_rope_bf16))
        scores = tl.dot(q_lora, tl.trans(kv_bf16), scores)

        m_j = tl.max(scores, axis=1)
        m_new = tl.maximum(m_i, m_j)
        alpha = tl.exp(m_i - m_new)
        p = tl.exp(scores - m_new[:, None])
        l_i = alpha * l_i + tl.sum(p, axis=1)
        acc = acc * alpha[:, None]

        acc = tl.dot(p.to(tl.bfloat16), kv_bf16, acc)
        m_i = m_new

    acc = acc * kv_scale

    norm_acc = acc / tl.maximum(l_i[:, None], 1e-12)
    lse = m_i + tl.log(tl.maximum(l_i, 1e-12))

    partial_base = (Partial_ptr
        + batch_id * (NUM_HEADS * NUM_SPLITS * BLOCK_V)
        + h_offs[:, None] * (NUM_SPLITS * BLOCK_V)
        + split_id * BLOCK_V
        + v_offs[None, :])
    tl.store(partial_base, norm_acc.to(tl.bfloat16))

    lse_base = (LSE_ptr
        + batch_id * (NUM_HEADS * NUM_SPLITS)
        + h_offs * NUM_SPLITS
        + split_id)
    tl.store(lse_base, lse)


@triton.jit
def _mla_reduce_32_1024(
    Partial_ptr, LSE_ptr, O_ptr,
    STRIDE_O_TOK: tl.constexpr, STRIDE_O_HEAD: tl.constexpr,
    BLOCK_V: tl.constexpr, NUM_SPLITS: tl.constexpr, NUM_HEADS: tl.constexpr,
    GRID_TOTAL: tl.constexpr,
):
    pid = tl.program_id(0)
    pid = remap_xcd_32_1024(pid, GRID_TOTAL)
    batch_id = pid // NUM_HEADS
    head_id = pid % NUM_HEADS
    v_offs = tl.arange(0, BLOCK_V)

    lse_base = LSE_ptr + batch_id * (NUM_HEADS * NUM_SPLITS) + head_id * NUM_SPLITS
    partial_base = Partial_ptr + batch_id * (NUM_HEADS * NUM_SPLITS * BLOCK_V) + head_id * (NUM_SPLITS * BLOCK_V)

    m_global = tl.full([1], float("-inf"), dtype=tl.float32)
    for s in range(NUM_SPLITS):
        lse_s = tl.load(lse_base + s)
        m_global = tl.maximum(m_global, lse_s)

    acc = tl.zeros([BLOCK_V], dtype=tl.float32)
    w_total = tl.zeros([1], dtype=tl.float32)
    for s in range(NUM_SPLITS):
        lse_s = tl.load(lse_base + s)
        w = tl.exp(lse_s - m_global)
        partial_s = tl.load(partial_base + s * BLOCK_V + v_offs).to(tl.float32)
        acc += w * partial_s
        w_total += w

    acc = acc / tl.maximum(w_total, 1e-12)

    # Hardcoded: q_tok = batch_id (seqlen=1)
    tl.store(O_ptr + batch_id * STRIDE_O_TOK + head_id * STRIDE_O_HEAD + v_offs, acc.to(tl.bfloat16))


_state_32_1024 = None

def _run_32_1024(q, kv_data, qo_indptr, kv_indptr, bs, kv_len):
    global _state_32_1024
    _NUM_SPLITS = 8
    _BLOCK_KV = 32
    _NUM_HEADS_PER_GROUP = 16
    _NUM_ITERS = 4  # 1024 / 8 / 32

    if _state_32_1024 is None:
        _state_32_1024 = {
            "partial": torch.empty((bs, NUM_HEADS, _NUM_SPLITS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
            "lse": torch.empty((bs, NUM_HEADS, _NUM_SPLITS), dtype=torch.float32, device="cuda"),
            "o": torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
        }

    kv_buf, kv_scale = kv_data["fp8"]

    stage1_grid = _NUM_SPLITS * bs  # 256
    reduce_grid = bs * NUM_HEADS  # 512

    _mla_stage1_32_1024[(stage1_grid,)](
        q, kv_buf, kv_scale,
        _state_32_1024["partial"], _state_32_1024["lse"],
        SM_SCALE,
        STRIDE_Q_TOK=9216, STRIDE_Q_HEAD=576,
        STRIDE_KV_TOK=576,
        BLOCK_KV=_BLOCK_KV, BLOCK_LORA=KV_LORA_RANK, BLOCK_ROPE=QK_ROPE_HEAD_DIM,
        BLOCK_V=V_HEAD_DIM,
        HEADS_PER_GROUP=_NUM_HEADS_PER_GROUP, NUM_SPLITS=_NUM_SPLITS, NUM_HEADS=NUM_HEADS,
        NUM_ITERS=_NUM_ITERS,
        KV_LEN=1024, TOKENS_PER_SPLIT=128,
        GRID_TOTAL=stage1_grid,
        num_warps=4, num_stages=1,
    )

    _mla_reduce_32_1024[(reduce_grid,)](
        _state_32_1024["partial"], _state_32_1024["lse"], _state_32_1024["o"],
        STRIDE_O_TOK=8192, STRIDE_O_HEAD=512,
        BLOCK_V=V_HEAD_DIM, NUM_SPLITS=_NUM_SPLITS, NUM_HEADS=NUM_HEADS,
        GRID_TOTAL=reduce_grid,
        num_warps=4,
    )

    return _state_32_1024["o"]

# ============================================================
# bs=32, kv_len=8192
# ============================================================

"""aiter a16w8 with page_size=8 ? scheduler sees 1/8 sequence length."""
import torch
import aiter
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1

_state_32_8192 = None

def _run_32_8192(q, kv_data, qo_indptr, kv_indptr, bs, kv_len):
    global _state_32_8192

    PAGE_SIZE = 8
    NUM_KV_SPLITS = 32
    nq, nkv, q_seq_len, dv = NUM_HEADS, NUM_KV_HEADS, 1, V_HEAD_DIM
    total_q = bs * q_seq_len

    if _state_32_8192 is None:
        pages_per_batch = kv_len // PAGE_SIZE
        total_pages = bs * pages_per_batch

        kv_indptr_pages = torch.zeros(bs + 1, dtype=torch.int32, device="cuda")
        for i in range(bs):
            kv_indptr_pages[i + 1] = kv_indptr_pages[i] + pages_per_batch

        kv_indices = torch.arange(total_pages, dtype=torch.int32, device="cuda")
        kv_last_page_len = torch.full((bs,), PAGE_SIZE, dtype=torch.int32, device="cuda")

        q_dtype, kv_dtype = torch.bfloat16, FP8_DTYPE

        info = get_mla_metadata_info_v1(
            bs, q_seq_len, nq, q_dtype, kv_dtype,
            is_sparse=False, fast_mode=False,
            num_kv_splits=NUM_KV_SPLITS, intra_batch_mode=False,
        )
        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_pages, 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=q_seq_len, uni_seqlen_qo=q_seq_len,
            fast_mode=False, max_split_per_batch=NUM_KV_SPLITS,
            intra_batch_mode=False, dtype_q=q_dtype, dtype_kv=kv_dtype,
        )

        num_partial = reduce_partial_map.size(0)
        _state_32_8192 = {
            "kv_indices": kv_indices,
            "kv_last_page_len": kv_last_page_len,
            "kv_indptr_pages": kv_indptr_pages,
            "work_metadata": 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,
            "logits": torch.empty((num_partial * q_seq_len, 1, nq, dv), dtype=torch.float32, device="cuda"),
            "attn_lse": torch.empty((num_partial * q_seq_len, 1, nq, 1), dtype=torch.float32, device="cuda"),
            "o": torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda"),
        }

    kv_buf, kv_scale = kv_data["fp8"]
    kv_4d = kv_buf.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_buf.shape[-1])

    aiter.mla_decode_stage1_asm_fwd(
        q.view(-1, NUM_HEADS, QK_HEAD_DIM),
        kv_4d,
        qo_indptr, _state_32_8192["kv_indptr_pages"],
        _state_32_8192["kv_indices"], _state_32_8192["kv_last_page_len"],
        None,
        _state_32_8192["work_metadata"], _state_32_8192["work_indptr"], _state_32_8192["work_info_set"],
        1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
        _state_32_8192["logits"], _state_32_8192["attn_lse"], _state_32_8192["o"],
        None, kv_scale,
    )

    aiter.mla_reduce_v1(
        _state_32_8192["logits"], _state_32_8192["attn_lse"],
        _state_32_8192["reduce_indptr"], _state_32_8192["reduce_final_map"], _state_32_8192["reduce_partial_map"],
        1, _state_32_8192["o"], None,
    )

    return _state_32_8192["o"]

# ============================================================
# bs=64, kv_len=1024
# ============================================================

"""bs=64, kv_len=1024 ? aiter a16w8 with page_size=2.

v22: Switch from custom Triton to aiter ASM with page_size=2.
"""
import torch
import aiter
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1


_state_64_1024 = None

def _run_64_1024(q, kv_data, qo_indptr, kv_indptr, bs, kv_len):
    global _state_64_1024

    PAGE_SIZE = 2
    NUM_KV_SPLITS = 32
    nq, nkv, q_seq_len, dv = NUM_HEADS, NUM_KV_HEADS, 1, V_HEAD_DIM
    total_q = bs * q_seq_len

    if _state_64_1024 is None:
        pages_per_batch = kv_len // PAGE_SIZE
        total_pages = bs * pages_per_batch

        kv_indptr_pages = torch.zeros(bs + 1, dtype=torch.int32, device="cuda")
        for i in range(bs):
            kv_indptr_pages[i + 1] = kv_indptr_pages[i] + pages_per_batch

        kv_indices = torch.arange(total_pages, dtype=torch.int32, device="cuda")
        kv_last_page_len = torch.full((bs,), PAGE_SIZE, dtype=torch.int32, device="cuda")

        q_dtype, kv_dtype = torch.bfloat16, FP8_DTYPE

        info = get_mla_metadata_info_v1(
            bs, q_seq_len, nq, q_dtype, kv_dtype,
            is_sparse=False, fast_mode=False,
            num_kv_splits=NUM_KV_SPLITS, intra_batch_mode=False,
        )
        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_pages, 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=q_seq_len, uni_seqlen_qo=q_seq_len,
            fast_mode=False, max_split_per_batch=NUM_KV_SPLITS,
            intra_batch_mode=False, dtype_q=q_dtype, dtype_kv=kv_dtype,
        )

        num_partial = reduce_partial_map.size(0)
        _state_64_1024 = {
            "kv_indices": kv_indices,
            "kv_last_page_len": kv_last_page_len,
            "kv_indptr_pages": kv_indptr_pages,
            "work_metadata": 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,
            "logits": torch.empty((num_partial * q_seq_len, 1, nq, dv), dtype=torch.float32, device="cuda"),
            "attn_lse": torch.empty((num_partial * q_seq_len, 1, nq, 1), dtype=torch.float32, device="cuda"),
            "o": torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda"),
        }

    kv_buf, kv_scale = kv_data["fp8"]
    kv_4d = kv_buf.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_buf.shape[-1])

    aiter.mla_decode_stage1_asm_fwd(
        q.view(-1, NUM_HEADS, QK_HEAD_DIM),
        kv_4d,
        qo_indptr, _state_64_1024["kv_indptr_pages"],
        _state_64_1024["kv_indices"], _state_64_1024["kv_last_page_len"],
        None,
        _state_64_1024["work_metadata"], _state_64_1024["work_indptr"], _state_64_1024["work_info_set"],
        1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
        _state_64_1024["logits"], _state_64_1024["attn_lse"], _state_64_1024["o"],
        None, kv_scale,
    )

    aiter.mla_reduce_v1(
        _state_64_1024["logits"], _state_64_1024["attn_lse"],
        _state_64_1024["reduce_indptr"], _state_64_1024["reduce_final_map"], _state_64_1024["reduce_partial_map"],
        1, _state_64_1024["o"], None,
    )

    return _state_64_1024["o"]

# ============================================================
# bs=64, kv_len=8192
# ============================================================

"""aiter a16w8 page_size=8 intra_batch_mode=True.

Hypothesis: with bs=64 all sequences same kv_len=8192, intra_batch_mode
gives scheduler uniform work distribution ? better CU utilization.
"""
import torch
import aiter
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1

_state_64_8192 = None

def _run_64_8192(q, kv_data, qo_indptr, kv_indptr, bs, kv_len):
    global _state_64_8192

    PAGE_SIZE = 8
    NUM_KV_SPLITS = 32
    nq, nkv, q_seq_len, dv = NUM_HEADS, NUM_KV_HEADS, 1, V_HEAD_DIM
    total_q = bs * q_seq_len

    if _state_64_8192 is None:
        pages_per_batch = kv_len // PAGE_SIZE
        total_pages = bs * pages_per_batch

        kv_indptr_pages = torch.zeros(bs + 1, dtype=torch.int32, device="cuda")
        for i in range(bs):
            kv_indptr_pages[i + 1] = kv_indptr_pages[i] + pages_per_batch

        kv_indices = torch.arange(total_pages, dtype=torch.int32, device="cuda")
        kv_last_page_len = torch.full((bs,), PAGE_SIZE, dtype=torch.int32, device="cuda")

        q_dtype, kv_dtype = torch.bfloat16, FP8_DTYPE

        info = get_mla_metadata_info_v1(
            bs, q_seq_len, nq, 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_pages, 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=q_seq_len, uni_seqlen_qo=q_seq_len,
            fast_mode=False, max_split_per_batch=NUM_KV_SPLITS,
            intra_batch_mode=True, dtype_q=q_dtype, dtype_kv=kv_dtype,
        )

        num_partial = reduce_partial_map.size(0)
        _state_64_8192 = {
            "kv_indices": kv_indices,
            "kv_last_page_len": kv_last_page_len,
            "kv_indptr_pages": kv_indptr_pages,
            "work_metadata": 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,
            "logits": torch.empty((num_partial * q_seq_len, 1, nq, dv), dtype=torch.float32, device="cuda"),
            "attn_lse": torch.empty((num_partial * q_seq_len, 1, nq, 1), dtype=torch.float32, device="cuda"),
            "o": torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda"),
        }

    kv_buf, kv_scale = kv_data["fp8"]
    kv_4d = kv_buf.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_buf.shape[-1])

    aiter.mla_decode_stage1_asm_fwd(
        q.view(-1, NUM_HEADS, QK_HEAD_DIM),
        kv_4d,
        qo_indptr, _state_64_8192["kv_indptr_pages"],
        _state_64_8192["kv_indices"], _state_64_8192["kv_last_page_len"],
        None,
        _state_64_8192["work_metadata"], _state_64_8192["work_indptr"], _state_64_8192["work_info_set"],
        1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
        _state_64_8192["logits"], _state_64_8192["attn_lse"], _state_64_8192["o"],
        None, kv_scale,
    )

    aiter.mla_reduce_v1(
        _state_64_8192["logits"], _state_64_8192["attn_lse"],
        _state_64_8192["reduce_indptr"], _state_64_8192["reduce_final_map"], _state_64_8192["reduce_partial_map"],
        1, _state_64_8192["o"], None,
    )

    return _state_64_8192["o"]

# ============================================================
# bs=256, kv_len=1024
# ============================================================

"""bs=256, kv_len=1024 ? aiter a16w8 with page_size=2.

v6: Fix page_size metadata ? kv_indptr in pages, page-based indices.
"""
import torch
import aiter
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1


_state_256_1024 = None

def _run_256_1024(q, kv_data, qo_indptr, kv_indptr, bs, kv_len):
    global _state_256_1024

    PAGE_SIZE = 2
    NUM_KV_SPLITS = 32
    MODE = "a16w8"
    nq, nkv, q_seq_len, dv = NUM_HEADS, NUM_KV_HEADS, 1, V_HEAD_DIM
    total_q = bs * q_seq_len

    if _state_256_1024 is None:
        pages_per_batch = kv_len // PAGE_SIZE  # 1024 / 2 = 512
        total_pages = bs * pages_per_batch

        kv_indptr_pages = torch.zeros(bs + 1, dtype=torch.int32, device="cuda")
        for i in range(bs):
            kv_indptr_pages[i + 1] = kv_indptr_pages[i] + pages_per_batch

        kv_indices = torch.arange(total_pages, dtype=torch.int32, device="cuda")
        kv_last_page_len = torch.full((bs,), PAGE_SIZE, dtype=torch.int32, device="cuda")

        q_dtype, kv_dtype = torch.bfloat16, FP8_DTYPE

        info = get_mla_metadata_info_v1(
            bs, q_seq_len, nq, q_dtype, kv_dtype,
            is_sparse=False, fast_mode=True,
            num_kv_splits=NUM_KV_SPLITS, intra_batch_mode=False,
        )
        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_pages, 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=q_seq_len, uni_seqlen_qo=q_seq_len,
            fast_mode=True, max_split_per_batch=NUM_KV_SPLITS,
            intra_batch_mode=False, dtype_q=q_dtype, dtype_kv=kv_dtype,
        )

        num_partial = reduce_partial_map.size(0)
        _state_256_1024 = {
            "kv_indices": kv_indices,
            "kv_last_page_len": kv_last_page_len,
            "kv_indptr_pages": kv_indptr_pages,
            "work_metadata": 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,
            "logits": torch.empty((num_partial * q_seq_len, 1, nq, dv), dtype=torch.float32, device="cuda"),
            "attn_lse": torch.empty((num_partial * q_seq_len, 1, nq, 1), dtype=torch.float32, device="cuda"),
            "o": torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda"),
        }

    kv_buf, kv_scale = kv_data["fp8"]
    kv_4d = kv_buf.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_buf.shape[-1])

    aiter.mla_decode_stage1_asm_fwd(
        q.view(-1, NUM_HEADS, QK_HEAD_DIM),
        kv_4d,
        qo_indptr, _state_256_1024["kv_indptr_pages"],
        _state_256_1024["kv_indices"], _state_256_1024["kv_last_page_len"],
        None,
        _state_256_1024["work_metadata"], _state_256_1024["work_indptr"], _state_256_1024["work_info_set"],
        1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
        _state_256_1024["logits"], _state_256_1024["attn_lse"], _state_256_1024["o"],
        None, kv_scale,
    )

    aiter.mla_reduce_v1(
        _state_256_1024["logits"], _state_256_1024["attn_lse"],
        _state_256_1024["reduce_indptr"], _state_256_1024["reduce_final_map"], _state_256_1024["reduce_partial_map"],
        1, _state_256_1024["o"], None,
    )

    return _state_256_1024["o"]

# ============================================================
# bs=256, kv_len=8192
# ============================================================

"""aiter a16w8 page_size=8 splits=64 intra_batch_mode=True.

Hypothesis: with bs=256 all sequences same length, intra_batch_mode gives
scheduler uniform work distribution ? better CU utilization.
"""
import torch
import aiter
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1

_state_256_8192 = None

def _run_256_8192(q, kv_data, qo_indptr, kv_indptr, bs, kv_len):
    global _state_256_8192

    PAGE_SIZE = 8
    NUM_KV_SPLITS = 64
    nq, nkv, q_seq_len, dv = NUM_HEADS, NUM_KV_HEADS, 1, V_HEAD_DIM
    total_q = bs * q_seq_len

    if _state_256_8192 is None:
        pages_per_batch = kv_len // PAGE_SIZE
        total_pages = bs * pages_per_batch

        kv_indptr_pages = torch.zeros(bs + 1, dtype=torch.int32, device="cuda")
        for i in range(bs):
            kv_indptr_pages[i + 1] = kv_indptr_pages[i] + pages_per_batch

        kv_indices = torch.arange(total_pages, dtype=torch.int32, device="cuda")
        kv_last_page_len = torch.full((bs,), PAGE_SIZE, dtype=torch.int32, device="cuda")

        q_dtype, kv_dtype = torch.bfloat16, FP8_DTYPE

        info = get_mla_metadata_info_v1(
            bs, q_seq_len, nq, q_dtype, kv_dtype,
            is_sparse=False, fast_mode=True,
            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_pages, 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=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=q_dtype, dtype_kv=kv_dtype,
        )

        num_partial = reduce_partial_map.size(0)
        _state_256_8192 = {
            "kv_indices": kv_indices,
            "kv_last_page_len": kv_last_page_len,
            "kv_indptr_pages": kv_indptr_pages,
            "work_metadata": 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,
            "logits": torch.empty((num_partial * q_seq_len, 1, nq, dv), dtype=torch.float32, device="cuda"),
            "attn_lse": torch.empty((num_partial * q_seq_len, 1, nq, 1), dtype=torch.float32, device="cuda"),
            "o": torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda"),
        }

    kv_buf, kv_scale = kv_data["fp8"]
    kv_4d = kv_buf.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_buf.shape[-1])

    aiter.mla_decode_stage1_asm_fwd(
        q.view(-1, NUM_HEADS, QK_HEAD_DIM),
        kv_4d,
        qo_indptr, _state_256_8192["kv_indptr_pages"],
        _state_256_8192["kv_indices"], _state_256_8192["kv_last_page_len"],
        None,
        _state_256_8192["work_metadata"], _state_256_8192["work_indptr"], _state_256_8192["work_info_set"],
        1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
        _state_256_8192["logits"], _state_256_8192["attn_lse"], _state_256_8192["o"],
        None, kv_scale,
    )

    aiter.mla_reduce_v1(
        _state_256_8192["logits"], _state_256_8192["attn_lse"],
        _state_256_8192["reduce_indptr"], _state_256_8192["reduce_final_map"], _state_256_8192["reduce_partial_map"],
        1, _state_256_8192["o"], None,
    )

    return _state_256_8192["o"]

# ============================================================
# Dispatch
# ============================================================

from task import input_t, output_t

_DISPATCH = {
    (4, 1024): _run_4_1024,
    (4, 8192): _run_4_8192,
    (32, 1024): _run_32_1024,
    (32, 8192): _run_32_8192,
    (64, 1024): _run_64_1024,
    (64, 8192): _run_64_8192,
    (256, 1024): _run_256_1024,
    (256, 8192): _run_256_8192,
}


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = config["batch_size"]
    kv_len = config["kv_seq_len"]
    return _DISPATCH[(bs, kv_len)](q, kv_data, qo_indptr, kv_indptr, bs, kv_len)
scrolls · 1091 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