Skip to content
KernelIndex
Search⌘K

submission 746469

chenxingqiang · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f8759a3cfaf68d074b30374cf698076912daf4ba7589f1088ae9c3de735a9a01
license declaredunknown
license concludedunknown
authorschenxingqiang
imported2026-08-15

Techniques

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

mmas = tl.dot(q_lora, tl.trans(k_lora)) + tl.dot(q_rope, tl.trans(k_rope))
num-warps = 4num_warps=4, num_stages=2,
online-softmaxm_new = tl.maximum(m_i, tl.max(s, axis=1))
stages = 2num_warps=4, num_stages=2,

Kernel source

submission.py590 lines
"""
Hybrid MLA decode: custom Triton for short KV, ASM for long KV.
"""
import os
import torch
import triton
import triton.language as tl
from task import input_t, output_t

from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.mla import mla_decode_fwd
from aiter.ops.triton.quant import dynamic_per_tensor_quant_fp8_i8
from aiter.ops.attention import mla_decode_stage1_asm_fwd, mla_reduce_v1
from aiter.utility.fp4_utils import mxfp4_to_f32, e8m0_to_f32

KV_LORA_RANK = 512
QK_ROPE_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_DIM
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8

_CACHE: dict = {}
_ASM_CACHE: dict = {}
_MXFP4_CACHE: dict = {}
_CACHE_LIMIT = 32


@triton.jit
def _decode_mxfp4_nibbles(x):
    mag = x & 0x7
    sign = x & 0x8
    out = tl.where(mag == 0, 0.0, 0.0)
    out = tl.where(mag == 1, 0.5, out)
    out = tl.where(mag == 2, 1.0, out)
    out = tl.where(mag == 3, 1.5, out)
    out = tl.where(mag == 4, 2.0, out)
    out = tl.where(mag == 5, 3.0, out)
    out = tl.where(mag == 6, 4.0, out)
    out = tl.where(mag == 7, 6.0, out)
    return tl.where(sign != 0, -out, out)


@triton.jit
def _mla_decode_attn(
    Q, KV, Out,
    Partial_O, Partial_M, Partial_L,
    qo_indptr, kv_indptr,
    kv_scale_ptr,
    sm_scale: tl.constexpr,
    stride_q_s, stride_q_h,
    stride_kv_s,
    stride_o_s, stride_o_h,
    NHEADS: tl.constexpr,
    D_LORA: tl.constexpr,
    D_ROPE: tl.constexpr,
    BLOCK_KV: tl.constexpr,
    NUM_SPLITS: tl.constexpr,
):
    pid = tl.program_id(0)
    pid_b = pid // NUM_SPLITS
    pid_s = pid % NUM_SPLITS

    qo_start = tl.load(qo_indptr + pid_b)
    kv_start = tl.load(kv_indptr + pid_b)
    kv_end = tl.load(kv_indptr + pid_b + 1)
    kv_len = kv_end - kv_start

    chunk = tl.cdiv(kv_len, NUM_SPLITS)
    my_start = kv_start + pid_s * chunk
    my_end = tl.minimum(kv_start + (pid_s + 1) * chunk, kv_end)
    my_len = my_end - my_start

    kv_scale = tl.load(kv_scale_ptr).to(tl.float32)
    qk_combined = kv_scale * sm_scale

    offs_h = tl.arange(0, NHEADS)
    offs_lora = tl.arange(0, D_LORA)
    offs_rope = tl.arange(0, D_ROPE)

    q_lora = tl.load(
        Q + qo_start * stride_q_s + offs_h[:, None] * stride_q_h + offs_lora[None, :]
    )
    q_rope = tl.load(
        Q + qo_start * stride_q_s + offs_h[:, None] * stride_q_h + D_LORA + offs_rope[None, :]
    )

    m_i = tl.full([NHEADS], value=float("-inf"), dtype=tl.float32)
    l_i = tl.zeros([NHEADS], dtype=tl.float32)
    acc = tl.zeros([NHEADS, D_LORA], dtype=tl.float32)

    offs_kv = tl.arange(0, BLOCK_KV)

    for kv_off in range(0, my_len, BLOCK_KV):
        valid = (kv_off + offs_kv) < my_len
        base = my_start + kv_off

        k_lora = tl.load(
            KV + (base + offs_kv[:, None]) * stride_kv_s + offs_lora[None, :],
            mask=valid[:, None], other=0.0,
        ).to(tl.bfloat16)

        k_rope = tl.load(
            KV + (base + offs_kv[:, None]) * stride_kv_s + D_LORA + offs_rope[None, :],
            mask=valid[:, None], other=0.0,
        ).to(tl.bfloat16)

        s = tl.dot(q_lora, tl.trans(k_lora)) + tl.dot(q_rope, tl.trans(k_rope))
        s *= qk_combined
        s = tl.where(valid[None, :], s, float("-inf"))

        m_new = tl.maximum(m_i, tl.max(s, axis=1))
        alpha = tl.exp(m_i - m_new)
        p = tl.exp(s - m_new[:, None])

        acc = acc * alpha[:, None] + tl.dot(p.to(tl.bfloat16), k_lora)
        l_i = l_i * alpha + tl.sum(p, axis=1)
        m_i = m_new

    if NUM_SPLITS == 1:
        acc = acc * (kv_scale / l_i[:, None])
        tl.store(
            Out + qo_start * stride_o_s + offs_h[:, None] * stride_o_h + offs_lora[None, :],
            acc.to(tl.bfloat16),
        )
    else:
        part_base = pid * NHEADS
        tl.store(
            Partial_O + pid * NHEADS * D_LORA + offs_h[:, None] * D_LORA + offs_lora[None, :],
            acc,
        )
        tl.store(Partial_M + part_base + offs_h, m_i)
        tl.store(Partial_L + part_base + offs_h, l_i)


@triton.jit
def _mla_decode_attn_mxfp4(
    Q, KV, KV_SCALES, Out,
    Partial_O, Partial_M, Partial_L,
    qo_indptr, kv_indptr,
    sm_scale: tl.constexpr,
    stride_q_s, stride_q_h,
    stride_kv_s, stride_scale_s,
    stride_o_s, stride_o_h,
    NHEADS: tl.constexpr,
    D_LORA: tl.constexpr,
    D_ROPE: tl.constexpr,
    BLOCK_KV: tl.constexpr,
    NUM_SPLITS: tl.constexpr,
):
    pid = tl.program_id(0)
    pid_b = pid // NUM_SPLITS
    pid_s = pid % NUM_SPLITS

    qo_start = tl.load(qo_indptr + pid_b)
    kv_start = tl.load(kv_indptr + pid_b)
    kv_end = tl.load(kv_indptr + pid_b + 1)
    kv_len = kv_end - kv_start

    chunk = tl.cdiv(kv_len, NUM_SPLITS)
    my_start = kv_start + pid_s * chunk
    my_end = tl.minimum(kv_start + (pid_s + 1) * chunk, kv_end)
    my_len = my_end - my_start

    offs_h = tl.arange(0, NHEADS)
    offs_lora = tl.arange(0, D_LORA)
    offs_rope = tl.arange(0, D_ROPE)

    q_lora = tl.load(
        Q + qo_start * stride_q_s + offs_h[:, None] * stride_q_h + offs_lora[None, :]
    )
    q_rope = tl.load(
        Q + qo_start * stride_q_s + offs_h[:, None] * stride_q_h + D_LORA + offs_rope[None, :]
    )

    m_i = tl.full([NHEADS], value=float("-inf"), dtype=tl.float32)
    l_i = tl.zeros([NHEADS], dtype=tl.float32)
    acc = tl.zeros([NHEADS, D_LORA], dtype=tl.float32)

    offs_kv = tl.arange(0, BLOCK_KV)
    packed_lora_idx = offs_lora // 2
    packed_rope_idx = (D_LORA + offs_rope) // 2
    scale_lora_idx = offs_lora // 32
    scale_rope_idx = (D_LORA + offs_rope) // 32
    lora_is_low = (offs_lora % 2) == 0
    rope_is_low = ((D_LORA + offs_rope) % 2) == 0

    for kv_off in range(0, my_len, BLOCK_KV):
        valid = (kv_off + offs_kv) < my_len
        base = my_start + kv_off

        packed_lora = tl.load(
            KV + (base + offs_kv[:, None]) * stride_kv_s + packed_lora_idx[None, :],
            mask=valid[:, None], other=0,
        )
        packed_rope = tl.load(
            KV + (base + offs_kv[:, None]) * stride_kv_s + packed_rope_idx[None, :],
            mask=valid[:, None], other=0,
        )

        nib_lora = tl.where(lora_is_low[None, :], packed_lora & 0xF, packed_lora >> 4)
        nib_rope = tl.where(rope_is_low[None, :], packed_rope & 0xF, packed_rope >> 4)

        scales_lora = tl.load(
            KV_SCALES + (base + offs_kv[:, None]) * stride_scale_s + scale_lora_idx[None, :],
            mask=valid[:, None], other=127,
        )
        scales_rope = tl.load(
            KV_SCALES + (base + offs_kv[:, None]) * stride_scale_s + scale_rope_idx[None, :],
            mask=valid[:, None], other=127,
        )

        k_lora = _decode_mxfp4_nibbles(nib_lora.to(tl.uint8)) * tl.exp2(scales_lora.to(tl.float32) - 127.0)
        k_rope = _decode_mxfp4_nibbles(nib_rope.to(tl.uint8)) * tl.exp2(scales_rope.to(tl.float32) - 127.0)

        s = tl.dot(q_lora, tl.trans(k_lora.to(tl.bfloat16))) + tl.dot(q_rope, tl.trans(k_rope.to(tl.bfloat16)))
        s *= sm_scale
        s = tl.where(valid[None, :], s, float("-inf"))

        m_new = tl.maximum(m_i, tl.max(s, axis=1))
        alpha = tl.exp(m_i - m_new)
        p = tl.exp(s - m_new[:, None])

        acc = acc * alpha[:, None] + tl.dot(p.to(tl.bfloat16), k_lora.to(tl.bfloat16))
        l_i = l_i * alpha + tl.sum(p, axis=1)
        m_i = m_new

    if NUM_SPLITS == 1:
        acc = acc / l_i[:, None]
        tl.store(
            Out + qo_start * stride_o_s + offs_h[:, None] * stride_o_h + offs_lora[None, :],
            acc.to(tl.bfloat16),
        )
    else:
        part_base = pid * NHEADS
        tl.store(
            Partial_O + pid * NHEADS * D_LORA + offs_h[:, None] * D_LORA + offs_lora[None, :],
            acc,
        )
        tl.store(Partial_M + part_base + offs_h, m_i)
        tl.store(Partial_L + part_base + offs_h, l_i)


@triton.jit
def _mla_decode_reduce(
    Partial_O, Partial_M, Partial_L,
    Out, qo_indptr, kv_scale_ptr,
    stride_o_s, stride_o_h,
    NHEADS: tl.constexpr,
    D_LORA: tl.constexpr,
    NUM_SPLITS: tl.constexpr,
    BLOCK_V: tl.constexpr,
):
    pid_b = tl.program_id(0)
    pid_h = tl.program_id(1)
    qo_start = tl.load(qo_indptr + pid_b)
    kv_scale = tl.load(kv_scale_ptr).to(tl.float32)

    offs_v = tl.arange(0, BLOCK_V)
    mask_v = offs_v < D_LORA

    m_global = float("-inf")
    l_global = 0.0
    acc = tl.zeros([BLOCK_V], dtype=tl.float32)

    for s in range(NUM_SPLITS):
        part_idx = pid_b * NUM_SPLITS + s
        m_s = tl.load(Partial_M + part_idx * NHEADS + pid_h)
        l_s = tl.load(Partial_L + part_idx * NHEADS + pid_h)
        po = tl.load(
            Partial_O + part_idx * NHEADS * D_LORA + pid_h * D_LORA + offs_v,
            mask=mask_v, other=0.0,
        )

        m_new = tl.maximum(m_global, m_s)
        alpha = tl.exp(m_global - m_new)
        beta = tl.exp(m_s - m_new)

        acc = acc * alpha + po * beta
        l_global = l_global * alpha + l_s * beta
        m_global = m_new

    acc = acc * (kv_scale / l_global)
    tl.store(
        Out + qo_start * stride_o_s + pid_h * stride_o_h + offs_v,
        acc.to(tl.bfloat16),
        mask=mask_v,
    )


def _get_num_splits(batch_size, kv_seq_len):
    target_blocks = 256
    blocks_per_batch = max(1, target_blocks // batch_size)
    max_splits = max(1, kv_seq_len // 64)
    return min(blocks_per_batch, max_splits)


def _select_asm_splits(batch_size, kv_seq_len):
    if kv_seq_len >= 8192:
        if batch_size >= 64:
            return 64
        return 32 if batch_size >= 16 else 16
    if kv_seq_len >= 4096:
        return 32
    return 16


def _select_mxfp4_splits(batch_size, kv_seq_len):
    if kv_seq_len >= 8192:
        if batch_size >= 64:
            return 32
        return 16
    if kv_seq_len >= 4096:
        return 16
    return 8


def _build_asm_cache(batch_size, kv_seq_len, num_kv_splits, nq, nkv, dq, dv,
                     q_dtype, kv_dtype, qo_indptr, kv_indptr, device):
    total_kv = batch_size * kv_seq_len
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)
    kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
    kv_gran = 128 if kv_seq_len >= 8192 else (64 if kv_seq_len >= 4096 else 16)
    info = get_mla_metadata_info_v1(
        batch_size, 1, 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=device) for s, t in info]
    wm, wi, wis, ri, rfm, rpm = work
    get_mla_metadata_v1(
        qo_indptr, kv_indptr, kv_last_page_len,
        nq // nkv, nkv, True, wm, wis, wi, ri, rfm, rpm,
        page_size=PAGE_SIZE, kv_granularity=kv_gran,
        max_seqlen_qo=1, uni_seqlen_qo=1,
        fast_mode=False, max_split_per_batch=num_kv_splits,
        intra_batch_mode=True, dtype_q=q_dtype, dtype_kv=kv_dtype,
    )
    n_partial = rpm.size(0)
    return {
        "kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,
        "wm": wm, "wi": wi, "wis": wis, "ri": ri, "rfm": rfm, "rpm": rpm,
        "logits": torch.empty((n_partial, 1, nq, dv), dtype=torch.float32, device=device),
        "attn_lse": torch.empty((n_partial, 1, nq, 1), dtype=torch.float32, device=device),
        "q_fp8": torch.empty((batch_size * nq, dq), dtype=FP8_DTYPE, device=device),
        "q_scale": torch.empty((1,), dtype=torch.float32, device=device),
        "o": torch.empty((batch_size, nq, dv), dtype=torch.bfloat16, device=device),
    }


def _run_triton_decode(q, kv_buffer_fp8, kv_scale, qo_indptr, kv_indptr,
                       batch_size, kv_seq_len, nq, dv, block_kv, num_splits, cache_tag):
    cache_key = (cache_tag, batch_size, kv_seq_len, num_splits, block_kv)
    c = _CACHE.get(cache_key)
    if c is None:
        o = torch.empty((batch_size, nq, dv), dtype=torch.bfloat16, device=q.device)
        if num_splits > 1:
            tp = batch_size * num_splits
            po = torch.empty((tp, nq, KV_LORA_RANK), dtype=torch.float32, device=q.device)
            pm = torch.empty((tp, nq), dtype=torch.float32, device=q.device)
            pl = torch.empty((tp, nq), dtype=torch.float32, device=q.device)
        else:
            po = pm = pl = None
        c = {"o": o, "po": po, "pm": pm, "pl": pl}
        if len(_CACHE) >= _CACHE_LIMIT:
            _CACHE.clear()
        _CACHE[cache_key] = c

    kv_flat = kv_buffer_fp8.view(-1, kv_buffer_fp8.shape[-1])
    _mla_decode_attn[(batch_size * num_splits,)](
        q, kv_flat, c["o"], c["po"], c["pm"], c["pl"],
        qo_indptr, kv_indptr, kv_scale, SM_SCALE,
        q.stride(0), q.stride(1), kv_flat.stride(0),
        c["o"].stride(0), c["o"].stride(1),
        NHEADS=nq, D_LORA=KV_LORA_RANK, D_ROPE=QK_ROPE_DIM,
        BLOCK_KV=block_kv, NUM_SPLITS=num_splits,
        num_warps=4, num_stages=2,
    )
    if num_splits > 1:
        _mla_decode_reduce[(batch_size, nq)](
            c["po"], c["pm"], c["pl"], c["o"],
            qo_indptr, kv_scale,
            c["o"].stride(0), c["o"].stride(1),
            NHEADS=nq, D_LORA=KV_LORA_RANK,
            NUM_SPLITS=num_splits, BLOCK_V=KV_LORA_RANK,
            num_warps=4,
        )
    return c["o"]


_FWD_CACHE: dict = {}


def _run_mla_decode_fwd(q, kv_buffer_fp8, kv_scale, qo_indptr, kv_indptr,
                        batch_size, kv_seq_len, nq, nkv, dq, dv):
    """Use high-level mla_decode_fwd (same as reference) for best perf."""
    num_kv_splits = 32 if kv_seq_len >= 4096 else 16
    cache_key = ("fwd", batch_size, kv_seq_len, num_kv_splits)
    c = _FWD_CACHE.get(cache_key)
    device = q.device

    if c is None:
        total_kv = batch_size * kv_seq_len
        kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)
        kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)

        info = get_mla_metadata_info_v1(
            batch_size, 1, nq, FP8_DTYPE, kv_buffer_fp8.dtype,
            is_sparse=False, fast_mode=False,
            num_kv_splits=num_kv_splits, intra_batch_mode=True,
        )
        work = [torch.empty(s, dtype=t, device=device) for s, t in info]
        wm, wi, wis, ri, rfm, rpm = work

        get_mla_metadata_v1(
            qo_indptr, kv_indptr, kv_last_page_len,
            nq // nkv, nkv, True, wm, wis, wi, ri, rfm, rpm,
            page_size=PAGE_SIZE, kv_granularity=16,
            max_seqlen_qo=1, uni_seqlen_qo=1,
            fast_mode=False, max_split_per_batch=num_kv_splits,
            intra_batch_mode=True, dtype_q=FP8_DTYPE, dtype_kv=kv_buffer_fp8.dtype,
        )

        c = {
            "kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,
            "wm": wm, "wi": wi, "wis": wis, "ri": ri, "rfm": rfm, "rpm": rpm,
            "q_fp8": torch.empty((batch_size * nq, dq), dtype=FP8_DTYPE, device=device),
            "q_scale": torch.empty((1,), dtype=torch.float32, device=device),
            "o": torch.empty((batch_size, nq, dv), dtype=torch.bfloat16, device=device),
        }
        if len(_FWD_CACHE) >= _CACHE_LIMIT:
            _FWD_CACHE.clear()
        _FWD_CACHE[cache_key] = c

    q_2d = q.view(-1, dq)
    if not q_2d.is_contiguous():
        q_2d = q_2d.contiguous()
    dynamic_per_tensor_quant_fp8_i8(c["q_fp8"], q_2d, c["q_scale"])
    kv_4d = kv_buffer_fp8.view(kv_buffer_fp8.shape[0], PAGE_SIZE, nkv, kv_buffer_fp8.shape[-1])

    mla_decode_fwd(
        c["q_fp8"].view(batch_size, nq, dq),
        kv_4d,
        c["o"],
        qo_indptr, kv_indptr,
        c["kv_indices"], c["kv_last_page_len"],
        1, page_size=PAGE_SIZE, nhead_kv=nkv, sm_scale=SM_SCALE,
        logit_cap=0.0, num_kv_splits=num_kv_splits,
        q_scale=c["q_scale"], kv_scale=kv_scale,
        intra_batch_mode=True,
        work_meta_data=c["wm"], work_indptr=c["wi"], work_info_set=c["wis"],
        reduce_indptr=c["ri"], reduce_final_map=c["rfm"], reduce_partial_map=c["rpm"],
    )
    return c["o"]


def _run_asm_decode(q, kv_buffer_fp8, kv_scale, qo_indptr, kv_indptr,
                    batch_size, kv_seq_len, nq, nkv, dq, dv):
    num_kv_splits = _select_asm_splits(batch_size, kv_seq_len)
    cache_key = ("asm", batch_size, kv_seq_len, num_kv_splits)
    c = _ASM_CACHE.get(cache_key)
    if c is None:
        c = _build_asm_cache(
            batch_size, kv_seq_len, num_kv_splits, nq, nkv, dq, dv,
            FP8_DTYPE, kv_buffer_fp8.dtype,
            qo_indptr, kv_indptr, q.device,
        )
        if len(_ASM_CACHE) >= _CACHE_LIMIT:
            _ASM_CACHE.clear()
        _ASM_CACHE[cache_key] = c

    q_2d = q.view(-1, dq)
    if not q_2d.is_contiguous():
        q_2d = q_2d.contiguous()
    dynamic_per_tensor_quant_fp8_i8(c["q_fp8"], q_2d, c["q_scale"])
    kv_4d = kv_buffer_fp8.view(kv_buffer_fp8.shape[0], PAGE_SIZE, nkv, kv_buffer_fp8.shape[-1])

    mla_decode_stage1_asm_fwd(
        c["q_fp8"].view(batch_size, nq, dq), kv_4d,
        qo_indptr, kv_indptr,
        c["kv_indices"], c["kv_last_page_len"],
        None, c["wm"], c["wi"], c["wis"],
        1, PAGE_SIZE, nkv, SM_SCALE,
        c["logits"], c["attn_lse"], c["o"],
        c["q_scale"], kv_scale,
    )
    mla_reduce_v1(
        c["logits"], c["attn_lse"],
        c["ri"], c["rfm"], c["rpm"],
        1, c["o"], None,
    )
    return c["o"]


def _dequantize_mxfp4_kv(kv_buffer_mxfp4, kv_scale_mxfp4):
    total_kv = kv_buffer_mxfp4.shape[0]
    key = (
        kv_buffer_mxfp4.data_ptr(),
        kv_scale_mxfp4.data_ptr(),
        total_kv,
        kv_buffer_mxfp4.device,
    )
    cached = _MXFP4_CACHE.get(key)
    if cached is not None:
        return cached

    num_blocks = QK_HEAD_DIM // 32
    kv_fp32 = mxfp4_to_f32(kv_buffer_mxfp4.view(total_kv, QK_HEAD_DIM // 2))
    scale_f32 = e8m0_to_f32(kv_scale_mxfp4)[:total_kv, :num_blocks]
    kv_fp32 = kv_fp32.view(total_kv, num_blocks, 32) * scale_f32.unsqueeze(-1)
    kv_bf16 = kv_fp32.view(total_kv, 1, QK_HEAD_DIM).to(torch.bfloat16)

    if len(_MXFP4_CACHE) >= 4:
        _MXFP4_CACHE.clear()
    _MXFP4_CACHE[key] = kv_bf16
    return kv_bf16


def _run_triton_decode_mxfp4(q, kv_buffer_mxfp4, kv_scale_mxfp4, qo_indptr, kv_indptr,
                             batch_size, kv_seq_len, nq, dv, num_splits, cache_tag):
    cache_key = (cache_tag, batch_size, kv_seq_len, num_splits)
    c = _CACHE.get(cache_key)
    if c is None:
        o = torch.empty((batch_size, nq, dv), dtype=torch.bfloat16, device=q.device)
        if num_splits > 1:
            tp = batch_size * num_splits
            po = torch.empty((tp, nq, KV_LORA_RANK), dtype=torch.float32, device=q.device)
            pm = torch.empty((tp, nq), dtype=torch.float32, device=q.device)
            pl = torch.empty((tp, nq), dtype=torch.float32, device=q.device)
        else:
            po = pm = pl = None
        c = {"o": o, "po": po, "pm": pm, "pl": pl}
        if len(_CACHE) >= _CACHE_LIMIT:
            _CACHE.clear()
        _CACHE[cache_key] = c

    # Cast to uint8 for Triton (FP4 dtypes not supported in pointer canonicalization)
    kv_flat = kv_buffer_mxfp4.view(torch.uint8).view(-1, kv_buffer_mxfp4.shape[-1])
    kv_scales = kv_scale_mxfp4.view(torch.uint8).view(kv_scale_mxfp4.shape[0], -1)
    _mla_decode_attn_mxfp4[(batch_size * num_splits,)](
        q, kv_flat, kv_scales, c["o"], c["po"], c["pm"], c["pl"],
        qo_indptr, kv_indptr, SM_SCALE,
        q.stride(0), q.stride(1),
        kv_flat.stride(0), kv_scales.stride(0),
        c["o"].stride(0), c["o"].stride(1),
        NHEADS=nq, D_LORA=KV_LORA_RANK, D_ROPE=QK_ROPE_DIM,
        BLOCK_KV=128, NUM_SPLITS=num_splits,
        num_warps=4, num_stages=2,
    )
    if num_splits > 1:
        unit_scale = torch.ones((1,), dtype=torch.float32, device=q.device)
        _mla_decode_reduce[(batch_size, nq)](
            c["po"], c["pm"], c["pl"], c["o"],
            qo_indptr, unit_scale,
            c["o"].stride(0), c["o"].stride(1),
            NHEADS=nq, D_LORA=KV_LORA_RANK,
            NUM_SPLITS=num_splits, BLOCK_V=KV_LORA_RANK,
            num_warps=4,
        )
    return c["o"]


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

    batch_size = config["batch_size"]
    nq = config["num_heads"]
    nkv = config["num_kv_heads"]
    dq = config["qk_head_dim"]
    dv = config["v_head_dim"]
    kv_seq_len = config["kv_seq_len"]

    kv_buffer_fp8, kv_scale = kv_data["fp8"]

    # Short KV: Use custom Triton with FP8 KV
    if kv_seq_len <= 2048:
        num_splits = _get_num_splits(batch_size, kv_seq_len)
        return _run_triton_decode(
            q, kv_buffer_fp8, kv_scale, qo_indptr, kv_indptr,
            batch_size, kv_seq_len, nq, dv,
            block_kv=64, num_splits=num_splits, cache_tag="triton-short",
        )

    # Long KV: ASM decode path
    return _run_asm_decode(
        q, kv_buffer_fp8, kv_scale, qo_indptr, kv_indptr,
        batch_size, kv_seq_len, nq, nkv, dq, dv,
    )
scrolls · 590 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