Skip to content
KernelIndex
Search⌘K

submission 590963

gwokhou · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v4.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-590963?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
1.13ms
#728 of 766
2026-03-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a61fb8eccef7ec61087a9e95187900937f869f1e02e49cb8ec56445dbb659129
license declaredunknown
license concludedunknown
authorsgwokhou
imported2026-08-26

Techniques

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

fp4Submission v4: inlines mla_decode_fwd and MXFP4 stage1 from aiter (no import of mla_decode_fwd).
num-warps = 4num_warps=4,
persistent-kernelDecode only — persistent mode with get_mla_metadata_v1.
split-kfor split_kv_id in range(0, num_valid_kv_splits):
stages = 2num_stages=2,
tile-k = 32BLOCK_K = 32
tile-n = 64BLOCK_N = 64

Kernel source

submission_v4.py1119 lines
"""
Reference implementation for MLA (Multi-head Latent Attention) decode kernel.
Submission v4: inlines mla_decode_fwd and MXFP4 stage1 from aiter (no import of mla_decode_fwd).

Uses the same aiter MLA API; mla_decode_fwd and its children call chain
(get_meta_param, _fwd_kernel_stage2_asm, mla_decode_stage1_mxfp4) are copied here.
DeepSeek R1 forward_absorb MLA: absorbed q (576), compressed kv_buffer (576),
output v_head_dim = kv_lora_rank = 512.

Decode only — persistent mode with get_mla_metadata_v1.
"""

from __future__ import annotations
import functools

import aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.jit.utils.chip_info import get_cu_num, get_gfx
from aiter.ops.triton.utils._triton import arch_info
from aiter.utility.fp4_utils import dynamic_mxfp4_quant, e8m0_to_f32, mxfp4_to_f32
from task import input_t, output_t
import torch
import triton
import triton.language as tl

# ---------------------------------------------------------------------------
# DeepSeek R1 latent MQA constants (forward_absorb path)
# ---------------------------------------------------------------------------
TOTAL_NUM_HEADS = 128
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)

PAGE_SIZE = 1
NUM_KV_SPLITS = 32

FP8_DTYPE = aiter_dtypes.fp8
MXFP4_DTYPE = aiter_dtypes.fp4x2

QKV_DTYPE = "mxfp4"

# MXFP4 block size (must match fp4_utils.dynamic_mxfp4_quant)
MXFP4_BLOCK_SIZE = 32


# ---------------------------------------------------------------------------
# FP8 quantization
# ---------------------------------------------------------------------------


def quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """
    Dynamic per-tensor FP8 quantization (following sglang scaled_fp8_quant).

    Args:
        tensor: bf16 tensor to quantize

    Returns:
        (fp8_tensor, scale) where scale is a scalar float32 tensor.
        Dequantize: fp8_tensor.to(bf16) * scale
    """
    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)


# ---------------------------------------------------------------------------
# MXFP4 quantization (aiter native: block-32)
# ---------------------------------------------------------------------------
def quantize_mxfp4(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    orig_shape = tensor.shape
    B, M, N = orig_shape
    tensor_2d = tensor.reshape(B * M, N)
    fp4_data_2d, scale_e8m0 = dynamic_mxfp4_quant(tensor_2d)
    fp4_data = fp4_data_2d.view(B, M, N // 2)
    return fp4_data, scale_e8m0


def dequantize_mxfp4(
    fp4_data: torch.Tensor,
    scale_e8m0: torch.Tensor,
    orig_shape: tuple,
    dtype: torch.dtype = torch.bfloat16,
) -> torch.Tensor:
    """
    Dequantize MXFP4 tensor using aiter utilities.

    Note: dynamic_mxfp4_quant may pad both row and block dimensions in scale_e8m0.
    We trim scales to match the actual data dimensions.

    Args:
        fp4_data:   packed FP4 data, shape [B, M, N//2] in fp4x2 or uint8
        scale_e8m0: E8M0 block scale factors (possibly padded) in fp8_e8m0
        orig_shape: original (B, M, N) for reshaping
        dtype:      output dtype

    Returns:
        Dequantized tensor of shape orig_shape.
    """
    B, M, N = orig_shape
    num_rows = B * M
    block_size = 32
    num_blocks = N // block_size
    fp4_data_2d = fp4_data.reshape(num_rows, N // 2)
    float_vals = mxfp4_to_f32(fp4_data_2d)
    scale_f32 = e8m0_to_f32(scale_e8m0)
    scale_f32 = scale_f32[:num_rows, :num_blocks]
    float_vals_blocked = float_vals.view(num_rows, num_blocks, block_size)
    scaled = float_vals_blocked * scale_f32.unsqueeze(-1)
    return scaled.view(B, M, N).to(dtype)


# ---------------------------------------------------------------------------
# Inlined from aiter.mla: stage2 kernel + get_meta_param + mla_decode_fwd
# ---------------------------------------------------------------------------
@triton.jit
def _fwd_kernel_stage2_asm(
    Mid_O,
    Mid_lse,
    O,
    qo_indptr,
    kv_indptr,
    num_kv_splits_indptr,
    stride_mid_ob: tl.int64,
    stride_mid_oh: tl.int64,
    stride_mid_os: tl.int64,
    stride_obs: tl.int64,
    stride_oh: tl.int64,
    MAYBE_FINAL_OUT: tl.constexpr,
    BATCH_NUM: tl.constexpr,
    BLOCK_DV: tl.constexpr,
    Lv: tl.constexpr,
    mgc: tl.constexpr,
):
    cur_batch = tl.program_id(0)
    cur_head = tl.program_id(1)
    cur_qo_start = tl.load(qo_indptr + cur_batch)
    cur_qo_end = tl.load(qo_indptr + cur_batch + 1)
    cur_split_start = tl.load(num_kv_splits_indptr + cur_batch)
    cur_split_end = tl.load(num_kv_splits_indptr + cur_batch + 1)
    num_max_kv_splits = tl.load(num_kv_splits_indptr + BATCH_NUM)
    cur_kv_seq_len = tl.load(kv_indptr + cur_batch + 1) - tl.load(kv_indptr + cur_batch)

    offs_d = tl.arange(0, BLOCK_DV)
    mask_d = offs_d < Lv

    offs_logic = cur_qo_start * stride_mid_ob + cur_head * stride_mid_oh
    offs_v = offs_logic * Lv + offs_d
    num_valid_kv_splits = tl.minimum(
        cur_split_end - cur_split_start, tl.cdiv(cur_kv_seq_len, mgc)
    )
    FINAL_OUT = MAYBE_FINAL_OUT and num_max_kv_splits == BATCH_NUM

    for cur_qo in range(cur_qo_start, cur_qo_end):
        if FINAL_OUT:
            input_ptr = Mid_O.to(tl.pointer_type(O.type.element_ty))
            out = tl.load(
                input_ptr
                + Lv * (cur_qo * stride_mid_os + cur_head * stride_mid_oh)
                + offs_d,
                mask=mask_d,
                other=0.0,
            )
            tl.store(
                O + cur_qo * stride_obs + cur_head * stride_oh + offs_d,
                out,
                mask=mask_d,
            )
        else:
            e_sum = 0.0
            e_max = -float("inf")
            acc = tl.zeros((BLOCK_DV,), dtype=tl.float32)
            for split_kv_id in range(0, num_valid_kv_splits):
                tv = tl.load(
                    Mid_O + offs_v + split_kv_id * stride_mid_os * Lv,
                    mask=mask_d,
                    other=0.0,
                )
                tlogic = tl.load(Mid_lse + offs_logic + split_kv_id * stride_mid_os)
                n_e_max = tl.maximum(tlogic, e_max)
                old_scale = tl.exp(e_max - n_e_max)
                acc *= old_scale
                exp_logic = tl.exp(tlogic - n_e_max)
                acc += exp_logic * tv
                e_sum = e_sum * old_scale + exp_logic
                e_max = n_e_max
            offs_logic += stride_mid_ob
            offs_v += stride_mid_ob * Lv
            tl.store(
                O + cur_qo * stride_obs + cur_head * stride_oh + offs_d,
                acc / e_sum,
                mask=mask_d,
            )


@functools.lru_cache()
def get_meta_param(num_kv_splits, bs, total_kv, nhead, max_seqlen_q, dtype):
    if num_kv_splits is None:
        cu_num = get_cu_num()
        avg_kv = total_kv / bs
        overhead = 84.1
        tmp = [
            (
                bs
                * i
                / ((bs * i + cu_num - 1) // cu_num * cu_num)
                * avg_kv
                / (avg_kv + overhead * i),
                i,
            )
            for i in range(1, 17)
        ]
        num_kv_splits = sorted(tmp, key=lambda x: x[0], reverse=True)[0][1]

    get_block_n_fp8 = {
        16: 128,
        32: 128,
        48: 64,
        64: 64,
        128: 32,
        256: 32,
        384: 32,
        512: 32,
    }
    if dtype == aiter_dtypes.fp8:
        min_block_n = get_block_n_fp8[int(nhead * max_seqlen_q)]
        num_kv_splits = min(
            num_kv_splits, int(total_kv / bs + min_block_n - 1) // min_block_n
        )

    num_kv_splits_indptr = torch.arange(
        0, (bs + 1) * num_kv_splits, num_kv_splits, dtype=torch.int, device="cuda"
    )
    return num_kv_splits, num_kv_splits_indptr


def mla_decode_fwd(
    q,
    kv_buffer,
    o,
    qo_indptr,
    kv_indptr,
    kv_indices,
    kv_last_page_lens,
    max_seqlen_q,
    page_size=1,
    nhead_kv=1,
    sm_scale=None,
    logit_cap=0.0,
    num_kv_splits=None,
    num_kv_splits_indptr=None,
    work_meta_data=None,
    work_indptr=None,
    work_info_set=None,
    reduce_indptr=None,
    reduce_final_map=None,
    reduce_partial_map=None,
    q_scale=None,
    kv_scale=None,
    intra_batch_mode=False,
    return_logits=False,
    return_lse=False,
    use_mxfp4=False,
):
    device = q.device
    assert logit_cap <= 0, f"{logit_cap=} is not support yet"
    if kv_buffer.dtype != torch.uint8:
        _, _, _, qk_head_dim = kv_buffer.shape
    else:
        _, _, qk_head_dim = q.shape

    if sm_scale is None:
        sm_scale = 1.0 / (qk_head_dim**0.5)

    ori_total_s, ori_nhead, ori_v_head_dim = o.shape
    total_s, nhead, v_head_dim = o.shape
    bs = qo_indptr.shape[0] - 1
    total_kv = kv_indices.shape[0]

    persistent_mode = work_meta_data is not None
    io_transformed = False

    if not persistent_mode:
        if num_kv_splits is None or num_kv_splits_indptr is None:
            num_kv_splits, num_kv_splits_indptr = get_meta_param(
                num_kv_splits, bs, total_kv, nhead, max_seqlen_q, q.dtype
            )

        mgc = 64 if max_seqlen_q == 1 and nhead == 16 else 16
        MAYBE_FINAL_OUT = True
        if nhead == 16 and max_seqlen_q == 1:
            MAYBE_FINAL_OUT = False

        logits = (
            o.view((total_s, num_kv_splits, nhead, v_head_dim))
            if (
                num_kv_splits == 1
                and (
                    q.dtype == aiter_dtypes.fp8
                    or (q.dtype == aiter_dtypes.bf16 and max_seqlen_q == 4)
                )
            )
            else torch.empty(
                (total_s, num_kv_splits, nhead, v_head_dim),
                dtype=torch.float32,
                device=device,
            )
        )
        attn_lse = torch.empty(
            (total_s, num_kv_splits, nhead, 1), dtype=torch.float32, device=device
        )
        final_lse = torch.empty((total_s, nhead), dtype=torch.float32, device=device)

        use_mxfp4_path = (
            use_mxfp4
            and get_gfx() == "gfx950"
            and arch_info.is_fp4_avail()
            and q.dtype == aiter_dtypes.bf16
            and kv_buffer.dtype == aiter_dtypes.bf16
            and qk_head_dim % 32 == 0
            and v_head_dim % 32 == 0
        )
        if use_mxfp4_path:
            mla_decode_stage1_mxfp4(
                q,
                kv_buffer,
                qo_indptr,
                kv_indptr,
                kv_indices,
                kv_last_page_lens,
                num_kv_splits_indptr,
                max_seqlen_q,
                page_size,
                nhead_kv,
                sm_scale,
                logits,
                attn_lse,
                o,
            )
        else:
            aiter.mla_decode_stage1_asm_fwd(
                q,
                kv_buffer,
                qo_indptr,
                kv_indptr,
                kv_indices,
                kv_last_page_lens,
                num_kv_splits_indptr,
                None,
                None,
                None,
                max_seqlen_q,
                page_size,
                nhead_kv,
                sm_scale,
                logits,
                attn_lse,
                o,
                q_scale,
                kv_scale,
            )

        if num_kv_splits == 1 and (
            use_mxfp4_path
            or q.dtype == aiter_dtypes.fp8
            or (q.dtype == aiter_dtypes.bf16 and max_seqlen_q == 4)
            or (
                q.dtype == aiter_dtypes.bf16
                and kv_buffer.dtype == aiter_dtypes.bf16
                and nhead in [32, 64]
            )
        ):
            return logits.view(total_s, nhead, v_head_dim), attn_lse

        Lv = v_head_dim
        BLOCK_DV = triton.next_power_of_2(Lv)
        grid = (bs, nhead)
        extra_kargs = {"waves_per_eu": 4}
        _fwd_kernel_stage2_asm[grid](
            logits,
            attn_lse,
            o,
            qo_indptr,
            kv_indptr,
            num_kv_splits_indptr,
            attn_lse.stride(0),
            attn_lse.stride(2),
            attn_lse.stride(1),
            o.stride(0),
            o.stride(1),
            MAYBE_FINAL_OUT=MAYBE_FINAL_OUT,
            BATCH_NUM=bs,
            BLOCK_DV=BLOCK_DV,
            Lv=Lv,
            mgc=mgc,
            num_warps=4,
            num_stages=2,
            **extra_kargs,
        )
    else:
        if num_kv_splits is None:
            num_kv_splits = get_cu_num()
        if (
            nhead == 16
            or (
                nhead == 128
                and q.dtype == aiter_dtypes.fp8
                and kv_buffer.dtype == aiter_dtypes.fp8
            )
            or (
                get_gfx() == "gfx950"
                and nhead == 32
                and q.dtype == aiter_dtypes.fp8
                and kv_buffer.dtype == aiter_dtypes.fp8
                and max_seqlen_q == 4
            )
        ):
            pass
        elif nhead in range(32, 128 + 1, 16) and persistent_mode:
            total_s = ori_total_s * (ori_nhead // 16)
            nhead = 16
            q = q.view(total_s, nhead, -1)
            o = o.view(total_s, nhead, -1)
            io_transformed = True
        else:
            assert False, f"{nhead=} and {max_seqlen_q=} not supported"

        logits = torch.empty(
            (reduce_partial_map.size(0) * max_seqlen_q, 1, nhead, v_head_dim),
            dtype=torch.float32,
            device=device,
        )
        attn_lse = torch.empty(
            (reduce_partial_map.size(0) * max_seqlen_q, 1, nhead, 1),
            dtype=torch.float32,
            device=device,
        )
        final_lse = (
            torch.empty((total_s, nhead), dtype=torch.float32, device=device)
            if return_lse
            else None
        )

        aiter.mla_decode_stage1_asm_fwd(
            q,
            kv_buffer,
            qo_indptr,
            kv_indptr,
            kv_indices,
            kv_last_page_lens,
            num_kv_splits_indptr,
            work_meta_data,
            work_indptr,
            work_info_set,
            max_seqlen_q,
            page_size,
            nhead_kv,
            sm_scale,
            logits,
            attn_lse,
            o,
            q_scale,
            kv_scale,
        )
        aiter.mla_reduce_v1(
            logits,
            attn_lse,
            reduce_indptr,
            reduce_final_map,
            reduce_partial_map,
            max_seqlen_q,
            o,
            final_lse,
        )

    if io_transformed:
        if return_logits:
            logits = logits.view(-1, 1, ori_nhead, v_head_dim)
        q = q.view(ori_total_s, ori_nhead, -1)
        o = o.view(ori_total_s, ori_nhead, -1)

    return logits, final_lse


# ---------------------------------------------------------------------------
# Inlined from aiter.ops.triton.attention.mla_decode_stage1_mxfp4
# ---------------------------------------------------------------------------
@triton.jit
def _qkt_mxfp4_kernel(
    Q_fp4_ptr,
    Q_scale_ptr,
    K_fp4_ptr,
    K_scale_ptr,
    scores_ptr,
    qo_indptr_ptr,
    kv_indptr_ptr,
    stride_q_s,
    stride_q_h,
    stride_q_k,
    stride_q_scale_s,
    stride_q_scale_k,
    stride_k_kv,
    stride_k_k,
    stride_k_scale_kv,
    stride_k_scale_k,
    stride_scores_s,
    stride_scores_h,
    stride_scores_n,
    total_s: tl.int32,
    nhead: tl.int32,
    qk_head_dim: tl.int32,
    max_kv_len: tl.int32,
    num_kv_splits: tl.int32,
    sm_scale: tl.float32,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    batch = tl.program_id(0)
    head = tl.program_id(1)
    split = tl.program_id(2)

    qo_start = tl.load(qo_indptr_ptr + batch)
    kv_start_b = tl.load(kv_indptr_ptr + batch)
    kv_end_b = tl.load(kv_indptr_ptr + batch + 1)
    kv_len_b = kv_end_b - kv_start_b

    kv_chunk_len = kv_len_b // num_kv_splits
    kv_start = kv_start_b + split * kv_chunk_len
    kv_end = kv_start + kv_chunk_len
    if split == num_kv_splits - 1:
        kv_end = kv_end_b
    kv_chunk_len = kv_end - kv_start

    if kv_chunk_len <= 0:
        return

    q_idx = qo_start
    if q_idx >= total_s:
        return

    scale_k_blocks = qk_head_dim // MXFP4_BLOCK_SIZE
    base_scores = (
        scores_ptr
        + q_idx * stride_scores_s
        + head * stride_scores_h
        + split * max_kv_len * stride_scores_n
    )

    for n_start in range(0, kv_chunk_len, BLOCK_N):
        n_end = tl.minimum(n_start + BLOCK_N, kv_chunk_len)
        n_size = n_end - n_start
        offs_n = tl.arange(0, BLOCK_N)

        acc = tl.zeros((1, BLOCK_N), dtype=tl.float32)

        for k_block in range(0, scale_k_blocks):
            k_start = k_block * MXFP4_BLOCK_SIZE
            offs_k = tl.arange(0, BLOCK_K // 2)

            q_ptrs = (
                Q_fp4_ptr
                + q_idx * stride_q_s
                + head * stride_q_h
                + (k_start // 2) * stride_q_k
            )
            q_load = tl.load(
                q_ptrs,
                mask=offs_k[None, :] < (qk_head_dim // 2 - k_start // 2),
                other=0,
            )
            q_scale_ptrs = (
                Q_scale_ptr
                + (q_idx * nhead + head) * stride_q_scale_s
                + k_block * stride_q_scale_k
            )
            q_s = tl.load(q_scale_ptrs)
            q_s_bc = tl.broadcast_to(q_s, (1, 1))

            k_base = (
                K_fp4_ptr
                + (kv_start + n_start) * stride_k_kv
                + (k_start // 2) * stride_k_k
            )
            mask_n = offs_n < n_size
            mask_k = offs_k < (qk_head_dim // 2 - k_start // 2)
            k_load = tl.load(
                k_base + offs_k[:, None] * stride_k_k + offs_n[None, :] * stride_k_kv,
                mask=mask_k[:, None] & mask_n[None, :],
                other=0,
            )
            k_scale_ptrs = (
                K_scale_ptr
                + (kv_start + n_start) * stride_k_scale_kv
                + k_block * stride_k_scale_k
            )
            k_s = tl.load(
                k_scale_ptrs + offs_n * stride_k_scale_kv, mask=mask_n, other=127
            )
            k_s_bc = tl.broadcast_to(k_s[None, :], (BLOCK_K // 2, BLOCK_N))

            acc = tl.dot_scaled(
                q_load,
                q_s_bc,
                "e2m1",
                k_load,
                k_s_bc,
                "e2m1",
                acc,
            )

        acc = acc * sm_scale
        tl.store(
            base_scores + n_start * stride_scores_n + offs_n * stride_scores_n,
            acc[0, :],
            mask=offs_n < n_size,
        )


@triton.jit
def _softmax_lse_pv_f32_kernel(
    scores_ptr,
    V_fp32_ptr,
    logits_ptr,
    lse_ptr,
    o_ptr,
    qo_indptr_ptr,
    kv_indptr_ptr,
    stride_scores_s,
    stride_scores_h,
    stride_scores_n,
    stride_v_kv,
    stride_v_d,
    stride_logits_s,
    stride_logits_h,
    stride_logits_d,
    stride_lse_s,
    stride_lse_h,
    stride_lse_split,
    stride_logits_split,
    stride_o_s,
    stride_o_h,
    stride_o_d,
    total_s: tl.int32,
    nhead: tl.int32,
    v_head_dim: tl.int32,
    max_kv_len: tl.int32,
    num_kv_splits: tl.int32,
    BLOCK_D: tl.constexpr,
):
    batch = tl.program_id(0)
    head = tl.program_id(1)
    split = tl.program_id(2)

    qo_start = tl.load(qo_indptr_ptr + batch)
    kv_start_b = tl.load(kv_indptr_ptr + batch)
    kv_end_b = tl.load(kv_indptr_ptr + batch + 1)
    kv_len_b = kv_end_b - kv_start_b

    kv_chunk_len = kv_len_b // num_kv_splits
    kv_start = kv_start_b + split * kv_chunk_len
    kv_end = kv_start + kv_chunk_len
    if split == num_kv_splits - 1:
        kv_end = kv_end_b
    kv_chunk_len = kv_end - kv_start

    if kv_chunk_len <= 0:
        return

    q_idx = qo_start
    if q_idx >= total_s:
        return

    base_scores = (
        scores_ptr
        + q_idx * stride_scores_s
        + head * stride_scores_h
        + split * max_kv_len * stride_scores_n
    )

    m_i = -float("inf")
    l_i = 0.0
    BLOCK_N = 64
    for n_start in range(0, kv_chunk_len, BLOCK_N):
        n_end = tl.minimum(n_start + BLOCK_N, kv_chunk_len)
        offs_n = tl.arange(0, BLOCK_N)
        mask_n = offs_n < (n_end - n_start)
        s = tl.load(
            base_scores + (n_start + offs_n) * stride_scores_n,
            mask=mask_n,
            other=-float("inf"),
        )
        m_ij = tl.maximum(m_i, tl.max(s, axis=0))
        p = tl.exp(s - m_ij)
        p = tl.where(mask_n, p, 0.0)
        alpha = tl.exp(m_i - m_ij)
        l_i = l_i * alpha + tl.sum(p, axis=0)
        m_i = m_ij

    lse = m_i + tl.log(l_i)

    acc_o = tl.zeros((BLOCK_D,), dtype=tl.float32)
    offs_d = tl.arange(0, BLOCK_D)

    for n_start in range(0, kv_chunk_len, BLOCK_N):
        n_end = tl.minimum(n_start + BLOCK_N, kv_chunk_len)
        offs_n = tl.arange(0, BLOCK_N)
        mask_n = offs_n < (n_end - n_start)
        s = tl.load(
            base_scores + (n_start + offs_n) * stride_scores_n,
            mask=mask_n,
            other=0.0,
        )
        p = tl.exp(s - m_i) / l_i
        p = tl.where(mask_n, p, 0.0)

        for ni in range(BLOCK_N):
            if n_start + ni >= kv_chunk_len:
                break
            p_val = tl.load(base_scores + (n_start + ni) * stride_scores_n)
            p_val = tl.exp(p_val - m_i) / l_i
            v_row = tl.load(
                V_fp32_ptr
                + (kv_start + n_start + ni) * stride_v_kv
                + offs_d * stride_v_d,
                mask=offs_d < v_head_dim,
                other=0.0,
            )
            acc_o += p_val * v_row

    lse_off = q_idx * stride_lse_s + split * stride_lse_split + head * stride_lse_h
    tl.store(lse_ptr + lse_off, lse)

    logits_off = (
        q_idx * stride_logits_s + split * stride_logits_split + head * stride_logits_h
    )
    tl.store(
        logits_ptr + logits_off + offs_d * stride_logits_d,
        acc_o,
        mask=offs_d < v_head_dim,
    )

    if num_kv_splits == 1:
        tl.store(
            o_ptr + q_idx * stride_o_s + head * stride_o_h + offs_d * stride_o_d,
            acc_o,
            mask=offs_d < v_head_dim,
        )


def _dequant_mxfp4_to_f32(
    x_fp4: torch.Tensor, scale_e8m0: torch.Tensor, dim: int
) -> torch.Tensor:
    from aiter.utility import fp4_utils

    vals = fp4_utils.mxfp4_to_f32(x_fp4)
    scale = scale_e8m0.view(torch.uint8).to(torch.float32)
    scale = torch.pow(2.0, 127.0 - scale)
    if scale.dim() == 2:
        scale = scale[:, : dim // MXFP4_BLOCK_SIZE].repeat_interleave(
            MXFP4_BLOCK_SIZE, dim=1
        )
    return vals * scale


def _gather_kv_from_paged(
    kv_buffer: torch.Tensor,
    kv_indices: torch.Tensor,
    page_size: int,
    nhead_kv: int,
    qk_head_dim: int,
    v_head_dim: int,
) -> tuple[torch.Tensor, torch.Tensor]:
    total_kv = kv_indices.shape[0]
    if kv_buffer.dim() == 4:
        num_page, ps, nk, feat = kv_buffer.shape
        kv_flat = kv_buffer.reshape(num_page * ps, nk, feat)
        flat_idx = (
            kv_indices * page_size
            + torch.arange(total_kv, device=kv_indices.device, dtype=kv_indices.dtype)
            % page_size
        )
        kv_gathered = kv_flat[flat_idx]
    else:
        kv_flat = kv_buffer
        kv_gathered = kv_flat[kv_indices]
    if kv_gathered.shape[-1] >= qk_head_dim + v_head_dim:
        K = kv_gathered[..., :qk_head_dim].contiguous()
        V = kv_gathered[..., qk_head_dim : qk_head_dim + v_head_dim].contiguous()
    else:
        K = kv_gathered.contiguous()
        V = kv_gathered[..., :v_head_dim].contiguous()
    return K, V


def mla_decode_stage1_mxfp4(
    q: torch.Tensor,
    kv_buffer: torch.Tensor,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    kv_indices: torch.Tensor,
    kv_last_page_lens: torch.Tensor,
    num_kv_splits_indptr: torch.Tensor,
    max_seqlen_q: int,
    page_size: int,
    nhead_kv: int,
    sm_scale: float,
    logits: torch.Tensor,
    attn_lse: torch.Tensor,
    o: torch.Tensor,
) -> None:
    assert arch_info.is_fp4_avail(), "MXFP4 requires gfx950"
    device = q.device
    total_s, nhead, qk_head_dim = q.shape
    v_head_dim = o.shape[-1]
    bs = qo_indptr.shape[0] - 1
    total_kv = kv_indices.shape[0]
    num_kv_splits = (num_kv_splits_indptr[1] - num_kv_splits_indptr[0]).item()

    assert qk_head_dim % MXFP4_BLOCK_SIZE == 0
    assert v_head_dim % MXFP4_BLOCK_SIZE == 0

    q_flat = q.reshape(-1, qk_head_dim).to(torch.bfloat16)
    Q_fp4, Q_scale = dynamic_mxfp4_quant(q_flat)
    Q_fp4 = Q_fp4.reshape(total_s, nhead, -1)
    scale_n_q = qk_head_dim // MXFP4_BLOCK_SIZE
    Q_scale = Q_scale.reshape(-1, scale_n_q)[: total_s * nhead].contiguous()

    K_gather, V_gather = _gather_kv_from_paged(
        kv_buffer, kv_indices, page_size, nhead_kv, qk_head_dim, v_head_dim
    )
    K_flat = K_gather.reshape(-1, qk_head_dim).to(torch.bfloat16)
    V_flat = V_gather.reshape(-1, v_head_dim).to(torch.bfloat16)
    K_fp4, K_scale = dynamic_mxfp4_quant(K_flat)
    V_fp4, V_scale = dynamic_mxfp4_quant(V_flat)

    V_fp32 = _dequant_mxfp4_to_f32(V_fp4, V_scale, v_head_dim)

    max_kv_len = (kv_indptr[1:] - kv_indptr[:-1]).max().item()
    max_kv_chunk = (max_kv_len + num_kv_splits - 1) // num_kv_splits

    scores = torch.empty(
        (total_s, nhead, num_kv_splits * max_kv_chunk),
        dtype=torch.float32,
        device=device,
    )
    scores.fill_(-1e9)

    BLOCK_N = 64
    BLOCK_K = 32

    grid_qkt = (bs, nhead, num_kv_splits)
    _qkt_mxfp4_kernel[grid_qkt](
        Q_fp4,
        Q_scale,
        K_fp4,
        K_scale,
        scores,
        qo_indptr,
        kv_indptr,
        stride_q_s=Q_fp4.stride(0),
        stride_q_h=Q_fp4.stride(1),
        stride_q_k=Q_fp4.stride(2),
        stride_q_scale_s=Q_scale.stride(0),
        stride_q_scale_k=Q_scale.stride(1),
        stride_k_kv=K_fp4.stride(0),
        stride_k_k=K_fp4.stride(1),
        stride_k_scale_kv=K_scale.stride(0),
        stride_k_scale_k=K_scale.stride(1),
        stride_scores_s=scores.stride(0),
        stride_scores_h=scores.stride(1),
        stride_scores_n=scores.stride(2),
        total_s=total_s,
        nhead=nhead,
        qk_head_dim=qk_head_dim,
        max_kv_len=max_kv_chunk,
        num_kv_splits=num_kv_splits,
        sm_scale=sm_scale,
        BLOCK_N=BLOCK_N,
        BLOCK_K=BLOCK_K,
        num_warps=4,
    )

    BLOCK_D = triton.next_power_of_2(v_head_dim)
    _softmax_lse_pv_f32_kernel[grid_qkt](
        scores,
        V_fp32,
        logits,
        attn_lse,
        o,
        qo_indptr,
        kv_indptr,
        stride_scores_s=scores.stride(0),
        stride_scores_h=scores.stride(1),
        stride_scores_n=scores.stride(2),
        stride_v_kv=V_fp32.stride(0),
        stride_v_d=V_fp32.stride(1),
        stride_logits_s=logits.stride(0),
        stride_logits_h=logits.stride(2),
        stride_logits_d=logits.stride(3),
        stride_lse_s=attn_lse.stride(0),
        stride_lse_h=attn_lse.stride(2),
        stride_lse_split=attn_lse.stride(1),
        stride_logits_split=logits.stride(1),
        stride_o_s=o.stride(0),
        stride_o_h=o.stride(1),
        stride_o_d=o.stride(2),
        total_s=total_s,
        nhead=nhead,
        v_head_dim=v_head_dim,
        max_kv_len=max_kv_chunk,
        num_kv_splits=num_kv_splits,
        BLOCK_D=BLOCK_D,
        num_warps=4,
    )


# ---------------------------------------------------------------------------
# Persistent mode metadata and wrapper (calls local mla_decode_fwd)
# ---------------------------------------------------------------------------
def _make_mla_decode_metadata(
    batch_size: int,
    max_q_len: int,
    nhead: int,
    nhead_kv: int,
    q_dtype: torch.dtype,
    kv_dtype: torch.dtype,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    kv_last_page_len: torch.Tensor,
    num_kv_splits: int = 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 _aiter_mla_decode(
    q: torch.Tensor,
    kv_buffer: torch.Tensor,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    config: dict,
    q_scale: torch.Tensor | None = None,
    kv_scale: torch.Tensor | None = None,
) -> torch.Tensor:
    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_indices = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")

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

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

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

    o = torch.empty((q.shape[0], nq, dv), dtype=torch.bfloat16, device="cuda")
    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


# ---------------------------------------------------------------------------
# Strategy A: mxfp4 KV (dequant then attend); fp8 / bf16 strategies
# ---------------------------------------------------------------------------
def _mla_decode_strategy_a(
    q: torch.Tensor,
    kv_buffer_mxfp4: torch.Tensor,
    kv_scale_mxfp4: torch.Tensor,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    config: dict,
) -> torch.Tensor:
    kv_orig_shape = (kv_buffer_mxfp4.shape[0], kv_buffer_mxfp4.shape[1], QK_HEAD_DIM)
    kv_bf16 = dequantize_mxfp4(kv_buffer_mxfp4, kv_scale_mxfp4, kv_orig_shape)
    return _aiter_mla_decode(
        q,
        kv_bf16,
        qo_indptr,
        kv_indptr,
        config,
        q_scale=None,
        kv_scale=None,
    )


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

    if QKV_DTYPE == "mxfp4":
        kv_buffer_mxfp4, kv_scale_mxfp4 = kv_data["mxfp4"]
        return _mla_decode_strategy_a(
            q,
            kv_buffer_mxfp4,
            kv_scale_mxfp4,
            qo_indptr,
            kv_indptr,
            config,
        )
    elif QKV_DTYPE == "fp8":
        q_input, q_scale = quantize_fp8(q)
        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,
        )
    elif QKV_DTYPE == "bf16":
        q_input, q_scale = q, None
        kv_input, kv_scale = kv_data["bf16"], None
        return _aiter_mla_decode(
            q_input,
            kv_input,
            qo_indptr,
            kv_indptr,
            config,
            q_scale=q_scale,
            kv_scale=kv_scale,
        )
    else:
        raise ValueError(f"Invalid QKV_DTYPE: {QKV_DTYPE}")
scrolls · 1119 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