Skip to content
KernelIndex
Search⌘K

submission 643908

Maxwell Cipher · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

mla_v13.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-643908?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
4.15ms
#742 of 766
2026-03-27

Reported · How evidence levels are derived →

Source and license

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

Techniques

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

fp4"""v13: Multi-head fused HIP MLA decode with MXFP4 KV dequant on MI355X.
shared-memory__shared__ float q_lds[NHEADS][QK_DIM];

Kernel source

mla_v13.py637 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""v13: Multi-head fused HIP MLA decode with MXFP4 KV dequant on MI355X.

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.

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).

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)

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)

# ---------------------------------------------------------------------------
# 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);
}

__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);
}

// =========================================================================
// 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);
}

// =========================================================================
// 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
// =========================================================================

__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;

    // Base index for this batch element's heads in the flattened bh dimension
    int bh_base = batch_id * NHEADS;

    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;
    }

    // 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;
    }

    // ---- 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,
    )
    _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)
# ---------------------------------------------------------------------------

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)
    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,
    )
    work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
    wm, wi, wis, ri, rfm, rpm = work
    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,
    )
    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,
    }


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"]

    kv_fp8, kv_scale = kv_data["fp8"]
    total_kv = int(kv_indptr[-1].item())

    q_fp8, q_scale = _quantize_fp8(q)

    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)

    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"],
    )
    return o


# ---------------------------------------------------------------------------
# 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"]

    # ---- Try HIP kernel path (MXFP4 dequant, multi-head fusion) ----
    if _HIP_AVAILABLE and qsl == 1:
        kv_fp4x2, kv_scales_e8m0 = kv_data["mxfp4"]

        # 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)
scrolls · 637 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