Skip to content
KernelIndex
Search⌘K

submission 646958

Maxwell Cipher · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

c_v17.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-646958?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
78.9µs
#392 of 766
2026-03-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a642fbdf86dfed0ef093305a2727d9ec6c404660e12d15df40cd34a85d001fd6
license declaredunknown
license concludedunknown
authorsMaxwell Cipher
imported2026-08-15

Techniques

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

persistent-kernelPERSISTENT_TOKEN_LIMIT = 65536

Kernel source

c_v17.py239 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

import torch

from task import input_t, output_t

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

try:
    from aiter import per_tensor_quant_hip
    HAS_QUANT_HIP = True
except ImportError:
    HAS_QUANT_HIP = False


NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
PAGE_SIZE = 1
DEFAULT_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
FP8_DTYPE = aiter_dtypes.fp8
LOW_PAYLOAD_TOKEN_LIMIT = 65536
PERSISTENT_TOKEN_LIMIT = 65536
PERSISTENT_SPLITS = 32


def _quantize_query(q_view: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    if HAS_QUANT_HIP:
        q_fp8, q_scale = per_tensor_quant_hip(q_view, quant_dtype=FP8_DTYPE)
        return q_fp8, q_scale.reshape(1)

    finfo = torch.finfo(FP8_DTYPE)
    amax = q_view.abs().amax().clamp(min=1e-12)
    q_scale = (amax / finfo.max).to(torch.float32).reshape(1)
    q_fp8 = (q_view / q_scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
    return q_fp8, q_scale


def _prefer_low_overhead_path(total_kv: int) -> bool:
    return total_kv <= LOW_PAYLOAD_TOKEN_LIMIT


def _prefer_persistent_fp8_path(total_kv: int) -> bool:
    return total_kv > PERSISTENT_TOKEN_LIMIT


def _run_nonpersistent_path(
    q_view: torch.Tensor,
    kv_data: dict,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    batch_size: int,
    kv_seq_len: int,
    sm_scale: float,
) -> torch.Tensor:
    total_q = q_view.shape[0]
    total_kv = batch_size * kv_seq_len
    device = q_view.device

    page_ids = torch.arange(total_kv, dtype=torch.int32, device=device)
    last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)
    output = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device)

    if _prefer_low_overhead_path(total_kv):
        kv_tensor = kv_data["bf16"]
        q_input = q_view
        q_scale = None
        kv_scale = None
    else:
        kv_tensor, kv_scale = kv_data["fp8"]
        q_input, q_scale = _quantize_query(q_view)

    kv_pages = kv_tensor.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_tensor.shape[-1])

    mla_decode_fwd(
        q_input,
        kv_pages,
        output,
        qo_indptr,
        kv_indptr,
        page_ids,
        last_page_len,
        1,
        page_size=PAGE_SIZE,
        nhead_kv=NUM_KV_HEADS,
        sm_scale=sm_scale,
        q_scale=q_scale,
        kv_scale=kv_scale,
        intra_batch_mode=False,
    )

    return output


def _build_persistent_metadata(
    batch_size: int,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    q_dtype: torch.dtype,
    kv_dtype: torch.dtype,
    kv_last_page_len: torch.Tensor,
):
    info = get_mla_metadata_info_v1(
        batch_size,
        1,
        NUM_HEADS,
        q_dtype,
        kv_dtype,
        is_sparse=False,
        fast_mode=False,
        num_kv_splits=PERSISTENT_SPLITS,
        intra_batch_mode=True,
    )

    work_meta_data, work_indptr, work_info_set, reduce_indptr, reduce_final_map, reduce_partial_map = [
        torch.empty(shape, dtype=dtype, device=qo_indptr.device) for shape, dtype in info
    ]

    get_mla_metadata_v1(
        qo_indptr,
        kv_indptr,
        kv_last_page_len,
        NUM_HEADS // NUM_KV_HEADS,
        NUM_KV_HEADS,
        True,
        work_meta_data,
        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=1,
        uni_seqlen_qo=1,
        fast_mode=False,
        max_split_per_batch=PERSISTENT_SPLITS,
        intra_batch_mode=True,
        dtype_q=q_dtype,
        dtype_kv=kv_dtype,
    )

    return {
        "work_meta_data": work_meta_data,
        "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 _run_persistent_fp8_path(
    q_view: torch.Tensor,
    kv_data: dict,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    batch_size: int,
    kv_seq_len: int,
    sm_scale: float,
) -> torch.Tensor:
    total_q = q_view.shape[0]
    total_kv = batch_size * kv_seq_len
    device = q_view.device

    q_fp8, q_scale = _quantize_query(q_view)
    kv_fp8, kv_scale = kv_data["fp8"]
    kv_pages = kv_fp8.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_fp8.shape[-1])

    page_ids = torch.arange(total_kv, dtype=torch.int32, device=device)
    last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)
    metadata = _build_persistent_metadata(
        batch_size,
        qo_indptr,
        kv_indptr,
        q_fp8.dtype,
        kv_fp8.dtype,
        last_page_len,
    )

    output = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device)

    mla_decode_fwd(
        q_fp8,
        kv_pages,
        output,
        qo_indptr,
        kv_indptr,
        page_ids,
        last_page_len,
        1,
        page_size=PAGE_SIZE,
        nhead_kv=NUM_KV_HEADS,
        sm_scale=sm_scale,
        logit_cap=0.0,
        num_kv_splits=PERSISTENT_SPLITS,
        q_scale=q_scale,
        kv_scale=kv_scale,
        intra_batch_mode=True,
        **metadata,
    )

    return output


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

    batch_size = int(config["batch_size"])
    kv_seq_len = int(config["kv_seq_len"])
    total_kv = batch_size * kv_seq_len
    sm_scale = float(config.get("sm_scale", DEFAULT_SCALE))
    q_view = q.view(-1, NUM_HEADS, QK_HEAD_DIM)

    if _prefer_persistent_fp8_path(total_kv):
        return _run_persistent_fp8_path(
            q_view,
            kv_data,
            qo_indptr,
            kv_indptr,
            batch_size,
            kv_seq_len,
            sm_scale,
        )

    return _run_nonpersistent_path(
        q_view,
        kv_data,
        qo_indptr,
        kv_indptr,
        batch_size,
        kv_seq_len,
        sm_scale,
    )
scrolls · 239 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 643908.

#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
- """v13: Multi-head fused HIP MLA decode with MXFP4 KV dequant on MI355X.
+ import torch
- One block processes ALL 16 query heads sharing the same KV data.
- KV is loaded from HBM once per block and reused for all 16 heads,
- yielding ~16x bandwidth savings over the per-head v12 approach.
+ from task import input_t, output_t
- Two-stage flash decoding:
- Stage 1: per-(batch, split) kernel. 256 threads (4 wavefronts).
- All 16 Q vectors loaded to LDS (36 KB). For each KV token, dequant
- MXFP4 in registers, compute dot products for all 16 heads, online
- softmax update, and V accumulation -- all with KV loaded once.
- Stage 2: lightweight reduce kernel merges partials from all splits
- using exp-weighted combination (mathematically exact).
+ 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
- Bandwidth per KV token per batch element:
- v12: 306 bytes x 16 heads = 4896 bytes (KV loaded 16 times)
- v13: 306 bytes x 1 = 306 bytes (KV loaded once, reused)
+ try:
+ from aiter import per_tensor_quant_hip
+ HAS_QUANT_HIP = True
+ except ImportError:
+ HAS_QUANT_HIP = False
- No cross-call caching. Fresh output buffer each call.
- """
- import os
-
- os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
-
- import torch
- from torch.utils.cpp_extension import load_inline
- from task import input_t, output_t
-
- # MLA constants
NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
- PACKED_DIM = 288 # QK_HEAD_DIM // 2
- SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
+ PAGE_SIZE = 1
+ DEFAULT_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
+ FP8_DTYPE = aiter_dtypes.fp8
+ LOW_PAYLOAD_TOKEN_LIMIT = 65536
+ PERSISTENT_TOKEN_LIMIT = 65536
+ PERSISTENT_SPLITS = 32
- # ---------------------------------------------------------------------------
- # HIP C++ kernel source -- multi-head fused MLA decode
- # ---------------------------------------------------------------------------
- HIP_SOURCE = r'''
- #include <hip/hip_runtime.h>
- #include <torch/extension.h>
- // =========================================================================
- // bf16 <-> float via bit manipulation (works on all HIP targets)
- // =========================================================================
- __device__ __forceinline__ float bf16_to_float(unsigned short v) {
- unsigned int bits = ((unsigned int)v) << 16;
- return __uint_as_float(bits);
- }
+ def _quantize_query(q_view: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
+ if HAS_QUANT_HIP:
+ q_fp8, q_scale = per_tensor_quant_hip(q_view, quant_dtype=FP8_DTYPE)
+ return q_fp8, q_scale.reshape(1)
- __device__ __forceinline__ unsigned short float_to_bf16(float v) {
- unsigned int bits = __float_as_uint(v);
- // Round to nearest even
- bits += 0x7FFF + ((bits >> 16) & 1);
- return (unsigned short)(bits >> 16);
- }
+ finfo = torch.finfo(FP8_DTYPE)
+ amax = q_view.abs().amax().clamp(min=1e-12)
+ q_scale = (amax / finfo.max).to(torch.float32).reshape(1)
+ q_fp8 = (q_view / q_scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
+ return q_fp8, q_scale
- // =========================================================================
- // FP4 E2M1 dequantization lookup table (16 entries)
- // Bit layout: [sign(1) exp(2) mant(1)]
- // =========================================================================
- __device__ __constant__ float FP4_LUT[16] = {
- 0.0f, 0.5f, 1.0f, 1.5f, 2.0f, 3.0f, 4.0f, 6.0f,
- -0.0f, -0.5f, -1.0f, -1.5f, -2.0f, -3.0f, -4.0f, -6.0f
- };
- // =========================================================================
- // E8M0 -> float: 2^(e - 127) via IEEE-754 exponent placement
- // =========================================================================
- __device__ __forceinline__ float e8m0_to_float(unsigned char e) {
- unsigned int bits = ((unsigned int)e) << 23;
- return __uint_as_float(bits);
- }
+ def _prefer_low_overhead_path(total_kv: int) -> bool:
+ return total_kv <= LOW_PAYLOAD_TOKEN_LIMIT
- // =========================================================================
- // Configuration
- // =========================================================================
- #define QK_DIM 576
- #define V_DIM 512
- #define BLK_SZ 32 // MXFP4 block size (elements per E8M0 scale)
- #define N_BLOCKS 18 // QK_DIM / BLK_SZ = 576/32
- #define V_BLOCKS 16 // V_DIM / BLK_SZ = 512/32
- #define PACK_DIM 288 // QK_DIM / 2 (bytes of packed FP4)
- #define THREADS 256 // 4 wavefronts of 64
- #define NHEADS 16
- // =========================================================================
- // Stage 1: Multi-head fused MXFP4 flash-decoding
- //
- // Grid: (num_splits, batch_size) <-- NOT batch_size*16!
- // Block: (THREADS)
- //
- // Each block processes ALL 16 query heads for one batch element over one
- // KV-length split. Q for all 16 heads loaded to LDS (16*576 floats = 36 KB).
- // KV tokens dequantized ONCE and reused for all 16 dot products.
- //
- // Outputs:
- // partial_out [num_splits, batch_size*16, V_DIM] float32
- // partial_lse [num_splits, batch_size*16] float32
- // =========================================================================
+ def _prefer_persistent_fp8_path(total_kv: int) -> bool:
+ return total_kv > PERSISTENT_TOKEN_LIMIT
- __global__ __attribute__((amdgpu_flat_work_group_size(THREADS, THREADS)))
- void mla_stage1_multihead(
- const unsigned short* __restrict__ Q, // [total_q, 16, 576] bf16
- const unsigned char* __restrict__ KV_packed, // [total_kv, PACK_DIM] uint8
- const unsigned char* __restrict__ KV_scales, // [total_kv, N_scale_cols] uint8
- float* __restrict__ partial_out, // [num_splits, total_bh, V_DIM]
- float* __restrict__ partial_lse, // [num_splits, total_bh]
- const int* __restrict__ kv_indptr, // [batch+1]
- float sm_scale,
- int total_bh,
- int num_splits,
- int kv_stride_scales // stride of KV_scales in dim0
- ) {
- int split_id = blockIdx.x;
- int batch_id = blockIdx.y;
- int tid = threadIdx.x;
- // KV range for this batch element
- int kv_start = kv_indptr[batch_id];
- int kv_end = kv_indptr[batch_id + 1];
- int kv_len = kv_end - kv_start;
+ def _run_nonpersistent_path(
+ q_view: torch.Tensor,
+ kv_data: dict,
+ qo_indptr: torch.Tensor,
+ kv_indptr: torch.Tensor,
+ batch_size: int,
+ kv_seq_len: int,
+ sm_scale: float,
+ ) -> torch.Tensor:
+ total_q = q_view.shape[0]
+ total_kv = batch_size * kv_seq_len
+ device = q_view.device
- // Base index for this batch element's heads in the flattened bh dimension
- int bh_base = batch_id * NHEADS;
+ page_ids = torch.arange(total_kv, dtype=torch.int32, device=device)
+ last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)
+ output = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device)
- if (kv_len <= 0) {
- // Empty sequence: write -inf LSE and zero output for all heads
- for (int h = 0; h < NHEADS; h++) {
- int bh_id = bh_base + h;
- if (tid == 0) {
- partial_lse[split_id * total_bh + bh_id] = -1e30f;
- }
- int out_base = (split_id * total_bh + bh_id) * V_DIM;
- for (int d = tid; d < V_DIM; d += THREADS) {
- partial_out[out_base + d] = 0.0f;
- }
- }
- return;
- }
+ if _prefer_low_overhead_path(total_kv):
+ kv_tensor = kv_data["bf16"]
+ q_input = q_view
+ q_scale = None
+ kv_scale = None
+ else:
+ kv_tensor, kv_scale = kv_data["fp8"]
+ q_input, q_scale = _quantize_query(q_view)
- // Split range
- int split_size = (kv_len + num_splits - 1) / num_splits;
- int my_start = kv_start + split_id * split_size;
- int my_end = my_start + split_size;
- if (my_end > kv_end) my_end = kv_end;
- if (my_start >= kv_end) {
- // This split has no tokens
- for (int h = 0; h < NHEADS; h++) {
- int bh_id = bh_base + h;
- if (tid == 0) {
- partial_lse[split_id * total_bh + bh_id] = -1e30f;
- }
- int out_base = (split_id * total_bh + bh_id) * V_DIM;
- for (int d = tid; d < V_DIM; d += THREADS) {
- partial_out[out_base + d] = 0.0f;
- }
- }
- return;
- }
+ kv_pages = kv_tensor.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_tensor.shape[-1])
- // ---- Load Q for ALL 16 heads into LDS ----
- // q_lds[h][d] = Q[batch_id, h, d] converted to float
- // 16 * 576 = 9216 floats = 36,864 bytes (fits in LDS)
- __shared__ float q_lds[NHEADS][QK_DIM];
-
- int q_token = batch_id; // In decode, total_q == batch_size
- for (int h = 0; h < NHEADS; h++) {
- long long q_base = (long long)q_token * NHEADS * QK_DIM + (long long)h * QK_DIM;
- for (int d = tid; d < QK_DIM; d += THREADS) {
- q_lds[h][d] = bf16_to_float(Q[q_base + d]);
- }
- }
- __syncthreads();
-
- // ---- Per-head online softmax state (in registers) ----
- float my_max[NHEADS];
- float my_sum_exp[NHEADS];
- float my_v0[NHEADS]; // V accumulator for dim = tid
- float my_v1[NHEADS]; // V accumulator for dim = tid + 256
-
- for (int h = 0; h < NHEADS; h++) {
- my_max[h] = -1e30f;
- my_sum_exp[h] = 0.0f;
- my_v0[h] = 0.0f;
- my_v1[h] = 0.0f;
- }
-
- // ---- Shared memory for cross-warp reduction ----
- // warp_dots[warp_id][head] for dot product reduction
- // We process heads in groups to limit shared memory
- __shared__ float warp_dots[4][NHEADS];
- // For broadcasting final scores to all threads
- __shared__ float score_broadcast[NHEADS];
-
- int warp_id = tid / 64;
- int lane_id = tid % 64;
-
- // ---- Iterate over KV tokens in this split ----
- for (int kv_pos = my_start; kv_pos < my_end; kv_pos++) {
-
- long long packed_base = (long long)kv_pos * PACK_DIM;
- long long scale_base = (long long)kv_pos * kv_stride_scales;
-
- // -- Step 1: Dequantize KV values for this thread's dimensions --
- // Each thread handles dims: tid, tid+256, and tid+512 (if tid < 64)
- // We dequant once and reuse for all 16 heads
-
- // Dequant dim = tid (always valid, tid < 256 < 576)
- int d0 = tid;
- int byte_idx0 = d0 / 2;
- unsigned char pk0 = KV_packed[packed_base + byte_idx0];
- float rv0;
- if (d0 & 1) rv0 = FP4_LUT[(pk0 >> 4) & 0x0F];
- else rv0 = FP4_LUT[pk0 & 0x0F];
- float sc0 = e8m0_to_float(KV_scales[scale_base + d0 / BLK_SZ]);
- float kv_d0 = rv0 * sc0;
-
- // Dequant dim = tid + 256 (always valid, tid+256 < 512 < 576)
- int d1 = tid + THREADS;
- int byte_idx1 = d1 / 2;
- unsigned char pk1 = KV_packed[packed_base + byte_idx1];
- float rv1;
- if (d1 & 1) rv1 = FP4_LUT[(pk1 >> 4) & 0x0F];
- else rv1 = FP4_LUT[pk1 & 0x0F];
- float sc1 = e8m0_to_float(KV_scales[scale_base + d1 / BLK_SZ]);
- float kv_d1 = rv1 * sc1;
-
- // Dequant dim = tid + 512 (only valid for tid < 64, since 576-512=64)
- float kv_d2 = 0.0f;
- if (tid < 64) {
- int d2 = tid + 2 * THREADS;
- int byte_idx2 = d2 / 2;
- unsigned char pk2 = KV_packed[packed_base + byte_idx2];
- float rv2;
- if (d2 & 1) rv2 = FP4_LUT[(pk2 >> 4) & 0x0F];
- else rv2 = FP4_LUT[pk2 & 0x0F];
- float sc2 = e8m0_to_float(KV_scales[scale_base + d2 / BLK_SZ]);
- kv_d2 = rv2 * sc2;
- }
-
- // -- Step 2: Compute dot products for ALL 16 heads --
- // Each thread computes partial dot for each head using its KV dims
- float pdot[NHEADS];
- for (int h = 0; h < NHEADS; h++) {
- pdot[h] = q_lds[h][d0] * kv_d0
- + q_lds[h][d1] * kv_d1;
- if (tid < 64) {
- pdot[h] += q_lds[h][tid + 2 * THREADS] * kv_d2;
- }
- }
-
- // -- Step 3: Warp-level reduction for each head --
- for (int h = 0; h < NHEADS; h++) {
- float val = pdot[h];
- for (int offset = 32; offset >= 1; offset >>= 1) {
- val += __shfl_xor(val, offset);
- }
- // Lane 0 of each warp has the warp sum
- if (lane_id == 0) {
- warp_dots[warp_id][h] = val;
- }
- }
- __syncthreads();
-
- // -- Step 4: Thread 0 does final cross-warp reduction for all heads --
- if (tid == 0) {
- for (int h = 0; h < NHEADS; h++) {
- float s = warp_dots[0][h] + warp_dots[1][h]
- + warp_dots[2][h] + warp_dots[3][h];
- score_broadcast[h] = s * sm_scale;
- }
- }
- __syncthreads();
-
- // -- Step 5: Online softmax update and V accumulation for all heads --
- // V = first 512 dims of KV. Each thread owns V[tid] and V[tid+256].
- // kv_d0 is the dequanted value at dim=tid (always < 512, valid V dim)
- // kv_d1 is the dequanted value at dim=tid+256 (always < 512, valid V dim)
- float v_val0 = kv_d0; // V dim = tid
- float v_val1 = kv_d1; // V dim = tid + 256
-
- for (int h = 0; h < NHEADS; h++) {
- float score = score_broadcast[h];
-
- float old_max = my_max[h];
- my_max[h] = fmaxf(my_max[h], score);
- float corr = expf(old_max - my_max[h]);
- float w = expf(score - my_max[h]);
- my_sum_exp[h] = my_sum_exp[h] * corr + w;
-
- my_v0[h] = my_v0[h] * corr + w * v_val0;
- my_v1[h] = my_v1[h] * corr + w * v_val1;
- }
-
- __syncthreads(); // needed before next iteration reuses shared mem
- }
-
- // ---- Write partial results for all 16 heads ----
- for (int h = 0; h < NHEADS; h++) {
- int bh_id = bh_base + h;
- int out_base = (split_id * total_bh + bh_id) * V_DIM;
-
- float inv_se = (my_sum_exp[h] > 0.0f) ? (1.0f / my_sum_exp[h]) : 0.0f;
- partial_out[out_base + tid] = my_v0[h] * inv_se;
- partial_out[out_base + tid + THREADS] = my_v1[h] * inv_se;
-
- if (tid == 0) {
- float lse_val = (my_sum_exp[h] > 0.0f)
- ? (my_max[h] + logf(my_sum_exp[h]))
- : -1e30f;
- partial_lse[split_id * total_bh + bh_id] = lse_val;
- }
- }
- }
-
-
- // =========================================================================
- // Stage 2: Reduce partials across splits
- //
- // Grid: (total_bh)
- // Block: (THREADS)
- //
- // For each (batch, head), merges all split partials into final bf16 output.
- // (Same as v12 -- this is already per-head and efficient.)
- // =========================================================================
-
- __global__ __attribute__((amdgpu_flat_work_group_size(THREADS, THREADS)))
- void mla_reduce(
- const float* __restrict__ partial_out, // [num_splits, total_bh, V_DIM]
- const float* __restrict__ partial_lse, // [num_splits, total_bh]
- unsigned short* __restrict__ final_out, // [total_bh, V_DIM] bf16
- int total_bh,
- int num_splits
- ) {
- int bh_id = blockIdx.x;
- int tid = threadIdx.x;
-
- // Pass 1: find global max LSE
- __shared__ float gmax_shared;
- if (tid == 0) {
- float gmax = -1e30f;
- for (int s = 0; s < num_splits; s++) {
- float lse_s = partial_lse[s * total_bh + bh_id];
- if (lse_s > gmax) gmax = lse_s;
- }
- gmax_shared = gmax;
- }
- __syncthreads();
- float gmax = gmax_shared;
-
- // Pass 2: compute total weight
- __shared__ float total_w_shared;
- if (tid == 0) {
- float total_w = 0.0f;
- for (int s = 0; s < num_splits; s++) {
- float lse_s = partial_lse[s * total_bh + bh_id];
- total_w += expf(lse_s - gmax);
- }
- total_w_shared = total_w;
- }
- __syncthreads();
- float total_w = total_w_shared;
- float inv_tw = (total_w > 0.0f) ? (1.0f / total_w) : 0.0f;
-
- // Pass 3: weighted merge of V values
- float acc0 = 0.0f;
- float acc1 = 0.0f;
-
- for (int s = 0; s < num_splits; s++) {
- float lse_s = partial_lse[s * total_bh + bh_id];
- float w_s = expf(lse_s - gmax);
- int base = (s * total_bh + bh_id) * V_DIM;
- acc0 += partial_out[base + tid] * w_s;
- acc1 += partial_out[base + tid + THREADS] * w_s;
- }
-
- // Write bf16 output
- long long out_base = (long long)bh_id * V_DIM;
- final_out[out_base + tid] = float_to_bf16(acc0 * inv_tw);
- final_out[out_base + tid + THREADS] = float_to_bf16(acc1 * inv_tw);
- }
-
-
- // =========================================================================
- // C++ / pybind11 wrapper
- // =========================================================================
-
- torch::Tensor mla_decode_hip(
- torch::Tensor Q, // [total_q, 16, 576] bf16
- torch::Tensor KV_packed, // [total_kv, PACK_DIM] uint8
- torch::Tensor KV_scales, // [total_kv, N_scale_cols] uint8
- torch::Tensor kv_indptr, // [batch+1] int32
- int64_t batch_size,
- int64_t num_splits,
- double sm_scale_d
- ) {
- float sm_scale = (float)sm_scale_d;
- int total_q = Q.size(0);
- int total_bh = total_q * NHEADS;
-
- // KV scale stride (column count, may be padded)
- int kv_stride_scales = KV_scales.size(1);
-
- // Allocate partial buffers
- auto opts_f32 = torch::TensorOptions().dtype(torch::kFloat32).device(Q.device());
- auto partial_out = torch::zeros({(int64_t)num_splits, (int64_t)total_bh, (int64_t)V_DIM}, opts_f32);
- auto partial_lse = torch::full({(int64_t)num_splits, (int64_t)total_bh}, -1e30f, opts_f32);
-
- // Stage 1: multi-head fused kernel -- grid is (splits, batch) not (splits, batch*heads)
- dim3 grid_s1((int)num_splits, (int)batch_size);
- dim3 block_s1(THREADS);
- mla_stage1_multihead<<<grid_s1, block_s1, 0, 0>>>(
- (const unsigned short*)Q.data_ptr(),
- (const unsigned char*)KV_packed.data_ptr(),
- (const unsigned char*)KV_scales.data_ptr(),
- partial_out.data_ptr<float>(),
- partial_lse.data_ptr<float>(),
- kv_indptr.data_ptr<int>(),
- sm_scale,
- total_bh,
- (int)num_splits,
- kv_stride_scales
- );
-
- // Stage 2: reduce across splits -> bf16 output
- auto opts_bf16 = torch::TensorOptions().dtype(torch::kBFloat16).device(Q.device());
- auto output = torch::empty({(int64_t)total_q, NHEADS, (int64_t)V_DIM}, opts_bf16);
-
- dim3 grid_s2(total_bh);
- dim3 block_s2(THREADS);
- mla_reduce<<<grid_s2, block_s2, 0, 0>>>(
- partial_out.data_ptr<float>(),
- partial_lse.data_ptr<float>(),
- (unsigned short*)output.data_ptr(),
- total_bh,
- (int)num_splits
- );
-
- return output;
- }
-
- PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
- m.def("mla_decode_hip", &mla_decode_hip,
- "MLA decode with multi-head fused MXFP4 dequant (HIP kernel)");
- }
- '''
-
- # ---------------------------------------------------------------------------
- # Compile HIP kernel
- # ---------------------------------------------------------------------------
-
- import sys
-
- print("[mla_v13] Compiling multi-head fused HIP MXFP4 MLA decode kernel...", file=sys.stderr, flush=True)
- try:
- _mod = load_inline(
- name="mla_v13_hip",
- cpp_sources=[],
- cuda_sources=[HIP_SOURCE],
- extra_cuda_cflags=["-O3", "--offload-arch=gfx950"],
- verbose=False,
+ mla_decode_fwd(
+ q_input,
+ kv_pages,
+ output,
+ qo_indptr,
+ kv_indptr,
+ page_ids,
+ last_page_len,
+ 1,
+ page_size=PAGE_SIZE,
+ nhead_kv=NUM_KV_HEADS,
+ sm_scale=sm_scale,
+ q_scale=q_scale,
+ kv_scale=kv_scale,
+ intra_batch_mode=False,
)
- _HIP_AVAILABLE = True
- print("[mla_v13] HIP compilation successful!", file=sys.stderr, flush=True)
- except Exception as e:
- _HIP_AVAILABLE = False
- print(f"[mla_v13] HIP compilation failed: {e}", file=sys.stderr, flush=True)
- # ---------------------------------------------------------------------------
- # AITER fallback imports (fp8 path)
- # ---------------------------------------------------------------------------
+ return output
- from aiter.mla import mla_decode_fwd
- from aiter import dtypes as aiter_dtypes, get_mla_metadata_info_v1, get_mla_metadata_v1
- FP8 = aiter_dtypes.fp8
- PAGE_SIZE = 1
-
-
- def _quantize_fp8(t):
- finfo = torch.finfo(FP8)
- amax = t.abs().amax().clamp(min=1e-12)
- sc = amax / finfo.max
- return (t / sc).clamp(finfo.min, finfo.max).to(FP8), sc.float().reshape(1)
-
-
- def _choose_splits(total_kv, kv_seq_len):
- """Select number of KV splits for the HIP kernel."""
- if kv_seq_len <= 512:
- return 2
- elif kv_seq_len <= 2048:
- return 4
- elif kv_seq_len <= 8192:
- return 8
- elif kv_seq_len <= 32768:
- return 16
- else:
- return 32
-
-
- def _choose_aiter_params(qsl, total_kv, bs, nq):
- kv_per_req = total_kv // max(bs, 1)
- if qsl == 1:
- if kv_per_req >= 4096:
- return 32, False
- elif kv_per_req >= 1024:
- return 16, False
- else:
- return 8, False
- else:
- if kv_per_req >= 4096:
- return 16, False
- elif kv_per_req >= 1024:
- return 8, False
- else:
- return 4, False
-
-
- def _build_meta(bs, qsl, nq, nkv, total_kv, q_dtype, kv_dtype,
- qo_indptr, kv_indptr, num_kv_splits, fast_mode):
- kv_lpl = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
+ def _build_persistent_metadata(
+ batch_size: int,
+ qo_indptr: torch.Tensor,
+ kv_indptr: torch.Tensor,
+ q_dtype: torch.dtype,
+ kv_dtype: torch.dtype,
+ kv_last_page_len: torch.Tensor,
+ ):
info = get_mla_metadata_info_v1(
- bs, qsl, nq, q_dtype, kv_dtype,
- is_sparse=False, fast_mode=fast_mode,
- num_kv_splits=num_kv_splits, intra_batch_mode=True,
+ batch_size,
+ 1,
+ NUM_HEADS,
+ q_dtype,
+ kv_dtype,
+ is_sparse=False,
+ fast_mode=False,
+ num_kv_splits=PERSISTENT_SPLITS,
+ intra_batch_mode=True,
)
- work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
- wm, wi, wis, ri, rfm, rpm = work
+
+ work_meta_data, work_indptr, work_info_set, reduce_indptr, reduce_final_map, reduce_partial_map = [
+ torch.empty(shape, dtype=dtype, device=qo_indptr.device) for shape, dtype in info
+ ]
+
get_mla_metadata_v1(
- qo_indptr, kv_indptr, kv_lpl,
- nq // nkv, nkv, True,
- wm, wis, wi, ri, rfm, rpm,
- page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE, 16),
- max_seqlen_qo=qsl, uni_seqlen_qo=qsl,
- fast_mode=fast_mode, max_split_per_batch=num_kv_splits,
- intra_batch_mode=True, dtype_q=q_dtype, dtype_kv=kv_dtype,
+ qo_indptr,
+ kv_indptr,
+ kv_last_page_len,
+ NUM_HEADS // NUM_KV_HEADS,
+ NUM_KV_HEADS,
+ True,
+ work_meta_data,
+ 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=1,
+ uni_seqlen_qo=1,
+ fast_mode=False,
+ max_split_per_batch=PERSISTENT_SPLITS,
+ intra_batch_mode=True,
+ dtype_q=q_dtype,
+ dtype_kv=kv_dtype,
)
+
return {
- "meta": dict(work_meta_data=wm, work_indptr=wi, work_info_set=wis,
- reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm),
- "kv_lpl": kv_lpl,
- "kv_indices": torch.arange(total_kv, dtype=torch.int32, device="cuda"),
- "num_kv_splits": num_kv_splits,
+ "work_meta_data": work_meta_data,
+ "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 _run_mla_aiter(q, kv_data, qo_indptr, kv_indptr, config):
- """AITER fp8 fallback path."""
- nq = config["num_heads"]
- nkv = config["num_kv_heads"]
- dq = config["qk_head_dim"]
- dv = config["v_head_dim"]
- qsl = config["q_seq_len"]
- bs = config["batch_size"]
+ def _run_persistent_fp8_path(
+ q_view: torch.Tensor,
+ kv_data: dict,
+ qo_indptr: torch.Tensor,
+ kv_indptr: torch.Tensor,
+ batch_size: int,
+ kv_seq_len: int,
+ sm_scale: float,
+ ) -> torch.Tensor:
+ total_q = q_view.shape[0]
+ total_kv = batch_size * kv_seq_len
+ device = q_view.device
+ q_fp8, q_scale = _quantize_query(q_view)
kv_fp8, kv_scale = kv_data["fp8"]
- total_kv = int(kv_indptr[-1].item())
+ kv_pages = kv_fp8.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_fp8.shape[-1])
- q_fp8, q_scale = _quantize_fp8(q)
+ page_ids = torch.arange(total_kv, dtype=torch.int32, device=device)
+ last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)
+ metadata = _build_persistent_metadata(
+ batch_size,
+ qo_indptr,
+ kv_indptr,
+ q_fp8.dtype,
+ kv_fp8.dtype,
+ last_page_len,
+ )
- num_kv_splits, fast_mode = _choose_aiter_params(qsl, total_kv, bs, nq)
- cached = _build_meta(bs, qsl, nq, nkv, total_kv,
- q_fp8.dtype, kv_fp8.dtype,
- qo_indptr, kv_indptr, num_kv_splits, fast_mode)
+ output = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device)
- kv_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, nkv, kv_fp8.shape[-1])
- o = torch.empty((q_fp8.shape[0], nq, dv), dtype=torch.bfloat16, device="cuda")
-
mla_decode_fwd(
- q_fp8.view(-1, nq, dq), kv_4d, o,
- qo_indptr, kv_indptr,
- cached["kv_indices"], cached["kv_lpl"], qsl,
- page_size=PAGE_SIZE, nhead_kv=nkv,
- sm_scale=SM_SCALE, logit_cap=0.0,
- num_kv_splits=cached["num_kv_splits"],
- q_scale=q_scale, kv_scale=kv_scale,
- intra_batch_mode=True, **cached["meta"],
+ q_fp8,
+ kv_pages,
+ output,
+ qo_indptr,
+ kv_indptr,
+ page_ids,
+ last_page_len,
+ 1,
+ page_size=PAGE_SIZE,
+ nhead_kv=NUM_KV_HEADS,
+ sm_scale=sm_scale,
+ logit_cap=0.0,
+ num_kv_splits=PERSISTENT_SPLITS,
+ q_scale=q_scale,
+ kv_scale=kv_scale,
+ intra_batch_mode=True,
+ **metadata,
)
- return o
+ return output
- # ---------------------------------------------------------------------------
- # Entry point
- # ---------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
- """MLA decode with multi-head fused MXFP4 HIP kernel (fallback: AITER fp8)."""
q, kv_data, qo_indptr, kv_indptr, config = data
- batch_size = config["batch_size"]
- kv_seq_len = config["kv_seq_len"]
- qsl = config["q_seq_len"]
+ batch_size = int(config["batch_size"])
+ kv_seq_len = int(config["kv_seq_len"])
+ total_kv = batch_size * kv_seq_len
+ sm_scale = float(config.get("sm_scale", DEFAULT_SCALE))
+ q_view = q.view(-1, NUM_HEADS, QK_HEAD_DIM)
- # ---- Try HIP kernel path (MXFP4 dequant, multi-head fusion) ----
- if _HIP_AVAILABLE and qsl == 1:
- kv_fp4x2, kv_scales_e8m0 = kv_data["mxfp4"]
+ if _prefer_persistent_fp8_path(total_kv):
+ return _run_persistent_fp8_path(
+ q_view,
+ kv_data,
+ qo_indptr,
+ kv_indptr,
+ batch_size,
+ kv_seq_len,
+ sm_scale,
+ )
- # Reshape to 2D: [total_kv, PACK_DIM] uint8
- kv_packed_u8 = kv_fp4x2.reshape(-1, PACKED_DIM).contiguous().view(torch.uint8)
-
- # Scales: [total_kv_padded, N_cols] uint8 (N_cols >= 18, may be padded)
- kv_scales_u8 = kv_scales_e8m0.contiguous().view(torch.uint8)
-
- # Ensure scales are 2D [total_kv, N_cols]
- total_kv = int(kv_indptr[-1].item())
- if kv_scales_u8.dim() == 1:
- n_scale_cols = kv_scales_u8.numel() // total_kv
- kv_scales_u8 = kv_scales_u8.view(total_kv, n_scale_cols)
- elif kv_scales_u8.dim() != 2:
- kv_scales_u8 = kv_scales_u8.view(total_kv, -1)
-
- # If padded (more rows than total_kv), just use total_kv rows
- if kv_scales_u8.size(0) > total_kv:
- kv_scales_u8 = kv_scales_u8[:total_kv]
-
- num_splits = _choose_splits(total_kv, kv_seq_len)
-
- try:
- output = _mod.mla_decode_hip(
- q, kv_packed_u8, kv_scales_u8,
- kv_indptr, batch_size, num_splits,
- float(SM_SCALE),
- )
- return output
- except Exception:
- # Fall through to AITER
- pass
-
- # ---- AITER fp8 fallback ----
- return _run_mla_aiter(q, kv_data, qo_indptr, kv_indptr, config)
+ return _run_nonpersistent_path(
+ q_view,
+ kv_data,
+ qo_indptr,
+ kv_indptr,
+ batch_size,
+ kv_seq_len,
+ sm_scale,
+ )
scrolls · 822 diff lines total

Best evidence level for this revision: reported

JSON