Skip to content
KernelIndex
Search⌘K

submission 665420

Barry_zhang · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v0011b.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-665420?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
166.6µs
#541 of 766
2026-03-29

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:9bd8573a3c57a16d0a5404df5bd5e8459533c5323c62fd0090de2627e9d844cc
license declaredunknown
license concludedunknown
authorsBarry_zhang
imported2026-08-15

Techniques

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

fp4Custom HIP kernel for MLA decode attention with MXFP4 KV cache on MI355X (gfx950).
online-softmaxfloat m_new = fmaxf(m_old, tile_max);
shared-memory__shared__ uint16_t q_lds[NUM_HEADS * QK_DIM]; // 16 * 576
split-k__global__ void mla_splitk_reduce_kernel(
tile-n = 16constexpr int BLOCK_N = 16; // KV tile size

Kernel source

submission_v0011b.py735 lines
"""
Custom HIP kernel for MLA decode attention with MXFP4 KV cache on MI355X (gfx950).

v0011: MFMA V accumulation replacing scalar V.
  - Phase D: 4 warps each handle 128 V dims via 8 MFMA 16x16x16 bf16_1k tiles
  - Phase C: threads 0-15 compute softmax + write bf16 weights to weight_lds
  - LDS: q_lds(18KB) + kv_lds(18KB) + score_lds(1KB) + weight_lds(512B) + softmax(192B) = ~38KB
  - Remove scalar V accumulation (acc_v, head_id, head_lane, etc.)
"""

from __future__ import annotations
from typing import Any
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
import torch
from task import input_t, output_t

# ---------------------------------------------------------------------------
# HIP kernel source (compiled as .hip / cuda_sources)
# ---------------------------------------------------------------------------

CUDA_SOURCE = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <cstdint>
#include <cfloat>

// =========================================================================
// Constants
// =========================================================================
constexpr int NUM_HEADS = 16;
constexpr int BLOCK_SIZE_ATT = 256;            // 4 warps x 64 threads
constexpr int QK_DIM = 576;
constexpr int V_DIM = 512;
constexpr int PACKED_KV_BYTES = 288;           // 576 / 2
constexpr int MX_BLOCK_SIZE = 32;
constexpr int NUM_MX_BLOCKS = 18;             // 576 / 32
constexpr int BLOCK_N = 16;                    // KV tile size
constexpr int MFMA_M = 16;
constexpr int MFMA_N = 16;
constexpr int MFMA_K = 16;
constexpr int WARP_SIZE = 64;
constexpr int K_ITERS = QK_DIM / MFMA_K;      // 36

// LOG2E for fast exp via exp2
constexpr float LOG2E_VAL = 1.4426950408889634f;

// =========================================================================
// FP4 E2M1 dequantization LUT (16 entries)
// =========================================================================
__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
};

// =========================================================================
// Helper: convert bf16 bits to float
// =========================================================================
__device__ __forceinline__ float bf16_to_float(uint16_t val) {
    union { float f; uint32_t u; } converter;
    converter.u = ((uint32_t)val) << 16;
    return converter.f;
}

// =========================================================================
// Helper: convert float to bf16 bits (round to nearest even)
// =========================================================================
__device__ __forceinline__ uint16_t float_to_bf16(float val) {
    union { float f; uint32_t u; } converter;
    converter.f = val;
    uint32_t bits = converter.u;
    uint32_t lsb = (bits >> 16) & 1;
    uint32_t rounding_bias = 0x7FFF + lsb;
    bits += rounding_bias;
    return (uint16_t)(bits >> 16);
}

// =========================================================================
// Main attention kernel (Split-K) with LDS tiling, MFMA score + MFMA V
//
// Grid: (num_splits * batch_size, 1, 1)
// Block: (256, 1, 1)  -- 4 warps x 64 threads
//
// Each block handles one (batch_item, split) pair for ALL 16 heads.
//
// LDS budget:
//   q_lds:       16 * 576 * 2 = 18,432 bytes
//   kv_lds:      16 * 576 * 2 = 18,432 bytes
//   score_lds:   16 * 16  * 4 = 1,024  bytes
//   weight_lds:  16 * 16  * 2 = 512    bytes
//   softmax_m:   16 * 4       = 64     bytes
//   softmax_l:   16 * 4       = 64     bytes
//   softmax_corr:16 * 4       = 64     bytes
//   Total: ~38,592 bytes ≈ 38 KB → floor(160KB / 38KB) = 4 blocks/CU
// =========================================================================
__global__ void mla_mxfp4_attention_kernel(
    const uint16_t* __restrict__ q,          // (total_q, 16, 576) bf16
    const uint8_t*  __restrict__ kv_buffer,  // (total_kv, 288) packed fp4x2
    const uint8_t*  __restrict__ kv_scale,   // (total_kv, scale_stride) E8M0
    float*          __restrict__ partial_out, // (num_splits*batch_size, 16, 512) fp32
    float*          __restrict__ partial_lse, // (num_splits*batch_size, 16) fp32
    uint16_t*       __restrict__ final_out,   // (total_q, 16, 512) bf16
    int batch_size,
    int kv_seq_len,
    int num_splits,
    int scale_stride,
    float sm_scale
) {
    typedef float __attribute__((ext_vector_type(4))) float4_t;
    typedef short __attribute__((ext_vector_type(4))) short4_t;

    int block_id = blockIdx.x;
    int split_id = block_id / batch_size;
    int batch_id = block_id % batch_size;

    int tid = threadIdx.x;
    int warp_id = tid / WARP_SIZE;    // 0..3
    int lane_id = tid % WARP_SIZE;    // 0..63

    // MFMA lane mapping
    int m_block = lane_id / 16;  // 0..3 — which group of 4 rows this lane handles
    int n_col = lane_id % 16;    // 0..15 — which column

    // KV range for this split
    int kv_per_split = (kv_seq_len + num_splits - 1) / num_splits;
    int kv_start = split_id * kv_per_split;
    int kv_end = kv_start + kv_per_split;
    if (kv_end > kv_seq_len) kv_end = kv_seq_len;

    int q_offset = batch_id;  // decode: total_q = batch_size, q_seq_len=1
    int kv_base = batch_id * kv_seq_len;

    // -----------------------------------------------------------------
    // LDS declarations
    // -----------------------------------------------------------------
    __shared__ uint16_t q_lds[NUM_HEADS * QK_DIM];          // 16 * 576
    __shared__ uint16_t kv_lds[BLOCK_N * QK_DIM];           // 16 * 576
    __shared__ float score_lds[BLOCK_N * NUM_HEADS];         // 16 * 16
    __shared__ uint16_t weight_lds[BLOCK_N * NUM_HEADS];     // 16 * 16 bf16
    __shared__ float softmax_m[NUM_HEADS];                   // per-head running max
    __shared__ float softmax_l[NUM_HEADS];                   // per-head running sum
    __shared__ float softmax_corr[NUM_HEADS];                // per-head correction factor

    // -----------------------------------------------------------------
    // Step 1: Cooperatively load Q into LDS
    // -----------------------------------------------------------------
    const uint16_t* q_batch = q + (int64_t)q_offset * NUM_HEADS * QK_DIM;
    for (int i = tid; i < NUM_HEADS * QK_DIM; i += BLOCK_SIZE_ATT) {
        q_lds[i] = q_batch[i];
    }

    // Initialize softmax state in LDS
    if (tid < NUM_HEADS) {
        softmax_m[tid] = -FLT_MAX;
        softmax_l[tid] = 0.0f;
        softmax_corr[tid] = 1.0f;
    }
    __syncthreads();

    // -----------------------------------------------------------------
    // MFMA V accumulators: each warp handles 128 V dims (8 tiles of 16)
    // Each lane holds float4 for 4 heads (m_block*4 + {0,1,2,3})
    // -----------------------------------------------------------------
    float4_t v_acc[8];
    #pragma unroll
    for (int i = 0; i < 8; i++) {
        v_acc[i][0] = 0.0f;
        v_acc[i][1] = 0.0f;
        v_acc[i][2] = 0.0f;
        v_acc[i][3] = 0.0f;
    }

    // Per-head online softmax state in registers (for 4 heads in this lane's m_block)
    float head_m[4] = {-FLT_MAX, -FLT_MAX, -FLT_MAX, -FLT_MAX};
    float head_l[4] = {0.0f, 0.0f, 0.0f, 0.0f};

    // -----------------------------------------------------------------
    // Step 2: Tile loop over KV tokens (BLOCK_N=16 per tile)
    // -----------------------------------------------------------------
    for (int tile_start = kv_start; tile_start < kv_end; tile_start += BLOCK_N) {
        int tile_end = tile_start + BLOCK_N;
        if (tile_end > kv_end) tile_end = kv_end;
        int tile_len = tile_end - tile_start;

        // =============================================================
        // Phase A: Prefetch + Dequant MXFP4 into kv_lds
        // Two-pass: first load all raw data, then process
        // This separates memory latency from compute for better pipelining
        // =============================================================
        {
            constexpr int MAX_LOADS = 5;  // ceil(1152 / 256) = 5
            int total_u32 = tile_len * (PACKED_KV_BYTES / 4);  // 16 * 72 = 1152

            // Pass 1: Prefetch raw KV data + scale bytes into registers
            uint32_t raw_kv[MAX_LOADS];
            uint8_t  raw_s0[MAX_LOADS];
            uint8_t  raw_s1[MAX_LOADS];
            int      raw_token[MAX_LOADS];
            int      raw_dim[MAX_LOADS];
            int num_loads = 0;

            for (int i = tid; i < total_u32; i += BLOCK_SIZE_ATT) {
                int token_in_tile = i / (PACKED_KV_BYTES / 4);
                int u32_in_token = i % (PACKED_KV_BYTES / 4);
                int byte_in_token = u32_in_token * 4;
                int token_idx = kv_base + tile_start + token_in_tile;
                int dim_base = byte_in_token * 2;

                // Prefetch: issue all global loads back-to-back
                raw_kv[num_loads] = *(const uint32_t*)(kv_buffer + (int64_t)token_idx * PACKED_KV_BYTES + byte_in_token);
                int blk0 = dim_base / MX_BLOCK_SIZE;
                int blk1 = (dim_base + 7) / MX_BLOCK_SIZE;
                raw_s0[num_loads] = kv_scale[(int64_t)token_idx * scale_stride + blk0];
                raw_s1[num_loads] = (blk1 != blk0) ? kv_scale[(int64_t)token_idx * scale_stride + blk1] : raw_s0[num_loads];
                raw_token[num_loads] = token_in_tile;
                raw_dim[num_loads] = dim_base;
                num_loads++;
            }

            // Pass 2: Dequant from registers to kv_lds (no global loads)
            for (int li = 0; li < num_loads; li++) {
                uint32_t packed4 = raw_kv[li];
                int token_in_tile = raw_token[li];
                int dim_base = raw_dim[li];
                int blk0 = dim_base / MX_BLOCK_SIZE;
                float s0 = exp2f((float)raw_s0[li] - 127.0f);
                float s1 = (raw_s1[li] != raw_s0[li]) ? exp2f((float)raw_s1[li] - 127.0f) : s0;

                #pragma unroll
                for (int j = 0; j < 4; j++) {
                    uint8_t byte_val = (packed4 >> (j * 8)) & 0xFF;
                    int d0 = dim_base + j * 2;
                    float scale = (d0 / MX_BLOCK_SIZE == blk0) ? s0 : s1;

                    kv_lds[token_in_tile * QK_DIM + d0] = float_to_bf16(FP4_LUT[byte_val & 0x0F] * scale);
                    kv_lds[token_in_tile * QK_DIM + d0 + 1] = float_to_bf16(FP4_LUT[byte_val >> 4] * scale);
                }
            }
        }
        __syncthreads();

        // =============================================================
        // Phase B: MFMA score computation — warp 0 only
        // 16 heads x 16 tokens, one MFMA chunk
        // =============================================================
        if (warp_id == 0) {
            int m = lane_id % 16;
            int k_sub = lane_id / 16;  // 0..3

            // Initialize accumulator
            float4_t score_acc = {0.0f, 0.0f, 0.0f, 0.0f};

            // K-loop: 576 dims in steps of 16
            for (int k = 0; k < QK_DIM; k += MFMA_K) {
                int k_offset = k + k_sub * 4;

                // Load 4 bf16 from Q for A matrix
                short4_t a_val;
                a_val[0] = (short)q_lds[m * QK_DIM + k_offset];
                a_val[1] = (short)q_lds[m * QK_DIM + k_offset + 1];
                a_val[2] = (short)q_lds[m * QK_DIM + k_offset + 2];
                a_val[3] = (short)q_lds[m * QK_DIM + k_offset + 3];

                // Load 4 bf16 from K for B matrix
                short4_t b_val;
                int token_in_tile = lane_id % 16;
                if (token_in_tile < tile_len) {
                    b_val[0] = (short)kv_lds[token_in_tile * QK_DIM + k_offset];
                    b_val[1] = (short)kv_lds[token_in_tile * QK_DIM + k_offset + 1];
                    b_val[2] = (short)kv_lds[token_in_tile * QK_DIM + k_offset + 2];
                    b_val[3] = (short)kv_lds[token_in_tile * QK_DIM + k_offset + 3];
                } else {
                    b_val[0] = 0; b_val[1] = 0; b_val[2] = 0; b_val[3] = 0;
                }

                // MFMA: S += Q * K^T
                score_acc = __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(
                    a_val, b_val, score_acc, 0, 0, 0);
            }

            // Apply sm_scale
            score_acc[0] *= sm_scale;
            score_acc[1] *= sm_scale;
            score_acc[2] *= sm_scale;
            score_acc[3] *= sm_scale;

            // Write scores to score_lds[token][head]
            // Output mapping: lane l holds C[m_block*4+{0,1,2,3}, n_col]
            // where n_col = lane_id % 16, m_block = lane_id / 16
            int sc_n_col = lane_id % 16;
            int sc_m_block = lane_id / 16;

            if (sc_n_col < tile_len) {
                score_lds[sc_n_col * NUM_HEADS + sc_m_block * 4 + 0] = score_acc[0];
                score_lds[sc_n_col * NUM_HEADS + sc_m_block * 4 + 1] = score_acc[1];
                score_lds[sc_n_col * NUM_HEADS + sc_m_block * 4 + 2] = score_acc[2];
                score_lds[sc_n_col * NUM_HEADS + sc_m_block * 4 + 3] = score_acc[3];
            }
        }
        __syncthreads();

        // =============================================================
        // Phase C: Softmax + Weight preparation (threads 0-15 only)
        // =============================================================
        if (tid < NUM_HEADS) {
            int h = tid;
            float tile_max = -FLT_MAX;
            float scores[BLOCK_N];
            for (int n = 0; n < tile_len; n++) {
                scores[n] = score_lds[n * NUM_HEADS + h];
                tile_max = fmaxf(tile_max, scores[n]);
            }
            float m_old = softmax_m[h];
            float m_new = fmaxf(m_old, tile_max);
            float correction = exp2f((m_old - m_new) * LOG2E_VAL);

            // Update running state
            float l_old = softmax_l[h] * correction;
            float l_new = l_old;

            // Compute attention weights and write to weight_lds
            for (int n = 0; n < tile_len; n++) {
                float w = exp2f((scores[n] - m_new) * LOG2E_VAL);
                l_new += w;
                weight_lds[n * NUM_HEADS + h] = float_to_bf16(w);
            }
            // Zero-pad remaining tokens
            for (int n = tile_len; n < BLOCK_N; n++) {
                weight_lds[n * NUM_HEADS + h] = 0;
            }

            softmax_m[h] = m_new;
            softmax_l[h] = l_new;
            softmax_corr[h] = correction;
        }
        __syncthreads();

        // =============================================================
        // Phase D: V MFMA — all 4 warps, each handles 128 V dims
        // =============================================================
        {
            // Read correction for 4 heads in this lane's m_block
            float corr[4];
            corr[0] = softmax_corr[m_block * 4 + 0];
            corr[1] = softmax_corr[m_block * 4 + 1];
            corr[2] = softmax_corr[m_block * 4 + 2];
            corr[3] = softmax_corr[m_block * 4 + 3];

            // Apply correction to all V accumulators
            #pragma unroll
            for (int i = 0; i < 8; i++) {
                v_acc[i][0] *= corr[0];
                v_acc[i][1] *= corr[1];
                v_acc[i][2] *= corr[2];
                v_acc[i][3] *= corr[3];
            }

            // V MFMA: 8 iterations over V dim chunks (each warp handles 128 V dims)
            int v_base = warp_id * 128;
            #pragma unroll
            for (int vi = 0; vi < 8; vi++) {
                int v_offset = v_base + vi * 16;
                if (v_offset >= V_DIM) break;

                // Load A matrix: attention weights[head, token]
                // MFMA A: lane l needs A[m=l%16, k_sub*4..k_sub*4+3]
                // = weight_lds[token * 16 + head] where token=(l/16)*4+j, head=l%16
                short4_t a_val;
                int k_base_a = (lane_id / 16) * 4;
                a_val[0] = (short)weight_lds[(k_base_a + 0) * NUM_HEADS + (lane_id % 16)];
                a_val[1] = (short)weight_lds[(k_base_a + 1) * NUM_HEADS + (lane_id % 16)];
                a_val[2] = (short)weight_lds[(k_base_a + 2) * NUM_HEADS + (lane_id % 16)];
                a_val[3] = (short)weight_lds[(k_base_a + 3) * NUM_HEADS + (lane_id % 16)];

                // Load B matrix: KV values[token, v_dim]
                // MFMA B: lane l needs B[n=l%16, k_sub*4..k_sub*4+3]
                // = kv_lds[token * QK_DIM + v_dim] where token=(l/16)*4+j, v_dim=v_offset+l%16
                short4_t b_val;
                int n_dim = v_offset + (lane_id % 16);
                int k_base_b = (lane_id / 16) * 4;
                if (n_dim < V_DIM) {
                    b_val[0] = (short)kv_lds[(k_base_b + 0) * QK_DIM + n_dim];
                    b_val[1] = (short)kv_lds[(k_base_b + 1) * QK_DIM + n_dim];
                    b_val[2] = (short)kv_lds[(k_base_b + 2) * QK_DIM + n_dim];
                    b_val[3] = (short)kv_lds[(k_base_b + 3) * QK_DIM + n_dim];
                } else {
                    b_val[0] = 0; b_val[1] = 0; b_val[2] = 0; b_val[3] = 0;
                }

                v_acc[vi] = __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(
                    a_val, b_val, v_acc[vi], 0, 0, 0);
            }
        }
        __syncthreads();
    }

    // =================================================================
    // Output: normalize and write results
    // =================================================================

    // Read final l_val for normalization
    float final_l[4];
    final_l[0] = softmax_l[m_block * 4 + 0];
    final_l[1] = softmax_l[m_block * 4 + 1];
    final_l[2] = softmax_l[m_block * 4 + 2];
    final_l[3] = softmax_l[m_block * 4 + 3];

    float inv_l[4];
    for (int i = 0; i < 4; i++)
        inv_l[i] = (final_l[i] > 0.0f) ? (1.0f / final_l[i]) : 0.0f;

    // Normalize
    #pragma unroll
    for (int i = 0; i < 8; i++) {
        v_acc[i][0] *= inv_l[0];
        v_acc[i][1] *= inv_l[1];
        v_acc[i][2] *= inv_l[2];
        v_acc[i][3] *= inv_l[3];
    }

    // Write output
    // MFMA output: lane l holds C[m_block*4+{0,1,2,3}, n_col] where n_col=l%16
    // For V: heads = m_block*4+{0,1,2,3}, v_dim = v_base + vi*16 + n_col
    int v_base_out = warp_id * 128;

    if (num_splits == 1) {
        for (int vi = 0; vi < 8; vi++) {
            int v_dim = v_base_out + vi * 16 + n_col;
            if (v_dim < V_DIM) {
                for (int h = 0; h < 4; h++) {
                    int head = m_block * 4 + h;
                    int64_t out_idx = ((int64_t)q_offset * NUM_HEADS + head) * V_DIM + v_dim;
                    final_out[out_idx] = float_to_bf16(v_acc[vi][h]);
                }
            }
        }
    } else {
        int split_batch_idx = split_id * batch_size + batch_id;
        for (int vi = 0; vi < 8; vi++) {
            int v_dim = v_base_out + vi * 16 + n_col;
            if (v_dim < V_DIM) {
                for (int h = 0; h < 4; h++) {
                    int head = m_block * 4 + h;
                    int64_t po_idx = ((int64_t)split_batch_idx * NUM_HEADS + head) * V_DIM + v_dim;
                    partial_out[po_idx] = v_acc[vi][h];
                }
            }
        }
        // Write LSE: lane with n_col==0 writes for each head in its m_block
        if (n_col == 0) {
            for (int h = 0; h < 4; h++) {
                int head = m_block * 4 + h;
                float m = softmax_m[head];
                float l = softmax_l[head];
                float lse = m + __logf(fmaxf(l, 1e-20f));
                int lse_idx = split_batch_idx * NUM_HEADS + head;
                partial_lse[lse_idx] = lse;
            }
        }
    }
}

// =========================================================================
// Split-K reduce kernel
// Grid: (batch_size, NUM_HEADS, 1), Block: (256, 1, 1)
// Each thread handles 2 V dims (512 / 256 = 2)
// =========================================================================
__global__ void mla_splitk_reduce_kernel(
    const float*    __restrict__ partial_out,
    const float*    __restrict__ partial_lse,
    uint16_t*       __restrict__ final_out,
    int batch_size,
    int num_splits
) {
    int batch_id = blockIdx.x;
    int head_id = blockIdx.y;
    int tid = threadIdx.x;

    constexpr int DIMS_PER_REDUCE_THREAD = 2;

    // Find global max LSE across splits
    float global_max = -FLT_MAX;
    for (int s = 0; s < num_splits; s++) {
        int split_batch_idx = s * batch_size + batch_id;
        float lse = partial_lse[split_batch_idx * NUM_HEADS + head_id];
        global_max = fmaxf(global_max, lse);
    }

    // Accumulate weighted outputs
    float acc[DIMS_PER_REDUCE_THREAD];
    #pragma unroll
    for (int i = 0; i < DIMS_PER_REDUCE_THREAD; i++) {
        acc[i] = 0.0f;
    }
    float total_weight = 0.0f;

    for (int s = 0; s < num_splits; s++) {
        int split_batch_idx = s * batch_size + batch_id;
        float lse = partial_lse[split_batch_idx * NUM_HEADS + head_id];
        float weight = exp2f((lse - global_max) * LOG2E_VAL);
        total_weight += weight;

        int64_t po_base = ((int64_t)split_batch_idx * NUM_HEADS + head_id) * V_DIM;
        #pragma unroll
        for (int i = 0; i < DIMS_PER_REDUCE_THREAD; i++) {
            int d = tid * DIMS_PER_REDUCE_THREAD + i;
            if (d < V_DIM) {
                acc[i] += weight * partial_out[po_base + d];
            }
        }
    }

    // Normalize and write bf16 output
    float inv_total = (total_weight > 0.0f) ? (1.0f / total_weight) : 0.0f;
    int64_t out_base = ((int64_t)batch_id * NUM_HEADS + head_id) * V_DIM;
    #pragma unroll
    for (int i = 0; i < DIMS_PER_REDUCE_THREAD; i++) {
        int d = tid * DIMS_PER_REDUCE_THREAD + i;
        if (d < V_DIM) {
            final_out[out_base + d] = float_to_bf16(acc[i] * inv_total);
        }
    }
}

// =========================================================================
// Torch C++ wrapper functions
// =========================================================================

void launch_mla_mxfp4_attention(
    torch::Tensor q,
    torch::Tensor kv_buffer,
    torch::Tensor kv_scale,
    torch::Tensor partial_out,
    torch::Tensor partial_lse,
    torch::Tensor final_out,
    int64_t batch_size,
    int64_t kv_seq_len,
    int64_t num_splits,
    int64_t scale_stride,
    double sm_scale
) {
    int grid_x = num_splits * batch_size;
    dim3 grid(grid_x, 1, 1);
    dim3 block(256, 1, 1);

    mla_mxfp4_attention_kernel<<<grid, block>>>(
        reinterpret_cast<const uint16_t*>(q.data_ptr()),
        reinterpret_cast<const uint8_t*>(kv_buffer.data_ptr()),
        reinterpret_cast<const uint8_t*>(kv_scale.data_ptr()),
        partial_out.data_ptr<float>(),
        partial_lse.data_ptr<float>(),
        reinterpret_cast<uint16_t*>(final_out.data_ptr()),
        (int)batch_size,
        (int)kv_seq_len,
        (int)num_splits,
        (int)scale_stride,
        (float)sm_scale
    );
}

void launch_mla_splitk_reduce(
    torch::Tensor partial_out,
    torch::Tensor partial_lse,
    torch::Tensor final_out,
    int64_t batch_size,
    int64_t num_splits
) {
    dim3 grid(batch_size, 16, 1);
    dim3 block(256, 1, 1);

    mla_splitk_reduce_kernel<<<grid, block>>>(
        partial_out.data_ptr<float>(),
        partial_lse.data_ptr<float>(),
        reinterpret_cast<uint16_t*>(final_out.data_ptr()),
        (int)batch_size,
        (int)num_splits
    );
}
"""

# ---------------------------------------------------------------------------
# Per-case split configs — tuned for BLOCK_N=16 and 4 blocks/CU target
# ---------------------------------------------------------------------------
SPLIT_CONFIGS = {
    (4, 1024): 64,
    (4, 8192): 64,
    (32, 1024): 32,
    (32, 8192): 64,
    (64, 1024): 16,
    (64, 8192): 32,
    (256, 1024): 4,
    (256, 8192): 16,
}

DEFAULT_SPLITS = 4

# ---------------------------------------------------------------------------
# Module-level caches
# ---------------------------------------------------------------------------
_module = None
_buffer_cache: dict[tuple, dict[str, torch.Tensor]] = {}


def _get_module():
    """Lazy-compile the HIP kernels via load_inline."""
    global _module
    if _module is not None:
        return _module

    from torch.utils.cpp_extension import load_inline

    _module = load_inline(
        name="mla_mxfp4_kernel_v0011b",
        cpp_sources=[
            """
void launch_mla_mxfp4_attention(
    torch::Tensor q,
    torch::Tensor kv_buffer,
    torch::Tensor kv_scale,
    torch::Tensor partial_out,
    torch::Tensor partial_lse,
    torch::Tensor final_out,
    int64_t batch_size,
    int64_t kv_seq_len,
    int64_t num_splits,
    int64_t scale_stride,
    double sm_scale);
void launch_mla_splitk_reduce(
    torch::Tensor partial_out,
    torch::Tensor partial_lse,
    torch::Tensor final_out,
    int64_t batch_size,
    int64_t num_splits);
"""
        ],
        cuda_sources=[CUDA_SOURCE],
        functions=["launch_mla_mxfp4_attention", "launch_mla_splitk_reduce"],
        extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3", "-w"],
        verbose=False,
    )
    return _module


def _get_buffers(
    batch_size: int,
    num_splits: int,
    device: torch.device,
) -> dict[str, torch.Tensor]:
    """Get or allocate cached buffers for partial outputs."""
    cache_key = (batch_size, num_splits, device)
    if cache_key in _buffer_cache:
        return _buffer_cache[cache_key]

    buffers: dict[str, torch.Tensor] = {}

    # Final output: (batch_size, 16, 512) bf16
    buffers["final_out"] = torch.empty(
        (batch_size, 16, 512), dtype=torch.bfloat16, device=device
    )

    if num_splits > 1:
        # Partial output: (num_splits * batch_size, 16, 512) fp32
        buffers["partial_out"] = torch.empty(
            (num_splits * batch_size, 16, 512), dtype=torch.float32, device=device
        )
        # Partial LSE: (num_splits * batch_size, 16) fp32
        buffers["partial_lse"] = torch.empty(
            (num_splits * batch_size, 16), dtype=torch.float32, device=device
        )
    else:
        # Dummy tensors (not used but needed for kernel launch signature)
        buffers["partial_out"] = torch.empty(1, dtype=torch.float32, device=device)
        buffers["partial_lse"] = torch.empty(1, dtype=torch.float32, device=device)

    _buffer_cache[cache_key] = buffers
    return buffers


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    """MLA decode attention with custom MXFP4 HIP kernel."""
    q, kv_data, qo_indptr, kv_indptr, config = data

    batch_size = int(config["batch_size"])
    kv_seq_len = int(config["kv_seq_len"])
    sm_scale = float(config["sm_scale"])

    # Extract MXFP4 KV cache
    kv_buffer, kv_scale = kv_data["mxfp4"]

    # kv_buffer: (total_kv, 1, 288) uint8 -> flatten to (total_kv, 288)
    kv_buffer_flat = kv_buffer.reshape(-1, 288)

    # kv_scale: (total_kv, N_blocks) uint8, N_blocks may be padded (>= 18)
    scale_stride = int(kv_scale.size(1))  # may be > 18 due to padding

    # Determine number of splits
    num_splits = SPLIT_CONFIGS.get((batch_size, kv_seq_len), DEFAULT_SPLITS)

    # Get compiled module
    mod = _get_module()

    # Get or allocate buffers
    buffers = _get_buffers(batch_size, num_splits, q.device)

    # Ensure q is contiguous with shape (total_q, 16, 576)
    q_contig = q.contiguous()

    # Launch main attention kernel
    mod.launch_mla_mxfp4_attention(
        q_contig,
        kv_buffer_flat,
        kv_scale,
        buffers["partial_out"],
        buffers["partial_lse"],
        buffers["final_out"],
        batch_size,
        kv_seq_len,
        num_splits,
        scale_stride,
        sm_scale,
    )

    # Launch reduce kernel if needed
    if num_splits > 1:
        mod.launch_mla_splitk_reduce(
            buffers["partial_out"],
            buffers["partial_lse"],
            buffers["final_out"],
            batch_size,
            num_splits,
        )

    return buffers["final_out"]
scrolls · 735 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 608930.

- # gpumode leaderboard reference
"""
- Reference implementation for MLA (Multi-head Latent Attention) decode kernel.
+ Custom HIP kernel for MLA decode attention with MXFP4 KV cache on MI355X (gfx950).
- Uses aiter MLA kernels (mla_decode_fwd) as the reference.
- DeepSeek R1 forward_absorb MLA: absorbed q (576), compressed kv_buffer (576),
- output v_head_dim = kv_lora_rank = 512.
-
- The input provides:
- q: (total_q, 16, 576) bfloat16 — absorbed query
- kv_data: dict with KV cache in three formats:
- "bf16": Tensor (total_kv, 1, 576) bfloat16 — highest precision
- "fp8": (Tensor, Tensor) kv_buffer fp8 + scalar scale — per-tensor quantized
- "mxfp4": (Tensor, Tensor) kv_buffer fp4x2 + fp8_e8m0 — block-32 quantized
- The reference quantizes Q to fp8 on-the-fly inside ref_kernel.
-
- The reference kernel quantizes Q to fp8 on-the-fly and uses fp8 KV (a8w8 kernel),
- which is ~2-3x faster than bf16 on MI355X with negligible accuracy loss.
-
- Decode only — persistent mode with get_mla_metadata_v1.
+ v0011: MFMA V accumulation replacing scalar V.
+ - Phase D: 4 warps each handle 128 V dims via 8 MFMA 16x16x16 bf16_1k tiles
+ - Phase C: threads 0-15 compute softmax + write bf16 weights to weight_lds
+ - LDS: q_lds(18KB) + kv_lds(18KB) + score_lds(1KB) + weight_lds(512B) + softmax(192B) = ~38KB
+ - Remove scalar V accumulation (acc_v, head_id, head_lane, etc.)
"""
+ from __future__ import annotations
+ from typing import Any
+ import os
+ os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
import torch
- import torch.nn.functional as F
from task import input_t, output_t
- from utils import make_match_reference
- from aiter.mla import mla_decode_fwd
- from aiter import dtypes as aiter_dtypes
- from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
- from aiter.utility.fp4_utils import (
- dynamic_mxfp4_quant,
- mxfp4_to_f32,
- e8m0_to_f32,
- )
-
# ---------------------------------------------------------------------------
- # DeepSeek R1 latent MQA constants (forward_absorb path)
- # https://huggingface.co/deepseek-ai/DeepSeek-R1-0528/blob/main/config.json
+ # HIP kernel source (compiled as .hip / cuda_sources)
# ---------------------------------------------------------------------------
- NUM_HEADS = 16
- NUM_KV_HEADS = 1
- KV_LORA_RANK = 512
- QK_ROPE_HEAD_DIM = 64
- QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM # 576
- V_HEAD_DIM = KV_LORA_RANK # 512
- SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
- PAGE_SIZE = 1
- NUM_KV_SPLITS = 32
+ CUDA_SOURCE = r"""
+ #include <torch/extension.h>
+ #include <hip/hip_runtime.h>
+ #include <cstdint>
+ #include <cfloat>
- # FP8 dtype (platform-specific via aiter)
- FP8_DTYPE = aiter_dtypes.fp8
+ // =========================================================================
+ // Constants
+ // =========================================================================
+ constexpr int NUM_HEADS = 16;
+ constexpr int BLOCK_SIZE_ATT = 256; // 4 warps x 64 threads
+ constexpr int QK_DIM = 576;
+ constexpr int V_DIM = 512;
+ constexpr int PACKED_KV_BYTES = 288; // 576 / 2
+ constexpr int MX_BLOCK_SIZE = 32;
+ constexpr int NUM_MX_BLOCKS = 18; // 576 / 32
+ constexpr int BLOCK_N = 16; // KV tile size
+ constexpr int MFMA_M = 16;
+ constexpr int MFMA_N = 16;
+ constexpr int MFMA_K = 16;
+ constexpr int WARP_SIZE = 64;
+ constexpr int K_ITERS = QK_DIM / MFMA_K; // 36
- # Query dtype for the reference kernel: "fp8" or "bf16"
- Q_DTYPE = "fp8"
+ // LOG2E for fast exp via exp2
+ constexpr float LOG2E_VAL = 1.4426950408889634f;
- # KV cache dtype for the reference kernel: "fp8" or "bf16"
- KV_DTYPE = "fp8"
+ // =========================================================================
+ // FP4 E2M1 dequantization LUT (16 entries)
+ // =========================================================================
+ __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
+ };
+ // =========================================================================
+ // Helper: convert bf16 bits to float
+ // =========================================================================
+ __device__ __forceinline__ float bf16_to_float(uint16_t val) {
+ union { float f; uint32_t u; } converter;
+ converter.u = ((uint32_t)val) << 16;
+ return converter.f;
+ }
- # ---------------------------------------------------------------------------
- # FP8 quantization (sglang style: dynamic per-tensor)
- # ---------------------------------------------------------------------------
- def quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
- """
- Dynamic per-tensor FP8 quantization (following sglang scaled_fp8_quant).
+ // =========================================================================
+ // Helper: convert float to bf16 bits (round to nearest even)
+ // =========================================================================
+ __device__ __forceinline__ uint16_t float_to_bf16(float val) {
+ union { float f; uint32_t u; } converter;
+ converter.f = val;
+ uint32_t bits = converter.u;
+ uint32_t lsb = (bits >> 16) & 1;
+ uint32_t rounding_bias = 0x7FFF + lsb;
+ bits += rounding_bias;
+ return (uint16_t)(bits >> 16);
+ }
- Args:
- tensor: bf16 tensor to quantize
+ // =========================================================================
+ // Main attention kernel (Split-K) with LDS tiling, MFMA score + MFMA V
+ //
+ // Grid: (num_splits * batch_size, 1, 1)
+ // Block: (256, 1, 1) -- 4 warps x 64 threads
+ //
+ // Each block handles one (batch_item, split) pair for ALL 16 heads.
+ //
+ // LDS budget:
+ // q_lds: 16 * 576 * 2 = 18,432 bytes
+ // kv_lds: 16 * 576 * 2 = 18,432 bytes
+ // score_lds: 16 * 16 * 4 = 1,024 bytes
+ // weight_lds: 16 * 16 * 2 = 512 bytes
+ // softmax_m: 16 * 4 = 64 bytes
+ // softmax_l: 16 * 4 = 64 bytes
+ // softmax_corr:16 * 4 = 64 bytes
+ // Total: ~38,592 bytes ≈ 38 KB → floor(160KB / 38KB) = 4 blocks/CU
+ // =========================================================================
+ __global__ void mla_mxfp4_attention_kernel(
+ const uint16_t* __restrict__ q, // (total_q, 16, 576) bf16
+ const uint8_t* __restrict__ kv_buffer, // (total_kv, 288) packed fp4x2
+ const uint8_t* __restrict__ kv_scale, // (total_kv, scale_stride) E8M0
+ float* __restrict__ partial_out, // (num_splits*batch_size, 16, 512) fp32
+ float* __restrict__ partial_lse, // (num_splits*batch_size, 16) fp32
+ uint16_t* __restrict__ final_out, // (total_q, 16, 512) bf16
+ int batch_size,
+ int kv_seq_len,
+ int num_splits,
+ int scale_stride,
+ float sm_scale
+ ) {
+ typedef float __attribute__((ext_vector_type(4))) float4_t;
+ typedef short __attribute__((ext_vector_type(4))) short4_t;
- 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)
+ int block_id = blockIdx.x;
+ int split_id = block_id / batch_size;
+ int batch_id = block_id % batch_size;
+ int tid = threadIdx.x;
+ int warp_id = tid / WARP_SIZE; // 0..3
+ int lane_id = tid % WARP_SIZE; // 0..63
- # ---------------------------------------------------------------------------
- # MXFP4 quantization (aiter native: block-32, fp4x2 + fp8_e8m0 dtypes)
- # Uses aiter.utility.fp4_utils.dynamic_mxfp4_quant
- # ---------------------------------------------------------------------------
+ // MFMA lane mapping
+ int m_block = lane_id / 16; // 0..3 — which group of 4 rows this lane handles
+ int n_col = lane_id % 16; // 0..15 — which column
- def quantize_mxfp4(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
- """
- MXFP4 block-wise quantization using aiter's dynamic_mxfp4_quant.
+ // KV range for this split
+ int kv_per_split = (kv_seq_len + num_splits - 1) / num_splits;
+ int kv_start = split_id * kv_per_split;
+ int kv_end = kv_start + kv_per_split;
+ if (kv_end > kv_seq_len) kv_end = kv_seq_len;
- Block size = 32. Each block gets an E8M0 scale factor.
- Two FP4 E2M1 values are packed per byte.
+ int q_offset = batch_id; // decode: total_q = batch_size, q_seq_len=1
+ int kv_base = batch_id * kv_seq_len;
- Args:
- tensor: bf16 tensor of shape [B, M, N] (N must be divisible by 32)
+ // -----------------------------------------------------------------
+ // LDS declarations
+ // -----------------------------------------------------------------
+ __shared__ uint16_t q_lds[NUM_HEADS * QK_DIM]; // 16 * 576
+ __shared__ uint16_t kv_lds[BLOCK_N * QK_DIM]; // 16 * 576
+ __shared__ float score_lds[BLOCK_N * NUM_HEADS]; // 16 * 16
+ __shared__ uint16_t weight_lds[BLOCK_N * NUM_HEADS]; // 16 * 16 bf16
+ __shared__ float softmax_m[NUM_HEADS]; // per-head running max
+ __shared__ float softmax_l[NUM_HEADS]; // per-head running sum
+ __shared__ float softmax_corr[NUM_HEADS]; // per-head correction factor
- Returns:
- (fp4_data, scale_e8m0)
- - fp4_data: shape [B, M, N//2] in aiter_dtypes.fp4x2
- - scale_e8m0: shape [B*M, ceil(N/32)] padded, in aiter_dtypes.fp8_e8m0
- """
- orig_shape = tensor.shape # (B, M, N)
- B, M, N = orig_shape
+ // -----------------------------------------------------------------
+ // Step 1: Cooperatively load Q into LDS
+ // -----------------------------------------------------------------
+ const uint16_t* q_batch = q + (int64_t)q_offset * NUM_HEADS * QK_DIM;
+ for (int i = tid; i < NUM_HEADS * QK_DIM; i += BLOCK_SIZE_ATT) {
+ q_lds[i] = q_batch[i];
+ }
- # dynamic_mxfp4_quant expects 2D: (B*M, N)
- tensor_2d = tensor.reshape(B * M, N)
- fp4_data_2d, scale_e8m0 = dynamic_mxfp4_quant(tensor_2d)
+ // Initialize softmax state in LDS
+ if (tid < NUM_HEADS) {
+ softmax_m[tid] = -FLT_MAX;
+ softmax_l[tid] = 0.0f;
+ softmax_corr[tid] = 1.0f;
+ }
+ __syncthreads();
- # Reshape fp4_data back to 3D: (B, M, N//2)
- fp4_data = fp4_data_2d.view(B, M, N // 2)
+ // -----------------------------------------------------------------
+ // MFMA V accumulators: each warp handles 128 V dims (8 tiles of 16)
+ // Each lane holds float4 for 4 heads (m_block*4 + {0,1,2,3})
+ // -----------------------------------------------------------------
+ float4_t v_acc[8];
+ #pragma unroll
+ for (int i = 0; i < 8; i++) {
+ v_acc[i][0] = 0.0f;
+ v_acc[i][1] = 0.0f;
+ v_acc[i][2] = 0.0f;
+ v_acc[i][3] = 0.0f;
+ }
- return fp4_data, scale_e8m0
+ // Per-head online softmax state in registers (for 4 heads in this lane's m_block)
+ float head_m[4] = {-FLT_MAX, -FLT_MAX, -FLT_MAX, -FLT_MAX};
+ float head_l[4] = {0.0f, 0.0f, 0.0f, 0.0f};
+ // -----------------------------------------------------------------
+ // Step 2: Tile loop over KV tokens (BLOCK_N=16 per tile)
+ // -----------------------------------------------------------------
+ for (int tile_start = kv_start; tile_start < kv_end; tile_start += BLOCK_N) {
+ int tile_end = tile_start + BLOCK_N;
+ if (tile_end > kv_end) tile_end = kv_end;
+ int tile_len = tile_end - tile_start;
- 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.
+ // =============================================================
+ // Phase A: Prefetch + Dequant MXFP4 into kv_lds
+ // Two-pass: first load all raw data, then process
+ // This separates memory latency from compute for better pipelining
+ // =============================================================
+ {
+ constexpr int MAX_LOADS = 5; // ceil(1152 / 256) = 5
+ int total_u32 = tile_len * (PACKED_KV_BYTES / 4); // 16 * 72 = 1152
- Note: dynamic_mxfp4_quant may pad both row and block dimensions in scale_e8m0.
- We trim scales to match the actual data dimensions.
+ // Pass 1: Prefetch raw KV data + scale bytes into registers
+ uint32_t raw_kv[MAX_LOADS];
+ uint8_t raw_s0[MAX_LOADS];
+ uint8_t raw_s1[MAX_LOADS];
+ int raw_token[MAX_LOADS];
+ int raw_dim[MAX_LOADS];
+ int num_loads = 0;
- 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
+ for (int i = tid; i < total_u32; i += BLOCK_SIZE_ATT) {
+ int token_in_tile = i / (PACKED_KV_BYTES / 4);
+ int u32_in_token = i % (PACKED_KV_BYTES / 4);
+ int byte_in_token = u32_in_token * 4;
+ int token_idx = kv_base + tile_start + token_in_tile;
+ int dim_base = byte_in_token * 2;
- Returns:
- Dequantized tensor of shape orig_shape.
- """
- B, M, N = orig_shape
- num_rows = B * M
- block_size = 32
- num_blocks = N // block_size # actual blocks needed (e.g. 576/32 = 18)
+ // Prefetch: issue all global loads back-to-back
+ raw_kv[num_loads] = *(const uint32_t*)(kv_buffer + (int64_t)token_idx * PACKED_KV_BYTES + byte_in_token);
+ int blk0 = dim_base / MX_BLOCK_SIZE;
+ int blk1 = (dim_base + 7) / MX_BLOCK_SIZE;
+ raw_s0[num_loads] = kv_scale[(int64_t)token_idx * scale_stride + blk0];
+ raw_s1[num_loads] = (blk1 != blk0) ? kv_scale[(int64_t)token_idx * scale_stride + blk1] : raw_s0[num_loads];
+ raw_token[num_loads] = token_in_tile;
+ raw_dim[num_loads] = dim_base;
+ num_loads++;
+ }
- # Unpack FP4 to float32: mxfp4_to_f32 expects (..., N//2) -> (..., N)
- fp4_data_2d = fp4_data.reshape(num_rows, N // 2)
- float_vals = mxfp4_to_f32(fp4_data_2d) # (num_rows, N)
+ // Pass 2: Dequant from registers to kv_lds (no global loads)
+ for (int li = 0; li < num_loads; li++) {
+ uint32_t packed4 = raw_kv[li];
+ int token_in_tile = raw_token[li];
+ int dim_base = raw_dim[li];
+ int blk0 = dim_base / MX_BLOCK_SIZE;
+ float s0 = exp2f((float)raw_s0[li] - 127.0f);
+ float s1 = (raw_s1[li] != raw_s0[li]) ? exp2f((float)raw_s1[li] - 127.0f) : s0;
- # Convert E8M0 scales to float32 and trim padded dimensions
- scale_f32 = e8m0_to_f32(scale_e8m0) # (padded_rows, padded_blocks)
- scale_f32 = scale_f32[:num_rows, :num_blocks] # (num_rows, num_blocks)
+ #pragma unroll
+ for (int j = 0; j < 4; j++) {
+ uint8_t byte_val = (packed4 >> (j * 8)) & 0xFF;
+ int d0 = dim_base + j * 2;
+ float scale = (d0 / MX_BLOCK_SIZE == blk0) ? s0 : s1;
- # Apply block scales
- float_vals_blocked = float_vals.view(num_rows, num_blocks, block_size)
- scaled = float_vals_blocked * scale_f32.unsqueeze(-1)
+ kv_lds[token_in_tile * QK_DIM + d0] = float_to_bf16(FP4_LUT[byte_val & 0x0F] * scale);
+ kv_lds[token_in_tile * QK_DIM + d0 + 1] = float_to_bf16(FP4_LUT[byte_val >> 4] * scale);
+ }
+ }
+ }
+ __syncthreads();
- return scaled.view(B, M, N).to(dtype)
+ // =============================================================
+ // Phase B: MFMA score computation — warp 0 only
+ // 16 heads x 16 tokens, one MFMA chunk
+ // =============================================================
+ if (warp_id == 0) {
+ int m = lane_id % 16;
+ int k_sub = lane_id / 16; // 0..3
+ // Initialize accumulator
+ float4_t score_acc = {0.0f, 0.0f, 0.0f, 0.0f};
- # ---------------------------------------------------------------------------
- # Persistent mode metadata helpers
- # ---------------------------------------------------------------------------
+ // K-loop: 576 dims in steps of 16
+ for (int k = 0; k < QK_DIM; k += MFMA_K) {
+ int k_offset = k + k_sub * 4;
- 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
+ // Load 4 bf16 from Q for A matrix
+ short4_t a_val;
+ a_val[0] = (short)q_lds[m * QK_DIM + k_offset];
+ a_val[1] = (short)q_lds[m * QK_DIM + k_offset + 1];
+ a_val[2] = (short)q_lds[m * QK_DIM + k_offset + 2];
+ a_val[3] = (short)q_lds[m * QK_DIM + k_offset + 3];
- # Populate the metadata buffers
- get_mla_metadata_v1(
- qo_indptr, kv_indptr, kv_last_page_len,
- nhead // nhead_kv, # num_heads_per_head_k
- nhead_kv, # num_heads_k
- True, # is_causal
- 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,
- )
+ // Load 4 bf16 from K for B matrix
+ short4_t b_val;
+ int token_in_tile = lane_id % 16;
+ if (token_in_tile < tile_len) {
+ b_val[0] = (short)kv_lds[token_in_tile * QK_DIM + k_offset];
+ b_val[1] = (short)kv_lds[token_in_tile * QK_DIM + k_offset + 1];
+ b_val[2] = (short)kv_lds[token_in_tile * QK_DIM + k_offset + 2];
+ b_val[3] = (short)kv_lds[token_in_tile * QK_DIM + k_offset + 3];
+ } else {
+ b_val[0] = 0; b_val[1] = 0; b_val[2] = 0; b_val[3] = 0;
+ }
- 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,
+ // MFMA: S += Q * K^T
+ score_acc = __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(
+ a_val, b_val, score_acc, 0, 0, 0);
+ }
+
+ // Apply sm_scale
+ score_acc[0] *= sm_scale;
+ score_acc[1] *= sm_scale;
+ score_acc[2] *= sm_scale;
+ score_acc[3] *= sm_scale;
+
+ // Write scores to score_lds[token][head]
+ // Output mapping: lane l holds C[m_block*4+{0,1,2,3}, n_col]
+ // where n_col = lane_id % 16, m_block = lane_id / 16
+ int sc_n_col = lane_id % 16;
+ int sc_m_block = lane_id / 16;
+
+ if (sc_n_col < tile_len) {
+ score_lds[sc_n_col * NUM_HEADS + sc_m_block * 4 + 0] = score_acc[0];
+ score_lds[sc_n_col * NUM_HEADS + sc_m_block * 4 + 1] = score_acc[1];
+ score_lds[sc_n_col * NUM_HEADS + sc_m_block * 4 + 2] = score_acc[2];
+ score_lds[sc_n_col * NUM_HEADS + sc_m_block * 4 + 3] = score_acc[3];
+ }
+ }
+ __syncthreads();
+
+ // =============================================================
+ // Phase C: Softmax + Weight preparation (threads 0-15 only)
+ // =============================================================
+ if (tid < NUM_HEADS) {
+ int h = tid;
+ float tile_max = -FLT_MAX;
+ float scores[BLOCK_N];
+ for (int n = 0; n < tile_len; n++) {
+ scores[n] = score_lds[n * NUM_HEADS + h];
+ tile_max = fmaxf(tile_max, scores[n]);
+ }
+ float m_old = softmax_m[h];
+ float m_new = fmaxf(m_old, tile_max);
+ float correction = exp2f((m_old - m_new) * LOG2E_VAL);
+
+ // Update running state
+ float l_old = softmax_l[h] * correction;
+ float l_new = l_old;
+
+ // Compute attention weights and write to weight_lds
+ for (int n = 0; n < tile_len; n++) {
+ float w = exp2f((scores[n] - m_new) * LOG2E_VAL);
+ l_new += w;
+ weight_lds[n * NUM_HEADS + h] = float_to_bf16(w);
+ }
+ // Zero-pad remaining tokens
+ for (int n = tile_len; n < BLOCK_N; n++) {
+ weight_lds[n * NUM_HEADS + h] = 0;
+ }
+
+ softmax_m[h] = m_new;
+ softmax_l[h] = l_new;
+ softmax_corr[h] = correction;
+ }
+ __syncthreads();
+
+ // =============================================================
+ // Phase D: V MFMA — all 4 warps, each handles 128 V dims
+ // =============================================================
+ {
+ // Read correction for 4 heads in this lane's m_block
+ float corr[4];
+ corr[0] = softmax_corr[m_block * 4 + 0];
+ corr[1] = softmax_corr[m_block * 4 + 1];
+ corr[2] = softmax_corr[m_block * 4 + 2];
+ corr[3] = softmax_corr[m_block * 4 + 3];
+
+ // Apply correction to all V accumulators
+ #pragma unroll
+ for (int i = 0; i < 8; i++) {
+ v_acc[i][0] *= corr[0];
+ v_acc[i][1] *= corr[1];
+ v_acc[i][2] *= corr[2];
+ v_acc[i][3] *= corr[3];
+ }
+
+ // V MFMA: 8 iterations over V dim chunks (each warp handles 128 V dims)
+ int v_base = warp_id * 128;
+ #pragma unroll
+ for (int vi = 0; vi < 8; vi++) {
+ int v_offset = v_base + vi * 16;
+ if (v_offset >= V_DIM) break;
+
+ // Load A matrix: attention weights[head, token]
+ // MFMA A: lane l needs A[m=l%16, k_sub*4..k_sub*4+3]
+ // = weight_lds[token * 16 + head] where token=(l/16)*4+j, head=l%16
+ short4_t a_val;
+ int k_base_a = (lane_id / 16) * 4;
+ a_val[0] = (short)weight_lds[(k_base_a + 0) * NUM_HEADS + (lane_id % 16)];
+ a_val[1] = (short)weight_lds[(k_base_a + 1) * NUM_HEADS + (lane_id % 16)];
+ a_val[2] = (short)weight_lds[(k_base_a + 2) * NUM_HEADS + (lane_id % 16)];
+ a_val[3] = (short)weight_lds[(k_base_a + 3) * NUM_HEADS + (lane_id % 16)];
+
+ // Load B matrix: KV values[token, v_dim]
+ // MFMA B: lane l needs B[n=l%16, k_sub*4..k_sub*4+3]
+ // = kv_lds[token * QK_DIM + v_dim] where token=(l/16)*4+j, v_dim=v_offset+l%16
+ short4_t b_val;
+ int n_dim = v_offset + (lane_id % 16);
+ int k_base_b = (lane_id / 16) * 4;
+ if (n_dim < V_DIM) {
+ b_val[0] = (short)kv_lds[(k_base_b + 0) * QK_DIM + n_dim];
+ b_val[1] = (short)kv_lds[(k_base_b + 1) * QK_DIM + n_dim];
+ b_val[2] = (short)kv_lds[(k_base_b + 2) * QK_DIM + n_dim];
+ b_val[3] = (short)kv_lds[(k_base_b + 3) * QK_DIM + n_dim];
+ } else {
+ b_val[0] = 0; b_val[1] = 0; b_val[2] = 0; b_val[3] = 0;
+ }
+
+ v_acc[vi] = __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(
+ a_val, b_val, v_acc[vi], 0, 0, 0);
+ }
+ }
+ __syncthreads();
}
+ // =================================================================
+ // Output: normalize and write results
+ // =================================================================
+ // Read final l_val for normalization
+ float final_l[4];
+ final_l[0] = softmax_l[m_block * 4 + 0];
+ final_l[1] = softmax_l[m_block * 4 + 1];
+ final_l[2] = softmax_l[m_block * 4 + 2];
+ final_l[3] = softmax_l[m_block * 4 + 3];
+
+ float inv_l[4];
+ for (int i = 0; i < 4; i++)
+ inv_l[i] = (final_l[i] > 0.0f) ? (1.0f / final_l[i]) : 0.0f;
+
+ // Normalize
+ #pragma unroll
+ for (int i = 0; i < 8; i++) {
+ v_acc[i][0] *= inv_l[0];
+ v_acc[i][1] *= inv_l[1];
+ v_acc[i][2] *= inv_l[2];
+ v_acc[i][3] *= inv_l[3];
+ }
+
+ // Write output
+ // MFMA output: lane l holds C[m_block*4+{0,1,2,3}, n_col] where n_col=l%16
+ // For V: heads = m_block*4+{0,1,2,3}, v_dim = v_base + vi*16 + n_col
+ int v_base_out = warp_id * 128;
+
+ if (num_splits == 1) {
+ for (int vi = 0; vi < 8; vi++) {
+ int v_dim = v_base_out + vi * 16 + n_col;
+ if (v_dim < V_DIM) {
+ for (int h = 0; h < 4; h++) {
+ int head = m_block * 4 + h;
+ int64_t out_idx = ((int64_t)q_offset * NUM_HEADS + head) * V_DIM + v_dim;
+ final_out[out_idx] = float_to_bf16(v_acc[vi][h]);
+ }
+ }
+ }
+ } else {
+ int split_batch_idx = split_id * batch_size + batch_id;
+ for (int vi = 0; vi < 8; vi++) {
+ int v_dim = v_base_out + vi * 16 + n_col;
+ if (v_dim < V_DIM) {
+ for (int h = 0; h < 4; h++) {
+ int head = m_block * 4 + h;
+ int64_t po_idx = ((int64_t)split_batch_idx * NUM_HEADS + head) * V_DIM + v_dim;
+ partial_out[po_idx] = v_acc[vi][h];
+ }
+ }
+ }
+ // Write LSE: lane with n_col==0 writes for each head in its m_block
+ if (n_col == 0) {
+ for (int h = 0; h < 4; h++) {
+ int head = m_block * 4 + h;
+ float m = softmax_m[head];
+ float l = softmax_l[head];
+ float lse = m + __logf(fmaxf(l, 1e-20f));
+ int lse_idx = split_batch_idx * NUM_HEADS + head;
+ partial_lse[lse_idx] = lse;
+ }
+ }
+ }
+ }
+
+ // =========================================================================
+ // Split-K reduce kernel
+ // Grid: (batch_size, NUM_HEADS, 1), Block: (256, 1, 1)
+ // Each thread handles 2 V dims (512 / 256 = 2)
+ // =========================================================================
+ __global__ void mla_splitk_reduce_kernel(
+ const float* __restrict__ partial_out,
+ const float* __restrict__ partial_lse,
+ uint16_t* __restrict__ final_out,
+ int batch_size,
+ int num_splits
+ ) {
+ int batch_id = blockIdx.x;
+ int head_id = blockIdx.y;
+ int tid = threadIdx.x;
+
+ constexpr int DIMS_PER_REDUCE_THREAD = 2;
+
+ // Find global max LSE across splits
+ float global_max = -FLT_MAX;
+ for (int s = 0; s < num_splits; s++) {
+ int split_batch_idx = s * batch_size + batch_id;
+ float lse = partial_lse[split_batch_idx * NUM_HEADS + head_id];
+ global_max = fmaxf(global_max, lse);
+ }
+
+ // Accumulate weighted outputs
+ float acc[DIMS_PER_REDUCE_THREAD];
+ #pragma unroll
+ for (int i = 0; i < DIMS_PER_REDUCE_THREAD; i++) {
+ acc[i] = 0.0f;
+ }
+ float total_weight = 0.0f;
+
+ for (int s = 0; s < num_splits; s++) {
+ int split_batch_idx = s * batch_size + batch_id;
+ float lse = partial_lse[split_batch_idx * NUM_HEADS + head_id];
+ float weight = exp2f((lse - global_max) * LOG2E_VAL);
+ total_weight += weight;
+
+ int64_t po_base = ((int64_t)split_batch_idx * NUM_HEADS + head_id) * V_DIM;
+ #pragma unroll
+ for (int i = 0; i < DIMS_PER_REDUCE_THREAD; i++) {
+ int d = tid * DIMS_PER_REDUCE_THREAD + i;
+ if (d < V_DIM) {
+ acc[i] += weight * partial_out[po_base + d];
+ }
+ }
+ }
+
+ // Normalize and write bf16 output
+ float inv_total = (total_weight > 0.0f) ? (1.0f / total_weight) : 0.0f;
+ int64_t out_base = ((int64_t)batch_id * NUM_HEADS + head_id) * V_DIM;
+ #pragma unroll
+ for (int i = 0; i < DIMS_PER_REDUCE_THREAD; i++) {
+ int d = tid * DIMS_PER_REDUCE_THREAD + i;
+ if (d < V_DIM) {
+ final_out[out_base + d] = float_to_bf16(acc[i] * inv_total);
+ }
+ }
+ }
+
+ // =========================================================================
+ // Torch C++ wrapper functions
+ // =========================================================================
+
+ void launch_mla_mxfp4_attention(
+ torch::Tensor q,
+ torch::Tensor kv_buffer,
+ torch::Tensor kv_scale,
+ torch::Tensor partial_out,
+ torch::Tensor partial_lse,
+ torch::Tensor final_out,
+ int64_t batch_size,
+ int64_t kv_seq_len,
+ int64_t num_splits,
+ int64_t scale_stride,
+ double sm_scale
+ ) {
+ int grid_x = num_splits * batch_size;
+ dim3 grid(grid_x, 1, 1);
+ dim3 block(256, 1, 1);
+
+ mla_mxfp4_attention_kernel<<<grid, block>>>(
+ reinterpret_cast<const uint16_t*>(q.data_ptr()),
+ reinterpret_cast<const uint8_t*>(kv_buffer.data_ptr()),
+ reinterpret_cast<const uint8_t*>(kv_scale.data_ptr()),
+ partial_out.data_ptr<float>(),
+ partial_lse.data_ptr<float>(),
+ reinterpret_cast<uint16_t*>(final_out.data_ptr()),
+ (int)batch_size,
+ (int)kv_seq_len,
+ (int)num_splits,
+ (int)scale_stride,
+ (float)sm_scale
+ );
+ }
+
+ void launch_mla_splitk_reduce(
+ torch::Tensor partial_out,
+ torch::Tensor partial_lse,
+ torch::Tensor final_out,
+ int64_t batch_size,
+ int64_t num_splits
+ ) {
+ dim3 grid(batch_size, 16, 1);
+ dim3 block(256, 1, 1);
+
+ mla_splitk_reduce_kernel<<<grid, block>>>(
+ partial_out.data_ptr<float>(),
+ partial_lse.data_ptr<float>(),
+ reinterpret_cast<uint16_t*>(final_out.data_ptr()),
+ (int)batch_size,
+ (int)num_splits
+ );
+ }
+ """
+
# ---------------------------------------------------------------------------
- # Aiter reference kernel (decode only)
+ # Per-case split configs — tuned for BLOCK_N=16 and 4 blocks/CU target
# ---------------------------------------------------------------------------
+ SPLIT_CONFIGS = {
+ (4, 1024): 64,
+ (4, 8192): 64,
+ (32, 1024): 32,
+ (32, 8192): 64,
+ (64, 1024): 16,
+ (64, 8192): 32,
+ (256, 1024): 4,
+ (256, 8192): 16,
+ }
- 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:
- """
- MLA decode attention using aiter persistent-mode kernel.
+ DEFAULT_SPLITS = 4
- Supports multiple Q/KV dtype combinations:
- - Q_DTYPE="fp8": fp8 Q + fp8 KV (a8w8) — fastest on MI355X
- - Q_DTYPE="bf16": bf16 Q + bf16 KV (a16w16) — highest precision
+ # ---------------------------------------------------------------------------
+ # Module-level caches
+ # ---------------------------------------------------------------------------
+ _module = None
+ _buffer_cache: dict[tuple, dict[str, torch.Tensor]] = {}
- q: (total_q, num_heads, 576) fp8 or bf16
- kv_buffer: (total_kv, 1, 576) fp8 or bf16
- q_scale: scalar float32 (required for fp8 Q, None for bf16)
- kv_scale: scalar float32 (required for fp8 KV, None for bf16)
- """
- 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())
- # Reshape kv_buffer to 4D for aiter: (total_kv, page_size, nhead_kv, dim)
- kv_buffer_4d = kv_buffer.view(kv_buffer.shape[0], PAGE_SIZE, nkv, kv_buffer.shape[-1])
+ def _get_module():
+ """Lazy-compile the HIP kernels via load_inline."""
+ global _module
+ if _module is not None:
+ return _module
- max_q_len = q_seq_len
- kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
- 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,
+ from torch.utils.cpp_extension import load_inline
+
+ _module = load_inline(
+ name="mla_mxfp4_kernel_v0011b",
+ cpp_sources=[
+ """
+ void launch_mla_mxfp4_attention(
+ torch::Tensor q,
+ torch::Tensor kv_buffer,
+ torch::Tensor kv_scale,
+ torch::Tensor partial_out,
+ torch::Tensor partial_lse,
+ torch::Tensor final_out,
+ int64_t batch_size,
+ int64_t kv_seq_len,
+ int64_t num_splits,
+ int64_t scale_stride,
+ double sm_scale);
+ void launch_mla_splitk_reduce(
+ torch::Tensor partial_out,
+ torch::Tensor partial_lse,
+ torch::Tensor final_out,
+ int64_t batch_size,
+ int64_t num_splits);
+ """
+ ],
+ cuda_sources=[CUDA_SOURCE],
+ functions=["launch_mla_mxfp4_attention", "launch_mla_splitk_reduce"],
+ extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3", "-w"],
+ verbose=False,
)
+ return _module
- 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,
+
+ def _get_buffers(
+ batch_size: int,
+ num_splits: int,
+ device: torch.device,
+ ) -> dict[str, torch.Tensor]:
+ """Get or allocate cached buffers for partial outputs."""
+ cache_key = (batch_size, num_splits, device)
+ if cache_key in _buffer_cache:
+ return _buffer_cache[cache_key]
+
+ buffers: dict[str, torch.Tensor] = {}
+
+ # Final output: (batch_size, 16, 512) bf16
+ buffers["final_out"] = torch.empty(
+ (batch_size, 16, 512), dtype=torch.bfloat16, device=device
)
- return o
+ if num_splits > 1:
+ # Partial output: (num_splits * batch_size, 16, 512) fp32
+ buffers["partial_out"] = torch.empty(
+ (num_splits * batch_size, 16, 512), dtype=torch.float32, device=device
+ )
+ # Partial LSE: (num_splits * batch_size, 16) fp32
+ buffers["partial_lse"] = torch.empty(
+ (num_splits * batch_size, 16), dtype=torch.float32, device=device
+ )
+ else:
+ # Dummy tensors (not used but needed for kernel launch signature)
+ buffers["partial_out"] = torch.empty(1, dtype=torch.float32, device=device)
+ buffers["partial_lse"] = torch.empty(1, dtype=torch.float32, device=device)
+
+ _buffer_cache[cache_key] = buffers
+ return buffers
+
+
+ @torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
- """Reference MLA decode attention. Uses Q_DTYPE and KV_DTYPE to select kernel variant."""
+ """MLA decode attention with custom MXFP4 HIP kernel."""
q, kv_data, qo_indptr, kv_indptr, config = data
- # Resolve Q
- if Q_DTYPE == "fp8":
- q_input, q_scale = quantize_fp8(q)
- else:
- q_input, q_scale = q, None
+ batch_size = int(config["batch_size"])
+ kv_seq_len = int(config["kv_seq_len"])
+ sm_scale = float(config["sm_scale"])
- # Resolve KV
- if KV_DTYPE == "fp8":
- kv_buffer_fp8, kv_scale = kv_data["fp8"]
- kv_input = kv_buffer_fp8
- else:
- 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,
- )
No newline at end of file
+ # Extract MXFP4 KV cache
+ kv_buffer, kv_scale = kv_data["mxfp4"]
+
+ # kv_buffer: (total_kv, 1, 288) uint8 -> flatten to (total_kv, 288)
+ kv_buffer_flat = kv_buffer.reshape(-1, 288)
+
+ # kv_scale: (total_kv, N_blocks) uint8, N_blocks may be padded (>= 18)
+ scale_stride = int(kv_scale.size(1)) # may be > 18 due to padding
+
+ # Determine number of splits
+ num_splits = SPLIT_CONFIGS.get((batch_size, kv_seq_len), DEFAULT_SPLITS)
+
+ # Get compiled module
+ mod = _get_module()
+
+ # Get or allocate buffers
+ buffers = _get_buffers(batch_size, num_splits, q.device)
+
+ # Ensure q is contiguous with shape (total_q, 16, 576)
+ q_contig = q.contiguous()
+
+ # Launch main attention kernel
+ mod.launch_mla_mxfp4_attention(
+ q_contig,
+ kv_buffer_flat,
+ kv_scale,
+ buffers["partial_out"],
+ buffers["partial_lse"],
+ buffers["final_out"],
+ batch_size,
+ kv_seq_len,
+ num_splits,
+ scale_stride,
+ sm_scale,
+ )
+
+ # Launch reduce kernel if needed
+ if num_splits > 1:
+ mod.launch_mla_splitk_reduce(
+ buffers["partial_out"],
+ buffers["partial_lse"],
+ buffers["final_out"],
+ batch_size,
+ num_splits,
+ )
+
+ return buffers["final_out"]
scrolls · 976 diff lines total

Best evidence level for this revision: reported

JSON