submission 666982
Barry_zhang · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 764 lines, June 9 Researcher Reciprocity License v1.0.
submission_v0016.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-666982?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:6f485f027d9cfd1bc5477040e791f2de15125ecf905c04bd54625009752ce708
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).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_v0016.py764 lines
# Submission #666700
# ============================================================
# Leaderboard: amd-mixed-mla (id: 765)
# File: submission.py
# User ID: 67768263
# Submitted: 2026-03-29T19:58:14.612115Z
# Status: done
# 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).
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
"""
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 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)
// 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"])
# 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 · 764 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 665830.
+ # Submission #666700+ # ============================================================+ # Leaderboard: amd-mixed-mla (id: 765)+ # File: submission.py+ # User ID: 67768263+ # Submitted: 2026-03-29T19:58:14.612115Z+ # Status: done++ # 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).- 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.)+ 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"""from __future__ import annotations⋯ 27 unchanged linesconstexpr 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: 16x16x32constexpr int WARP_SIZE = 64;- constexpr int K_ITERS = QK_DIM / MFMA_K; // 36+ constexpr int K_ITERS = QK_DIM / MFMA_K_SCORE; // 18 (was 36 with K=16)+ // 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 exp2constexpr float LOG2E_VAL = 1.4426950408889634f;⋯ 33 unchanged lines// 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:+ // LDS budget (v0014b — no score_lds):// q_lds: 16 * 576 * 2 = 18,432 bytes- // kv_lds: 16 * 576 * 2 = 18,432 bytes- // score_lds: 16 * 16 * 4 = 1,024 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: ~38,592 bytes ≈ 38 KB → floor(160KB / 38KB) = 4 blocks/CU+ // 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⋯ 36 unchanged lines// 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 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⋯ 28 unchanged linesv_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};+ // 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)⋯ 4 unchanged linesint tile_len = tile_end - tile_start;// =============================================================- // Phase A: Vectorized MXFP4 dequant into kv_lds+ // Phase A: Vectorized MXFP4 dequant into kv_lds (padded stride)// =============================================================int total_u32 = tile_len * (PACKED_KV_BYTES / 4); // 16 * 72 = 1152for (int i = tid; i < total_u32; i += BLOCK_SIZE_ATT) {⋯ 15 unchanged lines_su.u = (uint32_t)kv_scale[(int64_t)token_idx * scale_stride + blk] << 23;float scale = _su.f;- 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);+ 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();⋯ 3 unchanged lines// 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 accumulatorfloat4_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;+ // 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 4 bf16 from Q for A matrix- short4_t a_val;+ // 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 4 bf16 from K for B matrix- short4_t b_val;+ // 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 * 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];+ 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;}- // MFMA: S += Q * K^T- score_acc = __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(+ // 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);}⋯ 3 unchanged linesscore_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;+ // =============================================================+ // 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;- 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];+ // 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;}- }- __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]);+ // 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));}- 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;+ // 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);+ }- // 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);+ // 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;}- // Zero-pad remaining tokens- for (int n = tile_len; n < BLOCK_N; n++) {- weight_lds[n * NUM_HEADS + h] = 0;++ // 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);}- softmax_m[h] = m_new;- softmax_l[h] = l_new;- softmax_corr[h] = correction;+ // 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();// =============================================================⋯ 16 unchanged linesv_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⋯ 1 unchanged linesint 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]+ // 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 * QK_DIM + v_dim] where token=(l/16)*4+j, v_dim=v_offset+l%16+ // = kv_lds[token * KV_STRIDE + v_dim] where token=(l/16)*4+j, v_dim=v_offset+l%16short4_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];+ 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_val, b_val, v_acc[vi], 0, 0, 0);+ a_cached, b_val, v_acc[vi], 0, 0, 0);}}__syncthreads();⋯ 187 unchanged lines# 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,+ (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,- (256, 8192): 16,+ (256, 1024): 4, # splits=2 was regression+ (256, 8192): 8, # v0014b sweep: 16→8 = -4% win}DEFAULT_SPLITS = 4⋯ 14 unchanged linesfrom torch.utils.cpp_extension import load_inline_module = load_inline(- name="mla_mxfp4_kernel_v0011c",+ name="mla_mxfp4_kernel_v0016",cpp_sources=["""void launch_mla_mxfp4_attention(⋯ 115 unchanged lines)return buffers["final_out"]+
scrolls · 388 diff lines total
Best evidence level for this revision: reported
JSON