Skip to content
KernelIndex
Search⌘K

submission 754939

Shubham Kumar · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_hybrid_dot_scaled_v1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-754939?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
115.4µs
#482 of 766
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2919e22b19281cdd8530589e9651ceafc317f66d7ac0ca13354f04e6a23e20c1
license declaredunknown
license concludedunknown
authorsShubham Kumar
imported2026-08-26

Techniques

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

fp4batch_size * kv_seq_len <= 65536 -> Triton dot_scaled (native MXFP4 MFMA)
mmac = tl.dot(pb, v_block.to(tl.bfloat16))
persistent-kernelotherwise -> aiter a8w8 (FP8 persistent kernel)
split-kTriton path: two-stage split-K with tl.dot_scaled for score computation

Kernel source

submission_hybrid_dot_scaled_v1.py725 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""
Hybrid MLA decode kernel: Triton dot_scaled for small batch, aiter a8w8 for large batch.

Dispatch logic:
  batch_size * kv_seq_len <= 65536 -> Triton dot_scaled (native MXFP4 MFMA)
  otherwise                        -> aiter a8w8 (FP8 persistent kernel)

Triton path: two-stage split-K with tl.dot_scaled for score computation
  and manual FP4 dequant for V accumulation.
Aiter path: cached metadata, cached kv_indices, pre-allocated output,
  per-config NUM_KV_SPLITS tuning.
"""

import torch
import torch.nn.functional as F
import triton
import triton.language as tl
from task import input_t, output_t
from utils import make_match_reference

# ---------------------------------------------------------------------------
# Aiter imports (triggers ~222s JIT build)
# ---------------------------------------------------------------------------
from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.utility.fp4_utils import (
    dynamic_mxfp4_quant,
    mxfp4_to_f32,
    e8m0_to_f32,
)

# ---------------------------------------------------------------------------
# Shared constants
# ---------------------------------------------------------------------------
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)

# ---------------------------------------------------------------------------
# Triton path constants
# ---------------------------------------------------------------------------
NSPLITS_SMALL = 8    # for kv_seq_len <= 1024 (grid=batch*8, optimal for bs=32)
NSPLITS_MED = 4      # for large batch + kv>1024 (grid=batch*4, less overhead)
NSPLITS_LARGE = 64   # for small batch + kv>1024 (grid=batch*64, max parallelism)

PADDED_DIM = 768     # 576 padded to 768 = 3 * 256
PADDED_BYTES = 384   # 768 / 2 = 384 packed fp4x2 bytes
PADDED_SCALES = 24   # 768 / 32 = 24 scale blocks

# ---------------------------------------------------------------------------
# Aiter path constants
# ---------------------------------------------------------------------------
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
Q_DTYPE = "fp8"
KV_DTYPE = "fp8"

# ---------------------------------------------------------------------------
# Hybrid dispatch threshold
# ---------------------------------------------------------------------------
TRITON_THRESHOLD = 65536   # batch_size * kv_seq_len <= this -> use Triton

# ===========================================================================
# TRITON KERNELS
# ===========================================================================

# ---------------------------------------------------------------------------
# FP4 E2M1 dequantization helper (used only in V phase)
# ---------------------------------------------------------------------------
@triton.jit
def _dq4(x):
    """Dequantize FP4 E2M1 nibble (0-15) to float32."""
    n = x.to(tl.int32)
    e = (n >> 1) & 3
    m = (n & 1).to(tl.float32)
    val = tl.where(e == 0,
                   m * 0.5,
                   (1.0 + m * 0.5) * tl.math.exp2((e - 1).to(tl.float32)))
    sg = 1.0 - ((n >> 3) & 1).to(tl.float32) * 2.0
    return sg * val


# ---------------------------------------------------------------------------
# Stage 1: Split-K partial attention with native MXFP4 dot_scaled for scores
# ---------------------------------------------------------------------------
@triton.jit
def mla_s1(
    Q,          # query tensor: (total_q, 16, 576) bf16, row-major (ORIGINAL)
    KV,         # kv_fp4x2:    (total_kv, 288) uint8 (packed fp4x2, ORIGINAL)
    Sc,         # kv_scales:   (total_kv, 18+) uint8 (e8m0, ORIGINAL)
    Ind,        # kv_indptr:   (batch+1,) int32
    PO,         # partial_out: (batch*nsplits*16, 512) float32
    PM,         # partial_max: (batch*nsplits*16,)     float32
    PS,         # partial_sum: (batch*nsplits*16,)     float32
    Q_pad,      # padded Q:    (total_q, 16, 768) bf16
    KV_pad,     # padded KV:   (total_kv, 384) uint8
    Sc_pad,     # padded scales: (total_kv, 24) uint8
    sm_scale,   # softmax scale factor
    nsplits: tl.constexpr,   # number of splits
    sq_b: tl.constexpr,      # Q stride for batch dim (= 16 * 576)
    sq_h: tl.constexpr,      # Q stride for head dim  (= 576)
    skv: tl.constexpr,       # KV stride for token dim (= 288)
    ssc_t: tl.constexpr,     # Scale stride for token dim
    ssc_b: tl.constexpr,     # Scale stride (unused, kept for compat)
    sq_b_pad: tl.constexpr,  # Padded Q stride for batch dim (= 16 * 768)
    sq_h_pad: tl.constexpr,  # Padded Q stride for head dim  (= 768)
    skv_pad: tl.constexpr,   # Padded KV stride for token dim (= 384)
    ssc_pad: tl.constexpr,   # Padded Scale stride for token dim (= 24)
    BN: tl.constexpr,        # tile size for KV tokens (64)
):
    bid = tl.program_id(0)   # batch index
    sid = tl.program_id(1)   # split index

    # Compute token range for this split
    ks = tl.load(Ind + bid)
    ke = tl.load(Ind + bid + 1)
    kl = ke - ks
    cs = tl.cdiv(kl, nsplits)
    ms = ks + sid * cs
    me = tl.minimum(ms + cs, ke)

    si = bid * nsplits + sid   # flat split index
    hr = tl.arange(0, 16)     # head range [0..15]
    nr = tl.arange(0, BN)     # token range within tile [0..BN-1]

    # Early exit for empty splits
    if ms >= me:
        tl.store(PM + si * 16 + hr, tl.full([16], float('-inf'), tl.float32))
        tl.store(PS + si * 16 + hr, tl.zeros([16], tl.float32))
        return

    # Running softmax state
    mi = tl.full([16], float('-inf'), tl.float32)    # max logit per head
    li = tl.zeros([16], tl.float32)                  # sum of exp per head

    # 16 V accumulators: each (16, 32) for 16 heads x 32 dims = one scale block
    a0  = tl.zeros([16, 32], tl.float32)
    a1  = tl.zeros([16, 32], tl.float32)
    a2  = tl.zeros([16, 32], tl.float32)
    a3  = tl.zeros([16, 32], tl.float32)
    a4  = tl.zeros([16, 32], tl.float32)
    a5  = tl.zeros([16, 32], tl.float32)
    a6  = tl.zeros([16, 32], tl.float32)
    a7  = tl.zeros([16, 32], tl.float32)
    a8  = tl.zeros([16, 32], tl.float32)
    a9  = tl.zeros([16, 32], tl.float32)
    a10 = tl.zeros([16, 32], tl.float32)
    a11 = tl.zeros([16, 32], tl.float32)
    a12 = tl.zeros([16, 32], tl.float32)
    a13 = tl.zeros([16, 32], tl.float32)
    a14 = tl.zeros([16, 32], tl.float32)
    a15 = tl.zeros([16, 32], tl.float32)

    # Iterate over KV tiles within this split's range
    for ts in range(ms, me, BN):
        tid = ts + nr            # absolute token indices for this tile
        vm = tid < me            # validity mask

        # ---- Score computation using tl.dot_scaled for native MXFP4 MFMA ----
        scores = tl.zeros([BN, 16], tl.float32)

        for ki in tl.static_range(3):
            # K tile: (BN, 128) packed fp4x2 from padded KV
            k_tile = tl.load(KV_pad + tid[:, None] * skv_pad + ki * 128 + tl.arange(0, 128)[None, :],
                             mask=vm[:, None], other=0)

            # K scales: (BN, 8) raw e8m0 for this BK=256 chunk
            k_sc = tl.load(Sc_pad + tid[:, None] * ssc_pad + ki * 8 + tl.arange(0, 8)[None, :],
                           mask=vm[:, None], other=127)

            # Q^T tile: (256, 16) bf16
            d_offs = ki * 256 + tl.arange(0, 256)
            h_offs = tl.arange(0, 16)
            q_t = tl.load(Q_pad + bid * sq_b_pad + h_offs[None, :] * sq_h_pad + d_offs[:, None],
                          mask=True, other=0.0).to(tl.bfloat16)

            # Native MXFP4 x BF16 dot product via hardware MFMA
            scores += tl.dot_scaled(k_tile, k_sc, "e2m1", q_t, None, "bf16")

        # Transpose scores: (BN, 16) -> (16, BN) for softmax per head
        sc = tl.trans(scores) * sm_scale

        # ---- Online softmax ----
        sc = tl.where(vm[None, :], sc, float('-inf'))

        # Per-head max for this tile
        tile_max = tl.max(sc, axis=1)
        new_mi = tl.maximum(mi, tile_max)

        # Correction factor for previous accumulators
        alpha = tl.math.exp2((mi - new_mi) * 1.4426950408889634)
        # Softmax weights for this tile
        p = tl.math.exp2((sc - new_mi[:, None]) * 1.4426950408889634)
        tile_sum = tl.sum(p, axis=1)

        # Update running state
        li = li * alpha + tile_sum
        mi = new_mi

        # Rescale previous accumulators
        a0  = a0  * alpha[:, None]
        a1  = a1  * alpha[:, None]
        a2  = a2  * alpha[:, None]
        a3  = a3  * alpha[:, None]
        a4  = a4  * alpha[:, None]
        a5  = a5  * alpha[:, None]
        a6  = a6  * alpha[:, None]
        a7  = a7  * alpha[:, None]
        a8  = a8  * alpha[:, None]
        a9  = a9  * alpha[:, None]
        a10 = a10 * alpha[:, None]
        a11 = a11 * alpha[:, None]
        a12 = a12 * alpha[:, None]
        a13 = a13 * alpha[:, None]
        a14 = a14 * alpha[:, None]
        a15 = a15 * alpha[:, None]

        # ---- V accumulation: p @ V for first 512 dims = 16 blocks of 32 ----
        pb = p.to(tl.bfloat16)

        for vb in tl.static_range(16):
            # Load packed fp4x2 for V block vb: (BN, 16) uint8
            v_raw = tl.load(KV + (tid[:, None] * skv + vb * 16 + tl.arange(0, 16)[None, :]),
                            mask=vm[:, None], other=0)

            # Unpack
            vlo = (v_raw & 0xF).to(tl.uint8)
            vhi = (v_raw >> 4).to(tl.uint8)

            # Dequant
            fvlo = _dq4(vlo)
            fvhi = _dq4(vhi)

            # Load e8m0 scale
            vsc_raw = tl.load(Sc + tid * ssc_t + vb, mask=vm, other=0)
            vsc_exp = (vsc_raw.to(tl.int32) - 127).to(tl.float32)
            vsc_f = tl.math.exp2(vsc_exp)

            fvlo = fvlo * vsc_f[:, None]
            fvhi = fvhi * vsc_f[:, None]

            # Interleave even/odd to reconstruct 32 contiguous values per token
            v_joined = tl.join(fvlo, fvhi)
            v_block = tl.reshape(v_joined, [BN, 32])

            # p @ V_block: (16, BN) @ (BN, 32) -> (16, 32)
            c = tl.dot(pb, v_block.to(tl.bfloat16))

            if vb == 0:  a0  += c.to(tl.float32)
            if vb == 1:  a1  += c.to(tl.float32)
            if vb == 2:  a2  += c.to(tl.float32)
            if vb == 3:  a3  += c.to(tl.float32)
            if vb == 4:  a4  += c.to(tl.float32)
            if vb == 5:  a5  += c.to(tl.float32)
            if vb == 6:  a6  += c.to(tl.float32)
            if vb == 7:  a7  += c.to(tl.float32)
            if vb == 8:  a8  += c.to(tl.float32)
            if vb == 9:  a9  += c.to(tl.float32)
            if vb == 10: a10 += c.to(tl.float32)
            if vb == 11: a11 += c.to(tl.float32)
            if vb == 12: a12 += c.to(tl.float32)
            if vb == 13: a13 += c.to(tl.float32)
            if vb == 14: a14 += c.to(tl.float32)
            if vb == 15: a15 += c.to(tl.float32)

    # ---- Store partial results ----
    tl.store(PM + si * 16 + hr, mi)
    tl.store(PS + si * 16 + hr, li)

    # Store partial output: (16 heads, 512 dims) as 16 blocks of 32
    po_rows = (si * 16 + hr)[:, None] * 512
    dr = tl.arange(0, 32)
    tl.store(PO + po_rows + 0  * 32 + dr[None, :], a0)
    tl.store(PO + po_rows + 1  * 32 + dr[None, :], a1)
    tl.store(PO + po_rows + 2  * 32 + dr[None, :], a2)
    tl.store(PO + po_rows + 3  * 32 + dr[None, :], a3)
    tl.store(PO + po_rows + 4  * 32 + dr[None, :], a4)
    tl.store(PO + po_rows + 5  * 32 + dr[None, :], a5)
    tl.store(PO + po_rows + 6  * 32 + dr[None, :], a6)
    tl.store(PO + po_rows + 7  * 32 + dr[None, :], a7)
    tl.store(PO + po_rows + 8  * 32 + dr[None, :], a8)
    tl.store(PO + po_rows + 9  * 32 + dr[None, :], a9)
    tl.store(PO + po_rows + 10 * 32 + dr[None, :], a10)
    tl.store(PO + po_rows + 11 * 32 + dr[None, :], a11)
    tl.store(PO + po_rows + 12 * 32 + dr[None, :], a12)
    tl.store(PO + po_rows + 13 * 32 + dr[None, :], a13)
    tl.store(PO + po_rows + 14 * 32 + dr[None, :], a14)
    tl.store(PO + po_rows + 15 * 32 + dr[None, :], a15)


# ---------------------------------------------------------------------------
# Stage 2: Reduce partial results across splits
# ---------------------------------------------------------------------------
@triton.jit
def mla_s2(
    PO,     # partial_out: (batch*nsplits*16, 512) float32
    PM,     # partial_max: (batch*nsplits*16,)     float32
    PS,     # partial_sum: (batch*nsplits*16,)     float32
    Out,    # output:      (batch, 16, 512)        bf16
    nsplits: tl.constexpr,
    so_b: tl.constexpr,   # output stride for batch (= 16 * 512)
    so_h: tl.constexpr,   # output stride for head  (= 512)
):
    bid = tl.program_id(0)   # batch index
    hid = tl.program_id(1)   # head index [0..15]

    # Find global max across all splits for this (batch, head)
    global_max = tl.full([], float('-inf'), tl.float32)
    for s in range(nsplits):
        si = bid * nsplits + s
        m = tl.load(PM + si * 16 + hid)
        global_max = tl.maximum(global_max, m)

    # Compute weighted sum with log-sum-exp correction
    global_sum = tl.zeros([], tl.float32)
    acc = tl.zeros([512], tl.float32)
    dr = tl.arange(0, 512)

    for s in range(nsplits):
        si = bid * nsplits + s
        m = tl.load(PM + si * 16 + hid)
        l = tl.load(PS + si * 16 + hid)

        # Correction weight: exp(m_split - m_global)
        alpha = tl.math.exp2((m - global_max) * 1.4426950408889634)
        w = alpha * l

        # Load partial output row
        po_base = (si * 16 + hid) * 512
        pv = tl.load(PO + po_base + dr)

        acc += alpha * pv
        global_sum += w

    # Normalize and store
    acc = acc / global_sum
    out_base = bid * so_b + hid * so_h
    tl.store(Out + out_base + dr, acc.to(tl.bfloat16))


# ===========================================================================
# TRITON PATH: buffer caches and entry point
# ===========================================================================

_triton_buf_cache: dict = {}


def _get_triton_buffers(batch_size: int, nsplits: int):
    """Return (or allocate) partial output/max/sum buffers for Triton path."""
    key = (batch_size, nsplits)
    if key not in _triton_buf_cache:
        total_splits = batch_size * nsplits
        total_rows = total_splits * NUM_HEADS
        po = torch.empty((total_rows, V_HEAD_DIM), dtype=torch.float32, device="cuda")
        pm = torch.empty((total_rows,), dtype=torch.float32, device="cuda")
        ps = torch.empty((total_rows,), dtype=torch.float32, device="cuda")
        _triton_buf_cache[key] = (po, pm, ps)
    return _triton_buf_cache[key]


_triton_out_cache: dict = {}


def _get_triton_output(batch_size: int):
    """Return (or allocate) output tensor for Triton path."""
    if batch_size not in _triton_out_cache:
        _triton_out_cache[batch_size] = torch.empty(
            (batch_size, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"
        )
    return _triton_out_cache[batch_size]


_triton_pad_cache: dict = {}


def _get_padded_tensors(q, kv_flat, sc_flat, batch_size, total_kv):
    """Pad Q (576->768), KV (288->384), scales (18->24) for dot_scaled."""
    key = (batch_size, total_kv)
    if key not in _triton_pad_cache:
        # Pad KV: (total_kv, 288) -> (total_kv, 384) uint8, zeros for padding
        kv_padded = torch.zeros((total_kv, PADDED_BYTES), dtype=torch.uint8, device="cuda")
        kv_padded[:, :288] = kv_flat

        # Pad scales: (total_kv, 18) -> (total_kv, 24) uint8, 127 = neutral e8m0 (2^0)
        sc_padded = torch.full((total_kv, PADDED_SCALES), 127, dtype=torch.uint8, device="cuda")
        sc_cols = sc_flat.shape[1]
        sc_padded[:, :sc_cols] = sc_flat[:, :sc_cols]

        # Pad Q: (batch, 16, 576) -> (batch, 16, 768) bf16, zeros for padding
        q_padded = torch.zeros((batch_size, NUM_HEADS, PADDED_DIM), dtype=torch.bfloat16, device="cuda")
        q_padded[:, :, :QK_HEAD_DIM] = q[:batch_size]

        _triton_pad_cache[key] = (q_padded, kv_padded, sc_padded)
    else:
        q_padded, kv_padded, sc_padded = _triton_pad_cache[key]
        # Update with current data (cache is for allocation reuse)
        kv_padded[:, :288] = kv_flat
        sc_cols = sc_flat.shape[1]
        sc_padded[:, :sc_cols] = sc_flat[:, :sc_cols]
        q_padded[:, :, :QK_HEAD_DIM] = q[:batch_size]

    return q_padded, kv_padded, sc_padded


def _triton_dot_scaled_path(q, kv_data, kv_indptr, config):
    """Triton dot_scaled MLA decode for small batch * kv_seq_len cases."""
    batch_size = config["batch_size"]
    kv_seq_len = config["kv_seq_len"]

    # Unpack MXFP4 KV cache
    kv_fp4x2, kv_scales = kv_data["mxfp4"]

    # View as uint8 for Triton pointer arithmetic
    kv_u8 = kv_fp4x2.view(torch.uint8)
    sc_u8 = kv_scales.view(torch.uint8)

    # Flatten KV to 2D: (total_kv, 288)
    kv_flat = kv_u8.reshape(kv_u8.shape[0], -1)
    sc_flat = sc_u8

    total_kv = kv_flat.shape[0]

    # Strides for original tensors
    skv = kv_flat.stride(0)
    ssc_t = sc_flat.stride(0)

    # Q strides
    sq_b = q.stride(0) * 1
    sq_h = q.stride(1)

    # Pad tensors for dot_scaled
    q_padded, kv_padded, sc_padded = _get_padded_tensors(q, kv_flat, sc_flat, batch_size, total_kv)

    # Padded strides
    sq_b_pad = q_padded.stride(0)
    sq_h_pad = q_padded.stride(1)
    skv_pad = kv_padded.stride(0)
    ssc_pad = sc_padded.stride(0)

    # Select nsplits: 3-way dispatch for optimal grid size
    if kv_seq_len <= 1024:
        nsplits = NSPLITS_SMALL
    elif batch_size >= 64:
        nsplits = NSPLITS_MED
    else:
        nsplits = NSPLITS_LARGE

    # Allocate/reuse buffers
    po, pm, ps = _get_triton_buffers(batch_size, nsplits)
    out = _get_triton_output(batch_size)

    BN = 64

    # Stage 1: compute partial attention
    grid_s1 = (batch_size, nsplits)
    mla_s1[grid_s1](
        q, kv_flat, sc_flat, kv_indptr,
        po, pm, ps,
        q_padded, kv_padded, sc_padded,
        SM_SCALE,
        nsplits=nsplits,
        sq_b=sq_b,
        sq_h=sq_h,
        skv=skv,
        ssc_t=ssc_t,
        ssc_b=0,
        sq_b_pad=sq_b_pad,
        sq_h_pad=sq_h_pad,
        skv_pad=skv_pad,
        ssc_pad=ssc_pad,
        BN=BN,
    )

    # Stage 2: reduce across splits
    grid_s2 = (batch_size, NUM_HEADS)
    mla_s2[grid_s2](
        po, pm, ps, out,
        nsplits=nsplits,
        so_b=out.stride(0),
        so_h=out.stride(1),
    )

    return out


# ===========================================================================
# AITER PATH: caches, helpers, and entry point
# ===========================================================================

# Per-config NUM_KV_SPLITS tuning table
_SPLIT_TABLE = {
    (4,   1024): 4,
    (4,   8192): 16,
    (32,  1024): 4,
    (32,  8192): 16,
    (64,  1024): 4,
    (64,  8192): 16,
    (256, 1024): 8,
    (256, 8192): 16,
}
_DEFAULT_NUM_KV_SPLITS = 32

# Metadata buffer cache
_aiter_metadata_cache: dict = {}

# kv_indices cache
_aiter_kv_indices_cache: dict = {}

# Output tensor cache
_aiter_output_cache: dict = {}


def quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """Dynamic per-tensor FP8 quantization."""
    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 _make_mla_decode_metadata(
    batch_size, max_q_len, nhead, nhead_kv,
    q_dtype, kv_dtype,
    qo_indptr, kv_indptr, kv_last_page_len,
    num_kv_splits=_DEFAULT_NUM_KV_SPLITS,
):
    """Allocate and populate work buffers for persistent mla_decode_fwd."""
    info = get_mla_metadata_info_v1(
        batch_size, max_q_len, nhead, 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,
        nhead // nhead_kv,
        nhead_kv,
        True,
        work_metadata, work_info_set, work_indptr,
        reduce_indptr, reduce_final_map, reduce_partial_map,
        page_size=PAGE_SIZE,
        kv_granularity=max(PAGE_SIZE, 16),
        max_seqlen_qo=max_q_len,
        uni_seqlen_qo=max_q_len,
        fast_mode=False,
        max_split_per_batch=num_kv_splits,
        intra_batch_mode=True,
        dtype_q=q_dtype,
        dtype_kv=kv_dtype,
    )

    return {
        "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,
    }


def _get_num_kv_splits(batch_size: int, kv_len_per_seq: int) -> int:
    return _SPLIT_TABLE.get((batch_size, kv_len_per_seq), _DEFAULT_NUM_KV_SPLITS)


def _get_cached_kv_indices(total_kv_len: int) -> torch.Tensor:
    if total_kv_len not in _aiter_kv_indices_cache:
        _aiter_kv_indices_cache[total_kv_len] = torch.arange(
            total_kv_len, dtype=torch.int32, device="cuda"
        )
    return _aiter_kv_indices_cache[total_kv_len]


def _get_cached_output(total_q: int, nq: int, dv: int) -> torch.Tensor:
    key = (total_q, nq, dv)
    cached = _aiter_output_cache.get(key)
    if cached is None:
        cached = torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda")
        _aiter_output_cache[key] = cached
    return cached


def _get_cached_metadata(
    batch_size, max_q_len, total_kv_len, nhead, nhead_kv,
    q_dtype, kv_dtype,
    qo_indptr, kv_indptr, kv_last_page_len,
    num_kv_splits,
):
    cache_key = (batch_size, max_q_len, total_kv_len, q_dtype, kv_dtype, num_kv_splits)
    if cache_key not in _aiter_metadata_cache:
        _aiter_metadata_cache[cache_key] = _make_mla_decode_metadata(
            batch_size, max_q_len, nhead, nhead_kv,
            q_dtype, kv_dtype,
            qo_indptr, kv_indptr, kv_last_page_len,
            num_kv_splits=num_kv_splits,
        )
    return _aiter_metadata_cache[cache_key]


def _aiter_mla_decode(
    q, kv_buffer, qo_indptr, kv_indptr, config,
    q_scale=None, kv_scale=None,
):
    """MLA decode attention using aiter persistent-mode kernel."""
    batch_size = config["batch_size"]
    nq = config["num_heads"]
    nkv = config["num_kv_heads"]
    dq = config["qk_head_dim"]
    dv = config["v_head_dim"]
    q_seq_len = config["q_seq_len"]
    total_kv_len = int(kv_indptr[-1].item())

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

    max_q_len = q_seq_len

    kv_len_per_seq = total_kv_len // batch_size
    num_kv_splits = _get_num_kv_splits(batch_size, kv_len_per_seq)

    kv_indices = _get_cached_kv_indices(total_kv_len)

    kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)

    meta = _get_cached_metadata(
        batch_size, max_q_len, total_kv_len, nq, nkv,
        q.dtype, kv_buffer.dtype,
        qo_indptr, kv_indptr, kv_last_page_len,
        num_kv_splits=num_kv_splits,
    )

    o = _get_cached_output(q.shape[0], nq, dv)

    mla_decode_fwd(
        q.view(-1, nq, dq),
        kv_buffer_4d,
        o,
        qo_indptr,
        kv_indptr,
        kv_indices,
        kv_last_page_len,
        max_q_len,
        page_size=PAGE_SIZE,
        nhead_kv=nkv,
        sm_scale=SM_SCALE,
        logit_cap=0.0,
        num_kv_splits=num_kv_splits,
        q_scale=q_scale,
        kv_scale=kv_scale,
        intra_batch_mode=True,
        **meta,
    )
    return o


def _aiter_path(q, kv_data, qo_indptr, kv_indptr, config):
    """Aiter a8w8 FP8 path for large batch cases."""
    # Quantize Q to FP8
    q_input, q_scale = quantize_fp8(q)

    # Use FP8 KV cache
    kv_buffer_fp8, kv_scale = kv_data["fp8"]

    return _aiter_mla_decode(
        q_input, kv_buffer_fp8, qo_indptr, kv_indptr, config,
        q_scale=q_scale, kv_scale=kv_scale,
    )


# ===========================================================================
# HYBRID DISPATCH
# ===========================================================================

_aiter_warmed = False


def _warmup_aiter(q, kv_data, qo_indptr, kv_indptr, config):
    """Force one aiter call to trigger all JIT compilation and kernel loading."""
    global _aiter_warmed
    if _aiter_warmed:
        return
    _aiter_warmed = True
    # Run aiter path once (result discarded) to populate JIT caches
    try:
        _aiter_path(q, kv_data, qo_indptr, kv_indptr, config)
        torch.cuda.synchronize()
    except Exception:
        pass  # If warmup fails, aiter path will fail naturally later


def custom_kernel(data: input_t) -> output_t:
    """
    Hybrid MLA decode: Triton dot_scaled for small workloads, aiter a8w8 for large.

    Dispatch threshold: batch_size * kv_seq_len <= 65536 -> Triton
    This covers: bs=4/kv=1k, bs=4/kv=8k, bs=32/kv=1k, bs=64/kv=1k
    Aiter handles: bs=32/kv=8k, bs=64/kv=8k, bs=256/kv=1k, bs=256/kv=8k
    """
    q, kv_data, qo_indptr, kv_indptr, config = data

    batch_size = config["batch_size"]
    kv_seq_len = config["kv_seq_len"]
    workload = batch_size * kv_seq_len

    # Warmup aiter on first call regardless of path
    # This prevents JIT spikes during benchmark timing
    if not _aiter_warmed:
        _warmup_aiter(q, kv_data, qo_indptr, kv_indptr, config)

    if workload <= TRITON_THRESHOLD:
        return _triton_dot_scaled_path(q, kv_data, kv_indptr, config)
    else:
        return _aiter_path(q, kv_data, qo_indptr, kv_indptr, config)
scrolls · 725 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