submission 692061
npip99 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1616 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-692061?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:17d80bc7a2b400a1de62748b6c7255cf3353b7d288241c3e2851ea1ec78e4933
license declaredunknown
license concludedunknown
authorsnpip99
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
kv_data["mxfp4"][0].data_ptr(),num-warps = 4
constexpr u32 NUM_WARPS = 4;shared-memory
__shared__ union {Kernel source
submission.py1616 lines
# See {filename}.hip for details
import os
import torch
from torch.utils.cpp_extension import load_inline
from typing import Any
from task import input_t, output_t
if "PYTORCH_ROCM_ARCH" not in os.environ:
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950:xnack-"
CUDA_PRELUDE = """
#ifndef __gfx950__
#define __gfx950__
#endif
"""
CPP_WRAPPER = """
void entry(
const uintptr_t q, // .............. (bs, 16, 576) bf16
const uintptr_t kv_indptr, // ...... (bs+1,) int32
const uintptr_t kv_bf16, // ........ (total_kv, 576) bf16
const uintptr_t kv_fp8, // ......... (total_kv, 576) fp8
const uintptr_t kv_fp8_scale, // ... (1,) f32
const uintptr_t kv_mxfp4, // ....... (total_kv, 288) fp4x2
const uintptr_t kv_mxfp4_scale, // . (total_kv, 18) uint8_t
uintptr_t out, // .................. (bs, 16, 512) bf16
uint32_t bs
);
"""
CUDA_SRC = r"""
#define CDIV(a, b) (((a) + (b) - 1) / (b))
using u8 = uint8_t;
using u16 = uint16_t;
using u16x32 = u16 __attribute__((ext_vector_type(32)));
using u32 = uint32_t;
using u32x4 = u32 __attribute__((ext_vector_type(4)));
using u32x8 = u32 __attribute__((ext_vector_type(8)));
using u64 = uint64_t;
using f32 = float;
using f32x4 = f32 __attribute__((ext_vector_type(4)));
using bf16x2 = __bf16 __attribute__((ext_vector_type(2)));
using bf16x4 = __bf16 __attribute__((ext_vector_type(4)));
using bf16x8 = __bf16 __attribute__((ext_vector_type(8)));
__device__ inline f32 fast_exp(f32 x) {
constexpr f32 LOG2E = 1.4426950408889634f;
return __builtin_amdgcn_exp2f(x * LOG2E);
}
__device__ inline u16 f32_to_bf16(f32 v) {
return (u16)(__builtin_bit_cast(u32, v) >> 16);
}
__device__ inline float bf16_to_f32(uint16_t v) {
return __builtin_bit_cast(float, (uint32_t)v << 16);
}
// Quantize a group of 32 bf16 values -> E8M0 scale (biased uint8)
__device__ inline uint8_t bf16x32_to_scale_e8m0(const u16x32* row) {
uint32_t amax_u16x2_packed = 0;
const uint32_t* row32 = (const uint32_t*)row;
#pragma unroll
for (int i = 0; i < 16; i++) {
// abs both bf16 in one instruction
uint32_t v = row32[i] & 0x7FFF7FFFu;
// max both bf16 in one instruction
asm("v_pk_max_u16 %0, %0, %1" : "+v"(amax_u16x2_packed) : "v"(v));
}
uint16_t amax_u16 = max((uint16_t)amax_u16x2_packed, (uint16_t)(amax_u16x2_packed >> 16));
// Round exponent up when mantissa fraction >= 0.75
// (bf16 has 7 mantissa bits; 0x20 = 2^5 matches the f32 0x200000 = 2^21 rounding)
int biased_exp = ((amax_u16 + 0x20u) >> 7) & 0xFF;
int e8m0 = biased_exp - 2;
return (uint8_t)max(0, min(254, e8m0));
}
// Matches `fp4_utils.e8m0_to_f32`: scale_e8m0 << 23 as f32, with special cases
__device__ inline float e8m0_to_f32(uint8_t scale) {
return __builtin_bit_cast(float, (uint32_t)scale << 23);
}
__device__ inline __bf16 fp4_e2m1_to_bf16(u8 nibble, u8 scale) {
u8 s = (nibble >> 3) & 1;
u8 e = (nibble >> 1) & 3;
u8 m = nibble & 1;
f32 val = (e == 0) ? m * 0.5f : exp2f((f32)e - 1.0f) * (1.0f + m * 0.5f);
val = s ? -val : val;
u16 ret = f32_to_bf16(val * e8m0_to_f32(scale));
return *(const __bf16*)&ret;
}
__device__ inline uint8_t f32_to_fp4_e2m1(float v) {
// CAREFUL: We need to include the fix from https://github.com/ROCm/aiter/pull/975 / https://github.com/ROCm/aiter/pull/2249
// FP4 e2m1: EBITS=2, MBITS=1, EXP_BIAS=1
// denorm_exp = (127-1) + (23-1) + 1 = 149 => denorm_mask = 149 << 23 = 0x4A800000
// val_to_add = ((1-127) << 23) + ((1<<21)-1) = 0xC11FFFFF (int32 wrapping)
constexpr float FP4_MAX_NORMAL = 6.0f;
constexpr float FP4_MIN_NORMAL = 1.0f;
constexpr uint32_t DENORM_MASK_INT = 0x4A800000u;
constexpr float DENORM_MASK_FLOAT = __builtin_bit_cast(float, DENORM_MASK_INT);
constexpr uint32_t VAL_TO_ADD = 0xC11FFFFFu;
uint8_t s_bit = (uint8_t)((__builtin_bit_cast(uint32_t, v) >> 28) & 0x8u);
uint32_t abs_bits = __builtin_bit_cast(uint32_t, v) & 0x7FFFFFFFu;
float abs_v = __builtin_bit_cast(float, abs_bits);
uint8_t result;
if (abs_v >= FP4_MAX_NORMAL) {
// saturate branch
result = 0x7u;
} else if (abs_v < FP4_MIN_NORMAL) {
// denormal branch: float trick to extract low bits
uint32_t d = __builtin_bit_cast(uint32_t, abs_v + DENORM_MASK_FLOAT) - DENORM_MASK_INT;
result = (uint8_t)d;
} else {
// normal branch: bias-adjust exponent + round-to-nearest
uint32_t x = abs_bits;
uint32_t m_odd = (x >> 22) & 1u;
x = x + VAL_TO_ADD + m_odd;
result = (uint8_t)(x >> 22);
}
return s_bit | result;
}
// Take an f32 and its block's scaling term, and convert to fp4_e2m1
// The fp4_e2m1 will be in the low nibble of the return value
__device__ inline uint8_t f32_to_fp4_e2m1_scale(float v, uint8_t scale) {
float f32_scale = e8m0_to_f32((uint8_t)(254u - scale)); // 2^(-(exp_unb)) = 2^(127 - scale)
return f32_to_fp4_e2m1(v * f32_scale);
}
__device__ inline f32x4 mfma_bf16_16x16x32(bf16x8 a, bf16x8 b, f32x4 acc) {
#ifdef __gfx950__
// return __builtin_amdgcn_mfma_f32_16x16x32_bf16(a, b, acc, 0, 0, 0);
// AGPRs are bad, we can't multiply alpha into them. Pin VGPR.
asm("v_mfma_f32_16x16x32_bf16 %0, %1, %2, %0"
: "+v"(acc)
: "v"(a), "v"(b));
return acc;
#else
acc = __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(*(bf16x4*)&a, *(bf16x4*)&b, acc, 0, 0, 0);
acc = __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(*((bf16x4*)&a+1), *((bf16x4*)&b+1), acc, 0, 0, 0);
return acc;
#endif
}
__device__ inline void global_to_lds_1024b(const u8* global_base, u32 lds_base, u32 lane_id) {
#ifdef __gfx950__
asm volatile(
"s_mov_b32 m0, %0\n\t"
"global_load_lds_dwordx4 %1, off\n\t"
:: "s"(lds_base), "v"((const void*)(global_base + lane_id * 16))
: "memory", "m0"
);
#else
#pragma unroll
for (u32 sub = 0; sub < 4; sub++) {
asm volatile(
"s_mov_b32 m0, %0\n\t"
"global_load_lds_dword %1, off\n\t"
:: "s"(lds_base + sub * 256), "v"((const void*)(global_base + sub * 256 + lane_id * 4))
: "memory", "m0"
);
}
#endif
}
__device__ inline void global_to_lds_256b(const u8* global_base, u32 lds_base, u32 lane_id) {
asm volatile(
"s_mov_b32 m0, %0\n\t"
"global_load_lds_dword %1, off\n\t"
:: "s"(lds_base), "v"((const void*)(global_base + lane_id * 4))
: "memory", "m0"
);
}
__device__ inline f32 dpp_reduce_max_16(f32 v) {
#define DPP_MOV(x, ctrl) __builtin_bit_cast(f32, \
__builtin_amdgcn_mov_dpp(__builtin_bit_cast(u32, x), ctrl, 0xF, 0xF, false))
v = fmaxf(v, DPP_MOV(v, 0xB1)); // quad XOR 1
v = fmaxf(v, DPP_MOV(v, 0x4E)); // quad XOR 2
v = fmaxf(v, DPP_MOV(v, 0x124)); // row_ror:4
v = fmaxf(v, DPP_MOV(v, 0x128)); // row_ror:8
return v;
#undef DPP_MOV
}
__device__ inline f32 dpp_reduce_sum_16(f32 v) {
#define DPP_MOV(x, ctrl) __builtin_bit_cast(f32, \
__builtin_amdgcn_mov_dpp(__builtin_bit_cast(u32, x), ctrl, 0xF, 0xF, false))
v += DPP_MOV(v, 0xB1);
v += DPP_MOV(v, 0x4E);
v += DPP_MOV(v, 0x124);
v += DPP_MOV(v, 0x128);
return v;
#undef DPP_MOV
}
// Constants
constexpr f32 NEG_INF = -1e30f;
constexpr u32 THREADS_PER_WARP = 64;
#ifdef __gfx950__
constexpr u32 NUM_WARPS = 4;
#else
constexpr u32 NUM_WARPS = 2;
#endif
constexpr u32 THREADS_PER_BLOCK = THREADS_PER_WARP * NUM_WARPS;
constexpr u32 SCALE_GROUP_SIZE = 32;
// DeepSeek Parameters
constexpr u32 N_HEADS = 16;
constexpr u32 QK_HEAD_DIM = 576;
constexpr u32 V_HEAD_DIM = 512;
__global__ __launch_bounds__(THREADS_PER_BLOCK, 1)
void kernel(
const u16* __restrict__ q,
const u32* __restrict__ kv_indptr,
const u16* __restrict__ kv_bf16,
const u8* __restrict__ kv_fp8, const f32* __restrict__ kv_fp8_scale,
const u8* __restrict__ kv_mxfp4, const u8* __restrict__ kv_mxfp4_scale,
u16* __restrict__ out,
int bs
) {
u32 batch_idx = __builtin_amdgcn_readfirstlane(blockIdx.x);
u32 warp_id = __builtin_amdgcn_readfirstlane(threadIdx.x / THREADS_PER_WARP);
u32 lane_id = threadIdx.x % THREADS_PER_WARP;
u32 lane_col = lane_id % 16;
u32 lane_rowgroup = lane_id / 16; // rows \in 4*lane_rowgroup + {0,1,2,3}
// Pre-load Q into registers (quantization deferred to after LDS prefetch)
const u16* q_batch_item = q + batch_idx * N_HEADS * QK_HEAD_DIM;
constexpr u32 SCORE_MFMA_DOT_DIM = 128;
constexpr u32 SCORE_MFMA_ITERS = CDIV(QK_HEAD_DIM, SCORE_MFMA_DOT_DIM);
u32x4 q_data[SCORE_MFMA_ITERS][4];
#pragma unroll
for (u32 k = 0; k < SCORE_MFMA_ITERS; k++) {
u32 dim_base = k * SCORE_MFMA_DOT_DIM + lane_rowgroup * SCALE_GROUP_SIZE;
const u16* q_chunk = q_batch_item + lane_col * QK_HEAD_DIM + dim_base;
q_data[k][0] = *(const u32x4*)(q_chunk);
q_data[k][1] = *(const u32x4*)(q_chunk + 8);
q_data[k][2] = *(const u32x4*)(q_chunk + 16);
q_data[k][3] = *(const u32x4*)(q_chunk + 24);
}
u32 kv_start = kv_indptr[batch_idx];
u32 kv_end = kv_indptr[batch_idx + 1];
f32 softmax_max[4] = {NEG_INF, NEG_INF, NEG_INF, NEG_INF};
f32 softmax_denom[4] = {};
// KV Tiles
constexpr u32 KV_TILE_DIM = 32;
constexpr u32 PREFETCH = 2;
constexpr u32 KV_SCALE_STRIDE = CDIV(QK_HEAD_DIM / SCALE_GROUP_SIZE, 8) * 8;
// Shared Memory
__shared__ union {
// kv iterations
struct {
u8 kv_data[NUM_WARPS][PREFETCH][KV_TILE_DIM][QK_HEAD_DIM / 2];
u8 kv_scale[NUM_WARPS][PREFETCH][KV_TILE_DIM][KV_SCALE_STRIDE];
u16 weights[NUM_WARPS][N_HEADS][KV_TILE_DIM];
} tile;
// merge before write
struct {
f32 warp_max[NUM_WARPS][N_HEADS];
f32 warp_denom[NUM_WARPS][N_HEADS];
u16 values[NUM_WARPS][N_HEADS][V_HEAD_DIM];
} merge;
} lds;
// (u32 kv_tile_start = kv_start; kv_tile_start < kv_end; kv_tile_start += KV_TILE_DIM) {
constexpr u32 KV_DATA_BYTES = KV_TILE_DIM * (QK_HEAD_DIM / 2);
constexpr u32 KV_SCALE_BYTES = KV_TILE_DIM * KV_SCALE_STRIDE;
// These stay in SGPRs
u32 lds_data_buf0 = (u32)(uintptr_t)&lds.tile.kv_data[warp_id][0];
u32 lds_data_buf1 = (u32)(uintptr_t)&lds.tile.kv_data[warp_id][1];
u32 lds_scale_buf0 = (u32)(uintptr_t)&lds.tile.kv_scale[warp_id][0];
u32 lds_scale_buf1 = (u32)(uintptr_t)&lds.tile.kv_scale[warp_id][1];
auto load_lds = [&](uint32_t kv_iter) -> void {
u32 kv_tile_start = kv_start + kv_iter * KV_TILE_DIM;
u32 kv_buf = kv_iter % 2;
u32 lds_data_off = kv_buf ? lds_data_buf1 : lds_data_buf0;
u32 lds_scale_off = kv_buf ? lds_scale_buf1 : lds_scale_buf0;
constexpr u32 KV_DATA_U128 = KV_DATA_BYTES / 16;
static_assert(KV_DATA_U128 % THREADS_PER_WARP == 0);
constexpr u32 KV_DATA_ITERS = KV_DATA_U128 / THREADS_PER_WARP;
const u8* kv_data_src = (const u8*)(kv_mxfp4 + (u64)kv_tile_start * (QK_HEAD_DIM / 2));
#pragma unroll
for (u32 i = 0; i < KV_DATA_ITERS; i++) {
global_to_lds_1024b(kv_data_src + i * 1024, lds_data_off + i * 1024, lane_id);
}
constexpr u32 KV_SCALE_U32 = KV_SCALE_BYTES / 4;
constexpr u32 KV_SCALE_ITERS = CDIV(KV_SCALE_U32, THREADS_PER_WARP);
const u8* kv_scale_src = kv_mxfp4_scale + (u64)kv_tile_start * KV_SCALE_STRIDE;
#pragma unroll
for (u32 i = 0; i < KV_SCALE_ITERS; i++) {
global_to_lds_256b(kv_scale_src + i * 256, lds_scale_off + i * 256, lane_id);
}
};
u32 kv_range = kv_end - kv_start;
u32 kv_num_iters = kv_range / KV_TILE_DIM / NUM_WARPS;
if (kv_range % (NUM_WARPS * KV_TILE_DIM) != 0 || kv_num_iters < PREFETCH - 1) {
__builtin_trap();
}
u32 warp_kv_offset = warp_id * kv_num_iters;
#pragma unroll
for (u32 i = 0; i < PREFETCH - 1; i++) {
load_lds(warp_kv_offset + i);
}
constexpr u32 VALUE_MFMA_V_TILE_DIM = 16;
constexpr u32 VALUE_MFMA_DOT_DIM = 32;
constexpr u32 NUM_VALUE_TILES = V_HEAD_DIM / VALUE_MFMA_V_TILE_DIM;
f32x4 value_out_lanes[NUM_VALUE_TILES] = {};
__builtin_amdgcn_sched_barrier(0);
// ========================================
// Distributed Q quantization across warps
// Each warp quantizes k iterations where k % NUM_WARPS == warp_id, shares via LDS
// ========================================
__shared__ u32 q_lds[SCORE_MFMA_ITERS][64][8];
__shared__ u8 q_scale_lds[SCORE_MFMA_ITERS][64];
u32x8 q_lanes[SCORE_MFMA_ITERS];
u8 q_scale_lanes[SCORE_MFMA_ITERS];
#pragma unroll
for (u32 k = 0; k < SCORE_MFMA_ITERS; k++) {
if (k % NUM_WARPS != warp_id) continue;
u32 dim_base = k * SCORE_MFMA_DOT_DIM + lane_rowgroup * SCALE_GROUP_SIZE;
if (k < SCORE_MFMA_ITERS - 1 || dim_base + SCALE_GROUP_SIZE <= QK_HEAD_DIM) {
q_scale_lanes[k] = bf16x32_to_scale_e8m0((const u16x32*)&q_data[k]);
f32 s = e8m0_to_f32(q_scale_lanes[k]);
#pragma unroll
for (u32 r = 0; r < 4; r++) {
#ifdef __gfx950__
const bf16x2* v2 = (const bf16x2*)&q_data[k][r];
q_lanes[k][r] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(0, v2[0], s, 0);
q_lanes[k][r] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q_lanes[k][r], v2[1], s, 1);
q_lanes[k][r] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q_lanes[k][r], v2[2], s, 2);
q_lanes[k][r] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q_lanes[k][r], v2[3], s, 3);
#else
const u16* q_r = (const u16*)&q_data[k][r];
q_lanes[k][r] = 0;
#pragma unroll
for (u32 i = 0; i < 8; i++) {
u8 nibble = f32_to_fp4_e2m1_scale(bf16_to_f32(q_r[i]), q_scale_lanes[k]);
q_lanes[k][r] |= (u32)(nibble & 0xF) << (4 * i);
}
#endif
}
} else {
q_lanes[k] = {};
q_scale_lanes[k] = 0;
}
// Write to LDS
#pragma unroll
for (u32 r = 0; r < 8; r++) {
q_lds[k][lane_id][r] = q_lanes[k][r];
}
q_scale_lds[k][lane_id] = q_scale_lanes[k];
}
__syncthreads();
// Read back iterations this warp didn't compute
#pragma unroll
for (u32 k = 0; k < SCORE_MFMA_ITERS; k++) {
if (k % NUM_WARPS == warp_id) continue;
#pragma unroll
for (u32 r = 0; r < 8; r++) {
q_lanes[k][r] = q_lds[k][lane_id][r];
}
q_scale_lanes[k] = q_scale_lds[k][lane_id];
}
__builtin_amdgcn_sched_barrier(0);
for (u32 kv_iter = 0; kv_iter < kv_num_iters; kv_iter++) {
// ==========
// Cooperative load of KV Cache MXFP4 tile to LDS
// ==========
// Wait for the prefetched LDS to land
// asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
static_assert(PREFETCH <= 2); // Need custom vmcnt flag for PREFETCH > 2
__builtin_amdgcn_s_waitcnt(0x3f70);
// Prefetch the next LDS
u32 kv_buf = kv_iter % 2;
u32 kv_prefetch_iter = kv_iter + PREFETCH - 1;
if (kv_prefetch_iter < kv_num_iters) {
load_lds(warp_kv_offset + kv_prefetch_iter);
}
// ==========
// scores = per-head query @ key * 1/sqrt(d)
// scores = Q(N_HEADS,QK_HEAD_DIM) @ K(32,QK_HEAD_DIM)^T * 1/sqrt(QK_HEAD_DIM)
// ==========
constexpr u32 SCORE_MFMA_KV_TILE_DIM = 16;
static_assert(KV_TILE_DIM % SCORE_MFMA_KV_TILE_DIM == 0);
constexpr u32 MFMA_PER_KV_TILE_DIM = KV_TILE_DIM / SCORE_MFMA_KV_TILE_DIM;
f32x4 scores[MFMA_PER_KV_TILE_DIM] = {};
#pragma unroll
for (u32 score_mfma_idx = 0; score_mfma_idx < SCORE_MFMA_ITERS; score_mfma_idx++) {
#pragma unroll
for (u32 kv_mfma_idx = 0; kv_mfma_idx < MFMA_PER_KV_TILE_DIM; kv_mfma_idx++) {
u32 tok = kv_mfma_idx * 16 + lane_col;
u32 kv_byte_offset = score_mfma_idx * (SCORE_MFMA_DOT_DIM / 2) + lane_rowgroup * 16;
u32x8 kv_reg;
u8 kv_scale = 0;
if (score_mfma_idx < SCORE_MFMA_ITERS - 1 || kv_byte_offset + 16 <= QK_HEAD_DIM / 2) {
*(u32x4*)&kv_reg = *(const u32x4*)&lds.tile.kv_data[warp_id][kv_buf][tok][kv_byte_offset];
kv_scale = lds.tile.kv_scale[warp_id][kv_buf][tok][score_mfma_idx * (SCORE_MFMA_DOT_DIM / SCALE_GROUP_SIZE) + lane_rowgroup];
} else {
*(u32x4*)&kv_reg = {};
}
#ifdef __gfx950__
scores[kv_mfma_idx] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
q_lanes[score_mfma_idx],
kv_reg,
scores[kv_mfma_idx],
4, 4,
0, q_scale_lanes[score_mfma_idx],
0, kv_scale
);
#else
#pragma unroll
for (u32 sub = 0; sub < 4; sub++) {
bf16x8 q_bf16, kv_bf16_lane;
#pragma unroll
for (u32 i = 0; i < 8; i++) {
u32 qi = sub * 8 + i;
u8 q_nibble = (q_lanes[score_mfma_idx][qi / 8] >> (4 * (qi % 8))) & 0xF;
q_bf16[i] = fp4_e2m1_to_bf16(q_nibble, q_scale_lanes[score_mfma_idx]);
u8 kv_byte = ((const u8*)&kv_reg)[sub * 4 + i / 2];
u8 kv_nibble = (i % 2 == 0) ? (kv_byte & 0xF) : (kv_byte >> 4);
kv_bf16_lane[i] = fp4_e2m1_to_bf16(kv_nibble, kv_scale);
}
scores[kv_mfma_idx] = mfma_bf16_16x16x32(q_bf16, kv_bf16_lane, scores[kv_mfma_idx]);
}
#endif
}
}
// Scale by 1 / sqrt(QK_HEAD_DIM)
static_assert(QK_HEAD_DIM == 24 * 24); // Add constexpr sqrt to make this responsive w.r.t QK_HEAD_DIM
constexpr f32 SM_SCALE = 1.0f / 24.0f;
#pragma unroll
for (u32 i = 0; i < MFMA_PER_KV_TILE_DIM; i++) {
#pragma unroll
for (u32 j = 0; j < 4; j++) {
scores[i][j] *= SM_SCALE;
}
}
// ==========
// Online softmax
// ==========
// == MFMA thread-local KV-tile max reduction ==
// From the MFMA, each lane's (lane_col, lane_rowgroup) holds values from 4 different heads.
// We need to max-reduce over the MFMA KV Tiles, to get the per-lane max for each head
f32 lane_max[4] = {scores[0][0], scores[0][1], scores[0][2], scores[0][3]};
#pragma unroll
for (u32 i = 1; i < MFMA_PER_KV_TILE_DIM; i++) {
#pragma unroll
for (u32 j = 0; j < 4; j++) {
lane_max[j] = fmaxf(lane_max[j], scores[i][j]);
}
}
// == MFMA intra-lanegroup max reduction ==
// 16-consecutive lanes will share the same lane_rowgroup, so we reduce the max together.
// Each lane will then get the same true max of the KV_TILE_DIM rows, for its 4 heads.
#pragma unroll
for (u32 i = 0; i < 4; i++) {
lane_max[i] = dpp_reduce_max_16(lane_max[i]);
}
// == Update online state ==
// new_head_max is the new global max (max of accumulated kv_tile softmax_max[i] and current kv_tile lane_max[i]).
// alpha = exp(old_max - new_max) is the correction factor to rescale all previously accumulated values by
// We update previously accumulated softmax_denom and softmax_max by this alpha right now.
// - Updating previously accumulated value_out_lanes will be done later.
f32x4 alpha;
#pragma unroll
for (u32 i = 0; i < 4; i++) {
f32 new_head_max = fmaxf(softmax_max[i], lane_max[i]);
alpha[i] = fast_exp(softmax_max[i] - new_head_max);
softmax_denom[i] *= alpha[i];
softmax_max[i] = new_head_max;
}
// == Compute weights, thread-local KV-tile sum reduction ==
// weight = exp(score-score_max)
// For each score MFMA'ss output lane item, we calculate the weight for later value accumulation
// Weights are stored in LDS (for intra-warp permutation, value lanes are transposed)
// lane_sum is reduced across MFMA_PER_KV_TILE_DIM, for the 4 unique heads per lane.
f32 weights[MFMA_PER_KV_TILE_DIM][4];
f32 lane_sum[4] = {};
#pragma unroll
for (u32 i = 0; i < MFMA_PER_KV_TILE_DIM; i++) {
#pragma unroll
for (u32 j = 0; j < 4; j++) {
weights[i][j] = fast_exp(scores[i][j] - softmax_max[j]);
lds.tile.weights[warp_id][4 * lane_rowgroup + j][i * 16 + lane_col] = f32_to_bf16(weights[i][j]);
lane_sum[j] += weights[i][j];
}
}
// == accumulate per-lane softmax_denom ==
// We defer cross-lane reduction to after loop
#pragma unroll
for (u32 i = 0; i < 4; i++) {
softmax_denom[i] += lane_sum[i];
}
// ==========
// value_out += lds_weights(N_HEADS,KV_TILE_DIM) @ kv(KV_TILE_DIM,V_HEAD_DIM)^T
// ==========
// Read weights for MFMA value lane: [head=lcol, tokens lgrp*8..+7]
bf16x8 value_weight_lane = *(const bf16x8*)&lds.tile.weights[warp_id][lane_col][lane_rowgroup * 8];
// MFMA for weights * values
// We process in MFMA tiles of V_HEAD_DIM / VALUE_MFMA_V_TILE_DIM
constexpr u32 FEATURES_PER_DWORD = 8; // 4 bytes = 8 nibbles
constexpr u32 VALUEGROUP_ITERS = V_HEAD_DIM / (VALUE_MFMA_V_TILE_DIM * FEATURES_PER_DWORD); // 512/(16*8) = 4
static_assert(VALUEGROUP_ITERS == 4);
// Scales: all 4 valuegroups fall in the same scale group (lane_col * 16 bytes = lane_col * 32 features → scale_idx = lane_col)
u32 scale_base = (u32)(uintptr_t)&lds.tile.kv_scale[warp_id][kv_buf][lane_rowgroup * 8][lane_col];
u32 scale_reg[8];
__builtin_amdgcn_sched_barrier(0);
#pragma unroll
for (u32 row = 0; row < 8; row++) {
asm volatile("ds_read_u8 %0, %1 offset:%c2" : "=v"(scale_reg[row]) : "v"(scale_base), "n"(row * KV_SCALE_STRIDE));
}
__builtin_amdgcn_sched_barrier(0);
constexpr u32 VALUE_PREFETCH = 3;
u32 data_reg[VALUE_PREFETCH][8];
constexpr u32 VALUE_LOADS_PER_PRFETCH = 4;
auto loadValueLDS = [&](u32 vg_idx) {
u32 buf = vg_idx % VALUE_PREFETCH;
u32 byte_col_base = lane_col * 16 + vg_idx * 4;
u32 data_base = (u32)(uintptr_t)&lds.tile.kv_data[warp_id][kv_buf][lane_rowgroup * 8][byte_col_base];
u32 data_base_hi = data_base + 4 * (QK_HEAD_DIM / 2);
constexpr u32 QK_STRIDE_DWORD = (QK_HEAD_DIM / 2) / 4;
__builtin_amdgcn_sched_barrier(0);
asm volatile("ds_read2_b32 %0, %1 offset0:%c2 offset1:%c3"
: "=v"(*(u64*)&data_reg[buf][0])
: "v"(data_base), "n"((u32)0), "n"(QK_STRIDE_DWORD));
asm volatile("ds_read2_b32 %0, %1 offset0:%c2 offset1:%c3"
: "=v"(*(u64*)&data_reg[buf][2])
: "v"(data_base), "n"(QK_STRIDE_DWORD * 2), "n"(QK_STRIDE_DWORD * 3));
asm volatile("ds_read2_b32 %0, %1 offset0:%c2 offset1:%c3"
: "=v"(*(u64*)&data_reg[buf][4])
: "v"(data_base_hi), "n"((u32)0), "n"(QK_STRIDE_DWORD));
asm volatile("ds_read2_b32 %0, %1 offset0:%c2 offset1:%c3"
: "=v"(*(u64*)&data_reg[buf][6])
: "v"(data_base_hi), "n"(QK_STRIDE_DWORD * 2), "n"(QK_STRIDE_DWORD * 3));
__builtin_amdgcn_sched_barrier(0);
};
for (u32 valuegroup_idx = 0; valuegroup_idx < VALUE_PREFETCH - 1; valuegroup_idx++) {
loadValueLDS(valuegroup_idx);
}
#pragma unroll
for (u32 valuegroup_idx = 0; valuegroup_idx < VALUEGROUP_ITERS; valuegroup_idx++) {
// 16 lanes x 4 bytes = 64 bytes = 128 features per value group
// All 8 features in one dword share the same scale group
// (8 features < SCALE_GROUP_SIZE=32)
// Load from LDS
// 4 bytes (data) + 1 byte (scale) from each of 8 rows
__builtin_amdgcn_sched_barrier(0);
if (valuegroup_idx + VALUE_PREFETCH - 1 < VALUEGROUP_ITERS) {
loadValueLDS(valuegroup_idx + VALUE_PREFETCH - 1);
}
u32 inflight = (VALUEGROUP_ITERS - 1 - valuegroup_idx < VALUE_PREFETCH - 1)
? (VALUEGROUP_ITERS - 1 - valuegroup_idx)
: (VALUE_PREFETCH - 1);
static_assert(VALUE_PREFETCH <= 3); // Idk something weird happens here for >4. min(15, ..) doesn't fix it.
asm volatile("s_waitcnt lgkmcnt(%c0)" :: "n"(inflight * VALUE_LOADS_PER_PRFETCH) : "memory");
__builtin_amdgcn_sched_barrier(0);
u32 buf = valuegroup_idx % VALUE_PREFETCH;
// 4 byte columns, each producing 2 MFMAs (one for each nibble column)
#pragma unroll
for (u32 byte_col_offset = 0; byte_col_offset < 4; byte_col_offset++) {
// the value tile index for lo nibble / hi nibble
u32 lo_value_tile = valuegroup_idx * 8 + byte_col_offset * 2;
u32 hi_value_tile = lo_value_tile + 1;
value_out_lanes[lo_value_tile] *= alpha;
value_out_lanes[hi_value_tile] *= alpha;
bf16x8 value_lane_lo, value_lane_hi;
#pragma unroll
for (u32 row = 0; row < 8; row++) {
#ifdef __gfx950__
// f32 scale_f32 = e8m0_to_f32(scale);
// `e8m0_to_f32` requires an AND mask since it doesn't know ds_read_u8 will zero out the unused 3 bytes
f32 scale_f32 = __builtin_bit_cast(float, scale_reg[row] << 23);
u32 cvt;
switch (byte_col_offset) {
case 0: cvt = __builtin_bit_cast(u32, __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(data_reg[buf][row], scale_f32, 0)); break;
case 1: cvt = __builtin_bit_cast(u32, __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(data_reg[buf][row], scale_f32, 1)); break;
case 2: cvt = __builtin_bit_cast(u32, __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(data_reg[buf][row], scale_f32, 2)); break;
case 3: cvt = __builtin_bit_cast(u32, __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(data_reg[buf][row], scale_f32, 3)); break;
}
value_lane_lo[row] = __builtin_bit_cast(__bf16, (u16)cvt);
value_lane_hi[row] = __builtin_bit_cast(__bf16, (u16)(cvt >> 16));
#else
u8 scale = (u8)scale_reg[row];
u8 nibble_lo = (data_reg[buf][row] >> (byte_col_offset * 8)) & 0xF;
u8 nibble_hi = (data_reg[buf][row] >> (byte_col_offset * 8 + 4)) & 0xF;
value_lane_lo[row] = fp4_e2m1_to_bf16(nibble_lo, scale);
value_lane_hi[row] = fp4_e2m1_to_bf16(nibble_hi, scale);
#endif
}
value_out_lanes[lo_value_tile] = mfma_bf16_16x16x32(value_weight_lane, value_lane_lo, value_out_lanes[lo_value_tile]);
value_out_lanes[hi_value_tile] = mfma_bf16_16x16x32(value_weight_lane, value_lane_hi, value_out_lanes[hi_value_tile]);
}
}
}
// ==========
// Reduce online softmax parameters across all warps, and store results in LDS
// lds value_out[NUM_WARPS] = value_out_lanes * interwarp_max_correction / interwarp_denom
// ==========
// Reduce softmax_denom across lanes in this rowgroup (same value for lane_id / 16)
#pragma unroll
for (u32 i = 0; i < 4; i++) {
softmax_denom[i] = dpp_reduce_sum_16(softmax_denom[i]);
}
// == Store per-warp online softmax counters ==
// We've already reduced within a warp, so only the per-head leader `lane_col == 0` needs to write.
// There are 4 rowgroups, 4 heads each, 16 heads total.
if (lane_col == 0) {
#pragma unroll
for (u32 i = 0; i < 4; i++) {
lds.merge.warp_max[warp_id][4 * lane_rowgroup + i] = softmax_max[i];
lds.merge.warp_denom[warp_id][4 * lane_rowgroup + i] = softmax_denom[i];
}
}
__syncthreads();
// == per-head online softmax correction ==
// For the 4 heads, we get the max and denom by reducing the per-warp max/denom already written
// exp(warp_max-interwarp_max) will be the max correction on this warp
// We can also divide by the denominator, since inter-warp is the global data for this batch index.
f32 global_correction[4];
#pragma unroll
for (u32 i = 0; i < 4; i++) {
u32 head = 4 * lane_rowgroup + i;
f32 interwarp_max = NEG_INF;
#pragma unroll
for (u32 w = 0; w < NUM_WARPS; w++) {
interwarp_max = fmaxf(interwarp_max, lds.merge.warp_max[w][head]);
}
f32 interwarp_denom = 0;
#pragma unroll
for (u32 w = 0; w < NUM_WARPS; w++) {
interwarp_denom += lds.merge.warp_denom[w][head] * fast_exp(lds.merge.warp_max[w][head] - interwarp_max);
}
// interwarp denom, is also the global denom
global_correction[i] = fast_exp(softmax_max[i] - interwarp_max) / interwarp_denom;
}
// == per-head softmax'd values ==
// By adjusting our value_out_lanes by the correction, we can write the corrected results to LDS
// Interwarp reduction requires LDS communication.
#pragma unroll
for (u32 valuegroup_idx = 0; valuegroup_idx < 4; valuegroup_idx++) {
#pragma unroll
for (u32 value_tile_col_offset = 0; value_tile_col_offset < 8; value_tile_col_offset++) {
u32 value_tile_idx = valuegroup_idx * 8 + value_tile_col_offset;
#pragma unroll
for (u32 i = 0; i < 4; i++) {
u32 head = 4 * lane_rowgroup + i;
u32 vdim = lane_col * 32 + valuegroup_idx * 8 + value_tile_col_offset;
lds.merge.values[warp_id][head][vdim] = f32_to_bf16(value_out_lanes[value_tile_idx][i] * global_correction[i]);
}
}
}
__syncthreads();
// ==========
// Write to global
// global out = \sum_warp value_out
// ==========
// == Use LDS to reduce sum over warps, and write to global ==
// All threads can work together on reducing and writing to global
// The original lane assignments are irrelevant now, we just distribute the work evenly and in order.
u16* out_batch_item = out + batch_idx * N_HEADS * V_HEAD_DIM;
constexpr u32 TOTAL_ELEMS = N_HEADS * V_HEAD_DIM;
constexpr u32 CHUNK = 8;
constexpr u32 ITERS = TOTAL_ELEMS / (THREADS_PER_BLOCK * CHUNK);
static_assert(TOTAL_ELEMS % (THREADS_PER_BLOCK * CHUNK) == 0);
// LDS -> reg
u32x4 loaded[ITERS][NUM_WARPS];
#pragma unroll
for (u32 iter = 0; iter < ITERS; iter++) {
u32 base = (iter * THREADS_PER_BLOCK + threadIdx.x) * CHUNK;
#pragma unroll
for (u32 w = 0; w < NUM_WARPS; w++) {
loaded[iter][w] = *(const u32x4*)(&lds.merge.values[w][0][0] + base);
}
}
// reg -> reduce -> store
#pragma unroll
for (u32 iter = 0; iter < ITERS; iter++) {
f32 acc[CHUNK];
#pragma unroll
for (u32 i = 0; i < CHUNK; i++) {
u32 word = loaded[iter][0][i / 2];
acc[i] = bf16_to_f32((u16)(word >> (16 * (i & 1))));
}
#pragma unroll
for (u32 w = 1; w < NUM_WARPS; w++) {
#pragma unroll
for (u32 i = 0; i < CHUNK; i++) {
u32 word = loaded[iter][w][i / 2];
acc[i] += bf16_to_f32((u16)(word >> (16 * (i & 1))));
}
}
u32x4 out_packed;
#pragma unroll
for (u32 i = 0; i < 4; i++) {
out_packed[i] = (u32)f32_to_bf16(acc[i * 2]) | ((u32)f32_to_bf16(acc[i * 2 + 1]) << 16);
}
u32 base = (iter * THREADS_PER_BLOCK + threadIdx.x) * CHUNK;
*(u32x4*)(out_batch_item + base) = out_packed;
}
}
void entry(
const uintptr_t q, // .............. (bs, 16, 576) bf16
const uintptr_t kv_indptr, // ...... (bs+1,) int32
const uintptr_t kv_bf16, // ........ (total_kv, 576) bf16
const uintptr_t kv_fp8, // ......... (total_kv, 576) fp8
const uintptr_t kv_fp8_scale, // ... (1,) f32
const uintptr_t kv_mxfp4, // ....... (total_kv, 288) fp4x2
const uintptr_t kv_mxfp4_scale, // . (total_kv, 18) uint8_t
uintptr_t out, // .................. (bs, 16, 512) bf16
u32 bs
) {
kernel<<<bs, THREADS_PER_BLOCK>>>(
(const u16*)q,
(const u32*)kv_indptr,
(const u16*)kv_bf16,
(const u8*)kv_fp8, (const f32*)kv_fp8_scale,
(const u8*)kv_mxfp4, (const u8*)kv_mxfp4_scale,
(u16*)out,
bs
);
}
"""
CUDA_SRC_2 = r"""
#define CDIV(a, b) (((a) + (b) - 1) / (b))
using u8 = uint8_t;
using u16 = uint16_t;
using u16x32 = u16 __attribute__((ext_vector_type(32)));
using u32 = uint32_t;
using u32x4 = u32 __attribute__((ext_vector_type(4)));
using u32x8 = u32 __attribute__((ext_vector_type(8)));
using u64 = uint64_t;
using f32 = float;
using f32x4 = f32 __attribute__((ext_vector_type(4)));
using bf16x2 = __bf16 __attribute__((ext_vector_type(2)));
using bf16x4 = __bf16 __attribute__((ext_vector_type(4)));
using bf16x8 = __bf16 __attribute__((ext_vector_type(8)));
__device__ inline f32 fast_exp(f32 x) {
constexpr f32 LOG2E = 1.4426950408889634f;
return __builtin_amdgcn_exp2f(x * LOG2E);
}
__device__ inline u16 f32_to_bf16(f32 v) {
return (u16)(__builtin_bit_cast(u32, v) >> 16);
}
__device__ inline float bf16_to_f32(uint16_t v) {
return __builtin_bit_cast(float, (uint32_t)v << 16);
}
// Quantize a group of 32 bf16 values -> E8M0 scale (biased uint8)
__device__ inline uint8_t bf16x32_to_scale_e8m0(const u16x32* row) {
uint32_t amax_u16x2_packed = 0;
const uint32_t* row32 = (const uint32_t*)row;
#pragma unroll
for (int i = 0; i < 16; i++) {
// abs both bf16 in one instruction
uint32_t v = row32[i] & 0x7FFF7FFFu;
// max both bf16 in one instruction
asm("v_pk_max_u16 %0, %0, %1" : "+v"(amax_u16x2_packed) : "v"(v));
}
uint16_t amax_u16 = max((uint16_t)amax_u16x2_packed, (uint16_t)(amax_u16x2_packed >> 16));
// Round exponent up when mantissa fraction >= 0.75
// (bf16 has 7 mantissa bits; 0x20 = 2^5 matches the f32 0x200000 = 2^21 rounding)
int biased_exp = ((amax_u16 + 0x20u) >> 7) & 0xFF;
int e8m0 = biased_exp - 2;
return (uint8_t)max(0, min(254, e8m0));
}
// Matches `fp4_utils.e8m0_to_f32`: scale_e8m0 << 23 as f32, with special cases
__device__ inline float e8m0_to_f32(uint8_t scale) {
return __builtin_bit_cast(float, (uint32_t)scale << 23);
}
__device__ inline __bf16 fp4_e2m1_to_bf16(u8 nibble, u8 scale) {
u8 s = (nibble >> 3) & 1;
u8 e = (nibble >> 1) & 3;
u8 m = nibble & 1;
f32 val = (e == 0) ? m * 0.5f : exp2f((f32)e - 1.0f) * (1.0f + m * 0.5f);
val = s ? -val : val;
u16 ret = f32_to_bf16(val * e8m0_to_f32(scale));
return *(const __bf16*)&ret;
}
__device__ inline uint8_t f32_to_fp4_e2m1(float v) {
// CAREFUL: We need to include the fix from https://github.com/ROCm/aiter/pull/975 / https://github.com/ROCm/aiter/pull/2249
// FP4 e2m1: EBITS=2, MBITS=1, EXP_BIAS=1
// denorm_exp = (127-1) + (23-1) + 1 = 149 => denorm_mask = 149 << 23 = 0x4A800000
// val_to_add = ((1-127) << 23) + ((1<<21)-1) = 0xC11FFFFF (int32 wrapping)
constexpr float FP4_MAX_NORMAL = 6.0f;
constexpr float FP4_MIN_NORMAL = 1.0f;
constexpr uint32_t DENORM_MASK_INT = 0x4A800000u;
constexpr float DENORM_MASK_FLOAT = __builtin_bit_cast(float, DENORM_MASK_INT);
constexpr uint32_t VAL_TO_ADD = 0xC11FFFFFu;
uint8_t s_bit = (uint8_t)((__builtin_bit_cast(uint32_t, v) >> 28) & 0x8u);
uint32_t abs_bits = __builtin_bit_cast(uint32_t, v) & 0x7FFFFFFFu;
float abs_v = __builtin_bit_cast(float, abs_bits);
uint8_t result;
if (abs_v >= FP4_MAX_NORMAL) {
// saturate branch
result = 0x7u;
} else if (abs_v < FP4_MIN_NORMAL) {
// denormal branch: float trick to extract low bits
uint32_t d = __builtin_bit_cast(uint32_t, abs_v + DENORM_MASK_FLOAT) - DENORM_MASK_INT;
result = (uint8_t)d;
} else {
// normal branch: bias-adjust exponent + round-to-nearest
uint32_t x = abs_bits;
uint32_t m_odd = (x >> 22) & 1u;
x = x + VAL_TO_ADD + m_odd;
result = (uint8_t)(x >> 22);
}
return s_bit | result;
}
// Take an f32 and its block's scaling term, and convert to fp4_e2m1
// The fp4_e2m1 will be in the low nibble of the return value
__device__ inline uint8_t f32_to_fp4_e2m1_scale(float v, uint8_t scale) {
float f32_scale = e8m0_to_f32((uint8_t)(254u - scale)); // 2^(-(exp_unb)) = 2^(127 - scale)
return f32_to_fp4_e2m1(v * f32_scale);
}
__device__ inline f32x4 mfma_bf16_16x16x32(bf16x8 a, bf16x8 b, f32x4 acc) {
#ifdef __gfx950__
// return __builtin_amdgcn_mfma_f32_16x16x32_bf16(a, b, acc, 0, 0, 0);
// AGPRs are bad, we can't multiply alpha into them. Pin VGPR.
asm("v_mfma_f32_16x16x32_bf16 %0, %1, %2, %0"
: "+v"(acc)
: "v"(a), "v"(b));
return acc;
#else
acc = __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(*(bf16x4*)&a, *(bf16x4*)&b, acc, 0, 0, 0);
acc = __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(*((bf16x4*)&a+1), *((bf16x4*)&b+1), acc, 0, 0, 0);
return acc;
#endif
}
__device__ inline void global_to_lds_1024b(const u8* global_base, u8* lds_base, u32 lane_id) {
u32 lds_off = __builtin_amdgcn_readfirstlane((u32)(uintptr_t)lds_base);
#ifdef __gfx950__
asm volatile(
"s_mov_b32 m0, %0\n\t"
"global_load_lds_dwordx4 %1, off\n\t"
:: "s"(lds_off), "v"((const void*)(global_base + lane_id * 16))
: "memory", "m0"
);
#else
#pragma unroll
for (u32 sub = 0; sub < 4; sub++) {
asm volatile(
"s_mov_b32 m0, %0\n\t"
"global_load_lds_dword %1, off\n\t"
:: "s"(lds_off + sub * 256), "v"((const void*)(global_base + sub * 256 + lane_id * 4))
: "memory", "m0"
);
}
#endif
}
__device__ inline void global_to_lds_256b(const u8* global_base, u8* lds_base, u32 lane_id) {
u32 lds_off = __builtin_amdgcn_readfirstlane((u32)(uintptr_t)lds_base);
asm volatile(
"s_mov_b32 m0, %0\n\t"
"global_load_lds_dword %1, off\n\t"
:: "s"(lds_off), "v"((const void*)(global_base + lane_id * 4))
: "memory", "m0"
);
}
__device__ inline f32 dpp_reduce_max_16(f32 v) {
#define DPP_MOV(x, ctrl) __builtin_bit_cast(f32, \
__builtin_amdgcn_mov_dpp(__builtin_bit_cast(u32, x), ctrl, 0xF, 0xF, false))
v = fmaxf(v, DPP_MOV(v, 0xB1)); // quad XOR 1
v = fmaxf(v, DPP_MOV(v, 0x4E)); // quad XOR 2
v = fmaxf(v, DPP_MOV(v, 0x124)); // row_ror:4
v = fmaxf(v, DPP_MOV(v, 0x128)); // row_ror:8
return v;
#undef DPP_MOV
}
__device__ inline f32 dpp_reduce_sum_16(f32 v) {
#define DPP_MOV(x, ctrl) __builtin_bit_cast(f32, \
__builtin_amdgcn_mov_dpp(__builtin_bit_cast(u32, x), ctrl, 0xF, 0xF, false))
v += DPP_MOV(v, 0xB1);
v += DPP_MOV(v, 0x4E);
v += DPP_MOV(v, 0x124);
v += DPP_MOV(v, 0x128);
return v;
#undef DPP_MOV
}
// Constants
constexpr f32 NEG_INF = -1e30f;
constexpr u32 THREADS_PER_WARP = 64;
#ifdef __gfx950__
constexpr u32 NUM_WARPS = 4;
#else
constexpr u32 NUM_WARPS = 2;
#endif
constexpr u32 THREADS_PER_BLOCK = THREADS_PER_WARP * NUM_WARPS;
constexpr u32 SCALE_GROUP_SIZE = 32;
constexpr u32 KV_TILE_DIM = 32;
// DeepSeek Parameters
constexpr u32 N_HEADS = 16;
constexpr u32 QK_HEAD_DIM = 576;
constexpr u32 V_HEAD_DIM = 512;
__global__ __launch_bounds__(THREADS_PER_BLOCK, 1)
void kernel(
const u16* __restrict__ q,
const u32* __restrict__ kv_indptr,
const u16* __restrict__ kv_bf16,
const u8* __restrict__ kv_fp8, const f32* __restrict__ kv_fp8_scale,
const u8* __restrict__ kv_mxfp4, const u8* __restrict__ kv_mxfp4_scale,
u16* __restrict__ out,
int bs,
u32 num_splits,
u16* __restrict__ partial_values,
f32* __restrict__ partial_max,
f32* __restrict__ partial_denom,
u32* __restrict__ counter
) {
u32 batch_idx = blockIdx.x / num_splits;
u32 split_idx = blockIdx.x % num_splits;
u32 warp_id = threadIdx.x / THREADS_PER_WARP;
u32 lane_id = threadIdx.x % THREADS_PER_WARP;
u32 lane_col = lane_id % 16;
u32 lane_rowgroup = lane_id / 16; // rows \in 4*lane_rowgroup + {0,1,2,3}
// Pre-load and quantize Q into FP4 registers (constant across all KV tiles)
const u16* q_batch_item = q + batch_idx * N_HEADS * QK_HEAD_DIM;
constexpr u32 SCORE_MFMA_DOT_DIM = 128;
constexpr u32 SCORE_MFMA_ITERS = CDIV(QK_HEAD_DIM, SCORE_MFMA_DOT_DIM);
u32x8 q_lanes[SCORE_MFMA_ITERS];
u8 q_scale_lanes[SCORE_MFMA_ITERS];
#pragma unroll
for (u32 k = 0; k < SCORE_MFMA_ITERS; k++) {
u32 dim_base = k * SCORE_MFMA_DOT_DIM + lane_rowgroup * SCALE_GROUP_SIZE;
if (k < SCORE_MFMA_ITERS - 1 || dim_base + SCALE_GROUP_SIZE <= QK_HEAD_DIM) {
const u16* q_chunk = q_batch_item + lane_col * QK_HEAD_DIM + dim_base;
q_scale_lanes[k] = bf16x32_to_scale_e8m0((const u16x32*)q_chunk);
f32 s = e8m0_to_f32(q_scale_lanes[k]);
#pragma unroll
for (u32 r = 0; r < 4; r++) {
#ifdef __gfx950__
const bf16x2* v2 = (const bf16x2*)(q_chunk + r * 8);
q_lanes[k][r] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(0, v2[0], s, 0);
q_lanes[k][r] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q_lanes[k][r], v2[1], s, 1);
q_lanes[k][r] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q_lanes[k][r], v2[2], s, 2);
q_lanes[k][r] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q_lanes[k][r], v2[3], s, 3);
#else
q_lanes[k][r] = 0;
#pragma unroll
for (u32 i = 0; i < 8; i++) {
u8 nibble = f32_to_fp4_e2m1_scale(bf16_to_f32(q_chunk[r * 8 + i]), q_scale_lanes[k]);
q_lanes[k][r] |= (u32)(nibble & 0xF) << (4 * i);
}
#endif
}
} else {
q_lanes[k] = {};
q_scale_lanes[k] = 0;
}
}
u32 full_kv_start = kv_indptr[batch_idx];
u32 full_kv_end = kv_indptr[batch_idx + 1];
u32 full_kv_range = full_kv_end - full_kv_start;
u32 split_len = full_kv_range / num_splits;
u32 kv_start = full_kv_start + split_idx * split_len;
u32 kv_end = kv_start + split_len;
f32 softmax_max[4] = {NEG_INF, NEG_INF, NEG_INF, NEG_INF};
f32 softmax_denom[4] = {};
constexpr u32 VALUE_MFMA_V_TILE_DIM = 16;
constexpr u32 VALUE_MFMA_DOT_DIM = 32;
constexpr u32 NUM_VALUE_TILES = V_HEAD_DIM / VALUE_MFMA_V_TILE_DIM;
f32x4 value_out_lanes[NUM_VALUE_TILES] = {};
// KV Tiles
constexpr u32 PREFETCH = 2;
constexpr u32 KV_SCALE_STRIDE = CDIV(QK_HEAD_DIM / SCALE_GROUP_SIZE, 8) * 8;
// Shared Memory
__shared__ union {
// kv iterations
struct {
u8 kv_data[NUM_WARPS][PREFETCH][KV_TILE_DIM][QK_HEAD_DIM / 2];
u8 kv_scale[NUM_WARPS][PREFETCH][KV_TILE_DIM][KV_SCALE_STRIDE];
u16 weights[NUM_WARPS][N_HEADS][KV_TILE_DIM];
} tile;
// merge before write
struct {
f32 warp_max[NUM_WARPS][N_HEADS];
f32 warp_denom[NUM_WARPS][N_HEADS];
u16 values[NUM_WARPS][N_HEADS][V_HEAD_DIM];
} merge;
} lds;
u32 kv_range = kv_end - kv_start;
u32 kv_num_iters = kv_range / KV_TILE_DIM / NUM_WARPS;
if (kv_range % (NUM_WARPS * KV_TILE_DIM) != 0 || kv_num_iters < PREFETCH - 1) {
__builtin_trap();
}
u32 warp_kv_offset = warp_id * kv_num_iters;
// (u32 kv_tile_start = kv_start; kv_tile_start < kv_end; kv_tile_start += KV_TILE_DIM) {
auto load_lds = [&](uint32_t kv_iter) -> void {
u32 kv_tile_start = kv_start + (warp_kv_offset + kv_iter) * KV_TILE_DIM;
u32 kv_buf = kv_iter % 2;
constexpr u32 KV_DATA_BYTES = KV_TILE_DIM * (QK_HEAD_DIM / 2);
constexpr u32 KV_DATA_U128 = KV_DATA_BYTES / 16;
static_assert(KV_DATA_U128 % THREADS_PER_WARP == 0);
constexpr u32 KV_DATA_ITERS = KV_DATA_U128 / THREADS_PER_WARP;
const u8* kv_data_src = (const u8*)(kv_mxfp4 + (u64)kv_tile_start * (QK_HEAD_DIM / 2));
u8* lds_data_dst = (u8*)&lds.tile.kv_data[warp_id][kv_buf];
#pragma unroll
for (u32 i = 0; i < KV_DATA_ITERS; i++) {
global_to_lds_1024b(kv_data_src + i * 1024, lds_data_dst + i * 1024, lane_id);
}
constexpr u32 KV_SCALE_BYTES = KV_TILE_DIM * KV_SCALE_STRIDE;
constexpr u32 KV_SCALE_U32 = KV_SCALE_BYTES / 4;
constexpr u32 KV_SCALE_ITERS = CDIV(KV_SCALE_U32, THREADS_PER_WARP);
const u8* kv_scale_src = (const u8*)(kv_mxfp4_scale + (u64)kv_tile_start * KV_SCALE_STRIDE);
u8* lds_scale_data_dst = (u8*)&lds.tile.kv_scale[warp_id][kv_buf];
#pragma unroll
for (u32 i = 0; i < KV_SCALE_ITERS; i++) {
global_to_lds_256b(kv_scale_src + i * 256, lds_scale_data_dst + i * 256, lane_id);
}
};
#pragma unroll
for (u32 i = 0; i < PREFETCH - 1; i++) {
load_lds(i);
}
for (u32 kv_iter = 0; kv_iter < kv_num_iters; kv_iter++) {
// ==========
// Cooperative load of KV Cache MXFP4 tile to LDS
// ==========
// Wait for the prefetched LDS to land
// asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
static_assert(PREFETCH <= 2); // Need custom vmcnt flag for PREFETCH > 2
__builtin_amdgcn_s_waitcnt(0x3f70);
// Prefetch the next LDS
u32 kv_buf = kv_iter % 2;
u32 kv_prefetch_iter = kv_iter + PREFETCH - 1;
if (kv_prefetch_iter < kv_num_iters) {
load_lds(kv_prefetch_iter);
}
// ==========
// scores = per-head query @ key * 1/sqrt(d)
// scores = Q(N_HEADS,QK_HEAD_DIM) @ K(32,QK_HEAD_DIM)^T * 1/sqrt(QK_HEAD_DIM)
// ==========
constexpr u32 SCORE_MFMA_KV_TILE_DIM = 16;
static_assert(KV_TILE_DIM % SCORE_MFMA_KV_TILE_DIM == 0);
constexpr u32 MFMA_PER_KV_TILE_DIM = KV_TILE_DIM / SCORE_MFMA_KV_TILE_DIM;
f32x4 scores[MFMA_PER_KV_TILE_DIM] = {};
#pragma unroll
for (u32 score_mfma_idx = 0; score_mfma_idx < SCORE_MFMA_ITERS; score_mfma_idx++) {
#pragma unroll
for (u32 kv_mfma_idx = 0; kv_mfma_idx < MFMA_PER_KV_TILE_DIM; kv_mfma_idx++) {
u32 tok = kv_mfma_idx * 16 + lane_col;
u32 kv_byte_offset = score_mfma_idx * (SCORE_MFMA_DOT_DIM / 2) + lane_rowgroup * 16;
u32x8 kv_reg;
u8 kv_scale = 0;
if (score_mfma_idx < SCORE_MFMA_ITERS - 1 || kv_byte_offset + 16 <= QK_HEAD_DIM / 2) {
*(u32x4*)&kv_reg = *(const u32x4*)&lds.tile.kv_data[warp_id][kv_buf][tok][kv_byte_offset];
kv_scale = lds.tile.kv_scale[warp_id][kv_buf][tok][score_mfma_idx * (SCORE_MFMA_DOT_DIM / SCALE_GROUP_SIZE) + lane_rowgroup];
} else {
*(u32x4*)&kv_reg = {};
}
#ifdef __gfx950__
scores[kv_mfma_idx] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
q_lanes[score_mfma_idx],
kv_reg,
scores[kv_mfma_idx],
4, 4,
0, q_scale_lanes[score_mfma_idx],
0, kv_scale
);
#else
#pragma unroll
for (u32 sub = 0; sub < 4; sub++) {
bf16x8 q_bf16, kv_bf16_lane;
#pragma unroll
for (u32 i = 0; i < 8; i++) {
u32 qi = sub * 8 + i;
u8 q_nibble = (q_lanes[score_mfma_idx][qi / 8] >> (4 * (qi % 8))) & 0xF;
q_bf16[i] = fp4_e2m1_to_bf16(q_nibble, q_scale_lanes[score_mfma_idx]);
u8 kv_byte = ((const u8*)&kv_reg)[sub * 4 + i / 2];
u8 kv_nibble = (i % 2 == 0) ? (kv_byte & 0xF) : (kv_byte >> 4);
kv_bf16_lane[i] = fp4_e2m1_to_bf16(kv_nibble, kv_scale);
}
scores[kv_mfma_idx] = mfma_bf16_16x16x32(q_bf16, kv_bf16_lane, scores[kv_mfma_idx]);
}
#endif
}
}
// Scale by 1 / sqrt(QK_HEAD_DIM)
static_assert(QK_HEAD_DIM == 24 * 24); // Add constexpr sqrt to make this responsive w.r.t QK_HEAD_DIM
constexpr f32 SM_SCALE = 1.0f / 24.0f;
#pragma unroll
for (u32 i = 0; i < MFMA_PER_KV_TILE_DIM; i++) {
#pragma unroll
for (u32 j = 0; j < 4; j++) {
scores[i][j] *= SM_SCALE;
}
}
// ==========
// Online softmax
// ==========
// == MFMA thread-local KV-tile max reduction ==
// From the MFMA, each lane's (lane_col, lane_rowgroup) holds values from 4 different heads.
// We need to max-reduce over the MFMA KV Tiles, to get the per-lane max for each head
f32 lane_max[4] = {scores[0][0], scores[0][1], scores[0][2], scores[0][3]};
#pragma unroll
for (u32 i = 1; i < MFMA_PER_KV_TILE_DIM; i++) {
#pragma unroll
for (u32 j = 0; j < 4; j++) {
lane_max[j] = fmaxf(lane_max[j], scores[i][j]);
}
}
// == MFMA intra-lanegroup max reduction ==
// 16-consecutive lanes will share the same lane_rowgroup, so we reduce the max together.
// Each lane will then get the same true max of the KV_TILE_DIM rows, for its 4 heads.
#pragma unroll
for (u32 i = 0; i < 4; i++) {
lane_max[i] = dpp_reduce_max_16(lane_max[i]);
}
// == Update online state ==
// new_head_max is the new global max (max of accumulated kv_tile softmax_max[i] and current kv_tile lane_max[i]).
// alpha = exp(old_max - new_max) is the correction factor to rescale all previously accumulated values by
// We update previously accumulated softmax_denom and softmax_max by this alpha right now.
// - Updating previously accumulated value_out_lanes will be done later.
f32x4 alpha;
#pragma unroll
for (u32 i = 0; i < 4; i++) {
f32 new_head_max = fmaxf(softmax_max[i], lane_max[i]);
alpha[i] = fast_exp(softmax_max[i] - new_head_max);
softmax_denom[i] *= alpha[i];
softmax_max[i] = new_head_max;
}
// == Compute weights, thread-local KV-tile sum reduction ==
// weight = exp(score-score_max)
// For each score MFMA'ss output lane item, we calculate the weight for later value accumulation
// Weights are stored in LDS (for intra-warp permutation, value lanes are transposed)
// lane_sum is reduced across MFMA_PER_KV_TILE_DIM, for the 4 unique heads per lane.
f32 weights[MFMA_PER_KV_TILE_DIM][4];
f32 lane_sum[4] = {};
#pragma unroll
for (u32 i = 0; i < MFMA_PER_KV_TILE_DIM; i++) {
#pragma unroll
for (u32 j = 0; j < 4; j++) {
weights[i][j] = fast_exp(scores[i][j] - softmax_max[j]);
lds.tile.weights[warp_id][4 * lane_rowgroup + j][i * 16 + lane_col] = f32_to_bf16(weights[i][j]);
lane_sum[j] += weights[i][j];
}
}
// == accumulate per-lane softmax_denom ==
// We defer cross-lane reduction to after loop
#pragma unroll
for (u32 i = 0; i < 4; i++) {
softmax_denom[i] += lane_sum[i];
}
// ==========
// value_out += lds_weights(N_HEADS,KV_TILE_DIM) @ kv(KV_TILE_DIM,V_HEAD_DIM)^T
// ==========
// Read weights for MFMA value lane: [head=lcol, tokens lgrp*8..+7]
bf16x8 value_weight_lane = *(const bf16x8*)&lds.tile.weights[warp_id][lane_col][lane_rowgroup * 8];
// MFMA for weights * values
// We process in MFMA tiles of V_HEAD_DIM / VALUE_MFMA_V_TILE_DIM
constexpr u32 FEATURES_PER_DWORD = 8; // 4 bytes = 8 nibbles
constexpr u32 VALUEGROUP_ITERS = V_HEAD_DIM / (VALUE_MFMA_V_TILE_DIM * FEATURES_PER_DWORD); // 512/(16*8) = 4
static_assert(VALUEGROUP_ITERS == 4);
#pragma unroll
for (u32 valuegroup_idx = 0; valuegroup_idx < VALUEGROUP_ITERS; valuegroup_idx++) {
// 16 lanes x 4 bytes = 64 bytes = 128 features per value group
u32 byte_col_base = valuegroup_idx * 64 + lane_col * 4;
// All 8 features in one dword share the same scale group
// (8 features < SCALE_GROUP_SIZE=32)
u32 scale_idx = valuegroup_idx * 4 + lane_col / 4;
u32 data_base = (u32)(uintptr_t)&lds.tile.kv_data[warp_id][kv_buf][lane_rowgroup * 8][byte_col_base];
u32 scale_base = (u32)(uintptr_t)&lds.tile.kv_scale[warp_id][kv_buf][lane_rowgroup * 8][scale_idx];
// Load from LDS
// 4 bytes (data) + 1 byte (scale) from each of 8 rows
u32 data_reg[8];
u32 scale_reg[8];
__builtin_amdgcn_sched_barrier(0);
#pragma unroll
for (u32 row = 0; row < 8; row++) {
asm volatile(
"ds_read_b32 %0, %1 offset:%c2"
: "=v"(data_reg[row])
: "v"(data_base),
"n"(row * (QK_HEAD_DIM / 2))
);
asm volatile(
"ds_read_u8 %0, %1 offset:%c2"
: "=v"(scale_reg[row])
: "v"(scale_base),
"n"(row * KV_SCALE_STRIDE)
);
}
__builtin_amdgcn_sched_barrier(0);
asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");
__builtin_amdgcn_sched_barrier(0);
// 4 byte columns, each producing 2 MFMAs (one for each nibble column)
#pragma unroll
for (u32 byte_col_offset = 0; byte_col_offset < 4; byte_col_offset++) {
// the value tile index for lo nibble / hi nibble
u32 lo_value_tile = valuegroup_idx * 8 + byte_col_offset * 2;
u32 hi_value_tile = lo_value_tile + 1;
value_out_lanes[lo_value_tile] *= alpha;
value_out_lanes[hi_value_tile] *= alpha;
bf16x8 value_lane_lo, value_lane_hi;
#pragma unroll
for (u32 row = 0; row < 8; row++) {
u8 scale = (u8)scale_reg[row];
#ifdef __gfx950__
f32 scale_f32 = e8m0_to_f32(scale);
u32 cvt;
switch (byte_col_offset) {
case 0: cvt = __builtin_bit_cast(u32, __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(data_reg[row], scale_f32, 0)); break;
case 1: cvt = __builtin_bit_cast(u32, __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(data_reg[row], scale_f32, 1)); break;
case 2: cvt = __builtin_bit_cast(u32, __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(data_reg[row], scale_f32, 2)); break;
case 3: cvt = __builtin_bit_cast(u32, __builtin_amdgcn_cvt_scalef32_pk_bf16_fp4(data_reg[row], scale_f32, 3)); break;
}
value_lane_lo[row] = __builtin_bit_cast(__bf16, (u16)cvt);
value_lane_hi[row] = __builtin_bit_cast(__bf16, (u16)(cvt >> 16));
#else
u8 nibble_lo = (data_reg[row] >> (byte_col_offset * 8)) & 0xF;
u8 nibble_hi = (data_reg[row] >> (byte_col_offset * 8 + 4)) & 0xF;
value_lane_lo[row] = fp4_e2m1_to_bf16(nibble_lo, scale);
value_lane_hi[row] = fp4_e2m1_to_bf16(nibble_hi, scale);
#endif
}
value_out_lanes[lo_value_tile] = mfma_bf16_16x16x32(value_weight_lane, value_lane_lo, value_out_lanes[lo_value_tile]);
value_out_lanes[hi_value_tile] = mfma_bf16_16x16x32(value_weight_lane, value_lane_hi, value_out_lanes[hi_value_tile]);
}
}
}
// ==========
// Reduce online softmax parameters across all warps, and store results in LDS
// ==========
// Reduce softmax_denom across lanes in this rowgroup (same value for lane_id / 16)
#pragma unroll
for (u32 i = 0; i < 4; i++)
softmax_denom[i] = dpp_reduce_sum_16(softmax_denom[i]);
// == Store per-warp online softmax counters ==
// We've already reduced within a warp, so only the per-head leader `lane_col == 0` needs to write.
// There are 4 rowgroups, 4 heads each, 16 heads total.
if (lane_col == 0) {
#pragma unroll
for (u32 i = 0; i < 4; i++) {
lds.merge.warp_max[warp_id][4 * lane_rowgroup + i] = softmax_max[i];
lds.merge.warp_denom[warp_id][4 * lane_rowgroup + i] = softmax_denom[i];
}
}
__syncthreads();
f32 split_max[4];
f32 split_denom[4];
f32 warp_correction[4];
#pragma unroll
for (u32 i = 0; i < 4; i++) {
u32 head = 4 * lane_rowgroup + i;
split_max[i] = NEG_INF;
#pragma unroll
for (u32 w = 0; w < NUM_WARPS; w++) {
split_max[i] = fmaxf(split_max[i], lds.merge.warp_max[w][head]);
}
split_denom[i] = 0;
#pragma unroll
for (u32 w = 0; w < NUM_WARPS; w++) {
split_denom[i] += lds.merge.warp_denom[w][head] * fast_exp(lds.merge.warp_max[w][head] - split_max[i]);
}
warp_correction[i] = fast_exp(softmax_max[i] - split_max[i]);
}
#pragma unroll
for (u32 valuegroup_idx = 0; valuegroup_idx < 4; valuegroup_idx++) {
#pragma unroll
for (u32 value_tile_col_offset = 0; value_tile_col_offset < 8; value_tile_col_offset++) {
u32 value_tile_idx = valuegroup_idx * 8 + value_tile_col_offset;
#pragma unroll
for (u32 i = 0; i < 4; i++) {
u32 head = 4 * lane_rowgroup + i;
u32 vdim = valuegroup_idx * 128 + lane_col * 8 + value_tile_col_offset;
lds.merge.values[warp_id][head][vdim] = f32_to_bf16(value_out_lanes[value_tile_idx][i] * warp_correction[i]);
}
}
}
__syncthreads();
// ==========
// Write partials to global
// ==========
u32 split_base = batch_idx * num_splits + split_idx;
for (u32 idx = threadIdx.x; idx < N_HEADS * V_HEAD_DIM; idx += THREADS_PER_BLOCK) {
u32 head = idx / V_HEAD_DIM;
u32 vdim = idx % V_HEAD_DIM;
f32 sum = 0;
#pragma unroll
for (u32 w = 0; w < NUM_WARPS; w++) {
sum += bf16_to_f32(lds.merge.values[w][head][vdim]);
}
partial_values[split_base * N_HEADS * V_HEAD_DIM + idx] = f32_to_bf16(sum);
}
if (lane_col == 0 && warp_id == 0) {
#pragma unroll
for (u32 i = 0; i < 4; i++) {
u32 head = 4 * lane_rowgroup + i;
partial_max[split_base * N_HEADS + head] = split_max[i];
partial_denom[split_base * N_HEADS + head] = split_denom[i];
}
}
}
template<u32 NUM_SPLITS>
__global__ __launch_bounds__(256)
void reduce_splits(
const u16* __restrict__ partial_values,
const f32* __restrict__ partial_max,
const f32* __restrict__ partial_denom,
u16* __restrict__ out,
u32 bs
) {
constexpr u32 ELEMENTS = N_HEADS * V_HEAD_DIM;
u32 global_idx = blockIdx.x * 256 + threadIdx.x;
u32 total = bs * ELEMENTS;
if (global_idx >= total) return;
u32 batch_idx = global_idx / ELEMENTS;
u32 idx = global_idx % ELEMENTS;
u32 head = idx / V_HEAD_DIM;
// Load all data into registers first
f32 split_max[NUM_SPLITS];
f32 split_denom[NUM_SPLITS];
u16 split_val[NUM_SPLITS];
#pragma unroll
for (u32 s = 0; s < NUM_SPLITS; s++) {
u32 sb = batch_idx * NUM_SPLITS + s;
split_max[s] = partial_max[sb * N_HEADS + head];
}
#pragma unroll
for (u32 s = 0; s < NUM_SPLITS; s++) {
u32 sb = batch_idx * NUM_SPLITS + s;
split_denom[s] = partial_denom[sb * N_HEADS + head];
}
#pragma unroll
for (u32 s = 0; s < NUM_SPLITS; s++) {
u32 sb = batch_idx * NUM_SPLITS + s;
split_val[s] = partial_values[sb * ELEMENTS + idx];
}
// Compute
f32 global_max = NEG_INF;
#pragma unroll
for (u32 s = 0; s < NUM_SPLITS; s++)
global_max = fmaxf(global_max, split_max[s]);
f32 global_denom = 0;
#pragma unroll
for (u32 s = 0; s < NUM_SPLITS; s++)
global_denom += split_denom[s] * fast_exp(split_max[s] - global_max);
f32 val_sum = 0;
#pragma unroll
for (u32 s = 0; s < NUM_SPLITS; s++) {
f32 correction = fast_exp(split_max[s] - global_max) / global_denom;
val_sum += bf16_to_f32(split_val[s]) * correction;
}
out[global_idx] = f32_to_bf16(val_sum);
}
#define HIP_CALL(val) check((val), #val, __FILE__, __LINE__)
template <typename T> void check(T err, const char *const func, const char *const file, const int line) {
if (err != hipSuccess) {
fprintf(stderr, "HIP Runtime Error at: %s:%d\n", file, line);
fprintf(stderr, "%s %s\n", hipGetErrorString(err), func);
exit(1);
}
}
void entry(
const uintptr_t q, // .............. (bs, 16, 576) bf16
const uintptr_t kv_indptr, // ...... (bs+1,) int32
const uintptr_t kv_bf16, // ........ (total_kv, 576) bf16
const uintptr_t kv_fp8, // ......... (total_kv, 576) fp8
const uintptr_t kv_fp8_scale, // ... (1,) f32
const uintptr_t kv_mxfp4, // ....... (total_kv, 288) fp4x2
const uintptr_t kv_mxfp4_scale, // . (total_kv, 18) uint8_t
uintptr_t out, // .................. (bs, 16, 512) bf16
u32 bs
) {
// Reduction Buffers
constexpr u32 MAX_BS = 256;
constexpr u32 MAX_SPLITS = 256;
static u16* partial_values = nullptr;
static f32* partial_max = nullptr;
static f32* partial_denom = nullptr;
static u32* counter = nullptr;
if (!partial_values) {
HIP_CALL(hipMalloc(&partial_values, MAX_BS * MAX_SPLITS * N_HEADS * V_HEAD_DIM * sizeof(u16)));
HIP_CALL(hipMalloc(&partial_max, MAX_BS * MAX_SPLITS * N_HEADS * sizeof(f32)));
HIP_CALL(hipMalloc(&partial_denom, MAX_BS * MAX_SPLITS * N_HEADS * sizeof(f32)));
HIP_CALL(hipMalloc(&counter, MAX_BS * sizeof(u32)));
HIP_CALL(hipMemset(counter, 0, MAX_BS * sizeof(u32)));
}
// Get power of 2
u32 min_kvlen = ((const u32*)kv_indptr)[1] - ((const u32*)kv_indptr)[0];
u32 max_splits_kv = min_kvlen / (KV_TILE_DIM * NUM_WARPS);
u32 raw_num_splits = min(max_splits_kv, max(1u, 256u / bs));
u32 num_splits = 1u;
while (num_splits * 2 <= raw_num_splits) num_splits *= 2;
kernel<<<bs * num_splits, THREADS_PER_BLOCK>>>(
(const u16*)q,
(const u32*)kv_indptr,
(const u16*)kv_bf16,
(const u8*)kv_fp8, (const f32*)kv_fp8_scale,
(const u8*)kv_mxfp4, (const u8*)kv_mxfp4_scale,
(u16*)out,
bs,
num_splits,
partial_values,
partial_max,
partial_denom,
counter
);
u32 total = bs * N_HEADS * V_HEAD_DIM;
u32 reduce_blocks = CDIV(total, 256);
if (num_splits == 1) {
reduce_splits<1><<<reduce_blocks, 256>>>(partial_values, partial_max, partial_denom, (u16*)out, bs);
} else if (num_splits == 4) {
reduce_splits<4><<<reduce_blocks, 256>>>(partial_values, partial_max, partial_denom, (u16*)out, bs);
} else if (num_splits == 8) {
reduce_splits<8><<<reduce_blocks, 256>>>(partial_values, partial_max, partial_denom, (u16*)out, bs);
} else if (num_splits == 16) {
reduce_splits<16><<<reduce_blocks, 256>>>(partial_values, partial_max, partial_denom, (u16*)out, bs);
} else if (num_splits == 64) {
reduce_splits<64><<<reduce_blocks, 256>>>(partial_values, partial_max, partial_denom, (u16*)out, bs);
} else {
assert(false);
}
}
"""
class CompiledModule:
M: int
module: Any
out: torch.Tensor
def __init__(self, M: int) -> None:
CUDA_PRELUDE = """
#ifndef __gfx950__
#define __gfx950__
#endif
"""
cflags = ["--offload-arch=gfx950", "-std=c++20", "-O3", "-ffast-math", "-march=native", "-funroll-loops", "-fomit-frame-pointer"]
cflags.extend([f"-DM_DIM={M}"])
self.M = M
self.out = torch.empty((M, 16, 512), dtype=torch.bfloat16, device='cuda')
cuda_src = None
if M == 256:
cuda_src = CUDA_SRC
else:
cuda_src = CUDA_SRC_2
self.module = load_inline(
name=f"solution_{M}",
cpp_sources=[CPP_WRAPPER],
cuda_sources=[CUDA_PRELUDE + cuda_src],
functions=['entry'],
with_cuda=True,
verbose=False,
extra_cuda_cflags=cflags,
extra_cflags=cflags,
)
def inference(self, data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = qo_indptr.numel() - 1
self.module.entry(
q.data_ptr(),
kv_indptr.data_ptr(),
kv_data["bf16"].data_ptr(),
kv_data["fp8"][0].data_ptr(),
kv_data["fp8"][1].data_ptr(),
kv_data["mxfp4"][0].data_ptr(),
kv_data["mxfp4"][1].data_ptr(),
self.out.data_ptr(),
bs,
)
return self.out
_compiled_modules: dict[int, CompiledModule] = {}
def custom_kernel(data: input_t) -> output_t:
global _compiled_modules
q, kv_data, qo_indptr, kv_indptr, config = data
bs = qo_indptr.numel() - 1
if bs not in _compiled_modules:
_compiled_modules[bs] = CompiledModule(bs)
return _compiled_modules[bs].inference(data)
scrolls · 1616 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