submission 734135
willfisher · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 770 lines, June 9 Researcher Reciprocity License v1.0.
harry.jpeg.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-734135?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:129b13052addc70370781c18c39a00acb5ee4bb38dc64fd6eca044d68cbd9c0d
license declaredunknown
license concludedunknown
authorswillfisher
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
kv_buffer, kv_scales = kv_data["mxfp4"]num-warps = 8
constexpr int NUM_WARPS = 8;online-softmax
float running_max[HEADS_PER_LANE], running_sum[HEADS_PER_LANE];shared-memory
extern __shared__ char smem[];vector-width = float2
float2 f01 = __bfloat1622float2(bp[0]);Kernel source
harry.jpeg.py770 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
import os
if "PYTORCH_ROCM_ARCH" not in os.environ:
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
NUM_HEADS = 16
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM
V_HEAD_DIM = KV_LORA_RANK
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
CUDA_SRC = r"""
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <pybind11/pybind11.h>
constexpr int NUM_HEADS = 16;
constexpr int QK_HEAD_DIM = 576;
constexpr int V_HEAD_DIM = 512;
constexpr int KV_PACKED = 288;
constexpr int SCALE_COLS = 18;
constexpr int WAVEFRONT_SIZE = 64;
constexpr int NUM_WARPS = 8;
constexpr int BLOCK_THREADS = NUM_WARPS * WAVEFRONT_SIZE;
constexpr int HEADS_PER_LANE = 4;
constexpr float LOG2E = 1.4426950408889634f;
constexpr int FP4_K = 128; // dims per K chunk
constexpr int FP4_K_CHUNKS = 5; // ceil(576/128)
constexpr int V_K_CHUNKS = 4; // 512/128 — chunks that contain V data
constexpr int TILE_TOKENS = 128; // tokens per tile
// LDS layout: KV double-buf | V_fp8/Q reuse | score | attn_fp8 | stats
// KV loaded via coalesced DMA → LDS, then each thread reads to registers
constexpr int BUF_KV_BYTES = TILE_TOKENS * KV_PACKED; // 36864
constexpr int BUF_SC_BYTES = TILE_TOKENS * SCALE_COLS; // 2304
constexpr int BUF_TOTAL = BUF_KV_BYTES + BUF_SC_BYTES; // 39168
constexpr int V_FP8_STRIDE = V_HEAD_DIM + 8; // 520: 8-byte aligned, breaks bank conflicts (520/4%64=2)
constexpr int V_FP8_BYTES = TILE_TOKENS * V_FP8_STRIDE; // 66560
constexpr int SCORE_STRIDE = NUM_HEADS + 2; // 18 bf16/row: pad for bank conflicts (gcd(9,64)=1)
constexpr int SCORE_BYTES = TILE_TOKENS * SCORE_STRIDE * 2; // 4608
constexpr int ATTN_FP8_BYTES = TILE_TOKENS * NUM_HEADS + 136; // 2184 (fp8 attn + non-linear bank shift)
constexpr int STATS_BYTES = NUM_HEADS * 2 * 4; // 128
// Q reuses the V_fp8 area (Q written in prologue, cached to regs, then V_fp8 overwrites)
constexpr int OFF_KV0 = 0;
constexpr int OFF_KV1 = BUF_TOTAL; // 39168
constexpr int OFF_VFP8 = BUF_TOTAL * 2; // 78336 (also OFF_Q during prologue)
constexpr int OFF_Q = OFF_VFP8; // reuses V_fp8 area
constexpr int OFF_QSC = OFF_Q + NUM_HEADS * KV_PACKED; // 82944
constexpr int OFF_SCORE = OFF_VFP8 + V_FP8_BYTES; // 143872
constexpr int OFF_ATTN = OFF_SCORE + SCORE_BYTES; // 147968
constexpr int OFF_STATS = OFF_ATTN + ATTN_FP8_BYTES; // 150016
constexpr int TOTAL_LDS = OFF_STATS + STATS_BYTES; // 150144 (~147KB)
typedef __bf16 bf16x2 __attribute__((__vector_size__(2 * sizeof(__bf16))));
typedef float f32x4 __attribute__((__vector_size__(4 * sizeof(float))));
typedef int v4i32 __attribute__((__vector_size__(4 * sizeof(int))));
typedef int v2i32 __attribute__((__vector_size__(2 * sizeof(int))));
typedef int v8i32 __attribute__((__vector_size__(8 * sizeof(int))));
typedef uint32_t u32x4_vec __attribute__((__vector_size__(4 * sizeof(uint32_t))));
typedef __bf16 bf16x4 __attribute__((__vector_size__(4 * sizeof(__bf16))));
__device__ __forceinline__ v8i32 to_v8(v4i32 x) {
v8i32 r = {x[0], x[1], x[2], x[3], 0, 0, 0, 0};
return r;
}
// Fused DPP mirror butterfly: single-instruction v_max/v_add with DPP modifier
#define DPP_MAX(m, ctrl) asm volatile( \
"v_max_f32_dpp %0, %0, %0 " ctrl " row_mask:0xf bank_mask:0xf" : "+v"(m))
#define DPP_ADD(s, ctrl) asm volatile( \
"v_add_f32_dpp %0, %0, %0 " ctrl " row_mask:0xf bank_mask:0xf" : "+v"(s))
__device__ __forceinline__ float dpp_row_max(float m) {
DPP_MAX(m, "quad_perm:[1,0,3,2]");
DPP_MAX(m, "quad_perm:[2,3,0,1]");
DPP_MAX(m, "row_half_mirror");
DPP_MAX(m, "row_mirror");
return m;
}
__device__ __forceinline__ float dpp_row_sum(float s) {
DPP_ADD(s, "quad_perm:[1,0,3,2]");
DPP_ADD(s, "quad_perm:[2,3,0,1]");
DPP_ADD(s, "row_half_mirror");
DPP_ADD(s, "row_mirror");
return s;
}
using i32x4 = int32_t __attribute__((ext_vector_type(4)));
using i32x2 = int32_t __attribute__((ext_vector_type(2)));
using as3_uint32_ptr = uint32_t __attribute__((address_space(3)))*;
using as3_bf16x4_ptr = bf16x4 __attribute__((address_space(3)))*;
using as3_i32x2_ptr = i32x2 __attribute__((address_space(3)))*;
using as3_v2i32_ptr = v2i32 __attribute__((address_space(3)))*;
typedef __bf16 bf16x8 __attribute__((__vector_size__(8 * sizeof(__bf16))));
typedef short ev_short2 __attribute__((ext_vector_type(2)));
extern "C" __device__ __attribute__((const)) __bf16 llvm_exp2_bf16(__bf16) __asm("llvm.exp2.bf16");
extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
i32x4 rsrc, as3_uint32_ptr lds_ptr, int size, int voffset, int soffset, int offset, int aux)
__asm("llvm.amdgcn.raw.buffer.load.lds");
struct buffer_resource { uint64_t ptr; uint32_t range; uint32_t config; };
__device__ __forceinline__ i32x4 make_srsrc(const void* ptr, uint32_t range_bytes) {
buffer_resource rsrc = {reinterpret_cast<uint64_t>(ptr), range_bytes, 0x110000};
return *reinterpret_cast<const i32x4*>(&rsrc);
}
// ============================================================
// Grouped kernel: double-buffered KV LDS, per-warp V dequant
// ============================================================
template<int NUM_PARTIALS, int TILES_PER_BLOCK>
__global__ void __launch_bounds__(BLOCK_THREADS)
mla_grouped(
const __hip_bfloat16* __restrict__ q,
const uint8_t* __restrict__ kv_data,
const uint8_t* __restrict__ kv_scales,
const int32_t* __restrict__ qo_indptr,
const int32_t* __restrict__ kv_indptr,
__hip_bfloat16* __restrict__ partial_out,
float* __restrict__ partial_max,
float* __restrict__ partial_sum,
__hip_bfloat16* __restrict__ final_output,
float sm_scale
) {
const int batch_idx = blockIdx.x / NUM_PARTIALS;
const int partial_idx = blockIdx.x - batch_idx * NUM_PARTIALS;
const int tile_base = partial_idx * TILES_PER_BLOCK;
const int warp_id = threadIdx.x >> 6;
const int lane_id = threadIdx.x & 63;
const int lane_mod16 = lane_id & 15;
const int k_group = lane_id >> 4;
const int tok = warp_id * 16 + lane_mod16;
const int q_start = qo_indptr[batch_idx];
const int kv_start = kv_indptr[batch_idx];
const int kv_end = kv_indptr[batch_idx + 1];
extern __shared__ char smem[];
uint8_t* q_lds = reinterpret_cast<uint8_t*>(smem + OFF_Q);
uint8_t* qsc_lds = reinterpret_cast<uint8_t*>(smem + OFF_QSC);
uint8_t* v_fp8_lds = reinterpret_cast<uint8_t*>(smem + OFF_VFP8);
uint8_t* attn_fp8_lds = reinterpret_cast<uint8_t*>(smem + OFF_ATTN);
float* stats_lds = reinterpret_cast<float*>(smem + OFF_STATS);
// Helper: coalesced DMA load of full KV tile into LDS buffer
auto load_tile = [&](int hbm_start, int buf) {
uint8_t* dst_kv = reinterpret_cast<uint8_t*>(smem + (buf == 0 ? OFF_KV0 : OFF_KV1));
uint8_t* dst_sc = dst_kv + BUF_KV_BYTES;
constexpr int KV_LOADS = (BUF_KV_BYTES + 15) / 16;
i32x4 kv_srsrc = make_srsrc(kv_data + hbm_start * KV_PACKED, BUF_KV_BYTES);
for (int off = threadIdx.x; off < KV_LOADS; off += BLOCK_THREADS) {
int voff = off * 16;
llvm_amdgcn_raw_buffer_load_lds(kv_srsrc,
reinterpret_cast<as3_uint32_ptr>(reinterpret_cast<uintptr_t>(dst_kv) + voff),
16, voff, 0, 0, 0);
}
constexpr int SC_LOADS = (BUF_SC_BYTES + 15) / 16;
i32x4 sc_srsrc = make_srsrc(kv_scales + hbm_start * SCALE_COLS, BUF_SC_BYTES);
for (int off = threadIdx.x; off < SC_LOADS; off += BLOCK_THREADS) {
int voff = off * 16;
llvm_amdgcn_raw_buffer_load_lds(sc_srsrc,
reinterpret_cast<as3_uint32_ptr>(reinterpret_cast<uintptr_t>(dst_sc) + voff),
16, voff, 0, 0, 0);
}
};
// Issue tile 0 DMA (non-blocking, in flight during Q quantize)
load_tile(kv_start + tile_base * TILE_TOKENS, 0);
// ---- Q quantize (same as before) ----
{
const __hip_bfloat16* q_base = q + q_start * NUM_HEADS * QK_HEAD_DIM;
const int dc = warp_id;
if (dc < FP4_K_CHUNKS) {
const int qds = dc * FP4_K + (k_group << 5);
if (qds < QK_HEAD_DIM) {
__hip_bfloat16 qv[32];
const __hip_bfloat16* qp = q_base + lane_mod16 * QK_HEAD_DIM + qds;
#pragma unroll
for (int i = 0; i < 4; i++)
*reinterpret_cast<u32x4_vec*>(&qv[i*8]) = *reinterpret_cast<const u32x4_vec*>(qp + i*8);
uint16_t mx = 0;
#pragma unroll
for (int i = 0; i < 32; i++) { uint16_t b = *reinterpret_cast<const uint16_t*>(&qv[i]) & 0x7FFF; mx = max(mx, b); }
uint8_t e8 = 0;
if (mx != 0) { uint32_t bits = (uint32_t)mx << 16; bits = (bits + 0x200000u) & 0xFF800000u; e8 = (uint8_t)max(0, min(254, (int)(bits >> 23) - 2)); }
float sf = __uint_as_float((uint32_t)e8 << 23);
v4i32 ar;
#pragma unroll
for (int w = 0; w < 4; w++) {
unsigned int pk = 0;
pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk, __hip_bfloat162{qv[w*8+0],qv[w*8+1]}, sf, 0);
pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk, __hip_bfloat162{qv[w*8+2],qv[w*8+3]}, sf, 1);
pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk, __hip_bfloat162{qv[w*8+4],qv[w*8+5]}, sf, 2);
pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk, __hip_bfloat162{qv[w*8+6],qv[w*8+7]}, sf, 3);
ar[w] = (int)pk;
}
const int qb = qds >> 5;
*reinterpret_cast<u32x4_vec*>(q_lds + lane_mod16 * KV_PACKED + (qb << 4)) = *reinterpret_cast<const u32x4_vec*>(&ar);
qsc_lds[lane_mod16 * SCALE_COLS + qb] = e8;
}
}
}
// Running accumulators
f32x4 running_v[4];
float running_max[HEADS_PER_LANE], running_sum[HEADS_PER_LANE];
#pragma unroll
for (int i = 0; i < 4; i++) running_v[i] = {0,0,0,0};
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++) { running_max[h] = -INFINITY; running_sum[h] = 0; }
v8i32 q_v8[FP4_K_CHUNKS];
int q_scale_cached[FP4_K_CHUNKS];
const float sm_log2e = sm_scale * LOG2E;
using f16x2 = _Float16 __attribute__((ext_vector_type(2)));
// ============ TILE LOOP ============
for (int ti = 0; ti < TILES_PER_BLOCK; ti++) {
const int cur = ti & 1;
const int tile_start = kv_start + (tile_base + ti) * TILE_TOKENS;
const int tile_len = max(0, min(TILE_TOKENS, kv_end - tile_start));
// Wait for KV DMA + ensure all threads' LDS writes visible
asm volatile("s_waitcnt vmcnt(0)");
__syncthreads();
// Cache Q on first iteration (after sync ensures q_lds visible)
if (ti == 0) {
#pragma unroll
for (int d = 0; d < FP4_K_CHUNKS; d++) {
const int qb = d * (FP4_K / 32) + k_group;
v4i32 qr = *reinterpret_cast<const v4i32*>(q_lds + lane_mod16 * KV_PACKED + (qb << 4));
q_scale_cached[d] = (int)qsc_lds[lane_mod16 * SCALE_COLS + qb];
q_v8[d] = to_v8(qr);
}
if (k_group >= 2) { q_v8[4] = {0,0,0,0,0,0,0,0}; q_scale_cached[4] = 0; }
}
const uint8_t* kv_lds = reinterpret_cast<const uint8_t*>(smem + (cur == 0 ? OFF_KV0 : OFF_KV1));
const uint8_t* sc_lds = kv_lds + BUF_KV_BYTES;
const int tok_kv_base = tok * KV_PACKED + (k_group << 4);
const int tok_sc_base = tok * SCALE_COLS + k_group;
// Prefetch next tile (overlaps with entire tile compute)
if (ti + 1 < TILES_PER_BLOCK)
load_tile(kv_start + (tile_base + ti + 1) * TILE_TOKENS, 1 - cur);
// ---- Fused Score + V dequant, double-buffered LDS reads ----
f32x4 scores = {0,0,0,0};
const bool tok_valid = tok < tile_len;
// Prefetch chunk 0
v4i32 kv_next = {0,0,0,0};
int sc_next = 0;
if (tok_valid) {
kv_next = *reinterpret_cast<const v4i32*>(kv_lds + tok_kv_base);
sc_next = (int)sc_lds[tok_sc_base];
}
#pragma unroll
for (int d = 0; d < FP4_K_CHUNKS; d++) {
v4i32 kv = kv_next;
int sc = sc_next;
// Score MFMA first (fires on matrix core immediately)
scores = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
q_v8[d], to_v8(kv),
scores, 4, 4, 0, q_scale_cached[d], 0, sc);
// Prefetch chunk d+1 (LDS read during MFMA + dequant)
if (d + 1 < FP4_K_CHUNKS) {
const int next_dim = (d + 1) * (FP4_K / 2);
if (tok_valid && next_dim + (k_group << 4) + 16 <= KV_PACKED) {
kv_next = *reinterpret_cast<const v4i32*>(kv_lds + tok_kv_base + next_dim);
sc_next = (int)sc_lds[tok_sc_base + (d + 1) * (FP4_K / 32)];
} else {
kv_next = {0,0,0,0}; sc_next = 0;
}
}
// V dequant: fp4 → f16 → bf8 (VALU, overlaps with MFMA on matrix core)
if (d < V_K_CHUNKS) {
uint8_t* bf8_dst = v_fp8_lds + tok * V_FP8_STRIDE + d * FP4_K + (k_group << 5);
float block_sf = tok_valid ? __uint_as_float((uint32_t)sc << 23) : 0.0f;
ev_short2 ztmp = {0, 0};
#pragma unroll
for (int w = 0; w < 4; w++) {
uint32_t pk = (uint32_t)kv[w];
f16x2 h0 = __builtin_amdgcn_cvt_scalef32_pk_f16_fp4(pk, block_sf, 0);
f16x2 h1 = __builtin_amdgcn_cvt_scalef32_pk_f16_fp4(pk, block_sf, 1);
f16x2 h2 = __builtin_amdgcn_cvt_scalef32_pk_f16_fp4(pk, block_sf, 2);
f16x2 h3 = __builtin_amdgcn_cvt_scalef32_pk_f16_fp4(pk, block_sf, 3);
ev_short2 b01 = __builtin_amdgcn_cvt_scalef32_pk_bf8_f16(ztmp, h0, 1.0f, false);
ev_short2 b0123 = __builtin_amdgcn_cvt_scalef32_pk_bf8_f16(b01, h1, 1.0f, true);
ev_short2 b45 = __builtin_amdgcn_cvt_scalef32_pk_bf8_f16(ztmp, h2, 1.0f, false);
ev_short2 b4567 = __builtin_amdgcn_cvt_scalef32_pk_bf8_f16(b45, h3, 1.0f, true);
*reinterpret_cast<ev_short2*>(bf8_dst + w * 8) = b0123;
*reinterpret_cast<ev_short2*>(bf8_dst + w * 8 + 4) = b4567;
}
}
}
// ---- Register-based softmax: no score_lds ----
if (!tok_valid) { scores[0] = scores[1] = scores[2] = scores[3] = -INFINITY; }
else { scores[0] *= sm_log2e; scores[1] *= sm_log2e; scores[2] *= sm_log2e; scores[3] *= sm_log2e; }
// Warp-local max over 16 tokens (DPP mirror butterfly within k_group row)
float warp_max[HEADS_PER_LANE];
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++)
warp_max[h] = dpp_row_max(scores[h]);
// Warp-local exp + sum (DPP mirror butterfly, no * LOG2E — already in log2 space)
float attn_w[HEADS_PER_LANE], warp_sum[HEADS_PER_LANE];
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++) {
attn_w[h] = __builtin_amdgcn_exp2f(scores[h] - warp_max[h]);
warp_sum[h] = dpp_row_sum(attn_w[h]);
}
// Vectorized LDS write: 4 heads at once (float4)
float* scratch = reinterpret_cast<float*>(smem + OFF_SCORE);
if (lane_mod16 == 0) {
*reinterpret_cast<f32x4*>(&scratch[warp_id * NUM_HEADS + (k_group << 2)]) =
(f32x4){warp_max[0], warp_max[1], warp_max[2], warp_max[3]};
*reinterpret_cast<f32x4*>(&scratch[NUM_WARPS * NUM_HEADS + warp_id * NUM_HEADS + (k_group << 2)]) =
(f32x4){warp_sum[0], warp_sum[1], warp_sum[2], warp_sum[3]};
}
__syncthreads();
// Vectorized LDS reads: global max + corrected sum (float4 per warp)
float my_tm[HEADS_PER_LANE] = {-INFINITY, -INFINITY, -INFINITY, -INFINITY};
float my_ts[HEADS_PER_LANE] = {0, 0, 0, 0};
#pragma unroll
for (int w = 0; w < NUM_WARPS; w++) {
f32x4 wm = *reinterpret_cast<const f32x4*>(&scratch[w * NUM_HEADS + (k_group << 2)]);
my_tm[0] = fmaxf(my_tm[0], wm[0]); my_tm[1] = fmaxf(my_tm[1], wm[1]);
my_tm[2] = fmaxf(my_tm[2], wm[2]); my_tm[3] = fmaxf(my_tm[3], wm[3]);
}
#pragma unroll
for (int w = 0; w < NUM_WARPS; w++) {
f32x4 wm = *reinterpret_cast<const f32x4*>(&scratch[w * NUM_HEADS + (k_group << 2)]);
f32x4 ws = *reinterpret_cast<const f32x4*>(&scratch[NUM_WARPS * NUM_HEADS + w * NUM_HEADS + (k_group << 2)]);
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++)
my_ts[h] += ws[h] * __builtin_amdgcn_exp2f(wm[h] - my_tm[h]);
}
// Correct attn weights with global max + write fp8 to LDS
{
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++)
attn_w[h] *= __builtin_amdgcn_exp2f(warp_max[h] - my_tm[h]);
ev_short2 tmp = {0, 0};
ev_short2 fp8_lo = __builtin_amdgcn_cvt_scalef32_pk_fp8_f32(tmp, attn_w[0], attn_w[1], 1.0f, false);
ev_short2 fp8_all = __builtin_amdgcn_cvt_scalef32_pk_fp8_f32(fp8_lo, attn_w[2], attn_w[3], 1.0f, true);
int pad = ((tok & 16) >> 1) + ((tok & 32) << 2);
*reinterpret_cast<ev_short2*>(&attn_fp8_lds[tok * NUM_HEADS + pad + (k_group << 2)]) = fp8_all;
}
__syncthreads();
// ---- V multiply: fp8 attn × fp8 V, using transpose reads ----
// A = attn fp8 [16 heads × 128 tokens], B = V fp8 [16 dims × 128 tokens]
// 4 MFMAs per warp (512 dims / 16 per MFMA / 8 warps = 4 rounds)
f32x4 v_out[4];
#pragma unroll
for (int i = 0; i < 4; i++) v_out[i] = {0,0,0,0};
{
const int half = lane_mod16 & 1;
const int tok_in_grp = lane_mod16 >> 1;
const int half8 = half << 3;
const int kg16 = k_group << 4;
int tok_arr[4];
tok_arr[0] = kg16 + tok_in_grp;
tok_arr[1] = kg16 + 8 + tok_in_grp;
tok_arr[2] = 64 + kg16 + tok_in_grp;
tok_arr[3] = 64 + kg16 + 8 + tok_in_grp;
// Load A (attn fp8): 4 transpose reads, reused across all rounds
const uintptr_t attn_base = reinterpret_cast<uintptr_t>(attn_fp8_lds) + half8;
v8i32 a_reg;
{
v2i32* a_parts = reinterpret_cast<v2i32*>(&a_reg);
#pragma unroll
for (int c = 0; c < 4; c++) {
int t_c = tok_arr[c];
int pad_c = ((t_c & 16) >> 1) + ((t_c & 32) << 2);
a_parts[c] = __builtin_amdgcn_ds_read_tr8_b64_v2i32(
reinterpret_cast<as3_v2i32_ptr>(attn_base + t_c * NUM_HEADS + pad_c));
}
}
// V fp8: double-buffered B reads to hide LDS latency behind MFMA
const uintptr_t vfp8_base = reinterpret_cast<uintptr_t>(v_fp8_lds) + half8;
v8i32 b_reg[2];
// Pre-load round 0
{
const uintptr_t vb0 = vfp8_base + ((0 * 8 + warp_id) << 4);
v2i32* bp = reinterpret_cast<v2i32*>(&b_reg[0]);
#pragma unroll
for (int c = 0; c < 4; c++)
bp[c] = __builtin_amdgcn_ds_read_tr8_b64_v2i32(
reinterpret_cast<as3_v2i32_ptr>(vb0 + tok_arr[c] * V_FP8_STRIDE));
}
#pragma unroll
for (int r = 0; r < 4; r++) {
int cur = r & 1, nxt = (r + 1) & 1;
// Prefetch next round while MFMA executes current
if (r + 1 < 4) {
const uintptr_t vb_next = vfp8_base + (((r + 1) * 8 + warp_id) << 4);
v2i32* bp = reinterpret_cast<v2i32*>(&b_reg[nxt]);
#pragma unroll
for (int c = 0; c < 4; c++)
bp[c] = __builtin_amdgcn_ds_read_tr8_b64_v2i32(
reinterpret_cast<as3_v2i32_ptr>(vb_next + tok_arr[c] * V_FP8_STRIDE));
}
v_out[r] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a_reg, b_reg[cur], v_out[r], 0, 1, 0, 127, 0, 127);
}
}
// ---- Merge ----
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++) {
float nm = fmaxf(running_max[h], my_tm[h]);
float co = __builtin_amdgcn_exp2f(running_max[h] - nm);
float cn = __builtin_amdgcn_exp2f(my_tm[h] - nm);
running_sum[h] = running_sum[h] * co + my_ts[h] * cn;
running_max[h] = nm;
#pragma unroll
for (int i = 0; i < 4; i++)
running_v[i][h] = running_v[i][h] * co + v_out[i][h] * cn;
}
}
// ---- Write output ----
if (NUM_PARTIALS == 1) {
// Direct output: divide by sum, write final result (skip reduce kernel)
__hip_bfloat16* ob = final_output + (q_start * NUM_HEADS) * V_HEAD_DIM;
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++) {
float inv = (running_sum[h] > 0.0f) ? (1.0f / running_sum[h]) : 0.0f;
#pragma unroll
for (int i = 0; i < 4; i++) {
const int vd = (i*8 + warp_id)*16 + lane_mod16;
ob[((k_group<<2)+h)*V_HEAD_DIM + vd] = __float2bfloat16(running_v[i][h] * inv);
}
}
} else {
// Write partial for reduce kernel
const int pidx = batch_idx * NUM_PARTIALS + partial_idx;
__hip_bfloat16* ob = partial_out + pidx * (NUM_HEADS * V_HEAD_DIM);
#pragma unroll
for (int i = 0; i < 4; i++) {
const int vd = (i*8 + warp_id)*16 + lane_mod16;
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++)
ob[((k_group<<2)+h)*V_HEAD_DIM + vd] = __float2bfloat16(running_v[i][h]);
}
if (lane_mod16 == 0 && warp_id == 0) {
float* pm = partial_max + pidx * NUM_HEADS;
float* ps = partial_sum + pidx * NUM_HEADS;
#pragma unroll
for (int h = 0; h < HEADS_PER_LANE; h++) {
pm[(k_group<<2)+h] = running_max[h];
ps[(k_group<<2)+h] = running_sum[h];
}
}
}
}
// ============================================================
// Reduce kernel
// ============================================================
template<int NUM_PARTIALS>
__global__ void __launch_bounds__(WAVEFRONT_SIZE)
mla_reduce(
const __hip_bfloat16* __restrict__ partial_out,
const float* __restrict__ partial_max,
const float* __restrict__ partial_sum,
const int32_t* __restrict__ qo_indptr,
__hip_bfloat16* __restrict__ output
) {
const int bid = blockIdx.x, batch_idx = bid / NUM_HEADS, head_idx = bid - batch_idx * NUM_HEADS;
const int lane_id = threadIdx.x & 63, q_start = qo_indptr[batch_idx];
constexpr int ELEMS = V_HEAD_DIM / WAVEFRONT_SIZE; // 8
constexpr int PV_STRIDE = NUM_HEADS * V_HEAD_DIM;
typedef float ext_f32x4 __attribute__((ext_vector_type(4)));
typedef uint32_t ext_u32x4 __attribute__((ext_vector_type(4)));
float local_max = -INFINITY, local_sum = 0;
ext_f32x4 v_acc0 = {0,0,0,0}, v_acc1 = {0,0,0,0};
const int bo = batch_idx * NUM_PARTIALS;
const float* pm = partial_max + bo * NUM_HEADS + head_idx;
const float* ps = partial_sum + bo * NUM_HEADS + head_idx;
const __hip_bfloat16* pv_base = partial_out + bo * NUM_HEADS * V_HEAD_DIM
+ head_idx * V_HEAD_DIM + lane_id * ELEMS;
#pragma unroll
for (int p = 0; p < NUM_PARTIALS; p++) {
// Scalar loads for wave-uniform values (SALU path)
float cm = __builtin_amdgcn_readfirstlane(pm[p * NUM_HEADS]);
float cs = __builtin_amdgcn_readfirstlane(ps[p * NUM_HEADS]);
const __hip_bfloat16* pv = pv_base + p * PV_STRIDE;
// Vectorized bf16 load → native bf16x2→float2 conversion
const __hip_bfloat162* bp = reinterpret_cast<const __hip_bfloat162*>(pv);
float2 f01 = __bfloat1622float2(bp[0]);
float2 f23 = __bfloat1622float2(bp[1]);
float2 f45 = __bfloat1622float2(bp[2]);
float2 f67 = __bfloat1622float2(bp[3]);
ext_f32x4 v0 = {f01.x, f01.y, f23.x, f23.y};
ext_f32x4 v1 = {f45.x, f45.y, f67.x, f67.y};
// Wave-uniform branch (SALU)
if (cm >= local_max) {
float co = __builtin_amdgcn_exp2f(local_max - cm);
local_sum = local_sum * co + cs;
v_acc0 = v_acc0 * co + v0;
v_acc1 = v_acc1 * co + v1;
local_max = cm;
} else {
float cn = __builtin_amdgcn_exp2f(cm - local_max);
local_sum += cs * cn;
v_acc0 += v0 * cn;
v_acc1 += v1 * cn;
}
}
float inv = (local_sum > 0.0f) ? __builtin_amdgcn_rcpf(local_sum) : 0.0f;
v_acc0 *= inv;
v_acc1 *= inv;
// Pack f32 → bf16 via v_cvt_pk_bf16_f32 (2 f32 → 1 packed bf16x2 per instruction)
ext_u32x4 out_pack;
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[0]) : "v"(v_acc0[0]), "v"(v_acc0[1]));
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[1]) : "v"(v_acc0[2]), "v"(v_acc0[3]));
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[2]) : "v"(v_acc1[0]), "v"(v_acc1[1]));
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[3]) : "v"(v_acc1[2]), "v"(v_acc1[3]));
*reinterpret_cast<ext_u32x4*>(output + (q_start * NUM_HEADS + head_idx) * V_HEAD_DIM + lane_id * ELEMS) = out_pack;
}
// ============================================================
// Multi-wavefront reduce for high-partial shapes (e.g. 64 partials)
// Each wavefront handles NUM_PARTIALS/NWARPS partials, then LDS merge
// ============================================================
template<int NUM_PARTIALS, int NWARPS>
__global__ void __launch_bounds__(NWARPS * WAVEFRONT_SIZE)
mla_reduce_multi(
const __hip_bfloat16* __restrict__ partial_out,
const float* __restrict__ partial_max,
const float* __restrict__ partial_sum,
const int32_t* __restrict__ qo_indptr,
__hip_bfloat16* __restrict__ output
) {
const int bid = blockIdx.x, batch_idx = bid / NUM_HEADS, head_idx = bid - batch_idx * NUM_HEADS;
const int warp_id = threadIdx.x >> 6;
const int lane_id = threadIdx.x & 63, q_start = qo_indptr[batch_idx];
constexpr int ELEMS = V_HEAD_DIM / WAVEFRONT_SIZE;
constexpr int PV_STRIDE = NUM_HEADS * V_HEAD_DIM;
constexpr int PARTIALS_PER_WARP = NUM_PARTIALS / NWARPS;
typedef float ext_f32x4 __attribute__((ext_vector_type(4)));
typedef uint32_t ext_u32x4 __attribute__((ext_vector_type(4)));
float local_max = -INFINITY, local_sum = 0;
ext_f32x4 v_acc0 = {0,0,0,0}, v_acc1 = {0,0,0,0};
const int bo = batch_idx * NUM_PARTIALS;
const float* pm = partial_max + bo * NUM_HEADS + head_idx;
const float* ps = partial_sum + bo * NUM_HEADS + head_idx;
const __hip_bfloat16* pv_base = partial_out + bo * NUM_HEADS * V_HEAD_DIM
+ head_idx * V_HEAD_DIM + lane_id * ELEMS;
// Each warp processes its slice of partials
const int p_start = warp_id * PARTIALS_PER_WARP;
#pragma unroll
for (int p = p_start; p < p_start + PARTIALS_PER_WARP; p++) {
float cm = __builtin_amdgcn_readfirstlane(pm[p * NUM_HEADS]);
float cs = __builtin_amdgcn_readfirstlane(ps[p * NUM_HEADS]);
const __hip_bfloat162* bp = reinterpret_cast<const __hip_bfloat162*>(pv_base + p * PV_STRIDE);
float2 f01 = __bfloat1622float2(bp[0]);
float2 f23 = __bfloat1622float2(bp[1]);
float2 f45 = __bfloat1622float2(bp[2]);
float2 f67 = __bfloat1622float2(bp[3]);
ext_f32x4 v0 = {f01.x, f01.y, f23.x, f23.y};
ext_f32x4 v1 = {f45.x, f45.y, f67.x, f67.y};
if (cm >= local_max) {
float co = __builtin_amdgcn_exp2f(local_max - cm);
local_sum = local_sum * co + cs;
v_acc0 = v_acc0 * co + v0;
v_acc1 = v_acc1 * co + v1;
local_max = cm;
} else {
float cn = __builtin_amdgcn_exp2f(cm - local_max);
local_sum += cs * cn;
v_acc0 += v0 * cn;
v_acc1 += v1 * cn;
}
}
// === Cooperative pre-scaled LDS reduction ===
extern __shared__ char reduce_smem[];
float* smem_max = reinterpret_cast<float*>(reduce_smem);
float* smem_sum = smem_max + NWARPS;
float* smem_v = smem_sum + NWARPS;
// Step 1: Write max/sum (predicated — only lane 0, avoids 64-way bank conflict)
if (lane_id == 0) {
smem_max[warp_id] = local_max;
smem_sum[warp_id] = local_sum;
}
__syncthreads();
// Step 2: ALL warps compute global max (SALU via readfirstlane)
float global_max = -INFINITY;
#pragma unroll
for (int w = 0; w < NWARPS; w++)
global_max = fmaxf(global_max, __builtin_amdgcn_readfirstlane(smem_max[w]));
// Step 3: ALL warps pre-scale their own vectors + sum
float scale = __builtin_amdgcn_exp2f(local_max - global_max);
v_acc0 *= scale;
v_acc1 *= scale;
local_sum *= scale;
// Step 4: Write pre-scaled vectors + sum to LDS
float* wv_dst = &smem_v[warp_id * V_HEAD_DIM + lane_id * ELEMS];
*reinterpret_cast<ext_f32x4*>(wv_dst) = v_acc0;
*reinterpret_cast<ext_f32x4*>(wv_dst + 4) = v_acc1;
if (lane_id == 0) smem_sum[warp_id] = local_sum;
__syncthreads();
// Step 5: Warp 0 does pure branchless vector addition (no exp2f)
if (warp_id == 0) {
ext_f32x4 gv0 = {0,0,0,0}, gv1 = {0,0,0,0};
float gs = 0;
#pragma unroll
for (int w = 0; w < NWARPS; w++) {
const float* wv_base = &smem_v[w * V_HEAD_DIM + lane_id * ELEMS];
gv0 += *reinterpret_cast<const ext_f32x4*>(wv_base);
gv1 += *reinterpret_cast<const ext_f32x4*>(wv_base + 4);
gs += __builtin_amdgcn_readfirstlane(smem_sum[w]);
}
float inv = (gs > 0.0f) ? __builtin_amdgcn_rcpf(gs) : 0.0f;
gv0 *= inv;
gv1 *= inv;
ext_u32x4 out_pack;
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[0]) : "v"(gv0[0]), "v"(gv0[1]));
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[1]) : "v"(gv0[2]), "v"(gv0[3]));
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[2]) : "v"(gv1[0]), "v"(gv1[1]));
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(out_pack[3]) : "v"(gv1[2]), "v"(gv1[3]));
*reinterpret_cast<ext_u32x4*>(output + (q_start * NUM_HEADS + head_idx) * V_HEAD_DIM + lane_id * ELEMS) = out_pack;
}
}
// ============================================================
// Dispatch
// ============================================================
template<int NUM_PARTIALS, int TILES_PER_BLOCK>
void run_impl(uintptr_t q_ptr, uintptr_t kv_data_ptr, uintptr_t kv_scales_ptr,
uintptr_t qo_indptr_ptr, uintptr_t kv_indptr_ptr, uintptr_t output_ptr,
uintptr_t pout_ptr, uintptr_t pmax_ptr, uintptr_t psum_ptr,
float sm_scale, int batch_size) {
const auto* q_d = reinterpret_cast<const __hip_bfloat16*>(q_ptr);
const auto* kvd = reinterpret_cast<const uint8_t*>(kv_data_ptr);
const auto* kvs = reinterpret_cast<const uint8_t*>(kv_scales_ptr);
const auto* qoi = reinterpret_cast<const int32_t*>(qo_indptr_ptr);
const auto* kvi = reinterpret_cast<const int32_t*>(kv_indptr_ptr);
auto* out = reinterpret_cast<__hip_bfloat16*>(output_ptr);
auto* pout = reinterpret_cast<__hip_bfloat16*>(pout_ptr);
auto* pmax = reinterpret_cast<float*>(pmax_ptr);
auto* psum = reinterpret_cast<float*>(psum_ptr);
const int tb = batch_size * NUM_PARTIALS;
hipLaunchKernelGGL((mla_grouped<NUM_PARTIALS, TILES_PER_BLOCK>),
dim3(tb), dim3(BLOCK_THREADS), TOTAL_LDS, 0,
q_d, kvd, kvs, qoi, kvi, pout, pmax, psum, out, sm_scale);
if (NUM_PARTIALS > 1) {
if (NUM_PARTIALS >= 8) {
// Multi-wavefront reduce: 4 warps for 64 partials, 2 warps for 8
constexpr int REDUCE_WARPS = (NUM_PARTIALS >= 64) ? 16 : 2;
constexpr int REDUCE_THREADS = REDUCE_WARPS * WAVEFRONT_SIZE;
constexpr int REDUCE_LDS = REDUCE_WARPS * (V_HEAD_DIM * 4 + 2 * 4);
hipLaunchKernelGGL((mla_reduce_multi<NUM_PARTIALS, REDUCE_WARPS>),
dim3(batch_size * NUM_HEADS), dim3(REDUCE_THREADS), REDUCE_LDS, 0,
pout, pmax, psum, qoi, out);
} else {
hipLaunchKernelGGL((mla_reduce<NUM_PARTIALS>),
dim3(batch_size * NUM_HEADS), dim3(WAVEFRONT_SIZE), 0, 0,
pout, pmax, psum, qoi, out);
}
}
}
void run_s1(uintptr_t a, uintptr_t b, uintptr_t c, uintptr_t d, uintptr_t e, uintptr_t f,
uintptr_t po, uintptr_t pm, uintptr_t ps,
float g, int i, int max_kvseqlen) {
const int total_tiles = max_kvseqlen / 128;
const int tpb = max(1, i * total_tiles / 256);
const int np = total_tiles / tpb;
if (np == 1 && tpb == 8) run_impl<1, 8>(a, b, c, d, e, f, po, pm, ps, g, i);
else if (np == 1 && tpb == 64) run_impl<1, 64>(a, b, c, d, e, f, po, pm, ps, g, i);
else if (np == 4 && tpb == 2) run_impl<4, 2>(a, b, c, d, e, f, po, pm, ps, g, i);
else if (np == 4 && tpb == 16) run_impl<4, 16>(a, b, c, d, e, f, po, pm, ps, g, i);
else if (np == 8 && tpb == 1) run_impl<8, 1>(a, b, c, d, e, f, po, pm, ps, g, i);
else if (np == 8 && tpb == 8) run_impl<8, 8>(a, b, c, d, e, f, po, pm, ps, g, i);
else if (np == 64 && tpb == 1) run_impl<64, 1>(a, b, c, d, e, f, po, pm, ps, g, i);
else run_impl<1, 8>(a, b, c, d, e, f, po, pm, ps, g, i);
}
PYBIND11_MODULE(mla_new_kernel_module, m) { m.def("run_s1", &run_s1); }
"""
hip_module = load_inline(
name="mla_new_kernel_module",
cpp_sources="",
cuda_sources=CUDA_SRC,
with_cuda=True,
verbose=True,
extra_cuda_cflags=["-std=c++20", "-O3"],
no_implicit_headers=True,
)
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
kv_buffer, kv_scales = kv_data["mxfp4"]
total_q_len = q.size(0)
max_kvseqlen = kv_buffer.size(0) // batch_size
total_tiles = max_kvseqlen // 128
tpb = max(1, batch_size * total_tiles // 256)
np_ = total_tiles // tpb
tb = batch_size * np_
pout = torch.empty(tb, NUM_HEADS, V_HEAD_DIM, dtype=torch.bfloat16, device=q.device)
pmax = torch.empty(tb, NUM_HEADS, dtype=torch.float32, device=q.device)
psum = torch.empty(tb, NUM_HEADS, dtype=torch.float32, device=q.device)
output = torch.empty(total_q_len, NUM_HEADS, V_HEAD_DIM, dtype=torch.bfloat16, device=q.device)
hip_module.run_s1(
q.data_ptr(), kv_buffer.data_ptr(), kv_scales.data_ptr(),
qo_indptr.data_ptr(), kv_indptr.data_ptr(), output.data_ptr(),
pout.data_ptr(), pmax.data_ptr(), psum.data_ptr(),
SM_SCALE, batch_size, max_kvseqlen)
return output
scrolls · 770 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON