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
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.
fp4
Custom HIP kernel for MLA decode attention with MXFP4 KV cache on MI355X (gfx950).online-softmax
float m_new = fmaxf(m_old, tile_max);shared-memory
__shared__ uint16_t q_lds[NUM_HEADS * QK_DIM]; // 16 * 576split-k
__global__ void mla_splitk_reduce_kernel(tile-n = 16
constexpr int BLOCK_N = 16; // KV tile sizeKernel 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 Ffrom 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