Skip to content
KernelIndex
Search⌘K

submission 724364

Barry_zhang · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

040401-sub.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-724364?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
76.4µs
#382 of 766
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:992bc8c49109355de78998dae02f06d35061aff462b8c5553b6150d27cdae7f0
license declaredunknown
license concludedunknown
authorsBarry_zhang
imported2026-08-15

Techniques

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

persistent-kernelUses mla_decode_fwd non-persistent mode which handles all split logic internally.

Kernel source

040401-sub.py68 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""
v105: Use bf16 Q + bf16 KV (a16w16) to eliminate Q quantization overhead.

The per_tensor_quant_hip call costs ~5-10us. For small batch sizes (bs=4,32)
this is a significant fraction of total time. Using bf16 KV costs 2x bandwidth
but saves a kernel launch + quant compute.

Uses mla_decode_fwd non-persistent mode which handles all split logic internally.
The a16w16 kernel (mla_dec_stage1_bf16_a16w16_subQ16_mqa16) handles bf16+bf16
for qseqlen=1 non-persistent.

WARNING: page_size=1 EVERYWHERE.
"""

import torch
from task import input_t, output_t

from aiter.mla import mla_decode_fwd

# MLA constants
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)

_cache = {}


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

    batch_size = config["batch_size"]
    kv_seq_len = config["kv_seq_len"]
    q_total = q.shape[0]

    # bf16 path — no quantization needed
    kv_buffer_bf16 = kv_data["bf16"]
    q_bf16 = q.view(-1, NUM_HEADS, QK_HEAD_DIM)

    kv_buffer_4d = kv_buffer_bf16.view(-1, 1, NUM_KV_HEADS, kv_buffer_bf16.shape[-1])

    # Cache kv metadata per shape (constant across calls); allocate output fresh
    key = (batch_size, kv_seq_len)
    if key not in _cache:
        total_kv = batch_size * kv_seq_len
        kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
        kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")
        _cache[key] = (kv_indices, kv_last_page_len)

    kv_indices, kv_last_page_len = _cache[key]
    output = torch.empty((q_total, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")

    mla_decode_fwd(
        q_bf16, kv_buffer_4d, output,
        qo_indptr, kv_indptr,
        kv_indices, kv_last_page_len,
        1,  # max_seqlen_q
        page_size=1, nhead_kv=NUM_KV_HEADS, sm_scale=SM_SCALE,
        intra_batch_mode=False,
    )

    return output
scrolls · 68 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 666982.

- # Submission #666700
- # ============================================================
- # Leaderboard: amd-mixed-mla (id: 765)
- # File: submission.py
- # User ID: 67768263
- # Submitted: 2026-03-29T19:58:14.612115Z
- # Status: done
+ #!POPCORN leaderboard amd-mixed-mla
+ #!POPCORN gpu MI355X
- # Runs:
- # - benchmark on MI355X: passed (score: -) (2026-03-29T20:00:54.156512Z - 2026-03-29T20:05:29.246979Z)
-
- # Code:
- # ------------------------------------------------------------
"""
- Custom HIP kernel for MLA decode attention with MXFP4 KV cache on MI355X (gfx950).
+ v105: Use bf16 Q + bf16 KV (a16w16) to eliminate Q quantization overhead.
- v0014b: v0014a + fused Phase B+C (score MFMA → in-register shuffle softmax).
- - Eliminates score_lds entirely (~1KB LDS saved)
- - Reduces barriers from 4→3 per tile (25% fewer)
- - Softmax via __shfl_xor width=16 in warp 0 registers
- - KV_STRIDE = 584 for zero LDS bank conflicts
- - Cached A matrix in V MFMA loop
+ The per_tensor_quant_hip call costs ~5-10us. For small batch sizes (bs=4,32)
+ this is a significant fraction of total time. Using bf16 KV costs 2x bandwidth
+ but saves a kernel launch + quant compute.
+
+ Uses mla_decode_fwd non-persistent mode which handles all split logic internally.
+ The a16w16 kernel (mla_dec_stage1_bf16_a16w16_subQ16_mqa16) handles bf16+bf16
+ for qseqlen=1 non-persistent.
+
+ WARNING: page_size=1 EVERYWHERE.
"""
- 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)
- # ---------------------------------------------------------------------------
+ from aiter.mla import mla_decode_fwd
- CUDA_SOURCE = r"""
- #include <torch/extension.h>
- #include <hip/hip_runtime.h>
- #include <cstdint>
- #include <cfloat>
+ # MLA constants
+ 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)
- // =========================================================================
- // 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 MFMA_K_SCORE = 32; // gfx950 wide-K for score: 16x16x32
- constexpr int WARP_SIZE = 64;
- constexpr int K_ITERS = QK_DIM / MFMA_K_SCORE; // 18 (was 36 with K=16)
+ _cache = {}
- // LDS stride for kv_lds — padded to avoid bank conflicts
- // gfx950/CDNA4: 64 banks × 4 bytes. With QK_DIM=576, stride=576*2=1152B,
- // 1152/4=288, 288%64=32 → 8-way bank conflict. PAD=8 → stride=584*2=1168B,
- // 1168/4=292, 292%64=36 → all 16 tokens on different banks → ZERO conflicts.
- constexpr int KV_PAD = 8;
- constexpr int KV_STRIDE = QK_DIM + KV_PAD; // 584
- // 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
- //
- // LDS budget (v0014b — no score_lds):
- // q_lds: 16 * 576 * 2 = 18,432 bytes
- // kv_lds: 16 * 584 * 2 = 18,688 bytes (padded stride)
- // 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: ~37,824 bytes ≈ 37 KB → floor(160KB / 37KB) = 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 * KV_STRIDE]; // 16 * 584 (padded)
- __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 — UNUSED, state is in LDS
- // (removed head_m[4] and head_l[4] to save 8 VGPRs)
-
- // -----------------------------------------------------------------
- // 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: Vectorized MXFP4 dequant into kv_lds (padded stride)
- // =============================================================
- int total_u32 = tile_len * (PACKED_KV_BYTES / 4); // 16 * 72 = 1152
- 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;
-
- uint32_t packed4 = *(const uint32_t*)(kv_buffer + (int64_t)token_idx * PACKED_KV_BYTES + byte_in_token);
-
- #pragma unroll
- for (int j = 0; j < 4; j++) {
- uint8_t byte_val = (packed4 >> (j * 8)) & 0xFF;
- int d0 = dim_base + j * 2;
- int blk = d0 / MX_BLOCK_SIZE;
- // Fast E8M0→float: val→2^(val-127) = IEEE 754 with exponent=val
- union { float f; uint32_t u; } _su;
- _su.u = (uint32_t)kv_scale[(int64_t)token_idx * scale_stride + blk] << 23;
- float scale = _su.f;
-
- kv_lds[token_in_tile * KV_STRIDE + d0] = float_to_bf16(FP4_LUT[byte_val & 0x0F] * scale);
- kv_lds[token_in_tile * KV_STRIDE + 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) {
- typedef short __attribute__((ext_vector_type(8))) short8_t;
-
- 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 32 (gfx950 wide-K MFMA)
- for (int k = 0; k < QK_DIM; k += MFMA_K_SCORE) {
- int k_offset = k + k_sub * 8; // 8 elements per sub-group (32/4)
-
- // Load 8 bf16 from Q for A matrix
- short8_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];
- a_val[4] = (short)q_lds[m * QK_DIM + k_offset + 4];
- a_val[5] = (short)q_lds[m * QK_DIM + k_offset + 5];
- a_val[6] = (short)q_lds[m * QK_DIM + k_offset + 6];
- a_val[7] = (short)q_lds[m * QK_DIM + k_offset + 7];
-
- // Load 8 bf16 from K for B matrix (using padded stride)
- short8_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 * KV_STRIDE + k_offset];
- b_val[1] = (short)kv_lds[token_in_tile * KV_STRIDE + k_offset + 1];
- b_val[2] = (short)kv_lds[token_in_tile * KV_STRIDE + k_offset + 2];
- b_val[3] = (short)kv_lds[token_in_tile * KV_STRIDE + k_offset + 3];
- b_val[4] = (short)kv_lds[token_in_tile * KV_STRIDE + k_offset + 4];
- b_val[5] = (short)kv_lds[token_in_tile * KV_STRIDE + k_offset + 5];
- b_val[6] = (short)kv_lds[token_in_tile * KV_STRIDE + k_offset + 6];
- b_val[7] = (short)kv_lds[token_in_tile * KV_STRIDE + k_offset + 7];
- } else {
- b_val[0] = 0; b_val[1] = 0; b_val[2] = 0; b_val[3] = 0;
- b_val[4] = 0; b_val[5] = 0; b_val[6] = 0; b_val[7] = 0;
- }
-
- // gfx950 wide-K MFMA: S += Q * K^T (K=32 per instruction)
- score_acc = __builtin_amdgcn_mfma_f32_16x16x32_bf16(
- 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;
-
- // =============================================================
- // Fused Phase B+C: In-register softmax via warp shuffle
- //
- // MFMA output layout: lane l holds score_acc[0..3] for
- // heads (l/16)*4+{0,1,2,3} at token l%16
- // 4 groups of 16 lanes, each group handles 4 heads across 16 tokens
- // __shfl_xor with width=16 reduces within each 16-lane group
- // =============================================================
- int token = lane_id % 16;
- int sc_m_block = lane_id / 16;
-
- // Mask out-of-range tokens to -inf
- if (token >= tile_len) {
- score_acc[0] = -FLT_MAX;
- score_acc[1] = -FLT_MAX;
- score_acc[2] = -FLT_MAX;
- score_acc[3] = -FLT_MAX;
- }
-
- // 1. Tile-max reduction per component via butterfly shuffle
- float tile_max[4];
- #pragma unroll
- for (int c = 0; c < 4; c++) {
- tile_max[c] = score_acc[c];
- tile_max[c] = fmaxf(tile_max[c], __shfl_xor(tile_max[c], 1, 16));
- tile_max[c] = fmaxf(tile_max[c], __shfl_xor(tile_max[c], 2, 16));
- tile_max[c] = fmaxf(tile_max[c], __shfl_xor(tile_max[c], 4, 16));
- tile_max[c] = fmaxf(tile_max[c], __shfl_xor(tile_max[c], 8, 16));
- }
-
- // 2. Online softmax update — read LDS state (broadcast, no conflict)
- float m_old[4], m_new_local[4], correction_local[4];
- #pragma unroll
- for (int c = 0; c < 4; c++) {
- int head = sc_m_block * 4 + c;
- m_old[c] = softmax_m[head];
- m_new_local[c] = fmaxf(m_old[c], tile_max[c]);
- correction_local[c] = exp2f((m_old[c] - m_new_local[c]) * LOG2E_VAL);
- }
-
- // 3. Compute attention weights
- float w[4];
- #pragma unroll
- for (int c = 0; c < 4; c++) {
- w[c] = (token < tile_len)
- ? exp2f((score_acc[c] - m_new_local[c]) * LOG2E_VAL)
- : 0.0f;
- }
-
- // 4. Sum reduction via butterfly shuffle
- float sum_w[4];
- #pragma unroll
- for (int c = 0; c < 4; c++) {
- sum_w[c] = w[c];
- sum_w[c] += __shfl_xor(sum_w[c], 1, 16);
- sum_w[c] += __shfl_xor(sum_w[c], 2, 16);
- sum_w[c] += __shfl_xor(sum_w[c], 4, 16);
- sum_w[c] += __shfl_xor(sum_w[c], 8, 16);
- }
-
- // 5. Update running softmax state in LDS (one lane per group writes)
- if (token == 0) {
- #pragma unroll
- for (int c = 0; c < 4; c++) {
- int head = sc_m_block * 4 + c;
- float l_old = softmax_l[head] * correction_local[c];
- softmax_m[head] = m_new_local[c];
- softmax_l[head] = l_old + sum_w[c];
- softmax_corr[head] = correction_local[c];
- }
- }
-
- // 6. Write attention weights to weight_lds (all 64 lanes write)
- #pragma unroll
- for (int c = 0; c < 4; c++) {
- weight_lds[token * NUM_HEADS + sc_m_block * 4 + c] = float_to_bf16(w[c]);
- }
- }
- // Single barrier after fused B+C (was 2 barriers before)
- __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];
- }
-
- // Cache A matrix (attention weights) — same for all V dim chunks
- int k_base_a = (lane_id / 16) * 4;
- short4_t a_cached;
- a_cached[0] = (short)weight_lds[(k_base_a + 0) * NUM_HEADS + (lane_id % 16)];
- a_cached[1] = (short)weight_lds[(k_base_a + 1) * NUM_HEADS + (lane_id % 16)];
- a_cached[2] = (short)weight_lds[(k_base_a + 2) * NUM_HEADS + (lane_id % 16)];
- a_cached[3] = (short)weight_lds[(k_base_a + 3) * NUM_HEADS + (lane_id % 16)];
-
- // 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 B matrix: KV values[token, v_dim] (using padded stride)
- // MFMA B: lane l needs B[n=l%16, k_sub*4..k_sub*4+3]
- // = kv_lds[token * KV_STRIDE + 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) * KV_STRIDE + n_dim];
- b_val[1] = (short)kv_lds[(k_base_b + 1) * KV_STRIDE + n_dim];
- b_val[2] = (short)kv_lds[(k_base_b + 2) * KV_STRIDE + n_dim];
- b_val[3] = (short)kv_lds[(k_base_b + 3) * KV_STRIDE + 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_cached, 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): 16, # v0014b sweep: 64→16 = -21% win
- (4, 8192): 64, # 32 was regression, keep 64
- (32, 1024): 16, # v0014b sweep: 32→16 = -4% win
- (32, 8192): 64,
- (64, 1024): 16,
- (64, 8192): 32,
- (256, 1024): 4, # splits=2 was regression
- (256, 8192): 8, # v0014b sweep: 16→8 = -4% win
- }
-
- 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_v0016",
- 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"])
+ batch_size = config["batch_size"]
+ kv_seq_len = config["kv_seq_len"]
+ q_total = q.shape[0]
- # Extract MXFP4 KV cache
- kv_buffer, kv_scale = kv_data["mxfp4"]
+ # bf16 path — no quantization needed
+ kv_buffer_bf16 = kv_data["bf16"]
+ q_bf16 = q.view(-1, NUM_HEADS, QK_HEAD_DIM)
- # kv_buffer: (total_kv, 1, 288) uint8 -> flatten to (total_kv, 288)
- kv_buffer_flat = kv_buffer.reshape(-1, 288)
+ kv_buffer_4d = kv_buffer_bf16.view(-1, 1, NUM_KV_HEADS, kv_buffer_bf16.shape[-1])
- # 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
+ # Cache kv metadata per shape (constant across calls); allocate output fresh
+ key = (batch_size, kv_seq_len)
+ if key not in _cache:
+ total_kv = batch_size * kv_seq_len
+ kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
+ kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")
+ _cache[key] = (kv_indices, kv_last_page_len)
- # Determine number of splits
- num_splits = SPLIT_CONFIGS.get((batch_size, kv_seq_len), DEFAULT_SPLITS)
+ kv_indices, kv_last_page_len = _cache[key]
+ output = torch.empty((q_total, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
- # 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,
+ mla_decode_fwd(
+ q_bf16, kv_buffer_4d, output,
+ qo_indptr, kv_indptr,
+ kv_indices, kv_last_page_len,
+ 1, # max_seqlen_q
+ page_size=1, nhead_kv=NUM_KV_HEADS, sm_scale=SM_SCALE,
+ intra_batch_mode=False,
)
- # 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"]
-
+ return output
No newline at end of file
scrolls · 810 diff lines total

Best evidence level for this revision: reported

JSON