Skip to content
KernelIndex
Search⌘K

submission 687550

rosehulman. · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6bd9ba7b702ce27669bc2c54ade2da9a08d791e78eccbc56985a352bc1a59ff0
license declaredunknown
license concludedunknown
authorsrosehulman.
imported2026-08-15

Techniques

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

fp4MLA Decode — Aiter A8W8 primary with Triton MXFP4 for bandwidth-bound cases.
fp8scores += tl.dot_scaled(q_tile, None, "e4m3", k_t, k_scale, "e2m1")
mmascores += tl.dot(q_chunk, tl.trans(k_fp8.to(tl.bfloat16)))
num-warps = 4num_warps=4, num_stages=2,
online-softmaxm_new = tl.max(scores, axis=1)
split-kdef _mla_splitk_small(
stages = 2num_warps=4, num_stages=2,

Kernel source

submission.py1304 lines
"""
MLA Decode — Aiter A8W8 primary with Triton MXFP4 for bandwidth-bound cases.

Key design:
    - Aiter A8W8 remains the lowest-risk fast path for general configs.
    - Triton MXFP4 path targets large kvsl=8192 cases where HBM traffic dominates.
    - qsl>1: merge q_pos into head dimension on Triton paths to reuse KV reads.
    - Module-level warmup pre-compiles the Triton variants that dispatch may use.
"""

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

try:
    from aiter import dtypes as aiter_dtypes
    from aiter.mla import mla_decode_fwd
    from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
    import aiter as _aiter_mod
    import aiter.mla as _aiter_mla_mod
    _aiter_available = True
    _stage1_fn = getattr(_aiter_mla_mod, 'mla_decode_stage1_asm_fwd',
                         getattr(_aiter_mod, 'mla_decode_stage1_asm_fwd', None))
    _reduce_fn = getattr(_aiter_mla_mod, 'mla_reduce_v1',
                         getattr(_aiter_mod, 'mla_reduce_v1', None))
    _direct_available = _stage1_fn is not None and _reduce_fn is not None
except Exception:
    _aiter_available = False
    _direct_available = False

QK_DIM = 576
V_DIM = 512
FP8 = torch.float8_e4m3fn
LOG2E = tl.constexpr(1.4426950408889634)


@triton.jit
def _mla_fused(
    Q, KV_FP8,
    pO, pM, pL, Out,
    qo_indptr, kv_indptr, done_counter,
    sm_scale, kv_fp8_scale,
    stride_qt, stride_qh,
    stride_kv,
    stride_pO_s, stride_pO_t, stride_pO_h,
    stride_pm_s, stride_pm_t,
    stride_out_t, stride_out_h,
    gen_target, nheads, qseqlen, n_hg,
    NSPLITS: tl.constexpr,
    BS: tl.constexpr, HPB: tl.constexpr, DK: tl.constexpr,
):
    batch_id = tl.program_id(0)
    split_id = tl.program_id(1)
    hg_id = tl.program_id(2)

    h_start = hg_id * HPB
    h_off = h_start + tl.arange(0, HPB)
    h_mask = h_off < nheads

    q_s = tl.load(qo_indptr + batch_id)
    kv_s = tl.load(kv_indptr + batch_id)
    kv_e = tl.load(kv_indptr + batch_id + 1)
    kv_len = kv_e - kv_s

    chunk = tl.cdiv(kv_len, NSPLITS)
    c_s = kv_s + split_id * chunk
    c_e = tl.minimum(c_s + chunk, kv_e)

    s_idx = tl.arange(0, BS)
    dk_range = tl.arange(0, DK)

    N_FULL_K: tl.constexpr = 576 // DK
    REMAIN_K: tl.constexpr = 576 - N_FULL_K * DK
    N_FULL_V: tl.constexpr = 512 // DK

    combined_scale = sm_scale * kv_fp8_scale

    q_pos = 0
    while q_pos < qseqlen:
        qtok = q_s + q_pos
        mi = tl.full([HPB], float('-inf'), dtype=tl.float32)
        li = tl.zeros([HPB], dtype=tl.float32)
        a0 = tl.zeros([HPB, DK], dtype=tl.float32)
        a1 = tl.zeros([HPB, DK], dtype=tl.float32)
        a2 = tl.zeros([HPB, DK], dtype=tl.float32)
        a3 = tl.zeros([HPB, DK], dtype=tl.float32)

        for t_s in range(c_s, c_e, BS):
            t_e = tl.minimum(t_s + BS, c_e)
            smask = s_idx < (t_e - t_s)
            kv_ids = t_s + s_idx

            scores = tl.zeros([HPB, BS], dtype=tl.float32)
            for dc in range(N_FULL_K):
                q_ptrs = Q + qtok * stride_qt + h_off[:, None] * stride_qh + dc * DK + dk_range[None, :]
                q_chunk = tl.load(q_ptrs, mask=h_mask[:, None], other=0.0)
                k_ptrs = KV_FP8 + kv_ids[:, None] * stride_kv + dc * DK + dk_range[None, :]
                k_fp8 = tl.load(k_ptrs, mask=smask[:, None], other=0.0)
                scores += tl.dot(q_chunk, tl.trans(k_fp8.to(tl.bfloat16)))

            if REMAIN_K > 0:
                r_dk = tl.arange(0, 64)
                q_ptrs = Q + qtok * stride_qt + h_off[:, None] * stride_qh + N_FULL_K * DK + r_dk[None, :]
                q_chunk = tl.load(q_ptrs, mask=h_mask[:, None], other=0.0)
                k_ptrs = KV_FP8 + kv_ids[:, None] * stride_kv + N_FULL_K * DK + r_dk[None, :]
                k_fp8 = tl.load(k_ptrs, mask=smask[:, None], other=0.0)
                scores += tl.dot(q_chunk, tl.trans(k_fp8.to(tl.bfloat16)))

            scores = scores * combined_scale
            scores = tl.where(smask[None, :], scores, float('-inf'))
            m_new = tl.max(scores, axis=1)
            m_max = tl.maximum(mi, m_new)
            exp_old = tl.math.exp2((mi - m_max) * LOG2E)
            exp_s = tl.math.exp2((scores - m_max[:, None]) * LOG2E)
            l_tile = tl.sum(exp_s, axis=1)
            l_comb = li * exp_old + l_tile
            safe_l = tl.where(l_comb > 0.0, l_comb, 1.0)
            r = (li * exp_old / safe_l)[:, None]
            a0 *= r; a1 *= r; a2 *= r; a3 *= r
            probs = exp_s / safe_l[:, None]
            pb = probs.to(tl.bfloat16)

            for vc in range(N_FULL_V):
                v_ptrs = KV_FP8 + kv_ids[:, None] * stride_kv + vc * DK + dk_range[None, :]
                v_fp8 = tl.load(v_ptrs, mask=smask[:, None], other=0.0)
                c = tl.dot(pb, v_fp8.to(tl.bfloat16))
                if vc == 0:   a0 += c
                elif vc == 1: a1 += c
                elif vc == 2: a2 += c
                else:         a3 += c
            mi = m_max
            li = l_comb

        a0 *= kv_fp8_scale; a1 *= kv_fp8_scale; a2 *= kv_fp8_scale; a3 *= kv_fp8_scale
        po_base = pO + split_id * stride_pO_s + qtok * stride_pO_t
        for vc in range(N_FULL_V):
            if vc == 0:   w = a0
            elif vc == 1: w = a1
            elif vc == 2: w = a2
            else:         w = a3
            tl.store(po_base + h_off[:, None] * stride_pO_h + vc * DK + dk_range[None, :],
                     w, mask=h_mask[:, None])
        pm_off = pM + split_id * stride_pm_s + qtok * stride_pm_t + h_off
        pl_off = pL + split_id * stride_pm_s + qtok * stride_pm_t + h_off
        tl.store(pm_off, mi, mask=h_mask)
        tl.store(pl_off, li, mask=h_mask)
        q_pos += 1

    counter_idx = batch_id * n_hg + hg_id
    old_val = tl.atomic_add(done_counter + counter_idx, 1)
    is_last = (old_val + 1) == gen_target

    if is_last:
        N_FULL_V_R: tl.constexpr = 512 // DK
        q_pos2 = 0
        while q_pos2 < qseqlen:
            qtok = q_s + q_pos2
            gm = tl.full([HPB], float('-inf'), dtype=tl.float32)
            gl = tl.zeros([HPB], dtype=tl.float32)
            ra0 = tl.zeros([HPB, DK], dtype=tl.float32)
            ra1 = tl.zeros([HPB, DK], dtype=tl.float32)
            ra2 = tl.zeros([HPB, DK], dtype=tl.float32)
            ra3 = tl.zeros([HPB, DK], dtype=tl.float32)

            for s in range(NSPLITS):
                pm_val = tl.load(pM + s * stride_pm_s + qtok * stride_pm_t + h_off, mask=h_mask, other=float('-inf'))
                pl_val = tl.load(pL + s * stride_pm_s + qtok * stride_pm_t + h_off, mask=h_mask, other=0.0)
                m_new = tl.maximum(gm, pm_val)
                exp_old = tl.math.exp2((gm - m_new) * LOG2E)
                exp_new = tl.math.exp2((pm_val - m_new) * LOG2E)
                gl_new = gl * exp_old + pl_val * exp_new
                safe_gl = tl.where(gl_new > 0.0, gl_new, 1.0)
                r = (gl * exp_old / safe_gl)[:, None]
                ra0 *= r; ra1 *= r; ra2 *= r; ra3 *= r
                f = (pl_val * exp_new / safe_gl)[:, None]
                po_base = pO + s * stride_pO_s + qtok * stride_pO_t + h_off[:, None] * stride_pO_h
                for vc in range(N_FULL_V_R):
                    partial = tl.load(po_base + vc * DK + dk_range[None, :], mask=h_mask[:, None], other=0.0)
                    c = partial * f
                    if vc == 0:   ra0 += c
                    elif vc == 1: ra1 += c
                    elif vc == 2: ra2 += c
                    else:         ra3 += c
                gm = m_new
                gl = gl_new

            out_base = Out + qtok * stride_out_t + h_off[:, None] * stride_out_h
            for vc in range(N_FULL_V_R):
                if vc == 0:   w = ra0
                elif vc == 1: w = ra1
                elif vc == 2: w = ra2
                else:         w = ra3
                tl.store(out_base + vc * DK + dk_range[None, :], w.to(tl.bfloat16), mask=h_mask[:, None])
            q_pos2 += 1


@triton.jit
def _mla_nosplit(
    Q, KV_FP8,
    Out,
    qo_indptr, kv_indptr,
    sm_scale, kv_fp8_scale,
    stride_qt, stride_qh,
    stride_kv,
    stride_out_t, stride_out_h,
    nheads, qseqlen,
    BS: tl.constexpr, HPB: tl.constexpr, DK: tl.constexpr,
):
    batch_id = tl.program_id(0)
    hg_id = tl.program_id(1)
    h_start = hg_id * HPB
    h_off = h_start + tl.arange(0, HPB)
    h_mask = h_off < nheads
    q_s = tl.load(qo_indptr + batch_id)
    kv_s = tl.load(kv_indptr + batch_id)
    kv_e = tl.load(kv_indptr + batch_id + 1)
    s_idx = tl.arange(0, BS)
    dk_range = tl.arange(0, DK)

    N_FULL_K: tl.constexpr = 576 // DK
    REMAIN_K: tl.constexpr = 576 - N_FULL_K * DK
    N_FULL_V: tl.constexpr = 512 // DK

    combined_scale = sm_scale * kv_fp8_scale

    q_pos = 0
    while q_pos < qseqlen:
        qtok = q_s + q_pos
        mi = tl.full([HPB], float('-inf'), dtype=tl.float32)
        li = tl.zeros([HPB], dtype=tl.float32)
        a0 = tl.zeros([HPB, DK], dtype=tl.float32)
        a1 = tl.zeros([HPB, DK], dtype=tl.float32)
        a2 = tl.zeros([HPB, DK], dtype=tl.float32)
        a3 = tl.zeros([HPB, DK], dtype=tl.float32)
        for t_s in range(kv_s, kv_e, BS):
            t_e = tl.minimum(t_s + BS, kv_e)
            smask = s_idx < (t_e - t_s)
            kv_ids = t_s + s_idx

            scores = tl.zeros([HPB, BS], dtype=tl.float32)
            for dc in range(N_FULL_K):
                q_ptrs = Q + qtok * stride_qt + h_off[:, None] * stride_qh + dc * DK + dk_range[None, :]
                q_chunk = tl.load(q_ptrs, mask=h_mask[:, None], other=0.0)
                k_ptrs = KV_FP8 + kv_ids[:, None] * stride_kv + dc * DK + dk_range[None, :]
                k_fp8 = tl.load(k_ptrs, mask=smask[:, None], other=0.0)
                scores += tl.dot(q_chunk, tl.trans(k_fp8.to(tl.bfloat16)))

            if REMAIN_K > 0:
                r_dk = tl.arange(0, 64)
                q_ptrs = Q + qtok * stride_qt + h_off[:, None] * stride_qh + N_FULL_K * DK + r_dk[None, :]
                q_chunk = tl.load(q_ptrs, mask=h_mask[:, None], other=0.0)
                k_ptrs = KV_FP8 + kv_ids[:, None] * stride_kv + N_FULL_K * DK + r_dk[None, :]
                k_fp8 = tl.load(k_ptrs, mask=smask[:, None], other=0.0)
                scores += tl.dot(q_chunk, tl.trans(k_fp8.to(tl.bfloat16)))

            scores = scores * combined_scale
            scores = tl.where(smask[None, :], scores, float('-inf'))
            m_new = tl.max(scores, axis=1)
            m_max = tl.maximum(mi, m_new)
            exp_old = tl.math.exp2((mi - m_max) * LOG2E)
            exp_s = tl.math.exp2((scores - m_max[:, None]) * LOG2E)
            l_tile = tl.sum(exp_s, axis=1)
            l_comb = li * exp_old + l_tile
            safe_l = tl.where(l_comb > 0.0, l_comb, 1.0)
            r = (li * exp_old / safe_l)[:, None]
            a0 *= r; a1 *= r; a2 *= r; a3 *= r
            probs = exp_s / safe_l[:, None]
            pb = probs.to(tl.bfloat16)
            for vc in range(N_FULL_V):
                v_ptrs = KV_FP8 + kv_ids[:, None] * stride_kv + vc * DK + dk_range[None, :]
                v_fp8 = tl.load(v_ptrs, mask=smask[:, None], other=0.0)
                c = tl.dot(pb, v_fp8.to(tl.bfloat16))
                if vc == 0:   a0 += c
                elif vc == 1: a1 += c
                elif vc == 2: a2 += c
                else:         a3 += c
            mi = m_max
            li = l_comb
        out_base = Out + qtok * stride_out_t
        for vc in range(N_FULL_V):
            if vc == 0:   w = a0
            elif vc == 1: w = a1
            elif vc == 2: w = a2
            else:         w = a3
            w_scaled = (w * kv_fp8_scale).to(tl.bfloat16)
            tl.store(out_base + h_off[:, None] * stride_out_h + vc * DK + dk_range[None, :],
                     w_scaled, mask=h_mask[:, None])
        q_pos += 1


@triton.jit
def _mla_splitk_small(
    Q, KV_FP8,
    pO, pM, pL,
    qo_indptr, kv_indptr,
    sm_scale, kv_fp8_scale,
    stride_qt, stride_qh,
    stride_kv,
    stride_pO_s, stride_pO_t, stride_pO_h,
    stride_pm_s, stride_pm_t,
    nheads,
    NSPLITS: tl.constexpr,
    BS: tl.constexpr, HPB: tl.constexpr, QSEQLEN: tl.constexpr, DK: tl.constexpr,
):
    pid0 = tl.program_id(0)
    pid1 = tl.program_id(1)

    batch_id = pid0 // NSPLITS
    split_id = pid0 % NSPLITS
    h_start = pid1 * HPB
    h_off = h_start + tl.arange(0, HPB)
    h_mask = h_off < nheads

    q_s = tl.load(qo_indptr + batch_id)
    kv_s = tl.load(kv_indptr + batch_id)
    kv_e = tl.load(kv_indptr + batch_id + 1)
    kv_len = kv_e - kv_s

    chunk = tl.cdiv(kv_len, NSPLITS)
    c_s = kv_s + split_id * chunk
    c_e = tl.minimum(c_s + chunk, kv_e)

    s_idx = tl.arange(0, BS)
    dk_range = tl.arange(0, DK)

    N_FULL_K: tl.constexpr = 576 // DK
    REMAIN_K: tl.constexpr = 576 - N_FULL_K * DK
    N_FULL_V: tl.constexpr = 512 // DK

    for q_pos in range(QSEQLEN):
        qtok = q_s + q_pos
        mi = tl.full([HPB], float('-inf'), dtype=tl.float32)
        li = tl.zeros([HPB], dtype=tl.float32)

        a0 = tl.zeros([HPB, DK], dtype=tl.float32)
        a1 = tl.zeros([HPB, DK], dtype=tl.float32)
        a2 = tl.zeros([HPB, DK], dtype=tl.float32)
        a3 = tl.zeros([HPB, DK], dtype=tl.float32)

        for t_s in range(c_s, c_e, BS):
            t_e = tl.minimum(t_s + BS, c_e)
            smask = s_idx < (t_e - t_s)
            kv_ids = t_s + s_idx

            scores = tl.zeros([HPB, BS], dtype=tl.float32)
            for dc in range(N_FULL_K):
                q_ptrs = Q + qtok * stride_qt + h_off[:, None] * stride_qh + dc * DK + dk_range[None, :]
                q_chunk = tl.load(q_ptrs, mask=h_mask[:, None], other=0.0)
                k_ptrs = KV_FP8 + kv_ids[:, None] * stride_kv + dc * DK + dk_range[None, :]
                k_fp8 = tl.load(k_ptrs, mask=smask[:, None], other=0.0)
                scores += tl.dot(q_chunk, tl.trans(k_fp8.to(tl.bfloat16)))

            if REMAIN_K > 0:
                r_range = tl.arange(0, 64)
                q_ptrs = Q + qtok * stride_qt + h_off[:, None] * stride_qh + N_FULL_K * DK + r_range[None, :]
                q_chunk = tl.load(q_ptrs, mask=h_mask[:, None], other=0.0)
                k_ptrs = KV_FP8 + kv_ids[:, None] * stride_kv + N_FULL_K * DK + r_range[None, :]
                k_fp8 = tl.load(k_ptrs, mask=smask[:, None], other=0.0)
                scores += tl.dot(q_chunk, tl.trans(k_fp8.to(tl.bfloat16)))

            scores = scores * sm_scale
            scores = tl.where(smask[None, :], scores, float('-inf'))

            m_new = tl.max(scores, axis=1)
            m_max = tl.maximum(mi, m_new)
            exp_old = tl.math.exp2((mi - m_max) * LOG2E)
            exp_s = tl.math.exp2((scores - m_max[:, None]) * LOG2E)
            l_tile = tl.sum(exp_s, axis=1)
            l_comb = li * exp_old + l_tile
            safe_l = tl.where(l_comb > 0.0, l_comb, 1.0)

            r = (li * exp_old / safe_l)[:, None]
            a0 *= r; a1 *= r; a2 *= r; a3 *= r

            probs = exp_s / safe_l[:, None]
            pb = probs.to(tl.bfloat16)

            for vc in range(N_FULL_V):
                v_ptrs = KV_FP8 + kv_ids[:, None] * stride_kv + vc * DK + dk_range[None, :]
                v_fp8 = tl.load(v_ptrs, mask=smask[:, None], other=0.0)
                c = tl.dot(pb, v_fp8.to(tl.bfloat16))
                if vc == 0:
                    a0 += c
                elif vc == 1:
                    a1 += c
                elif vc == 2:
                    a2 += c
                else:
                    a3 += c

            mi = m_max
            li = l_comb

        a0 *= kv_fp8_scale; a1 *= kv_fp8_scale; a2 *= kv_fp8_scale; a3 *= kv_fp8_scale

        po_base = pO + split_id * stride_pO_s + qtok * stride_pO_t
        for vc in range(N_FULL_V):
            if vc == 0:
                w = a0
            elif vc == 1:
                w = a1
            elif vc == 2:
                w = a2
            else:
                w = a3
            tl.store(po_base + h_off[:, None] * stride_pO_h + vc * DK + dk_range[None, :],
                     w, mask=h_mask[:, None])

        pm_off = pM + split_id * stride_pm_s + qtok * stride_pm_t + h_off
        pl_off = pL + split_id * stride_pm_s + qtok * stride_pm_t + h_off
        tl.store(pm_off, mi, mask=h_mask)
        tl.store(pl_off, li, mask=h_mask)


@triton.jit
def _mla_reduce_small(
    pO, pM, pL, Out,
    stride_pO_s, stride_pm_s,
    NSPLITS: tl.constexpr,
    VDIM: tl.constexpr,
    BLK: tl.constexpr,
):
    th = tl.program_id(0)
    gm = tl.full([], float('-inf'), dtype=tl.float32)
    for s in range(NSPLITS):
        gm = tl.maximum(gm, tl.load(pM + s * stride_pm_s + th))
    gl = tl.zeros([], dtype=tl.float32)
    for s in range(NSPLITS):
        pm = tl.load(pM + s * stride_pm_s + th)
        pl = tl.load(pL + s * stride_pm_s + th)
        gl += pl * tl.math.exp2((pm - gm) * LOG2E)
    safe_gl = tl.where(gl > 0.0, gl, 1.0)
    d = tl.arange(0, BLK)
    for d_start in range(0, VDIM, BLK):
        acc = tl.zeros([BLK], dtype=tl.float32)
        for s in range(NSPLITS):
            pm = tl.load(pM + s * stride_pm_s + th)
            pl = tl.load(pL + s * stride_pm_s + th)
            w = pl * tl.math.exp2((pm - gm) * LOG2E)
            acc += tl.load(pO + s * stride_pO_s + th * VDIM + d_start + d) * w
        tl.store(Out + th * VDIM + d_start + d, (acc / safe_gl).to(tl.bfloat16))


# ═══════════════════════════════════════════════════════════════════════
#  MXFP4 kernels — ~1.9x bandwidth savings for large configs
# ═══════════════════════════════════════════════════════════════════════

@triton.jit
def _fp4e2m1_lookup(u):
    """Convert unsigned 3-bit fp4 magnitude (0-7) to float32."""
    return tl.where(u < 4,
               tl.where(u < 2,
                        tl.where(u == 0, 0.0, 0.5),
                        tl.where(u == 2, 1.0, 1.5)),
               tl.where(u < 6,
                        tl.where(u == 4, 2.0, 3.0),
                        tl.where(u == 6, 4.0, 6.0)))


@triton.jit
def _mla_hybrid_splitk(
    Q_FP8, KV_PACKED, KV_SCALE,
    pO, pM, pL,
    qo_indptr, kv_indptr,
    sm_scale,
    stride_qt, stride_qh,
    stride_kvp, stride_ks,
    stride_pO_s, stride_pO_t, stride_pO_h,
    stride_pm_s, stride_pm_t,
    nheads,
    NSPLITS: tl.constexpr,
    BS: tl.constexpr, HPB: tl.constexpr, DK: tl.constexpr,
):
    N_K_FULL: tl.constexpr = 576 // DK
    K_REM: tl.constexpr = 576 - N_K_FULL * DK
    K_REM_PACKED: tl.constexpr = K_REM // 2
    K_REM_NSB: tl.constexpr = K_REM // 32
    N_V: tl.constexpr = 512 // DK
    PKD: tl.constexpr = DK // 2
    NSB: tl.constexpr = DK // 32
    HALF_DK: tl.constexpr = DK // 2

    pid0 = tl.program_id(0)
    hg_id = tl.program_id(1)
    batch_id = pid0 // NSPLITS
    split_id = pid0 % NSPLITS

    h_start = hg_id * HPB
    h_off = h_start + tl.arange(0, HPB)
    h_mask = h_off < nheads

    q_s = tl.load(qo_indptr + batch_id)
    kv_s = tl.load(kv_indptr + batch_id)
    kv_e = tl.load(kv_indptr + batch_id + 1)
    kv_len = kv_e - kv_s

    chunk_size = tl.cdiv(kv_len, NSPLITS)
    c_s = kv_s + split_id * chunk_size
    c_e = tl.minimum(c_s + chunk_size, kv_e)

    s_idx = tl.arange(0, BS)
    dk_range = tl.arange(0, DK)
    pk_range = tl.arange(0, PKD)
    sb_range = tl.arange(0, NSB)
    half_range = tl.arange(0, HALF_DK)
    scale_map = pk_range // 16

    qtok = q_s
    mi = tl.full([HPB], float('-inf'), dtype=tl.float32)
    li = tl.zeros([HPB], dtype=tl.float32)
    ae0 = tl.zeros([HPB, HALF_DK], dtype=tl.float32)
    ao0 = tl.zeros([HPB, HALF_DK], dtype=tl.float32)
    ae1 = tl.zeros([HPB, HALF_DK], dtype=tl.float32)
    ao1 = tl.zeros([HPB, HALF_DK], dtype=tl.float32)
    ae2 = tl.zeros([HPB, HALF_DK], dtype=tl.float32)
    ao2 = tl.zeros([HPB, HALF_DK], dtype=tl.float32)
    ae3 = tl.zeros([HPB, HALF_DK], dtype=tl.float32)
    ao3 = tl.zeros([HPB, HALF_DK], dtype=tl.float32)

    for t_s in range(c_s, c_e, BS):
        t_e = tl.minimum(t_s + BS, c_e)
        smask = s_idx < (t_e - t_s)
        kv_ids = t_s + s_idx

        scores = tl.zeros([HPB, BS], dtype=tl.float32)

        for dc in range(N_K_FULL):
            q_ptrs = Q_FP8 + qtok * stride_qt + h_off[:, None] * stride_qh + dc * DK + dk_range[None, :]
            q_tile = tl.load(q_ptrs, mask=h_mask[:, None], other=0.0)

            k_ptrs = KV_PACKED + kv_ids[:, None] * stride_kvp + dc * PKD + pk_range[None, :]
            k_packed = tl.load(k_ptrs, mask=smask[:, None], other=0)
            k_t = tl.trans(k_packed)

            ks_ptrs = KV_SCALE + kv_ids[:, None] * stride_ks + dc * NSB + sb_range[None, :]
            k_scale = tl.load(ks_ptrs, mask=smask[:, None], other=127)

            scores += tl.dot_scaled(q_tile, None, "e4m3", k_t, k_scale, "e2m1")

        if K_REM > 0:
            rem_range = tl.arange(0, K_REM)
            rem_pk_range = tl.arange(0, K_REM_PACKED)
            rem_sb_range = tl.arange(0, K_REM_NSB)

            qr_ptrs = Q_FP8 + qtok * stride_qt + h_off[:, None] * stride_qh + N_K_FULL * DK + rem_range[None, :]
            qr = tl.load(qr_ptrs, mask=h_mask[:, None], other=0.0)

            kr_ptrs = KV_PACKED + kv_ids[:, None] * stride_kvp + N_K_FULL * PKD + rem_pk_range[None, :]
            kr = tl.load(kr_ptrs, mask=smask[:, None], other=0)
            kr_t = tl.trans(kr)

            ksr_ptrs = KV_SCALE + kv_ids[:, None] * stride_ks + N_K_FULL * NSB + rem_sb_range[None, :]
            kr_scale = tl.load(ksr_ptrs, mask=smask[:, None], other=127)

            scores += tl.dot_scaled(qr, None, "e4m3", kr_t, kr_scale, "e2m1")

        scores = scores * sm_scale
        scores = tl.where(smask[None, :], scores, float('-inf'))
        m_new = tl.max(scores, axis=1)
        m_max = tl.maximum(mi, m_new)
        exp_old = tl.math.exp2((mi - m_max) * LOG2E)
        exp_s = tl.math.exp2((scores - m_max[:, None]) * LOG2E)
        l_tile = tl.sum(exp_s, axis=1)
        l_comb = li * exp_old + l_tile
        safe_l = tl.where(l_comb > 0.0, l_comb, 1.0)
        r = (li * exp_old / safe_l)[:, None]
        ae0 *= r; ao0 *= r; ae1 *= r; ao1 *= r
        ae2 *= r; ao2 *= r; ae3 *= r; ao3 *= r
        probs = exp_s / safe_l[:, None]
        pb = probs.to(tl.bfloat16)

        for vc in range(N_V):
            vp_ptrs = KV_PACKED + kv_ids[:, None] * stride_kvp + vc * PKD + pk_range[None, :]
            vp = tl.load(vp_ptrs, mask=smask[:, None], other=0).to(tl.uint8)

            sc_col = vc * NSB + scale_map
            vs_ptrs = KV_SCALE + kv_ids[:, None] * stride_ks + sc_col[None, :]
            vs = tl.math.exp2(tl.load(vs_ptrs, mask=smask[:, None], other=127).to(tl.float32) - 127.0)

            vl = _fp4e2m1_lookup(vp & 7) * (1.0 - ((vp >> 3) & 1).to(tl.float32) * 2.0) * vs
            vh = _fp4e2m1_lookup((vp >> 4) & 7) * (1.0 - ((vp >> 7) & 1).to(tl.float32) * 2.0) * vs

            ce = tl.dot(pb, vl.to(tl.bfloat16))
            co = tl.dot(pb, vh.to(tl.bfloat16))
            if vc == 0:   ae0 += ce; ao0 += co
            elif vc == 1: ae1 += ce; ao1 += co
            elif vc == 2: ae2 += ce; ao2 += co
            else:         ae3 += ce; ao3 += co
        mi = m_max
        li = l_comb

    po_base = pO + split_id * stride_pO_s + qtok * stride_pO_t
    for vc in range(N_V):
        if vc == 0:   we, wo = ae0, ao0
        elif vc == 1: we, wo = ae1, ao1
        elif vc == 2: we, wo = ae2, ao2
        else:         we, wo = ae3, ao3
        even_offs = vc * DK + 2 * half_range
        odd_offs = even_offs + 1
        tl.store(po_base + h_off[:, None] * stride_pO_h + even_offs[None, :], we, mask=h_mask[:, None])
        tl.store(po_base + h_off[:, None] * stride_pO_h + odd_offs[None, :], wo, mask=h_mask[:, None])
    pm_off = pM + split_id * stride_pm_s + qtok * stride_pm_t + h_off
    pl_off = pL + split_id * stride_pm_s + qtok * stride_pm_t + h_off
    tl.store(pm_off, mi, mask=h_mask)
    tl.store(pl_off, li, mask=h_mask)


@triton.jit
def _mla_mxfp4_reduce(
    pO, pM, pL, Out,
    stride_pO_s, stride_pO_t, stride_pO_h,
    stride_pm_s, stride_pm_t,
    stride_out_t, stride_out_h,
    nheads,
    NSPLITS: tl.constexpr,
    HPB: tl.constexpr, DK: tl.constexpr,
):
    qtok = tl.program_id(0)
    hg_id = tl.program_id(1)
    h_start = hg_id * HPB
    h_off = h_start + tl.arange(0, HPB)
    h_mask = h_off < nheads
    dk_range = tl.arange(0, DK)

    N_FULL_V: tl.constexpr = 512 // DK

    gm = tl.full([HPB], float('-inf'), dtype=tl.float32)
    gl = tl.zeros([HPB], dtype=tl.float32)
    ra0 = tl.zeros([HPB, DK], dtype=tl.float32)
    ra1 = tl.zeros([HPB, DK], dtype=tl.float32)
    ra2 = tl.zeros([HPB, DK], dtype=tl.float32)
    ra3 = tl.zeros([HPB, DK], dtype=tl.float32)

    for s in range(NSPLITS):
        pm_val = tl.load(pM + s * stride_pm_s + qtok * stride_pm_t + h_off,
                         mask=h_mask, other=float('-inf'))
        pl_val = tl.load(pL + s * stride_pm_s + qtok * stride_pm_t + h_off,
                         mask=h_mask, other=0.0)
        m_new = tl.maximum(gm, pm_val)
        exp_old = tl.math.exp2((gm - m_new) * LOG2E)
        exp_new = tl.math.exp2((pm_val - m_new) * LOG2E)
        gl_new = gl * exp_old + pl_val * exp_new
        safe_gl = tl.where(gl_new > 0.0, gl_new, 1.0)
        r = (gl * exp_old / safe_gl)[:, None]
        ra0 *= r; ra1 *= r; ra2 *= r; ra3 *= r
        f = (pl_val * exp_new / safe_gl)[:, None]
        po_base = pO + s * stride_pO_s + qtok * stride_pO_t + h_off[:, None] * stride_pO_h
        for vc in range(N_FULL_V):
            partial = tl.load(po_base + vc * DK + dk_range[None, :],
                              mask=h_mask[:, None], other=0.0)
            c = partial * f
            if vc == 0:   ra0 += c
            elif vc == 1: ra1 += c
            elif vc == 2: ra2 += c
            else:         ra3 += c
        gm = m_new
        gl = gl_new

    out_base = Out + qtok * stride_out_t + h_off[:, None] * stride_out_h
    for vc in range(N_FULL_V):
        if vc == 0:   w = ra0
        elif vc == 1: w = ra1
        elif vc == 2: w = ra2
        else:         w = ra3
        tl.store(out_base + vc * DK + dk_range[None, :],
                 w.to(tl.bfloat16), mask=h_mask[:, None])


# ═══════════════════════════════════════════════════════════════════════
#  Warmup
# ═══════════════════════════════════════════════════════════════════════

_warmed = False

def _do_warmup():
    global _warmed
    if _warmed:
        return
    _warmed = True
    BS_W = 128; DK_W = 128

    qi_w = torch.tensor([0, 1], dtype=torch.int32, device="cuda")
    ki_w = torch.tensor([0, BS_W + 1], dtype=torch.int32, device="cuda")
    kv_fp8_w = torch.zeros((BS_W + 1, QK_DIM), dtype=FP8, device="cuda")
    kv_mx_w = torch.zeros((BS_W + 1, QK_DIM // 2), dtype=torch.uint8, device="cuda")
    kv_ms_w = torch.full((BS_W + 1, 24), 127, dtype=torch.uint8, device="cuda")

    for HPB_W in [16, 64]:
        for n_hg_v in [1, 2]:
            nh_v = n_hg_v * HPB_W
            q_w = torch.zeros((1, nh_v, QK_DIM), dtype=torch.bfloat16, device="cuda")
            qf_w = torch.zeros((1, nh_v, QK_DIM), dtype=FP8, device="cuda")
            o_w = torch.zeros((1, nh_v, V_DIM), dtype=torch.bfloat16, device="cuda")

            _mla_nosplit[(1, n_hg_v)](
                q_w, kv_fp8_w, o_w,
                qi_w, ki_w,
                1.0, 1.0,
                q_w.stride(0), q_w.stride(1),
                kv_fp8_w.stride(0),
                o_w.stride(0), o_w.stride(1),
                nh_v, 1,
                BS=BS_W, HPB=HPB_W, DK=DK_W,
                num_warps=4, num_stages=2,
            )
            for NS in [2, 4, 8, 16, 32, 64]:
                pO_w = torch.zeros((NS, 1, nh_v, V_DIM), dtype=torch.float32, device="cuda")
                pM_w = torch.zeros((NS, 1, nh_v), dtype=torch.float32, device="cuda")
                pL_w = torch.zeros((NS, 1, nh_v), dtype=torch.float32, device="cuda")
                ctr_w = torch.zeros(n_hg_v, dtype=torch.int32, device="cuda")
                _mla_fused[(1, NS, n_hg_v)](
                    q_w, kv_fp8_w,
                    pO_w, pM_w, pL_w, o_w,
                    qi_w, ki_w, ctr_w,
                    1.0, 1.0,
                    q_w.stride(0), q_w.stride(1),
                    kv_fp8_w.stride(0),
                    pO_w.stride(0), pO_w.stride(1), pO_w.stride(2),
                    pM_w.stride(0), pM_w.stride(1),
                    o_w.stride(0), o_w.stride(1),
                    NS, nh_v, 1, n_hg_v,
                    NSPLITS=NS, BS=BS_W, HPB=HPB_W, DK=DK_W,
                    num_warps=4, num_stages=2,
                )
            if HPB_W == 16:
                for NS in [4, 8]:
                    pO_w = torch.zeros((NS, 1, nh_v, V_DIM), dtype=torch.float32, device="cuda")
                    pM_w = torch.zeros((NS, 1, nh_v), dtype=torch.float32, device="cuda")
                    pL_w = torch.zeros((NS, 1, nh_v), dtype=torch.float32, device="cuda")
                    _mla_splitk_small[(NS, n_hg_v)](
                        q_w, kv_fp8_w,
                        pO_w, pM_w, pL_w,
                        qi_w, ki_w,
                        1.0, 1.0,
                        q_w.stride(0), q_w.stride(1),
                        kv_fp8_w.stride(0),
                        pO_w.stride(0), pO_w.stride(1), pO_w.stride(2),
                        pM_w.stride(0), pM_w.stride(1),
                        nh_v,
                        NSPLITS=NS, BS=BS_W, HPB=HPB_W, QSEQLEN=1, DK=DK_W,
                        num_warps=4, num_stages=2,
                    )
                    pO_flat_w = pO_w.view(NS, -1)
                    pM_flat_w = pM_w.view(NS, -1)
                    pL_flat_w = pL_w.view(NS, -1)
                    _mla_reduce_small[(nh_v,)](
                        pO_flat_w, pM_flat_w, pL_flat_w, o_w.view(-1),
                        pO_flat_w.stride(0), pM_flat_w.stride(0),
                        NSPLITS=NS, VDIM=V_DIM, BLK=128,
                        num_warps=2, num_stages=1,
                    )
            for NS in [4, 8, 16, 32]:
                pO_w = torch.zeros((NS, 1, nh_v, V_DIM), dtype=torch.float32, device="cuda")
                pM_w = torch.zeros((NS, 1, nh_v), dtype=torch.float32, device="cuda")
                pL_w = torch.zeros((NS, 1, nh_v), dtype=torch.float32, device="cuda")
                qf_w.zero_()
                _mla_hybrid_splitk[(NS, n_hg_v)](
                    qf_w, kv_mx_w, kv_ms_w,
                    pO_w, pM_w, pL_w,
                    qi_w, ki_w,
                    1.0,
                    qf_w.stride(0), qf_w.stride(1),
                    kv_mx_w.stride(0), kv_ms_w.stride(0),
                    pO_w.stride(0), pO_w.stride(1), pO_w.stride(2),
                    pM_w.stride(0), pM_w.stride(1),
                    nh_v,
                    NSPLITS=NS, BS=BS_W, HPB=HPB_W, DK=DK_W,
                    num_warps=4, num_stages=2,
                )
                _mla_mxfp4_reduce[(1, n_hg_v)](
                    pO_w, pM_w, pL_w, o_w,
                    pO_w.stride(0), pO_w.stride(1), pO_w.stride(2),
                    pM_w.stride(0), pM_w.stride(1),
                    o_w.stride(0), o_w.stride(1),
                    nh_v,
                    NSPLITS=NS, HPB=HPB_W, DK=DK_W,
                    num_warps=2, num_stages=1,
                )
    torch.cuda.synchronize()
    torch.cuda.empty_cache()


# ═══════════════════════════════════════════════════════════════════════
#  Aiter FP8 — lean dispatch with pre-computed args
# ═══════════════════════════════════════════════════════════════════════

_NS_TABLE = {
    (4, 1024): 16, (4, 8192): 32,
    (32, 1024): 16, (32, 8192): 32,
    (64, 1024): 16, (64, 8192): 64,
    (256, 1024): 64, (256, 8192): 64,
}

class _AiterEntry:
    __slots__ = ['qf_buf', 'qf_view', 'o', 'kvb4', 'kvi', 'klp',
                 'qs', 'kvs', 'meta', 'ns', 'sms', 'nkv', 'qsl',
                 'wk', 'logits', 'attn_lse', 'q_folded', 'o_folded',
                 'use_direct', 'direct_safe', 'q_holder', 'graph',
                 'qo_indptr_own', 'kv_indptr_own']

_aiter_table = {}

def _aiter_setup(q, kv_data, qo_indptr, kv_indptr, config, key):
    FP8_A = aiter_dtypes.fp8
    nh = config["num_heads"]; nkv = config["num_kv_heads"]
    dqk = config["qk_head_dim"]; dv = config["v_head_dim"]
    bs = config["batch_size"]; qsl = config["q_seq_len"]
    kvsl = config["kv_seq_len"]; sms = config["sm_scale"]

    qf_buf = torch.empty(q.shape, dtype=FP8_A, device="cuda")
    qs = torch.ones(1, dtype=torch.float32, device="cuda")

    kvb, kvs = kv_data["fp8"]
    tkv = bs * kvsl
    kvi = torch.arange(tkv, dtype=torch.int32, device="cuda")
    klp = torch.full((bs,), kvsl, dtype=torch.int32, device="cuda")
    kvb4 = kvb.view(kvb.shape[0], 1, nkv, kvb.shape[-1])

    qo_own = torch.arange(0, bs + 1, dtype=torch.int32, device="cuda") * qsl
    kv_own = torch.arange(0, (bs + 1) * kvsl, kvsl, dtype=torch.int32, device="cuda")

    ns = 64 if qsl == 1 else _NS_TABLE.get((bs, kvsl), 64)
    ibm = True

    info = get_mla_metadata_info_v1(
        bs, qsl, nh, qf_buf.dtype, kvb.dtype, is_sparse=False, fast_mode=False,
        num_kv_splits=ns, intra_batch_mode=ibm)
    wk = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
    get_mla_metadata_v1(
        qo_indptr, kv_indptr, klp, nh // nkv, nkv, True,
        wk[0], wk[2], wk[1], wk[3], wk[4], wk[5],
        page_size=1, kv_granularity=16,
        max_seqlen_qo=qsl, uni_seqlen_qo=qsl, fast_mode=False,
        max_split_per_batch=ns, intra_batch_mode=ibm,
        dtype_q=qf_buf.dtype, dtype_kv=kvb.dtype)

    total_q = q.shape[0]
    e = _AiterEntry()
    e.qf_buf = qf_buf
    e.qf_view = qf_buf.view(-1, nh, dqk)
    e.o = torch.empty((total_q, nh, dv), dtype=torch.bfloat16, device="cuda")
    e.kvb4 = kvb4
    e.kvi = kvi
    e.klp = klp
    e.qs = qs
    e.kvs = kvs
    e.meta = dict(
        work_meta_data=wk[0], work_indptr=wk[1], work_info_set=wk[2],
        reduce_indptr=wk[3], reduce_final_map=wk[4], reduce_partial_map=wk[5])
    e.ns = ns
    e.sms = sms
    e.nkv = nkv
    e.qsl = qsl
    e.wk = wk
    e.qo_indptr_own = qo_own
    e.kv_indptr_own = kv_own
    e.use_direct = False
    e.direct_safe = False
    e.graph = None
    e.q_holder = None

    if _direct_available:
        try:
            if qsl == 1:
                fold = max(nh // 16, 1)
                nh_f = 16 if nh >= 16 else nh
                total_s_f = total_q * fold
                e.q_folded = e.qf_view.view(total_s_f, nh_f, dqk)
                e.o_folded = e.o.view(total_s_f, nh_f, dv)
            else:
                raise RuntimeError(f"direct path unsupported for nh={nh}, qsl={qsl}")

            rpm_sz = wk[5].size(0)
            e.logits = torch.empty((rpm_sz * qsl, 1, nh_f, dv),
                                   dtype=torch.float32, device="cuda")
            e.attn_lse = torch.empty((rpm_sz * qsl, 1, nh_f, 1),
                                     dtype=torch.float32, device="cuda")
            e.q_holder = torch.empty_like(q)
            e.use_direct = True
            e.direct_safe = True

            # Warm up the direct path to trigger JIT/compilation
            e.q_holder.copy_(q)
            e.qf_buf.copy_(e.q_holder)
            _stage1_fn(
                e.q_folded, kvb4, qo_own, kv_own, kvi, klp,
                None, wk[0], wk[1], wk[2],
                qsl, 1, nkv, sms,
                e.logits, e.attn_lse, e.o_folded,
                qs, kvs)
            _reduce_fn(
                e.logits, e.attn_lse, wk[3], wk[4], wk[5],
                qsl, e.o_folded, None)
            torch.cuda.synchronize()
        except Exception as ex:
            print(f"[MLA] Direct/graph setup failed bs={bs},kvsl={kvsl}: {ex}", file=sys.stderr)
            e.use_direct = False
            e.graph = None

    _aiter_table[key] = e
    return e


def _aiter_refresh_inputs(e, kv_data):
    kvb, kvs = kv_data["fp8"]
    e.kvb4 = kvb.view(kvb.shape[0], 1, e.nkv, kvb.shape[-1])
    e.kvs = kvs


# ═══════════════════════════════════════════════════════════════════════
#  Host dispatch
# ═══════════════════════════════════════════════════════════════════════

_prev_buf_key = None
_prev_bufs = None
_prev_out_key = None
_prev_out = None
_counter_cache = {}
_gen_counter = {}
_qo_merged_cache = {}
_config_cache = {}
_mxfp4_config_cache = {}
_q_fp8_cache = {}
_kv_fp8_view_cache = {}
_kv_fp8_scale_cache = {}

_SMALL_TRITON_NS = {
    4: 8,
    32: 8,
    64: 4,
}


def _get_bufs(nsplits, total_q, nh):
    global _prev_buf_key, _prev_bufs
    key = (nsplits, total_q, nh)
    if key != _prev_buf_key:
        pO = torch.empty((nsplits, total_q, nh, V_DIM), dtype=torch.float32, device="cuda")
        pM = torch.empty((nsplits, total_q, nh), dtype=torch.float32, device="cuda")
        pL = torch.empty((nsplits, total_q, nh), dtype=torch.float32, device="cuda")
        _prev_buf_key = key
        _prev_bufs = (pO, pM, pL)
    return _prev_bufs

def _get_output(total_q, nh):
    global _prev_out_key, _prev_out
    key = (total_q, nh)
    if key != _prev_out_key:
        _prev_out_key = key
        _prev_out = torch.empty((total_q, nh, V_DIM), dtype=torch.bfloat16, device="cuda")
    return _prev_out

def _get_counter(bs, n_hg):
    key = (bs, n_hg)
    if key not in _counter_cache:
        _counter_cache[key] = torch.zeros(bs * n_hg, dtype=torch.int32, device="cuda")
        _gen_counter[key] = 0
    return _counter_cache[key]

def _get_qo_merged(bs):
    if bs not in _qo_merged_cache:
        _qo_merged_cache[bs] = torch.arange(bs + 1, dtype=torch.int32, device="cuda")
    return _qo_merged_cache[bs]


def _get_q_fp8(shape):
    key = tuple(shape)
    if key not in _q_fp8_cache:
        _q_fp8_cache[key] = torch.empty(shape, dtype=FP8, device="cuda")
    return _q_fp8_cache[key]


def _get_kv_fp8_view(kvb_fp8):
    key = kvb_fp8.data_ptr()
    view = _kv_fp8_view_cache.get(key)
    if view is None:
        view = kvb_fp8.reshape(-1, QK_DIM)
        _kv_fp8_view_cache[key] = view
    return view


def _get_kv_fp8_scale(kvs_fp8):
    key = kvs_fp8.data_ptr()
    scale = _kv_fp8_scale_cache.get(key)
    if scale is None:
        scale = kvs_fp8.float().item()
        _kv_fp8_scale_cache[key] = scale
    return scale


def _get_uint8_view(tensor):
    key = tensor.data_ptr()
    view = _config_cache.get(("view", key))
    if view is None:
        flat = tensor.reshape(-1, tensor.shape[-1])
        view = flat.view(torch.uint8).reshape(flat.shape[0], flat.shape[1])
        _config_cache[("view", key)] = view
    return view


def _should_use_mxfp4(config):
    bs = config["batch_size"]
    qsl = config["q_seq_len"]
    kvsl = config["kv_seq_len"]
    return qsl > 1 and kvsl >= 8192 and bs >= 32


def _should_use_small_triton(config):
    return (config["q_seq_len"] == 1 and
            config["kv_seq_len"] == 1024 and
            config["batch_size"] in _SMALL_TRITON_NS)


def _run_mxfp4(q, kv_data, qo_indptr, kv_indptr, config):
    nh = config["num_heads"]
    bs = config["batch_size"]
    qsl = config["q_seq_len"]
    kvsl = config["kv_seq_len"]
    sms = config["sm_scale"]
    total_q = q.shape[0]

    kv_mxfp4_raw, kv_scale_raw = kv_data["mxfp4"]
    kv_mxfp4 = _get_uint8_view(kv_mxfp4_raw)
    kv_scale = _get_uint8_view(kv_scale_raw)

    if qsl > 1:
        nh_eff = qsl * nh
        total_q_eff = bs
        q_for_kernel = q.reshape(bs, nh_eff, QK_DIM)
        qo_for_kernel = _get_qo_merged(bs)
    else:
        nh_eff = nh
        total_q_eff = total_q
        q_for_kernel = q
        qo_for_kernel = qo_indptr

    HPB = 64 if nh_eff >= 64 else 16
    DK = 128
    BS = 128
    n_hg = triton.cdiv(nh_eff, HPB)

    q_fp8 = _get_q_fp8(q_for_kernel.shape)
    q_fp8.copy_(q_for_kernel)

    cfg_key = (bs, nh_eff, kvsl, HPB)
    if cfg_key not in _mxfp4_config_cache:
        grid_per_split = bs * n_hg
        max_splits_kv = max(1, kvsl // BS)
        target_wgs = 304 * 2
        raw_splits = max(1, target_wgs // max(grid_per_split, 1))
        raw_splits = min(raw_splits, max_splits_kv)
        ns = 1
        for candidate in [1, 2, 4, 8, 16, 32, 64]:
            if candidate <= raw_splits:
                ns = candidate
        _mxfp4_config_cache[cfg_key] = ns
    nsplits = _mxfp4_config_cache[cfg_key]

    pO, pM, pL = _get_bufs(nsplits, total_q_eff, nh_eff)
    output = _get_output(total_q_eff, nh_eff)

    _mla_hybrid_splitk[(bs * nsplits, n_hg)](
        q_fp8, kv_mxfp4, kv_scale,
        pO, pM, pL,
        qo_for_kernel, kv_indptr,
        sms,
        q_fp8.stride(0), q_fp8.stride(1),
        kv_mxfp4.stride(0), kv_scale.stride(0),
        pO.stride(0), pO.stride(1), pO.stride(2),
        pM.stride(0), pM.stride(1),
        nh_eff,
        NSPLITS=nsplits, BS=BS, HPB=HPB, DK=DK,
        num_warps=4, num_stages=2,
    )

    _mla_mxfp4_reduce[(total_q_eff, n_hg)](
        pO, pM, pL, output,
        pO.stride(0), pO.stride(1), pO.stride(2),
        pM.stride(0), pM.stride(1),
        output.stride(0), output.stride(1),
        nh_eff,
        NSPLITS=nsplits, HPB=HPB, DK=DK,
        num_warps=2, num_stages=1,
    )

    if qsl > 1:
        return output.reshape(bs, qsl, nh, V_DIM).reshape(total_q, nh, V_DIM)
    return output


def _run_small_triton(q, kv_data, qo_indptr, kv_indptr, config):
    nh = config["num_heads"]
    bs = config["batch_size"]
    sms = config["sm_scale"]
    total_q = q.shape[0]

    kvb_fp8, kvs_fp8 = kv_data["fp8"]
    kv_fp8 = _get_kv_fp8_view(kvb_fp8)
    kv_fp8_scale_val = _get_kv_fp8_scale(kvs_fp8)
    effective_sms = sms * kv_fp8_scale_val

    HPB = 16
    DK = 128
    BS = 128
    n_hg = triton.cdiv(nh, HPB)
    nsplits = _SMALL_TRITON_NS[bs]
    output = _get_output(total_q, nh)

    if nsplits <= 1:
        _mla_nosplit[(bs, n_hg)](
            q, kv_fp8, output,
            qo_indptr, kv_indptr,
            sms, kv_fp8_scale_val,
            q.stride(0), q.stride(1),
            kv_fp8.stride(0),
            output.stride(0), output.stride(1),
            nh, 1,
            BS=BS, HPB=HPB, DK=DK,
            num_warps=4, num_stages=2,
        )
        return output

    pO, pM, pL = _get_bufs(nsplits, total_q, nh)
    _mla_splitk_small[(bs * nsplits, n_hg)](
        q, kv_fp8,
        pO, pM, pL,
        qo_indptr, kv_indptr,
        effective_sms, kv_fp8_scale_val,
        q.stride(0), q.stride(1),
        kv_fp8.stride(0),
        pO.stride(0), pO.stride(1), pO.stride(2),
        pM.stride(0), pM.stride(1),
        nh,
        NSPLITS=nsplits, BS=BS, HPB=HPB, QSEQLEN=1, DK=DK,
        num_warps=4, num_stages=2,
    )

    pO_flat = pO.view(nsplits, -1)
    pM_flat = pM.view(nsplits, -1)
    pL_flat = pL.view(nsplits, -1)
    _mla_reduce_small[(total_q * nh,)](
        pO_flat, pM_flat, pL_flat, output.view(-1),
        pO_flat.stride(0), pM_flat.stride(0),
        NSPLITS=nsplits, VDIM=V_DIM, BLK=128,
        num_warps=2, num_stages=1,
    )
    return output


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data

    if _should_use_mxfp4(config):
        try:
            _do_warmup()
            return _run_mxfp4(q, kv_data, qo_indptr, kv_indptr, config)
        except Exception as ex:
            print(f"[MLA] MXFP4 path failed, fallback: {ex}", file=sys.stderr)
            traceback.print_exc(file=sys.stderr)

    if _should_use_small_triton(config):
        try:
            _do_warmup()
            return _run_small_triton(q, kv_data, qo_indptr, kv_indptr, config)
        except Exception as ex:
            print(f"[MLA] Small Triton path failed, fallback: {ex}", file=sys.stderr)
            traceback.print_exc(file=sys.stderr)

    if _aiter_available:
        key = (
            config["batch_size"],
            config["kv_seq_len"],
            config["num_heads"],
            config["q_seq_len"],
        )
        e = _aiter_table.get(key)
        if e is None:
            e = _aiter_setup(q, kv_data, qo_indptr, kv_indptr, config, key)
        _aiter_refresh_inputs(e, kv_data)
        if e.graph is not None:
            e.q_holder.copy_(q)
            e.graph.replay()
            return e.o
        if e.use_direct:
            e.qf_buf.copy_(q)
            try:
                _stage1_fn(
                    e.q_folded, e.kvb4, e.qo_indptr_own, e.kv_indptr_own,
                    e.kvi, e.klp,
                    None, e.wk[0], e.wk[1], e.wk[2],
                    e.qsl, 1, e.nkv, e.sms,
                    e.logits, e.attn_lse, e.o_folded,
                    e.qs, e.kvs)
                _reduce_fn(
                    e.logits, e.attn_lse, e.wk[3], e.wk[4], e.wk[5],
                    e.qsl, e.o_folded, None)
                return e.o
            except Exception as ex:
                print(f"[MLA] Direct call failed, fallback: {ex}", file=sys.stderr)
                e.use_direct = False
        e.qf_buf.copy_(q)
        mla_decode_fwd(
            e.qf_view, e.kvb4, e.o, qo_indptr, kv_indptr, e.kvi, e.klp,
            e.qsl, page_size=1, nhead_kv=e.nkv, sm_scale=e.sms, logit_cap=0.0,
            num_kv_splits=e.ns, q_scale=e.qs, kv_scale=e.kvs,
            intra_batch_mode=True, **e.meta)
        return e.o

    _do_warmup()
    return _run_triton(q, kv_data, qo_indptr, kv_indptr, config)


def _run_triton(q, kv_data, qo_indptr, kv_indptr, config):
    nh = config["num_heads"]
    bs = config["batch_size"]
    qsl = config["q_seq_len"]
    kvsl = config["kv_seq_len"]
    sms = config["sm_scale"]
    total_q = q.shape[0]

    kvb_fp8, kvs_fp8 = kv_data["fp8"]
    kv_fp8 = _get_kv_fp8_view(kvb_fp8)
    kv_fp8_scale_val = _get_kv_fp8_scale(kvs_fp8)

    if qsl > 1:
        nh_eff = qsl * nh
        total_q_eff = bs
        q_for_kernel = q.reshape(bs, nh_eff, QK_DIM)
        qo_for_kernel = _get_qo_merged(bs)
        qsl_eff = 1
    else:
        nh_eff = nh
        total_q_eff = total_q
        q_for_kernel = q
        qo_for_kernel = qo_indptr
        qsl_eff = 1

    if qsl > 1 and nh_eff >= 64:
        HPB = 64
    else:
        HPB = 16

    DK = 128; BS = 128
    num_warps = 4; num_stages = 2
    n_hg = triton.cdiv(nh_eff, HPB)

    config_key = (bs, nh_eff, kvsl, HPB)
    if config_key not in _config_cache:
        grid_per_split = bs * n_hg
        max_splits_kv = max(1, kvsl // BS)
        target_wgs = 304 * 2
        raw_splits = max(1, target_wgs // max(grid_per_split, 1))
        raw_splits = min(raw_splits, max_splits_kv)
        NSPLITS = 1
        for s in [1, 2, 4, 8, 16, 32, 64]:
            if s <= raw_splits:
                NSPLITS = s
        _config_cache[config_key] = NSPLITS
    NSPLITS = _config_cache[config_key]

    output = _get_output(total_q_eff, nh_eff)

    if NSPLITS <= 1:
        _mla_nosplit[(bs, n_hg)](
            q_for_kernel, kv_fp8, output,
            qo_for_kernel, kv_indptr,
            sms, kv_fp8_scale_val,
            q_for_kernel.stride(0), q_for_kernel.stride(1),
            kv_fp8.stride(0),
            output.stride(0), output.stride(1),
            nh_eff, qsl_eff,
            BS=BS, HPB=HPB, DK=DK,
            num_warps=num_warps, num_stages=num_stages,
        )
    else:
        pO, pM, pL = _get_bufs(NSPLITS, total_q_eff, nh_eff)
        done_counter = _get_counter(bs, n_hg)
        ckey = (bs, n_hg)
        _gen_counter[ckey] += NSPLITS
        gen_target = _gen_counter[ckey]

        _mla_fused[(bs, NSPLITS, n_hg)](
            q_for_kernel, kv_fp8,
            pO, pM, pL, output,
            qo_for_kernel, kv_indptr,
            done_counter,
            sms, kv_fp8_scale_val,
            q_for_kernel.stride(0), q_for_kernel.stride(1),
            kv_fp8.stride(0),
            pO.stride(0), pO.stride(1), pO.stride(2),
            pM.stride(0), pM.stride(1),
            output.stride(0), output.stride(1),
            gen_target, nh_eff, qsl_eff, n_hg,
            NSPLITS=NSPLITS, BS=BS, HPB=HPB, DK=DK,
            num_warps=num_warps, num_stages=num_stages,
        )

    if qsl > 1:
        return output.reshape(bs, qsl, nh, V_DIM).reshape(total_q, nh, V_DIM)
    return output
scrolls · 1304 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