submission 749874
John Hahn · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 3953 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-749874?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:377d002038d9993ecf85b13eba5e1bbb612236879a8f29471794f780f774f42e
license declaredunknown
license concludedunknown
authorsJohn Hahn
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
constexpr int KV_FP4_STRIDE = 288; // QK_DIM/2 packed FP4 bytes per rowshared-memory
__shared__ LDS lds;vector-width = uint4
const uint4* p = reinterpret_cast<const uint4*>(&kv_lds[addr]);Kernel source
submission.py3953 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
v506: Based on v505. Add tp=4 (nhead=32) support with macro-based dispatch.
On first call per config: registers C++ state, caches KV pointers.
On subsequent calls with same data: dispatch_cached(key) skips all tensor
argument parsing, dict lookups, and data_ptr() calls. Only passes int64 key.
"""
from task import input_t, output_t
import torch
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
HIP_SOURCE = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDASt_QQ_.h>
#include <hip/hip_runtime.h>
#include <hip/hip_fp16.h>
#include <hip/hip_bfloat16.h>
#include <hip/amd_detail/amd_hip_fp8.h>
#include <cstdint>
#include <algorithm>
#include <unordered_map>
// ============================================================================
// MFMA intrinsic types
// ============================================================================
using mfma_acc_t = __attribute__((ext_vector_type(4))) float;
using mfma_input_t = __attribute__((ext_vector_type(8))) uint32_t;
using mfma_bf16_input_t = __attribute__((ext_vector_type(8))) short;
__device__ __forceinline__ mfma_acc_t
mfma_f32_16x16x128_fp8(mfma_input_t a, mfma_input_t b, mfma_acc_t c) {
return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a, b, c, 0, 0, 0, 0, 0, 0);
}
__device__ __forceinline__ uint32_t
pack_bf16x2(float a, float b) {
union { short2 s; uint32_t u; } r;
asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(r.u) : "v"(a), "v"(b));
return r.u;
}
__device__ __forceinline__ mfma_acc_t
mfma_f32_16x16x32_bf16(mfma_bf16_input_t a, mfma_bf16_input_t b, mfma_acc_t c) {
return __builtin_amdgcn_mfma_f32_16x16x32_bf16(a, b, c, 0, 0, 0);
}
__device__ __forceinline__ uint8_t
fp32_to_fp8_e4m3(float val) {
uint32_t packed = __builtin_amdgcn_cvt_pk_fp8_f32(val, val, 0u, false);
return (uint8_t)(packed & 0xFF);
}
// ============================================================================
// Constants
// ============================================================================
constexpr int QK_DIM = 576;
constexpr int V_DIM = 512;
constexpr int NTHREADS = 256;
constexpr int WAVESIZE = 64;
constexpr int FP8_MFMA_K = 128;
constexpr int QK_MFMAS = 5;
constexpr int QK_MFMAS_LATENT = 4; // First 4 iters cover dims 0-511 (latent)
constexpr int ROPE_DIM = 64; // Dims 512-575 (RoPE)
constexpr int HEAD_GROUP = 16;
constexpr int KV_TILE = 32;
constexpr int KV_HALF = 16;
constexpr int KV_FP8_STRIDE = 576;
constexpr int LDS_KV_STRIDE = 576; // Contiguous layout: matches KV_FP8_STRIDE for chunk-based loading (44% BW savings)
constexpr int KV_FP4_STRIDE = 288; // QK_DIM/2 packed FP4 bytes per row
constexpr int KV_SCALE_PER_ROW = 18; // QK_DIM/32 E8M0 block scales per row
constexpr int V_MFMAS_PER_WAVE = 8;
// my_scale includes LOG2E so we can use v_exp_f32 (exp2) directly, saving 1 VALU per exp call
constexpr float LOG2E_F = 1.4426950408889634f;
constexpr float LN2_F = 0.6931471805599453f;
// ============================================================================
// V accumulation helpers (vectorized 4-byte loads)
// ============================================================================
// Load 4 consecutive FP8 bytes from each of 8 KV positions
__device__ __forceinline__ void
load_v_4bytes(const uint8_t* kv_lds, int v_base4, int lane_group, uint32_t raw4[8]) {
int base_pos = lane_group * 8;
#pragma unroll
for (int i = 0; i < 8; i++) {
raw4[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(base_pos + i) * LDS_KV_STRIDE + v_base4]);
}
}
// Convert byte BYTE_K (0-3) from 8 raw4 values to BF16 B-operand
// __builtin_amdgcn_cvt_f32_fp8 requires compile-time constant byte index.
template <int BYTE_K>
__device__ __forceinline__ void
v_convert_bf16(const uint32_t raw4[8], uint32_t bp[4]) {
float vf0 = __builtin_amdgcn_cvt_f32_fp8(raw4[0], BYTE_K);
float vf1 = __builtin_amdgcn_cvt_f32_fp8(raw4[1], BYTE_K);
bp[0] = pack_bf16x2(vf0, vf1);
float vf2 = __builtin_amdgcn_cvt_f32_fp8(raw4[2], BYTE_K);
float vf3 = __builtin_amdgcn_cvt_f32_fp8(raw4[3], BYTE_K);
bp[1] = pack_bf16x2(vf2, vf3);
float vf4 = __builtin_amdgcn_cvt_f32_fp8(raw4[4], BYTE_K);
float vf5 = __builtin_amdgcn_cvt_f32_fp8(raw4[5], BYTE_K);
bp[2] = pack_bf16x2(vf4, vf5);
float vf6 = __builtin_amdgcn_cvt_f32_fp8(raw4[6], BYTE_K);
float vf7 = __builtin_amdgcn_cvt_f32_fp8(raw4[7], BYTE_K);
bp[3] = pack_bf16x2(vf6, vf7);
}
// Packed FP8→BF16 conversion: converts all 4 byte positions at once using cvt_pk_f32_fp8.
// Produces 4 B-operand arrays (bp0..bp3) for byte positions 0..3.
// Uses 16 cvt_pk + 16 pack = 32 VALU vs 32 cvt + 16 pack = 48 VALU for the scalar version.
typedef float v2f __attribute__((ext_vector_type(2)));
__device__ __forceinline__ void
v_convert_bf16_packed(const uint32_t raw4[8], uint32_t bp0[4], uint32_t bp1[4],
uint32_t bp2[4], uint32_t bp3[4]) {
#pragma unroll
for (int i = 0; i < 8; i += 2) {
v2f lo_a = __builtin_amdgcn_cvt_pk_f32_fp8(raw4[i], false); // bytes 0,1
v2f hi_a = __builtin_amdgcn_cvt_pk_f32_fp8(raw4[i], true); // bytes 2,3
v2f lo_b = __builtin_amdgcn_cvt_pk_f32_fp8(raw4[i+1], false);
v2f hi_b = __builtin_amdgcn_cvt_pk_f32_fp8(raw4[i+1], true);
bp0[i/2] = pack_bf16x2(lo_a[0], lo_b[0]); // byte 0
bp1[i/2] = pack_bf16x2(lo_a[1], lo_b[1]); // byte 1
bp2[i/2] = pack_bf16x2(hi_a[0], hi_b[0]); // byte 2
bp3[i/2] = pack_bf16x2(hi_a[1], hi_b[1]); // byte 3
}
}
// Direct FP8→BF16 via v_cvt_scalef32_pk_bf16_fp8 + v_perm_b32 transpose.
// Same 32 VALU but 4 fewer intermediate VGPRs (packed BF16 vs F32 pairs).
// Better for some shapes where register pressure is the bottleneck.
__device__ __forceinline__ void
v_convert_bf16_packed_direct(const uint32_t raw4[8], uint32_t bp0[4], uint32_t bp1[4],
uint32_t bp2[4], uint32_t bp3[4]) {
float one = 1.0f;
#pragma unroll
for (int i = 0; i < 8; i += 2) {
uint32_t pk_lo_a, pk_hi_a, pk_lo_b, pk_hi_b;
asm("v_cvt_scalef32_pk_bf16_fp8 %0, %1, %2"
: "=v"(pk_lo_a) : "v"(raw4[i]), "v"(one));
asm("v_cvt_scalef32_pk_bf16_fp8 %0, %1, %2 op_sel:[1,0,0]"
: "=v"(pk_hi_a) : "v"(raw4[i]), "v"(one));
asm("v_cvt_scalef32_pk_bf16_fp8 %0, %1, %2"
: "=v"(pk_lo_b) : "v"(raw4[i+1]), "v"(one));
asm("v_cvt_scalef32_pk_bf16_fp8 %0, %1, %2 op_sel:[1,0,0]"
: "=v"(pk_hi_b) : "v"(raw4[i+1]), "v"(one));
asm("v_perm_b32 %0, %1, %2, %3"
: "=v"(bp0[i/2]) : "v"(pk_lo_b), "v"(pk_lo_a), "s"(0x05040100u));
asm("v_perm_b32 %0, %1, %2, %3"
: "=v"(bp1[i/2]) : "v"(pk_lo_b), "v"(pk_lo_a), "s"(0x07060302u));
asm("v_perm_b32 %0, %1, %2, %3"
: "=v"(bp2[i/2]) : "v"(pk_hi_b), "v"(pk_hi_a), "s"(0x05040100u));
asm("v_perm_b32 %0, %1, %2, %3"
: "=v"(bp3[i/2]) : "v"(pk_hi_b), "v"(pk_hi_a), "s"(0x07060302u));
}
}
// Dispatch wrapper: selects conversion method based on template parameter
template <bool DIRECT_CVT>
__device__ __forceinline__ void
v_convert_bf16_dispatch(const uint32_t raw4[8], uint32_t bp0[4], uint32_t bp1[4],
uint32_t bp2[4], uint32_t bp3[4]) {
if constexpr (DIRECT_CVT)
v_convert_bf16_packed_direct(raw4, bp0, bp1, bp2, bp3);
else
v_convert_bf16_packed(raw4, bp0, bp1, bp2, bp3);
}
// Half-packed FP8→BF16: converts 2 byte positions at a time using cvt_pk_f32_fp8.
// Produces 2 B-operand arrays for either lo (bytes 0,1) or hi (bytes 2,3).
// Uses 8 cvt_pk + 8 pack = 16 VALU per call (vs 24 scalar for 2 byte positions).
// Total for all 4 byte positions: 32 VALU (same as full packed, but only 8 VGPRs for outputs).
template <bool HI>
__device__ __forceinline__ void
v_convert_bf16_half_packed(const uint32_t raw4[8], uint32_t bp0[4], uint32_t bp1[4]) {
#pragma unroll
for (int i = 0; i < 8; i += 2) {
v2f a = __builtin_amdgcn_cvt_pk_f32_fp8(raw4[i], HI);
v2f b = __builtin_amdgcn_cvt_pk_f32_fp8(raw4[i+1], HI);
bp0[i/2] = pack_bf16x2(a[0], b[0]);
bp1[i/2] = pack_bf16x2(a[1], b[1]);
}
}
// Helper macro to do convert + MFMA for one byte position
#define V_BYTE_MFMA(raw4, byte_k, m_idx, a_bf16) \
do { \
uint32_t _bp[4]; \
if constexpr ((byte_k) == 0) v_convert_bf16<0>((raw4), _bp); \
else if constexpr ((byte_k) == 1) v_convert_bf16<1>((raw4), _bp); \
else if constexpr ((byte_k) == 2) v_convert_bf16<2>((raw4), _bp); \
else v_convert_bf16<3>((raw4), _bp); \
mfma_bf16_input_t _b = assemble_bf16_input(_bp); \
mfma_acc_t _c = {v_acc[(m_idx)][0], v_acc[(m_idx)][1], v_acc[(m_idx)][2], v_acc[(m_idx)][3]}; \
mfma_acc_t _r = mfma_f32_16x16x32_bf16((a_bf16), _b, _c); \
v_acc[(m_idx)][0]=_r[0]; v_acc[(m_idx)][1]=_r[1]; v_acc[(m_idx)][2]=_r[2]; v_acc[(m_idx)][3]=_r[3]; \
} while(0)
// Wide LDS read: load 32 bytes as 2× uint4 into mfma_input_t (for QK frag_b)
// Requires 16-byte aligned address (guaranteed by LDS_KV_STRIDE being multiple of 16).
__device__ __forceinline__ void
load_frag_b_wide(const uint8_t* kv_lds, int addr, mfma_input_t& frag_b) {
const uint4* p = reinterpret_cast<const uint4*>(&kv_lds[addr]);
uint4 lo = p[0];
uint4 hi = p[1];
frag_b[0] = lo.x; frag_b[1] = lo.y; frag_b[2] = lo.z; frag_b[3] = lo.w;
frag_b[4] = hi.x; frag_b[5] = hi.y; frag_b[6] = hi.z; frag_b[7] = hi.w;
}
__device__ __forceinline__ mfma_bf16_input_t
assemble_bf16_input(const uint32_t bp[4]) {
// Zero-copy reinterpret: union avoids element-by-element short copies
union { uint32_t u[4]; mfma_bf16_input_t v; } cvt;
cvt.u[0] = bp[0]; cvt.u[1] = bp[1]; cvt.u[2] = bp[2]; cvt.u[3] = bp[3];
return cvt.v;
}
// ============================================================================
// FP8-to-BF16 MFMA input conversion (for RoPE BF16 MFMAs)
// ============================================================================
// Convert 2 uint32 (8 FP8 bytes) → mfma_bf16_input_t (8 BF16 values)
__device__ __forceinline__ mfma_bf16_input_t
fp8x8_to_bf16_input(uint32_t w0, uint32_t w1) {
float f0 = __builtin_amdgcn_cvt_f32_fp8(w0, 0);
float f1 = __builtin_amdgcn_cvt_f32_fp8(w0, 1);
float f2 = __builtin_amdgcn_cvt_f32_fp8(w0, 2);
float f3 = __builtin_amdgcn_cvt_f32_fp8(w0, 3);
float f4 = __builtin_amdgcn_cvt_f32_fp8(w1, 0);
float f5 = __builtin_amdgcn_cvt_f32_fp8(w1, 1);
float f6 = __builtin_amdgcn_cvt_f32_fp8(w1, 2);
float f7 = __builtin_amdgcn_cvt_f32_fp8(w1, 3);
uint32_t bp[4] = {
pack_bf16x2(f0, f1),
pack_bf16x2(f2, f3),
pack_bf16x2(f4, f5),
pack_bf16x2(f6, f7)
};
return assemble_bf16_input(bp);
}
// ============================================================================
// Q scale computation (replaces 4 PyTorch kernel launches)
// ============================================================================
// Dead: q_scale computation removed (BF16 fused path doesn't need it)
// Dead: bf16_to_fp8_kernel removed (BF16 fused path does inline conversion)
// ============================================================================
// LDS Layout
// ============================================================================
constexpr int LDS_BUF_SIZE = KV_TILE * LDS_KV_STRIDE;
struct LDS {
uint8_t kv_fp8[2][LDS_BUF_SIZE];
};
// OCC4 single-buffer LDS: 18,432 bytes < 32,768 (128KB/4)
// Extra 64B padding: unified FP8 QK loop reads 64B past buffer at t=4 (lane_groups 2,3)
// Q is zeroed there so 0×garbage=0, but reads must hit valid LDS
struct LDS_OCC4 {
uint8_t kv_fp8[LDS_BUF_SIZE];
uint8_t _pad_rope[64];
};
// OCC2 triple-buffer LDS: 3 × 18,432 = 55,296 bytes < 65,536 (128KB/2)
// Enables issuing loads 2 tiles ahead for 2× memory latency hiding
struct LDS_OCC2 {
uint8_t kv_fp8[3][LDS_BUF_SIZE];
};
// OCC1 5-buffer LDS for V2 kernel: 5 × 18,432 = 92,160 bytes < 163,840 (160KB at occ=1)
// 4 retained buffers for FP8 V MFMA (K=128 = 4×32 tokens), plus 1 prefetch buffer
constexpr int V2_NUM_BUFS = 5;
constexpr int V2_MEGATILE = 4;
constexpr int V2_MEGATILE_TOKENS = V2_MEGATILE * KV_TILE; // 128
struct LDS_V2 {
uint8_t kv_fp8[V2_NUM_BUFS][LDS_BUF_SIZE]; // 92,160 bytes
};
// FP4 path: double-buffered FP4 loads + scale cache (native FP4 MFMA, no FP8 buffer)
constexpr int LDS_FP4_BUF_SIZE = KV_TILE * KV_FP4_STRIDE + (WAVESIZE - 1) * 16; // 10,224
struct LDS_FP4 {
uint8_t fp4[2][LDS_FP4_BUF_SIZE]; // 20,448 bytes (double-buffered FP4)
uint8_t scale_cache[KV_TILE * KV_SCALE_PER_ROW]; // 576 bytes (E8M0 scales for current tile)
float wave_amax[4]; // 16 bytes
}; // Total: ~21,040 bytes (fits at occ=3: 53,333 limit)
// ============================================================================
// FP4 loading (async buffer_load_dwordx4...lds for MXFP4 data)
// ============================================================================
__device__ __forceinline__ void
issue_kv_fp4_loads(
int lds_buf_offset,
const uint8_t* fp4_src,
int tile_count,
int tid
) {
const int wave_id = tid / WAVESIZE;
const int lane_id = tid % WAVESIZE;
const int voff = lane_id * 16;
auto rsrc = __builtin_amdgcn_make_buffer_rsrc(
const_cast<uint8_t*>(fp4_src), 0, tile_count * KV_FP4_STRIDE, 0x20000);
int row = wave_id;
for (; row + 4 < tile_count; row += 8) {
int r0 = __builtin_amdgcn_readfirstlane(row);
int soff0 = __builtin_amdgcn_readfirstlane(r0 * KV_FP4_STRIDE);
int lds0 = __builtin_amdgcn_readfirstlane(lds_buf_offset + r0 * KV_FP4_STRIDE);
int soff1 = __builtin_amdgcn_readfirstlane((r0 + 4) * KV_FP4_STRIDE);
int lds1 = __builtin_amdgcn_readfirstlane(lds_buf_offset + (r0 + 4) * KV_FP4_STRIDE);
asm volatile(
"s_mov_b32 m0, %[lds0]\n"
"buffer_load_dwordx4 %[voff], %[rsrc], %[soff0] offen lds\n"
"s_mov_b32 m0, %[lds1]\n"
"buffer_load_dwordx4 %[voff], %[rsrc], %[soff1] offen lds"
:
: [rsrc] "s" (rsrc),
[voff] "v" (voff),
[soff0] "s" (soff0), [lds0] "s" (lds0),
[soff1] "s" (soff1), [lds1] "s" (lds1)
: "memory", "m0"
);
}
if (row < tile_count) {
int r0 = __builtin_amdgcn_readfirstlane(row);
int soff0 = __builtin_amdgcn_readfirstlane(r0 * KV_FP4_STRIDE);
int lds0 = __builtin_amdgcn_readfirstlane(lds_buf_offset + r0 * KV_FP4_STRIDE);
asm volatile(
"s_mov_b32 m0, %[lds0]\n"
"buffer_load_dwordx4 %[voff], %[rsrc], %[soff0] offen lds"
:
: [rsrc] "s" (rsrc),
[voff] "v" (voff),
[soff0] "s" (soff0), [lds0] "s" (lds0)
: "memory", "m0"
);
}
}
// ============================================================================
// Native FP4 MFMA helpers (no dequant — FP4 used directly in MFMA)
// ============================================================================
// Mixed FP8(A) × FP4(B) MFMA with hardware per-block E8M0 scaling on B
// scale_b: 4 packed E8M0 bytes, one per 32-element K block
__device__ __forceinline__ mfma_acc_t
mfma_f32_16x16x128_fp8_fp4_scaled(mfma_input_t a, mfma_input_t b, mfma_acc_t c, int scale_b) {
return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a, b, c, /*cbsz=*/0, /*blgp=*/4, /*opsel_a=*/0, /*scale_a=*/0, /*opsel_b=*/1, scale_b);
}
// Load FP4 B-fragment: 16 bytes into lower 4 dwords, zero upper 4
__device__ __forceinline__ void
load_frag_b_fp4(const uint8_t* fp4_lds, int addr, mfma_input_t& frag_b) {
const uint4* p = reinterpret_cast<const uint4*>(&fp4_lds[addr]);
uint4 lo = p[0];
frag_b[0] = lo.x; frag_b[1] = lo.y; frag_b[2] = lo.z; frag_b[3] = lo.w;
frag_b[4] = 0; frag_b[5] = 0; frag_b[6] = 0; frag_b[7] = 0;
}
// FP4 MFMA bytes per iteration: 128 FP4 elements × 0.5 bytes = 64 bytes
constexpr int FP4_MFMA_BYTES = 64;
// Hardware FP4->F32 conversion: preload FP4 data and E8M0 scales for V group
// Uses uint16_t loads (always 2-byte aligned since v_fp4_off = V_dim_index/2 is even)
__device__ __forceinline__ void
load_v_fp4_data(const uint8_t* fp4_lds, int v_fp4_off, int lane_group,
const uint8_t* scale_cache, int scale_block,
uint32_t fp4_raw[8], float scales[8]) {
int base_pos = lane_group * 8;
#pragma unroll
for (int i = 0; i < 8; i++) {
int pos = base_pos + i;
// Load only 2 FP4 bytes (4 FP4 values = 4 V dims) — 2-byte aligned
fp4_raw[i] = (uint32_t)*reinterpret_cast<const uint16_t*>(
&fp4_lds[pos * KV_FP4_STRIDE + v_fp4_off]);
uint8_t e8m0 = scale_cache[pos * KV_SCALE_PER_ROW + scale_block];
scales[i] = __uint_as_float((uint32_t)e8m0 << 23);
}
}
// Convert FP4 byte pair to BF16 MFMA inputs using hardware cvt instruction
// BYTE_IDX: 0 for V dims 0,1; 1 for V dims 2,3
// Produces bp_lo (lower nibble, even V dim) and bp_hi (upper nibble, odd V dim)
using floatx2_t = __attribute__((ext_vector_type(2))) float;
template <int BYTE_IDX>
__device__ __forceinline__ void
v_cvt_fp4_hw_pair(const uint32_t fp4_raw[8], const float scales[8],
uint32_t bp_lo[4], uint32_t bp_hi[4]) {
#pragma unroll
for (int i = 0; i < 8; i += 2) {
floatx2_t r0 = __builtin_amdgcn_cvt_scalef32_pk_f32_fp4(
fp4_raw[i], scales[i], BYTE_IDX);
floatx2_t r1 = __builtin_amdgcn_cvt_scalef32_pk_f32_fp4(
fp4_raw[i+1], scales[i+1], BYTE_IDX);
bp_lo[i/2] = pack_bf16x2(r0[0], r1[0]);
bp_hi[i/2] = pack_bf16x2(r0[1], r1[1]);
}
}
// ============================================================================
// Score extraction
// ============================================================================
// Score extraction: shuffle MFMA accumulators across lanes.
// IMPORTANT: Do NOT take address of acc arrays (pointer indirection causes
// compiler to use stack-spilled values instead of live MFMA registers).
// Pass individual floats to avoid the issue.
__device__ __forceinline__ float
extract_score_4(float a0, float a1, float a2, float a3,
int head, int src) {
float v0 = __shfl(a0, src);
float v1 = __shfl(a1, src);
float v2 = __shfl(a2, src);
float v3 = __shfl(a3, src);
int k = head % 4;
return (k == 0) ? v0 : (k == 1) ? v1 : (k == 2) ? v2 : v3;
}
// ============================================================================
// Cooperative KV tile load (async buffer_load_dwordx4...lds)
// ============================================================================
// Contiguous chunk-based async loading: HBM→LDS via buffer_load_dwordx4.
// With LDS_KV_STRIDE == KV_FP8_STRIDE == 576, rows are contiguous in both
// HBM and LDS. Instead of loading one row per instruction (64×16=1024 bytes,
// 44% wasted on the 576-byte rows), we load sequential 1024-byte chunks that
// span multiple rows. For 32 rows: ceil(32*576/1024) = 18 chunks instead of
// 32 instructions — 44% bandwidth reduction.
// OOB bytes (beyond tile_count*576) are zeroed by hardware and land in unused
// LDS slots beyond the valid rows.
__device__ __forceinline__ void
issue_kv_loads(
int lds_buf_offset,
const uint8_t* kv_src,
int tile_count,
int tid
) {
const int wave_id = tid / WAVESIZE;
const int lane_id = tid % WAVESIZE;
const int voff = lane_id * 16;
int total_bytes = tile_count * KV_FP8_STRIDE;
int num_chunks = (total_bytes + 1023) / 1024;
auto rsrc = __builtin_amdgcn_make_buffer_rsrc(
const_cast<uint8_t*>(kv_src), 0, total_bytes, 0x20000);
int chunk = wave_id;
for (; chunk + 4 < num_chunks; chunk += 8) {
int c0 = __builtin_amdgcn_readfirstlane(chunk);
int soff0 = __builtin_amdgcn_readfirstlane(c0 * 1024);
int lds0 = __builtin_amdgcn_readfirstlane(lds_buf_offset + c0 * 1024);
int c1 = __builtin_amdgcn_readfirstlane(c0 + 4);
int soff1 = __builtin_amdgcn_readfirstlane(c1 * 1024);
int lds1 = __builtin_amdgcn_readfirstlane(lds_buf_offset + c1 * 1024);
asm volatile(
"s_mov_b32 m0, %[lds0]\n"
"buffer_load_dwordx4 %[voff], %[rsrc], %[soff0] offen lds\n"
"s_mov_b32 m0, %[lds1]\n"
"buffer_load_dwordx4 %[voff], %[rsrc], %[soff1] offen lds"
:
: [rsrc] "s" (rsrc),
[voff] "v" (voff),
[soff0] "s" (soff0), [lds0] "s" (lds0),
[soff1] "s" (soff1), [lds1] "s" (lds1)
: "memory", "m0"
);
}
if (chunk < num_chunks) {
int c0 = __builtin_amdgcn_readfirstlane(chunk);
int soff0 = __builtin_amdgcn_readfirstlane(c0 * 1024);
int lds0 = __builtin_amdgcn_readfirstlane(lds_buf_offset + c0 * 1024);
asm volatile(
"s_mov_b32 m0, %[lds0]\n"
"buffer_load_dwordx4 %[voff], %[rsrc], %[soff0] offen lds"
:
: [rsrc] "s" (rsrc),
[voff] "v" (voff),
[soff0] "s" (soff0), [lds0] "s" (lds0)
: "memory", "m0"
);
}
}
__device__ __forceinline__ void wait_kv_loads() {
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
}
// Strided KV tile load: loads tile_count tokens from non-contiguous positions.
// Token i comes from kv_base[(first_token + i*stride) * QK_DIM] and is packed
// contiguously in lds.kv_fp8[buf_idx] for MFMA processing. Uses VGPR-mediated
// global_load → ds_write (not DMA). Same LDS layout as issue_kv_loads — callers
// don't see the difference. Synchronous: all LDS writes complete before return.
template <typename LDS_T>
__device__ __forceinline__ void
issue_kv_loads_strided(
LDS_T& lds,
int buf_idx,
const uint8_t* __restrict__ kv_base,
int first_token,
int stride,
int tile_count,
int tid
) {
constexpr int ITEMS_PER_TOKEN = QK_DIM / 16; // 576/16 = 36
const int total_items = tile_count * ITEMS_PER_TOKEN;
for (int item = tid; item < total_items; item += 256) {
int tok = item / ITEMS_PER_TOKEN;
int sub = item - tok * ITEMS_PER_TOKEN;
int byte_off = sub * 16;
int global_token = first_token + tok * stride;
uint4 val = *reinterpret_cast<const uint4*>(
kv_base + (size_t)global_token * QK_DIM + byte_off);
*reinterpret_cast<uint4*>(
&lds.kv_fp8[buf_idx][tok * LDS_KV_STRIDE + byte_off]) = val;
}
asm volatile("s_waitcnt vmcnt(0) lgkmcnt(0)" ::: "memory");
}
// ============================================================================
// COMMON: tile loop body (shared between FP8 and BF16 paths)
// ============================================================================
// Both paths use the same tile processing once q_cache is populated.
// Factored into a macro-like inline to avoid code duplication.
// ============================================================================
// Stage 1 — FP8 path (for kv=1024: pre-converted Q)
// ============================================================================
template <int NHEAD, bool DIRECT_CVT = false>
__global__ void
__launch_bounds__(256, 3)
mla_decode_stage1_fp8(
const uint8_t* __restrict__ q_fp8,
const float* __restrict__ q_scales,
const uint8_t* __restrict__ kv_fp8,
const float* __restrict__ kv_scale_ptr,
__hip_bfloat16* __restrict__ mid_o,
float* __restrict__ mid_lse,
__hip_bfloat16* __restrict__ output,
const int32_t* __restrict__ batch_map,
const int32_t* __restrict__ kv_indptr,
int batch_size,
int num_kv_splits
) {
constexpr float sm_scale = 1.0f / 24.0f;
const float kv_scale = *kv_scale_ptr;
constexpr int NUM_HEAD_GROUPS = (NHEAD + HEAD_GROUP - 1) / HEAD_GROUP;
const int q_pos = blockIdx.x;
const int split_id = blockIdx.y;
const int tid = threadIdx.x;
const int wave_id = tid / WAVESIZE;
const int lane_id = tid % WAVESIZE;
const int lane_group = lane_id / 16;
const int lane_col = lane_id % 16;
const int batch_id = batch_map[q_pos];
const int kv_start = kv_indptr[batch_id];
const int kv_end = kv_indptr[batch_id + 1];
const int total_kv_len = kv_end - kv_start;
int positions_per_split = (total_kv_len + num_kv_splits - 1) / num_kv_splits;
int split_start = split_id * positions_per_split;
int split_end = min(split_start + positions_per_split, total_kv_len);
__shared__ LDS lds;
for (int hg = 0; hg < NUM_HEAD_GROUPS; hg++) {
int h_start = hg * HEAD_GROUP;
int h_count = min(HEAD_GROUP, NHEAD - h_start);
float my_scale = 0.0f;
if (lane_col < h_count) {
int abs_h = h_start + lane_col;
my_scale = q_scales[q_pos * NHEAD + abs_h] * kv_scale * sm_scale * LOG2E_F;
}
float my_head_max = -1e30f;
float my_head_sum = 0.0f;
mfma_acc_t v_acc[V_MFMAS_PER_WAVE];
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++)
v_acc[m] = {0.0f, 0.0f, 0.0f, 0.0f};
if (split_start >= split_end) {
if (num_kv_splits > 1) {
__hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM
+ split_id * V_DIM;
float* ml = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits
+ split_id;
__hip_bfloat16 zero_bf16 = __float2bfloat16(0.0f);
#pragma unroll
for (int k = 0; k < 4; k++) {
int head = lane_group * 4 + k;
if (head < h_count) {
int abs_h = h_start + head;
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++) {
int v_idx = wave_id * 128 + (m / 4) * 64 + lane_col * 4 + (m % 4);
if (v_idx < V_DIM)
mo[abs_h * num_kv_splits * V_DIM + v_idx] = zero_bf16;
}
}
}
if (wave_id == 0 && lane_group == 0 && lane_col < h_count)
ml[(h_start + lane_col) * num_kv_splits] = -1e30f;
}
__syncthreads();
continue;
}
const uint8_t* q_base = q_fp8 + (size_t)(q_pos * NHEAD + h_start) * QK_DIM;
int first_count = min(KV_TILE, split_end - split_start);
issue_kv_loads(0, kv_fp8 + (size_t)(kv_start + split_start) * QK_DIM,
first_count, tid);
uint32_t q_cache[QK_MFMAS][8];
#pragma unroll
for (int t = 0; t < QK_MFMAS; t++) {
int q_byte = t * FP8_MFMA_K + lane_group * 32;
if (q_byte + 32 <= QK_DIM && lane_col < h_count) {
const uint32_t* ap = reinterpret_cast<const uint32_t*>(
q_base + (size_t)lane_col * QK_DIM + q_byte);
q_cache[t][0]=ap[0]; q_cache[t][1]=ap[1]; q_cache[t][2]=ap[2]; q_cache[t][3]=ap[3];
q_cache[t][4]=ap[4]; q_cache[t][5]=ap[5]; q_cache[t][6]=ap[6]; q_cache[t][7]=ap[7];
} else {
q_cache[t][0]=0; q_cache[t][1]=0; q_cache[t][2]=0; q_cache[t][3]=0;
q_cache[t][4]=0; q_cache[t][5]=0; q_cache[t][6]=0; q_cache[t][7]=0;
}
}
wait_kv_loads();
__syncthreads();
int buf = 0;
for (int tile_start = split_start; tile_start < split_end; tile_start += KV_TILE) {
int tile_count = min(KV_TILE, split_end - tile_start);
int next_start = tile_start + KV_TILE;
int has_next = (next_start < split_end);
if (__builtin_expect(has_next, 1)) {
int next_count = min(KV_TILE, split_end - next_start);
issue_kv_loads((buf ^ 1) * LDS_BUF_SIZE,
kv_fp8 + (size_t)(kv_start + next_start) * QK_DIM,
next_count, tid);
}
const uint8_t* kv_lds = lds.kv_fp8[buf];
mfma_acc_t acc_lo = {0.0f, 0.0f, 0.0f, 0.0f};
mfma_acc_t acc_hi = {0.0f, 0.0f, 0.0f, 0.0f};
// Unified FP8 QK scoring for all dims 0-575 (5 iterations)
// t=4: lane_groups 2,3 exceed QK_DIM — zero KV to avoid FP8 NaN (0x80)
asm volatile("s_setprio 3" :::);
#pragma unroll
for (int t = 0; t < QK_MFMAS; t++) {
mfma_input_t frag_q = {q_cache[t][0], q_cache[t][1], q_cache[t][2], q_cache[t][3],
q_cache[t][4], q_cache[t][5], q_cache[t][6], q_cache[t][7]};
int kv_byte = t * FP8_MFMA_K + lane_group * 32;
mfma_input_t frag_kv_lo, frag_kv_hi;
if (kv_byte + 32 <= QK_DIM) {
load_frag_b_wide(kv_lds, lane_col * LDS_KV_STRIDE + kv_byte, frag_kv_lo);
load_frag_b_wide(kv_lds, (lane_col + KV_HALF) * LDS_KV_STRIDE + kv_byte, frag_kv_hi);
} else {
frag_kv_lo[0]=0; frag_kv_lo[1]=0; frag_kv_lo[2]=0; frag_kv_lo[3]=0;
frag_kv_lo[4]=0; frag_kv_lo[5]=0; frag_kv_lo[6]=0; frag_kv_lo[7]=0;
frag_kv_hi[0]=0; frag_kv_hi[1]=0; frag_kv_hi[2]=0; frag_kv_hi[3]=0;
frag_kv_hi[4]=0; frag_kv_hi[5]=0; frag_kv_hi[6]=0; frag_kv_hi[7]=0;
}
acc_lo = mfma_f32_16x16x128_fp8(frag_kv_lo, frag_q, acc_lo);
acc_hi = mfma_f32_16x16x128_fp8(frag_kv_hi, frag_q, acc_hi);
}
asm volatile("s_setprio 0" ::: "memory");
// Issue V preloads for BOTH groups before softmax — eliminates
// LDS reads during V MFMAs for cleaner MFMA throughput
uint32_t raw4_cur[8], raw4_next_v[8];
int blo = lane_group * 4;
int bhi = lane_group * 4 + 16;
int v_off0 = wave_id * 128 + lane_col * 4;
int v_off1 = v_off0 + 64;
{
#pragma unroll
for (int i = 0; i < 4; i++) {
raw4_cur[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(blo + i) * LDS_KV_STRIDE + v_off0]);
raw4_cur[i+4] = *reinterpret_cast<const uint32_t*>(&kv_lds[(bhi + i) * LDS_KV_STRIDE + v_off0]);
raw4_next_v[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(blo + i) * LDS_KV_STRIDE + v_off1]);
raw4_next_v[i+4] = *reinterpret_cast<const uint32_t*>(&kv_lds[(bhi + i) * LDS_KV_STRIDE + v_off1]);
}
}
// Softmax: FMA fusion — max over raw acc, fmaf combines scale+sub
float partial_max = -1e30f;
partial_max = fmaxf(partial_max, fmaxf(acc_lo[0], acc_lo[1]));
partial_max = fmaxf(partial_max, fmaxf(acc_lo[2], acc_lo[3]));
partial_max = fmaxf(partial_max, fmaxf(acc_hi[0], acc_hi[1]));
partial_max = fmaxf(partial_max, fmaxf(acc_hi[2], acc_hi[3]));
partial_max = fmaxf(partial_max, __shfl_xor(partial_max, 16));
partial_max = fmaxf(partial_max, __shfl_xor(partial_max, 32));
float scaled_partial_max = partial_max * my_scale;
float old_max = my_head_max;
float new_max = fmaxf(old_max, scaled_partial_max);
my_head_max = new_max;
float neg_max = -new_max;
if (scaled_partial_max > old_max && old_max > -1e29f) {
float my_rescale = __builtin_amdgcn_exp2f(old_max - new_max);
my_head_sum *= my_rescale;
#pragma unroll
for (int k = 0; k < 4; k++) {
float r = __shfl(my_rescale, lane_group * 4 + k);
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++)
v_acc[m][k] *= r;
}
}
uint32_t a_packed[4];
float tile_sum_partial = 0.0f;
#pragma unroll
for (int j = 0; j < 4; j += 2) {
float es0 = __builtin_amdgcn_exp2f(fmaf(acc_lo[j], my_scale, neg_max));
float es1 = __builtin_amdgcn_exp2f(fmaf(acc_lo[j+1], my_scale, neg_max));
tile_sum_partial += es0 + es1;
a_packed[j / 2] = pack_bf16x2(es0, es1);
}
#pragma unroll
for (int j = 0; j < 4; j += 2) {
float es0 = __builtin_amdgcn_exp2f(fmaf(acc_hi[j], my_scale, neg_max));
float es1 = __builtin_amdgcn_exp2f(fmaf(acc_hi[j+1], my_scale, neg_max));
tile_sum_partial += es0 + es1;
a_packed[2 + j / 2] = pack_bf16x2(es0, es1);
}
tile_sum_partial += __shfl_xor(tile_sum_partial, 16);
tile_sum_partial += __shfl_xor(tile_sum_partial, 32);
my_head_sum += tile_sum_partial;
// Zero-copy reinterpret a_packed as mfma_bf16_input_t via union
union { uint32_t u[4]; mfma_bf16_input_t v; } a_cvt;
a_cvt.u[0] = a_packed[0]; a_cvt.u[1] = a_packed[1];
a_cvt.u[2] = a_packed[2]; a_cvt.u[3] = a_packed[3];
mfma_bf16_input_t a_bf16 = a_cvt.v;
// Vectorized V accumulation — both groups already preloaded
{
// Group 0: packed convert all 4 byte positions
asm volatile("s_setprio 3" :::);
uint32_t bp0[4], bp1[4], bp2[4], bp3[4];
v_convert_bf16_dispatch<DIRECT_CVT>(raw4_cur, bp0, bp1, bp2, bp3);
// Group 0 MFMAs
v_acc[0] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp0), v_acc[0]);
v_acc[1] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp1), v_acc[1]);
v_acc[2] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp2), v_acc[2]);
v_acc[3] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp3), v_acc[3]);
// Group 1: packed convert + 4 MFMAs (raw4_next_v already loaded before softmax)
uint32_t bp0b[4], bp1b[4], bp2b[4], bp3b[4];
v_convert_bf16_dispatch<DIRECT_CVT>(raw4_next_v, bp0b, bp1b, bp2b, bp3b);
v_acc[4] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp0b), v_acc[4]);
v_acc[5] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp1b), v_acc[5]);
v_acc[6] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp2b), v_acc[6]);
v_acc[7] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp3b), v_acc[7]);
asm volatile("s_setprio 0" :::);
}
if (__builtin_expect(has_next, 1)) wait_kv_loads();
__syncthreads();
buf ^= 1;
}
if (num_kv_splits == 1) {
#pragma unroll
for (int k = 0; k < 4; k++) {
int abs_h = h_start + lane_group * 4 + k;
float hs = __shfl(my_head_sum, lane_group * 4 + k);
float final_scale = (hs > 0.0f) ? kv_scale / hs : 0.0f;
__hip_bfloat16* o_ptr = output + (size_t)q_pos * NHEAD * V_DIM + abs_h * V_DIM;
#pragma unroll
for (int g = 0; g < 2; g++) {
int base_idx = wave_id * 128 + g * 64 + lane_col * 4;
uint32_t pk0 = pack_bf16x2(v_acc[g*4][k] * final_scale, v_acc[g*4+1][k] * final_scale);
uint32_t pk1 = pack_bf16x2(v_acc[g*4+2][k] * final_scale, v_acc[g*4+3][k] * final_scale);
*reinterpret_cast<uint64_t*>(&o_ptr[base_idx]) = ((uint64_t)pk1 << 32) | (uint64_t)pk0;
}
}
} else {
// Layout: [q_pos][head][split][v_dim] — stage2-friendly (contiguous splits per head)
__hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM
+ split_id * V_DIM;
float* ml = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits
+ split_id;
#pragma unroll
for (int k = 0; k < 4; k++) {
int abs_h = h_start + lane_group * 4 + k;
float hs = __shfl(my_head_sum, lane_group * 4 + k);
float final_scale = (hs > 0.0f) ? kv_scale / hs : 0.0f;
#pragma unroll
for (int g = 0; g < 2; g++) {
int base_idx = wave_id * 128 + g * 64 + lane_col * 4;
uint32_t pk0 = pack_bf16x2(v_acc[g*4][k] * final_scale, v_acc[g*4+1][k] * final_scale);
uint32_t pk1 = pack_bf16x2(v_acc[g*4+2][k] * final_scale, v_acc[g*4+3][k] * final_scale);
*reinterpret_cast<uint64_t*>(&mo[abs_h * num_kv_splits * V_DIM + base_idx]) = ((uint64_t)pk1 << 32) | (uint64_t)pk0;
}
}
if (wave_id == 0 && lane_group == 0) {
int abs_h = h_start + lane_col;
ml[abs_h * num_kv_splits] = (my_head_sum > 0.0f) ?
my_head_max * LN2_F + __logf(my_head_sum) : -1e30f;
}
}
__syncthreads();
}
}
// ============================================================================
// Stage 1 — FP8 OCC4 path (occupancy=4, single-buffer, trade SW double-buffer
// for HW wave interleaving across 4 WGs/CU)
// ============================================================================
template <int NHEAD>
__global__ void
__launch_bounds__(256, 4)
mla_decode_stage1_fp8_occ4(
const uint8_t* __restrict__ q_fp8,
const float* __restrict__ q_scales,
const uint8_t* __restrict__ kv_fp8,
const float* __restrict__ kv_scale_ptr,
__hip_bfloat16* __restrict__ mid_o,
float* __restrict__ mid_lse,
__hip_bfloat16* __restrict__ output,
const int32_t* __restrict__ batch_map,
const int32_t* __restrict__ kv_indptr,
int batch_size,
int num_kv_splits
) {
constexpr float sm_scale = 1.0f / 24.0f;
const float kv_scale = *kv_scale_ptr;
constexpr int NUM_HEAD_GROUPS = (NHEAD + HEAD_GROUP - 1) / HEAD_GROUP;
const int q_pos = blockIdx.x;
const int split_id = blockIdx.y;
const int tid = threadIdx.x;
const int wave_id = tid / WAVESIZE;
const int lane_id = tid % WAVESIZE;
const int lane_group = lane_id / 16;
const int lane_col = lane_id % 16;
const int batch_id = batch_map[q_pos];
const int kv_start = kv_indptr[batch_id];
const int kv_end = kv_indptr[batch_id + 1];
const int total_kv_len = kv_end - kv_start;
int positions_per_split = (total_kv_len + num_kv_splits - 1) / num_kv_splits;
int split_start = split_id * positions_per_split;
int split_end = min(split_start + positions_per_split, total_kv_len);
__shared__ LDS_OCC4 lds;
for (int hg = 0; hg < NUM_HEAD_GROUPS; hg++) {
int h_start = hg * HEAD_GROUP;
int h_count = min(HEAD_GROUP, NHEAD - h_start);
float my_scale = 0.0f;
if (lane_col < h_count) {
int abs_h = h_start + lane_col;
my_scale = q_scales[q_pos * NHEAD + abs_h] * kv_scale * sm_scale * LOG2E_F;
}
float my_head_max = -1e30f;
float my_head_sum = 0.0f;
mfma_acc_t v_acc[V_MFMAS_PER_WAVE];
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++)
v_acc[m] = {0.0f, 0.0f, 0.0f, 0.0f};
if (split_start >= split_end) {
if (num_kv_splits > 1) {
__hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM
+ split_id * V_DIM;
float* ml = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits
+ split_id;
__hip_bfloat16 zero_bf16 = __float2bfloat16(0.0f);
#pragma unroll
for (int k = 0; k < 4; k++) {
int head = lane_group * 4 + k;
if (head < h_count) {
int abs_h = h_start + head;
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++) {
int v_idx = wave_id * 128 + (m / 4) * 64 + lane_col * 4 + (m % 4);
if (v_idx < V_DIM)
mo[abs_h * num_kv_splits * V_DIM + v_idx] = zero_bf16;
}
}
}
if (wave_id == 0 && lane_group == 0 && lane_col < h_count)
ml[(h_start + lane_col) * num_kv_splits] = -1e30f;
}
__syncthreads();
continue;
}
const uint8_t* q_base = q_fp8 + (size_t)(q_pos * NHEAD + h_start) * QK_DIM;
// Pre-cache Q in VGPRs (same as occ=3 path)
uint32_t q_cache[QK_MFMAS][8];
#pragma unroll
for (int t = 0; t < QK_MFMAS; t++) {
int q_byte = t * FP8_MFMA_K + lane_group * 32;
if (q_byte + 32 <= QK_DIM && lane_col < h_count) {
const uint32_t* ap = reinterpret_cast<const uint32_t*>(
q_base + (size_t)lane_col * QK_DIM + q_byte);
q_cache[t][0]=ap[0]; q_cache[t][1]=ap[1]; q_cache[t][2]=ap[2]; q_cache[t][3]=ap[3];
q_cache[t][4]=ap[4]; q_cache[t][5]=ap[5]; q_cache[t][6]=ap[6]; q_cache[t][7]=ap[7];
} else {
q_cache[t][0]=0; q_cache[t][1]=0; q_cache[t][2]=0; q_cache[t][3]=0;
q_cache[t][4]=0; q_cache[t][5]=0; q_cache[t][6]=0; q_cache[t][7]=0;
}
}
// Overlapped tile loop: preload V into VGPRs → release LDS → issue next loads
// Softmax + V accumulation overlap with next tile's HBM→LDS loads
{
int first_tile_count = min(KV_TILE, split_end - split_start);
issue_kv_loads(0, kv_fp8 + (size_t)(kv_start + split_start) * QK_DIM,
first_tile_count, tid);
}
for (int tile_start = split_start; tile_start < split_end; tile_start += KV_TILE) {
int tile_count = min(KV_TILE, split_end - tile_start);
bool has_next = (tile_start + KV_TILE < split_end);
// Wait for current tile's loads to complete
wait_kv_loads();
__syncthreads();
const uint8_t* kv_lds = lds.kv_fp8;
mfma_acc_t acc_lo = {0.0f, 0.0f, 0.0f, 0.0f};
mfma_acc_t acc_hi = {0.0f, 0.0f, 0.0f, 0.0f};
// QK scoring: 5 FP8 MFMAs (reads LDS)
asm volatile("s_setprio 3" :::);
#pragma unroll
for (int t = 0; t < QK_MFMAS; t++) {
mfma_input_t frag_q = {q_cache[t][0], q_cache[t][1], q_cache[t][2], q_cache[t][3],
q_cache[t][4], q_cache[t][5], q_cache[t][6], q_cache[t][7]};
int kv_byte = t * FP8_MFMA_K + lane_group * 32;
mfma_input_t frag_kv_lo, frag_kv_hi;
if (kv_byte + 32 <= QK_DIM) {
load_frag_b_wide(kv_lds, lane_col * LDS_KV_STRIDE + kv_byte, frag_kv_lo);
load_frag_b_wide(kv_lds, (lane_col + KV_HALF) * LDS_KV_STRIDE + kv_byte, frag_kv_hi);
} else {
frag_kv_lo[0]=0; frag_kv_lo[1]=0; frag_kv_lo[2]=0; frag_kv_lo[3]=0;
frag_kv_lo[4]=0; frag_kv_lo[5]=0; frag_kv_lo[6]=0; frag_kv_lo[7]=0;
frag_kv_hi[0]=0; frag_kv_hi[1]=0; frag_kv_hi[2]=0; frag_kv_hi[3]=0;
frag_kv_hi[4]=0; frag_kv_hi[5]=0; frag_kv_hi[6]=0; frag_kv_hi[7]=0;
}
mfma_acc_t c_lo = {acc_lo[0], acc_lo[1], acc_lo[2], acc_lo[3]};
mfma_acc_t r_lo = mfma_f32_16x16x128_fp8(frag_kv_lo, frag_q, c_lo);
acc_lo[0]=r_lo[0]; acc_lo[1]=r_lo[1]; acc_lo[2]=r_lo[2]; acc_lo[3]=r_lo[3];
mfma_acc_t c_hi = {acc_hi[0], acc_hi[1], acc_hi[2], acc_hi[3]};
mfma_acc_t r_hi = mfma_f32_16x16x128_fp8(frag_kv_hi, frag_q, c_hi);
acc_hi[0]=r_hi[0]; acc_hi[1]=r_hi[1]; acc_hi[2]=r_hi[2]; acc_hi[3]=r_hi[3];
}
asm volatile("s_setprio 0" ::: "memory");
// Preload V data into VGPRs BEFORE releasing LDS (like occ3 pattern)
uint32_t raw4_cur[8], raw4_next[8];
{
int blo = lane_group * 4;
int bhi = lane_group * 4 + 16;
int v_off0 = wave_id * 128 + lane_col * 4;
int v_off1 = wave_id * 128 + 64 + lane_col * 4;
#pragma unroll
for (int i = 0; i < 4; i++) {
raw4_cur[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(blo + i) * LDS_KV_STRIDE + v_off0]);
raw4_cur[i+4] = *reinterpret_cast<const uint32_t*>(&kv_lds[(bhi + i) * LDS_KV_STRIDE + v_off0]);
raw4_next[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(blo + i) * LDS_KV_STRIDE + v_off1]);
raw4_next[i+4] = *reinterpret_cast<const uint32_t*>(&kv_lds[(bhi + i) * LDS_KV_STRIDE + v_off1]);
}
}
// LDS reads complete — release LDS and start next tile's loads (overlap!)
__syncthreads();
if (__builtin_expect(has_next, 1)) {
int next_start = tile_start + KV_TILE;
int next_count = min(KV_TILE, split_end - next_start);
issue_kv_loads(0, kv_fp8 + (size_t)(kv_start + next_start) * QK_DIM,
next_count, tid);
}
// Score extraction + softmax (pure VALU — overlaps with HBM loads)
float scores[8];
scores[0] = acc_lo[0]; scores[1] = acc_lo[1]; scores[2] = acc_lo[2]; scores[3] = acc_lo[3];
scores[4] = acc_hi[0]; scores[5] = acc_hi[1]; scores[6] = acc_hi[2]; scores[7] = acc_hi[3];
float partial_max = -1e30f;
#pragma unroll
for (int j = 0; j < 4; j++)
partial_max = fmaxf(partial_max, scores[j] * my_scale);
#pragma unroll
for (int j = 0; j < 4; j++)
partial_max = fmaxf(partial_max, scores[j+4] * my_scale);
partial_max = fmaxf(partial_max, __shfl_xor(partial_max, 16));
partial_max = fmaxf(partial_max, __shfl_xor(partial_max, 32));
float old_max = my_head_max;
float new_max = fmaxf(old_max, partial_max);
my_head_max = new_max;
if (partial_max > old_max && old_max > -1e29f) {
float my_rescale = __builtin_amdgcn_exp2f(old_max - new_max);
my_head_sum *= my_rescale;
#pragma unroll
for (int k = 0; k < 4; k++) {
float r = __shfl(my_rescale, lane_group * 4 + k);
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++)
v_acc[m][k] *= r;
}
}
uint32_t a_packed[4];
float tile_sum_partial = 0.0f;
#pragma unroll
for (int j = 0; j < 4; j += 2) {
float es0 = __builtin_amdgcn_exp2f(scores[j] * my_scale - my_head_max);
float es1 = __builtin_amdgcn_exp2f(scores[j+1] * my_scale - my_head_max);
tile_sum_partial += es0 + es1;
a_packed[j / 2] = pack_bf16x2(es0, es1);
}
#pragma unroll
for (int j = 0; j < 4; j += 2) {
float es0 = __builtin_amdgcn_exp2f(scores[j+4] * my_scale - my_head_max);
float es1 = __builtin_amdgcn_exp2f(scores[j+5] * my_scale - my_head_max);
tile_sum_partial += es0 + es1;
a_packed[2 + j / 2] = pack_bf16x2(es0, es1);
}
tile_sum_partial += __shfl_xor(tile_sum_partial, 16);
tile_sum_partial += __shfl_xor(tile_sum_partial, 32);
my_head_sum += tile_sum_partial;
// Zero-copy reinterpret a_packed as mfma_bf16_input_t via union
union { uint32_t u[4]; mfma_bf16_input_t v; } a_cvt;
a_cvt.u[0] = a_packed[0]; a_cvt.u[1] = a_packed[1];
a_cvt.u[2] = a_packed[2]; a_cvt.u[3] = a_packed[3];
mfma_bf16_input_t a_bf16 = a_cvt.v;
// V accumulation from preloaded VGPRs (packed convert, overlaps with HBM loads)
asm volatile("s_setprio 3" :::);
{
// Group 0: packed convert all 4 byte positions, then 4 MFMAs
{
uint32_t bp0[4], bp1[4], bp2[4], bp3[4];
v_convert_bf16_packed(raw4_cur, bp0, bp1, bp2, bp3);
v_acc[0] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp0), v_acc[0]);
v_acc[1] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp1), v_acc[1]);
v_acc[2] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp2), v_acc[2]);
v_acc[3] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp3), v_acc[3]);
}
// Group 1: packed convert + 4 MFMAs
{
uint32_t bp0[4], bp1[4], bp2[4], bp3[4];
v_convert_bf16_packed(raw4_next, bp0, bp1, bp2, bp3);
v_acc[4] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp0), v_acc[4]);
v_acc[5] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp1), v_acc[5]);
v_acc[6] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp2), v_acc[6]);
v_acc[7] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp3), v_acc[7]);
}
}
asm volatile("s_setprio 0" ::: "memory");
}
// Output (same as occ=3 but scalar stores to stay within occ4 VGPR budget)
if (num_kv_splits == 1) {
#pragma unroll
for (int k = 0; k < 4; k++) {
int abs_h = h_start + lane_group * 4 + k;
float hs = __shfl(my_head_sum, lane_group * 4 + k);
float final_scale = (hs > 0.0f) ? kv_scale / hs : 0.0f;
__hip_bfloat16* o_ptr = output + (size_t)q_pos * NHEAD * V_DIM + abs_h * V_DIM;
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++) {
int v_idx = wave_id * 128 + (m / 4) * 64 + lane_col * 4 + (m % 4);
o_ptr[v_idx] = __float2bfloat16(v_acc[m][k] * final_scale);
}
}
} else {
__hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM
+ split_id * V_DIM;
float* ml = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits
+ split_id;
#pragma unroll
for (int k = 0; k < 4; k++) {
int abs_h = h_start + lane_group * 4 + k;
float hs = __shfl(my_head_sum, lane_group * 4 + k);
float final_scale = (hs > 0.0f) ? kv_scale / hs : 0.0f;
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++) {
int v_idx = wave_id * 128 + (m / 4) * 64 + lane_col * 4 + (m % 4);
mo[abs_h * num_kv_splits * V_DIM + v_idx] = __float2bfloat16(v_acc[m][k] * final_scale);
}
}
if (wave_id == 0 && lane_group == 0) {
int abs_h = h_start + lane_col;
ml[abs_h * num_kv_splits] = (my_head_sum > 0.0f) ?
my_head_max * LN2_F + __logf(my_head_sum) : -1e30f;
}
}
__syncthreads();
}
}
// ============================================================================
// Stage 1 — FP8 OCC2 path (occupancy=2, triple-buffer, 2× memory latency hiding)
// Issues loads 2 tiles ahead: while computing tile N, tile N+1 is ready and
// tile N+2 is being loaded. vmcnt(4) keeps 1 tile of loads in flight.
// ============================================================================
template <int NHEAD>
__global__ void
__launch_bounds__(256, 2)
mla_decode_stage1_fp8_occ2(
const uint8_t* __restrict__ q_fp8,
const float* __restrict__ q_scales,
const uint8_t* __restrict__ kv_fp8,
const float* __restrict__ kv_scale_ptr,
__hip_bfloat16* __restrict__ mid_o,
float* __restrict__ mid_lse,
__hip_bfloat16* __restrict__ output,
const int32_t* __restrict__ batch_map,
const int32_t* __restrict__ kv_indptr,
int batch_size,
int num_kv_splits
) {
constexpr float sm_scale = 1.0f / 24.0f;
const float kv_scale = *kv_scale_ptr;
constexpr int NUM_HEAD_GROUPS = (NHEAD + HEAD_GROUP - 1) / HEAD_GROUP;
const int q_pos = blockIdx.x;
const int split_id = blockIdx.y;
const int tid = threadIdx.x;
const int wave_id = tid / WAVESIZE;
const int lane_id = tid % WAVESIZE;
const int lane_group = lane_id / 16;
const int lane_col = lane_id % 16;
const int batch_id = batch_map[q_pos];
const int kv_start = kv_indptr[batch_id];
const int kv_end = kv_indptr[batch_id + 1];
const int total_kv_len = kv_end - kv_start;
int positions_per_split = (total_kv_len + num_kv_splits - 1) / num_kv_splits;
int split_start = split_id * positions_per_split;
int split_end = min(split_start + positions_per_split, total_kv_len);
__shared__ LDS_OCC2 lds;
for (int hg = 0; hg < NUM_HEAD_GROUPS; hg++) {
int h_start = hg * HEAD_GROUP;
int h_count = min(HEAD_GROUP, NHEAD - h_start);
float my_scale = 0.0f;
if (lane_col < h_count) {
int abs_h = h_start + lane_col;
my_scale = q_scales[q_pos * NHEAD + abs_h] * kv_scale * sm_scale * LOG2E_F;
}
float my_head_max = -1e30f;
float my_head_sum = 0.0f;
mfma_acc_t v_acc[V_MFMAS_PER_WAVE];
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++)
v_acc[m] = {0.0f, 0.0f, 0.0f, 0.0f};
if (split_start >= split_end) {
if (num_kv_splits > 1) {
__hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM
+ split_id * V_DIM;
float* ml = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits
+ split_id;
__hip_bfloat16 zero_bf16 = __float2bfloat16(0.0f);
#pragma unroll
for (int k = 0; k < 4; k++) {
int head = lane_group * 4 + k;
if (head < h_count) {
int abs_h = h_start + head;
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++) {
int v_idx = wave_id * 128 + (m / 4) * 64 + lane_col * 4 + (m % 4);
if (v_idx < V_DIM)
mo[abs_h * num_kv_splits * V_DIM + v_idx] = zero_bf16;
}
}
}
if (wave_id == 0 && lane_group == 0 && lane_col < h_count)
ml[(h_start + lane_col) * num_kv_splits] = -1e30f;
}
__syncthreads();
continue;
}
const uint8_t* q_base = q_fp8 + (size_t)(q_pos * NHEAD + h_start) * QK_DIM;
// Count tiles for triple-buffer management
int num_tiles = (split_end - split_start + KV_TILE - 1) / KV_TILE;
// Load tile 0 into buf[0]
int first_count = min(KV_TILE, split_end - split_start);
issue_kv_loads(0, kv_fp8 + (size_t)(kv_start + split_start) * QK_DIM,
first_count, tid);
uint32_t q_cache[QK_MFMAS][8];
#pragma unroll
for (int t = 0; t < QK_MFMAS; t++) {
int q_byte = t * FP8_MFMA_K + lane_group * 32;
if (q_byte + 32 <= QK_DIM && lane_col < h_count) {
const uint32_t* ap = reinterpret_cast<const uint32_t*>(
q_base + (size_t)lane_col * QK_DIM + q_byte);
q_cache[t][0]=ap[0]; q_cache[t][1]=ap[1]; q_cache[t][2]=ap[2]; q_cache[t][3]=ap[3];
q_cache[t][4]=ap[4]; q_cache[t][5]=ap[5]; q_cache[t][6]=ap[6]; q_cache[t][7]=ap[7];
} else {
q_cache[t][0]=0; q_cache[t][1]=0; q_cache[t][2]=0; q_cache[t][3]=0;
q_cache[t][4]=0; q_cache[t][5]=0; q_cache[t][6]=0; q_cache[t][7]=0;
}
}
wait_kv_loads();
__syncthreads();
// Issue tile 1 loads async into buf[1] (if exists)
if (num_tiles > 1) {
int t1_start = split_start + KV_TILE;
int t1_count = min(KV_TILE, split_end - t1_start);
issue_kv_loads(1 * LDS_BUF_SIZE, kv_fp8 + (size_t)(kv_start + t1_start) * QK_DIM,
t1_count, tid);
}
// Triple-buffer tile loop
for (int tile_idx = 0; tile_idx < num_tiles; tile_idx++) {
int tile_start = split_start + tile_idx * KV_TILE;
int tile_count = min(KV_TILE, split_end - tile_start);
int buf = tile_idx % 3;
// Issue loads for tile_idx+2 into buf[(tile_idx+2)%3]
if (tile_idx + 2 < num_tiles) {
int t2_start = split_start + (tile_idx + 2) * KV_TILE;
int t2_count = min(KV_TILE, split_end - t2_start);
issue_kv_loads(((tile_idx + 2) % 3) * LDS_BUF_SIZE,
kv_fp8 + (size_t)(kv_start + t2_start) * QK_DIM,
t2_count, tid);
}
const uint8_t* kv_lds = lds.kv_fp8[buf];
mfma_acc_t acc_lo = {0.0f, 0.0f, 0.0f, 0.0f};
mfma_acc_t acc_hi = {0.0f, 0.0f, 0.0f, 0.0f};
// Unified FP8 QK scoring for all dims 0-575 (5 iterations)
asm volatile("s_setprio 3" :::);
#pragma unroll
for (int t = 0; t < QK_MFMAS; t++) {
mfma_input_t frag_q = {q_cache[t][0], q_cache[t][1], q_cache[t][2], q_cache[t][3],
q_cache[t][4], q_cache[t][5], q_cache[t][6], q_cache[t][7]};
int kv_byte = t * FP8_MFMA_K + lane_group * 32;
mfma_input_t frag_kv_lo, frag_kv_hi;
if (kv_byte + 32 <= QK_DIM) {
load_frag_b_wide(kv_lds, lane_col * LDS_KV_STRIDE + kv_byte, frag_kv_lo);
load_frag_b_wide(kv_lds, (lane_col + KV_HALF) * LDS_KV_STRIDE + kv_byte, frag_kv_hi);
} else {
frag_kv_lo[0]=0; frag_kv_lo[1]=0; frag_kv_lo[2]=0; frag_kv_lo[3]=0;
frag_kv_lo[4]=0; frag_kv_lo[5]=0; frag_kv_lo[6]=0; frag_kv_lo[7]=0;
frag_kv_hi[0]=0; frag_kv_hi[1]=0; frag_kv_hi[2]=0; frag_kv_hi[3]=0;
frag_kv_hi[4]=0; frag_kv_hi[5]=0; frag_kv_hi[6]=0; frag_kv_hi[7]=0;
}
acc_lo = mfma_f32_16x16x128_fp8(frag_kv_lo, frag_q, acc_lo);
acc_hi = mfma_f32_16x16x128_fp8(frag_kv_hi, frag_q, acc_hi);
}
asm volatile("s_setprio 0" ::: "memory");
// V preloads from LDS
uint32_t raw4_cur[8], raw4_next[8];
{
int blo = lane_group * 4;
int bhi = lane_group * 4 + 16;
int v_off0 = wave_id * 128 + lane_col * 4;
int v_off1 = wave_id * 128 + 64 + lane_col * 4;
#pragma unroll
for (int i = 0; i < 4; i++) {
raw4_cur[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(blo + i) * LDS_KV_STRIDE + v_off0]);
raw4_cur[i+4] = *reinterpret_cast<const uint32_t*>(&kv_lds[(bhi + i) * LDS_KV_STRIDE + v_off0]);
raw4_next[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(blo + i) * LDS_KV_STRIDE + v_off1]);
raw4_next[i+4] = *reinterpret_cast<const uint32_t*>(&kv_lds[(bhi + i) * LDS_KV_STRIDE + v_off1]);
}
}
// Score extraction
float scores[8];
scores[0] = acc_lo[0]; scores[1] = acc_lo[1]; scores[2] = acc_lo[2]; scores[3] = acc_lo[3];
scores[4] = acc_hi[0]; scores[5] = acc_hi[1]; scores[6] = acc_hi[2]; scores[7] = acc_hi[3];
float partial_max = -1e30f;
#pragma unroll
for (int j = 0; j < 4; j++) {
int pos = lane_group * 4 + j;
if (pos < tile_count) partial_max = fmaxf(partial_max, scores[j] * my_scale);
}
#pragma unroll
for (int j = 0; j < 4; j++) {
int pos = lane_group * 4 + 16 + j;
if (pos < tile_count) partial_max = fmaxf(partial_max, scores[j+4] * my_scale);
}
partial_max = fmaxf(partial_max, __shfl_xor(partial_max, 16));
partial_max = fmaxf(partial_max, __shfl_xor(partial_max, 32));
float old_max = my_head_max;
float new_max = fmaxf(old_max, partial_max);
my_head_max = new_max;
// Conditional rescale: skip when max unchanged (saves ~37 VALU per stable tile)
if (partial_max > old_max && old_max > -1e29f) {
float my_rescale = __builtin_amdgcn_exp2f(old_max - new_max);
my_head_sum *= my_rescale;
#pragma unroll
for (int k = 0; k < 4; k++) {
float r = __shfl(my_rescale, lane_group * 4 + k);
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++)
v_acc[m][k] *= r;
}
}
uint32_t a_packed[4];
float tile_sum_partial = 0.0f;
#pragma unroll
for (int j = 0; j < 4; j += 2) {
int pos0 = lane_group * 4 + j;
int pos1 = lane_group * 4 + j + 1;
float es0 = (pos0 < tile_count) ? __builtin_amdgcn_exp2f(scores[j] * my_scale - my_head_max) : 0.0f;
float es1 = (pos1 < tile_count) ? __builtin_amdgcn_exp2f(scores[j+1] * my_scale - my_head_max) : 0.0f;
tile_sum_partial += es0 + es1;
a_packed[j / 2] = pack_bf16x2(es0, es1);
}
#pragma unroll
for (int j = 0; j < 4; j += 2) {
int pos0 = lane_group * 4 + 16 + j;
int pos1 = lane_group * 4 + 16 + j + 1;
float es0 = (pos0 < tile_count) ? __builtin_amdgcn_exp2f(scores[j+4] * my_scale - my_head_max) : 0.0f;
float es1 = (pos1 < tile_count) ? __builtin_amdgcn_exp2f(scores[j+5] * my_scale - my_head_max) : 0.0f;
tile_sum_partial += es0 + es1;
a_packed[2 + j / 2] = pack_bf16x2(es0, es1);
}
tile_sum_partial += __shfl_xor(tile_sum_partial, 16);
tile_sum_partial += __shfl_xor(tile_sum_partial, 32);
my_head_sum += tile_sum_partial;
// Zero-copy reinterpret a_packed as mfma_bf16_input_t via union
union { uint32_t u[4]; mfma_bf16_input_t v; } a_cvt;
a_cvt.u[0] = a_packed[0]; a_cvt.u[1] = a_packed[1];
a_cvt.u[2] = a_packed[2]; a_cvt.u[3] = a_packed[3];
mfma_bf16_input_t a_bf16 = a_cvt.v;
// V accumulation (same as occ3)
{
asm volatile("s_waitcnt lgkmcnt(8)" ::: "memory");
asm volatile("s_setprio 3" :::);
{
uint32_t bp[4];
v_convert_bf16<0>(raw4_cur, bp);
v_acc[0] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp), v_acc[0]);
}
{
uint32_t bp[4];
v_convert_bf16<1>(raw4_cur, bp);
v_acc[1] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp), v_acc[1]);
}
{
uint32_t bp[4];
v_convert_bf16<2>(raw4_cur, bp);
v_acc[2] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp), v_acc[2]);
}
{
uint32_t bp[4];
v_convert_bf16<3>(raw4_cur, bp);
v_acc[3] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp), v_acc[3]);
}
asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");
{
uint32_t bp[4];
v_convert_bf16<0>(raw4_next, bp);
v_acc[4] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp), v_acc[4]);
}
{
uint32_t bp[4];
v_convert_bf16<1>(raw4_next, bp);
v_acc[5] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp), v_acc[5]);
}
{
uint32_t bp[4];
v_convert_bf16<2>(raw4_next, bp);
v_acc[6] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp), v_acc[6]);
}
{
uint32_t bp[4];
v_convert_bf16<3>(raw4_next, bp);
v_acc[7] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp), v_acc[7]);
}
asm volatile("s_setprio 0" ::: "memory");
}
// Triple-buffer wait: keep tile_idx+2 loads in flight, wait for tile_idx+1
// vmcnt(4) = min loads per wave per tile (waves 2,3 issue 4 loads for 18-chunk tile)
// For last tiles where no more loads are in flight, vmcnt(0) ensures all complete
if (tile_idx + 2 < num_tiles) {
asm volatile("s_waitcnt vmcnt(4)" ::: "memory");
} else if (tile_idx + 1 < num_tiles) {
wait_kv_loads();
}
__syncthreads();
}
if (num_kv_splits == 1) {
#pragma unroll
for (int k = 0; k < 4; k++) {
int head = lane_group * 4 + k;
if (head >= h_count) continue;
int abs_h = h_start + head;
float hs = __shfl(my_head_sum, head);
float final_scale = (hs > 0.0f) ? kv_scale / hs : 0.0f;
__hip_bfloat16* o_ptr = output + (size_t)q_pos * NHEAD * V_DIM + abs_h * V_DIM;
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++) {
int v_idx = wave_id * 128 + (m / 4) * 64 + lane_col * 4 + (m % 4);
if (v_idx < V_DIM)
o_ptr[v_idx] = __float2bfloat16(v_acc[m][k] * final_scale);
}
}
} else {
__hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM
+ split_id * V_DIM;
float* ml = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits
+ split_id;
#pragma unroll
for (int k = 0; k < 4; k++) {
int head = lane_group * 4 + k;
if (head >= h_count) continue;
int abs_h = h_start + head;
float hs = __shfl(my_head_sum, head);
float final_scale = (hs > 0.0f) ? kv_scale / hs : 0.0f;
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++) {
int v_idx = wave_id * 128 + (m / 4) * 64 + lane_col * 4 + (m % 4);
if (v_idx < V_DIM)
mo[abs_h * num_kv_splits * V_DIM + v_idx] = __float2bfloat16(v_acc[m][k] * final_scale);
}
}
if (wave_id == 0 && lane_group == 0 && lane_col < h_count) {
int abs_h = h_start + lane_col;
ml[abs_h * num_kv_splits] = (my_head_sum > 0.0f) ?
my_head_max * LN2_F + __logf(my_head_sum) : -1e30f;
}
}
__syncthreads();
}
}
// ============================================================================
// Stage 1 — BF16 path (for kv=8192: inline Q conversion, overlapped with KV load)
// ============================================================================
template <int NHEAD, bool DIRECT_CVT = false, bool PRELOAD_V = true, bool USE_SETPRIO = true>
__global__ void
__launch_bounds__(256, 3)
mla_decode_stage1_bf16(
const __hip_bfloat16* __restrict__ q_bf16,
const uint8_t* __restrict__ kv_fp8,
const float* __restrict__ kv_scale_ptr,
__hip_bfloat16* __restrict__ mid_o,
float* __restrict__ mid_lse,
__hip_bfloat16* __restrict__ output,
const int32_t* __restrict__ batch_map,
const int32_t* __restrict__ kv_indptr,
int batch_size,
int num_kv_splits
) {
constexpr float sm_scale = 1.0f / 24.0f;
const float kv_scale = *kv_scale_ptr;
constexpr int NUM_HEAD_GROUPS = (NHEAD + HEAD_GROUP - 1) / HEAD_GROUP;
const int q_pos = blockIdx.x;
const int split_id = blockIdx.y;
const int tid = threadIdx.x;
const int wave_id = tid / WAVESIZE;
const int lane_id = tid % WAVESIZE;
const int lane_group = lane_id / 16;
const int lane_col = lane_id % 16;
const int batch_id = batch_map[q_pos];
const int kv_start = kv_indptr[batch_id];
const int kv_end = kv_indptr[batch_id + 1];
const int total_kv_len = kv_end - kv_start;
int positions_per_split = (total_kv_len + num_kv_splits - 1) / num_kv_splits;
int split_start = split_id * positions_per_split;
int split_end = min(split_start + positions_per_split, total_kv_len);
__shared__ LDS lds;
for (int hg = 0; hg < NUM_HEAD_GROUPS; hg++) {
int h_start = hg * HEAD_GROUP;
int h_count = min(HEAD_GROUP, NHEAD - h_start);
float my_head_max = -1e30f;
float my_head_sum = 0.0f;
mfma_acc_t v_acc[V_MFMAS_PER_WAVE];
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++)
v_acc[m] = {0.0f, 0.0f, 0.0f, 0.0f};
if (split_start >= split_end) {
if (num_kv_splits > 1) {
__hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM
+ split_id * V_DIM;
float* ml = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits
+ split_id;
__hip_bfloat16 zero_bf16 = __float2bfloat16(0.0f);
#pragma unroll
for (int k = 0; k < 4; k++) {
int head = lane_group * 4 + k;
if (head < h_count) {
int abs_h = h_start + head;
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++) {
int v_idx = wave_id * 128 + (m / 4) * 64 + lane_col * 4 + (m % 4);
if (v_idx < V_DIM)
mo[abs_h * num_kv_splits * V_DIM + v_idx] = zero_bf16;
}
}
}
if (wave_id == 0 && lane_group == 0 && lane_col < h_count)
ml[(h_start + lane_col) * num_kv_splits] = -1e30f;
}
__syncthreads();
continue;
}
// ============================================================
// KEY OPTIMIZATION: Issue first KV tile load BEFORE Q conversion
// buffer_load...lds is async — Q conversion overlaps with it
// ============================================================
int first_tile = split_start;
int first_count = min(KV_TILE, split_end - first_tile);
issue_kv_loads(0, kv_fp8 + (size_t)(kv_start + first_tile) * QK_DIM,
first_count, tid);
// Inline Q BF16→FP8 conversion (no scaling — FP8 e4m3 range ±448 covers typical Q values)
const __hip_bfloat16* q_head_base = q_bf16 + (size_t)(q_pos * NHEAD + h_start) * QK_DIM;
float my_scale = 0.0f;
if (lane_col < h_count) {
my_scale = kv_scale * sm_scale * LOG2E_F;
}
// Convert BF16 → FP8 (unscaled, direct hw conversion)
uint32_t q_cache[QK_MFMAS][8];
#pragma unroll
for (int t = 0; t < QK_MFMAS; t++) {
int elem_start = t * FP8_MFMA_K + lane_group * 32;
if (elem_start + 32 <= QK_DIM && lane_col < h_count) {
const uint32_t* wp = reinterpret_cast<const uint32_t*>(
q_head_base + (size_t)lane_col * QK_DIM + elem_start);
#pragma unroll
for (int j = 0; j < 8; j++) {
uint32_t w0 = wp[j * 2];
uint32_t w1 = wp[j * 2 + 1];
uint32_t pk;
float s1 = 1.0f;
asm("v_cvt_scalef32_pk_fp8_bf16 %0, %1, %2"
: "=v"(pk) : "v"(w0), "v"(s1));
asm("v_cvt_scalef32_pk_fp8_bf16 %0, %1, %2 op_sel:[0,0,1]"
: "+v"(pk) : "v"(w1), "v"(s1));
q_cache[t][j] = pk;
}
} else {
q_cache[t][0]=0; q_cache[t][1]=0; q_cache[t][2]=0; q_cache[t][3]=0;
q_cache[t][4]=0; q_cache[t][5]=0; q_cache[t][6]=0; q_cache[t][7]=0;
}
}
// Now wait for the KV loads that were issued before Q conversion
wait_kv_loads();
__syncthreads();
int buf = 0;
for (int tile_start = first_tile; tile_start < split_end; tile_start += KV_TILE) {
int tile_count = min(KV_TILE, split_end - tile_start);
int next_start = tile_start + KV_TILE;
int has_next = (next_start < split_end);
if (__builtin_expect(has_next, 1)) {
int next_count = min(KV_TILE, split_end - next_start);
issue_kv_loads((buf ^ 1) * LDS_BUF_SIZE,
kv_fp8 + (size_t)(kv_start + next_start) * QK_DIM,
next_count, tid);
}
const uint8_t* kv_lds = lds.kv_fp8[buf];
mfma_acc_t acc_lo = {0.0f, 0.0f, 0.0f, 0.0f};
mfma_acc_t acc_hi = {0.0f, 0.0f, 0.0f, 0.0f};
// Transposed QK scoring: KV as A operand (rows=kv_pos), Q as B operand (cols=head)
// Result: acc[k] at lane = score for head=lane_col at kv_pos=lane_group*4+k
// Eliminates all 32 cross-lane shuffles for score extraction
if constexpr (USE_SETPRIO) asm volatile("s_setprio 3" :::);
#pragma unroll
for (int t = 0; t < QK_MFMAS; t++) {
mfma_input_t frag_q = {q_cache[t][0], q_cache[t][1], q_cache[t][2], q_cache[t][3],
q_cache[t][4], q_cache[t][5], q_cache[t][6], q_cache[t][7]};
int kv_byte = t * FP8_MFMA_K + lane_group * 32;
mfma_input_t frag_kv_lo, frag_kv_hi;
if (kv_byte + 32 <= QK_DIM) {
load_frag_b_wide(kv_lds, lane_col * LDS_KV_STRIDE + kv_byte, frag_kv_lo);
load_frag_b_wide(kv_lds, (lane_col + KV_HALF) * LDS_KV_STRIDE + kv_byte, frag_kv_hi);
} else {
frag_kv_lo[0]=0; frag_kv_lo[1]=0; frag_kv_lo[2]=0; frag_kv_lo[3]=0;
frag_kv_lo[4]=0; frag_kv_lo[5]=0; frag_kv_lo[6]=0; frag_kv_lo[7]=0;
frag_kv_hi[0]=0; frag_kv_hi[1]=0; frag_kv_hi[2]=0; frag_kv_hi[3]=0;
frag_kv_hi[4]=0; frag_kv_hi[5]=0; frag_kv_hi[6]=0; frag_kv_hi[7]=0;
}
// Transposed: KV as A (rows), Q as B (cols)
acc_lo = mfma_f32_16x16x128_fp8(frag_kv_lo, frag_q, acc_lo);
acc_hi = mfma_f32_16x16x128_fp8(frag_kv_hi, frag_q, acc_hi);
}
if constexpr (USE_SETPRIO) asm volatile("s_setprio 0" ::: "memory");
// V preload strategy: PRELOAD_V=true loads V from LDS before softmax
// (better MFMA throughput for large KV), PRELOAD_V=false defers to
// after softmax (lower register pressure for small batch)
uint32_t raw4_cur[8], raw4_next_v[8];
int blo = lane_group * 4;
int bhi = lane_group * 4 + 16;
int v_off0 = wave_id * 128 + lane_col * 4;
int v_off1 = v_off0 + 64;
if constexpr (PRELOAD_V) {
#pragma unroll
for (int i = 0; i < 4; i++) {
raw4_cur[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(blo + i) * LDS_KV_STRIDE + v_off0]);
raw4_cur[i+4] = *reinterpret_cast<const uint32_t*>(&kv_lds[(bhi + i) * LDS_KV_STRIDE + v_off0]);
raw4_next_v[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(blo + i) * LDS_KV_STRIDE + v_off1]);
raw4_next_v[i+4] = *reinterpret_cast<const uint32_t*>(&kv_lds[(bhi + i) * LDS_KV_STRIDE + v_off1]);
}
}
// Softmax: FMA fusion — max over raw acc, fmaf combines scale+sub
float partial_max = -1e30f;
partial_max = fmaxf(partial_max, fmaxf(acc_lo[0], acc_lo[1]));
partial_max = fmaxf(partial_max, fmaxf(acc_lo[2], acc_lo[3]));
partial_max = fmaxf(partial_max, fmaxf(acc_hi[0], acc_hi[1]));
partial_max = fmaxf(partial_max, fmaxf(acc_hi[2], acc_hi[3]));
partial_max = fmaxf(partial_max, __shfl_xor(partial_max, 16));
partial_max = fmaxf(partial_max, __shfl_xor(partial_max, 32));
float scaled_partial_max = partial_max * my_scale;
float old_max = my_head_max;
float new_max = fmaxf(old_max, scaled_partial_max);
my_head_max = new_max;
float neg_max = -new_max;
if (scaled_partial_max > old_max && old_max > -1e29f) {
float my_rescale = __builtin_amdgcn_exp2f(old_max - new_max);
my_head_sum *= my_rescale;
#pragma unroll
for (int k = 0; k < 4; k++) {
float r = __shfl(my_rescale, lane_group * 4 + k);
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++)
v_acc[m][k] *= r;
}
}
uint32_t a_packed[4];
float tile_sum_partial = 0.0f;
#pragma unroll
for (int j = 0; j < 4; j += 2) {
float es0 = __builtin_amdgcn_exp2f(fmaf(acc_lo[j], my_scale, neg_max));
float es1 = __builtin_amdgcn_exp2f(fmaf(acc_lo[j+1], my_scale, neg_max));
tile_sum_partial += es0 + es1;
a_packed[j / 2] = pack_bf16x2(es0, es1);
}
#pragma unroll
for (int j = 0; j < 4; j += 2) {
float es0 = __builtin_amdgcn_exp2f(fmaf(acc_hi[j], my_scale, neg_max));
float es1 = __builtin_amdgcn_exp2f(fmaf(acc_hi[j+1], my_scale, neg_max));
tile_sum_partial += es0 + es1;
a_packed[2 + j / 2] = pack_bf16x2(es0, es1);
}
tile_sum_partial += __shfl_xor(tile_sum_partial, 16);
tile_sum_partial += __shfl_xor(tile_sum_partial, 32);
my_head_sum += tile_sum_partial;
// Zero-copy reinterpret a_packed as mfma_bf16_input_t via union
union { uint32_t u[4]; mfma_bf16_input_t v; } a_cvt;
a_cvt.u[0] = a_packed[0]; a_cvt.u[1] = a_packed[1];
a_cvt.u[2] = a_packed[2]; a_cvt.u[3] = a_packed[3];
mfma_bf16_input_t a_bf16 = a_cvt.v;
// Vectorized V accumulation
{
// When PRELOAD_V=false, load V data from LDS here (after softmax)
if constexpr (!PRELOAD_V) {
#pragma unroll
for (int i = 0; i < 4; i++) {
raw4_cur[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(blo + i) * LDS_KV_STRIDE + v_off0]);
raw4_cur[i+4] = *reinterpret_cast<const uint32_t*>(&kv_lds[(bhi + i) * LDS_KV_STRIDE + v_off0]);
raw4_next_v[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(blo + i) * LDS_KV_STRIDE + v_off1]);
raw4_next_v[i+4] = *reinterpret_cast<const uint32_t*>(&kv_lds[(bhi + i) * LDS_KV_STRIDE + v_off1]);
}
}
// Group 0: packed convert all 4 byte positions
if constexpr (USE_SETPRIO) asm volatile("s_setprio 3" :::);
uint32_t bp0[4], bp1[4], bp2[4], bp3[4];
v_convert_bf16_dispatch<DIRECT_CVT>(raw4_cur, bp0, bp1, bp2, bp3);
// Group 0 MFMAs
v_acc[0] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp0), v_acc[0]);
v_acc[1] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp1), v_acc[1]);
v_acc[2] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp2), v_acc[2]);
v_acc[3] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp3), v_acc[3]);
// Group 1: packed convert + 4 MFMAs
uint32_t bp0b[4], bp1b[4], bp2b[4], bp3b[4];
v_convert_bf16_dispatch<DIRECT_CVT>(raw4_next_v, bp0b, bp1b, bp2b, bp3b);
v_acc[4] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp0b), v_acc[4]);
v_acc[5] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp1b), v_acc[5]);
v_acc[6] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp2b), v_acc[6]);
v_acc[7] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp3b), v_acc[7]);
if constexpr (USE_SETPRIO) asm volatile("s_setprio 0" :::);
}
if (__builtin_expect(has_next, 1)) wait_kv_loads();
__syncthreads();
buf ^= 1;
}
if (num_kv_splits == 1) {
#pragma unroll
for (int k = 0; k < 4; k++) {
int abs_h = h_start + lane_group * 4 + k;
float hs = __shfl(my_head_sum, lane_group * 4 + k);
float final_scale = (hs > 0.0f) ? kv_scale / hs : 0.0f;
__hip_bfloat16* o_ptr = output + (size_t)q_pos * NHEAD * V_DIM + abs_h * V_DIM;
#pragma unroll
for (int g = 0; g < 2; g++) {
int base_idx = wave_id * 128 + g * 64 + lane_col * 4;
uint32_t pk0 = pack_bf16x2(v_acc[g*4][k] * final_scale, v_acc[g*4+1][k] * final_scale);
uint32_t pk1 = pack_bf16x2(v_acc[g*4+2][k] * final_scale, v_acc[g*4+3][k] * final_scale);
*reinterpret_cast<uint64_t*>(&o_ptr[base_idx]) = ((uint64_t)pk1 << 32) | (uint64_t)pk0;
}
}
} else {
// Layout: [q_pos][head][split][v_dim] — stage2-friendly (contiguous splits per head)
__hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM
+ split_id * V_DIM;
float* ml = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits
+ split_id;
#pragma unroll
for (int k = 0; k < 4; k++) {
int abs_h = h_start + lane_group * 4 + k;
float hs = __shfl(my_head_sum, lane_group * 4 + k);
float final_scale = (hs > 0.0f) ? kv_scale / hs : 0.0f;
#pragma unroll
for (int g = 0; g < 2; g++) {
int base_idx = wave_id * 128 + g * 64 + lane_col * 4;
uint32_t pk0 = pack_bf16x2(v_acc[g*4][k] * final_scale, v_acc[g*4+1][k] * final_scale);
uint32_t pk1 = pack_bf16x2(v_acc[g*4+2][k] * final_scale, v_acc[g*4+3][k] * final_scale);
*reinterpret_cast<uint64_t*>(&mo[abs_h * num_kv_splits * V_DIM + base_idx]) = ((uint64_t)pk1 << 32) | (uint64_t)pk0;
}
}
if (wave_id == 0 && lane_group == 0) {
int abs_h = h_start + lane_col;
ml[abs_h * num_kv_splits] = (my_head_sum > 0.0f) ?
my_head_max * LN2_F + __logf(my_head_sum) : -1e30f;
}
}
__syncthreads();
}
}
// ============================================================================
// Stage 1 — BF16 fused PAGED path (page_size=KV_TILE=32 for exact alignment)
// ============================================================================
// Identical compute to mla_decode_stage1_bf16 but KV addressed via page tables.
// Only the address computation per tile changes — all MFMA/softmax/V is shared.
template <int NHEAD, bool DIRECT_CVT = true, bool PRELOAD_V = true, bool USE_SETPRIO = false>
__global__ void
__launch_bounds__(256, 3)
mla_decode_stage1_bf16_paged(
const __hip_bfloat16* __restrict__ q_bf16,
const uint8_t* __restrict__ kv_pages, // [total_pages, page_size, QK_DIM]
const float* __restrict__ kv_scale_ptr,
__hip_bfloat16* __restrict__ mid_o,
float* __restrict__ mid_lse,
__hip_bfloat16* __restrict__ output,
const int32_t* __restrict__ batch_map,
const int32_t* __restrict__ kv_indices, // logical->physical page map
const int32_t* __restrict__ paged_kv_indptr, // [batch+1] page ranges
const int32_t* __restrict__ kv_last_page_len, // valid tokens in last page
int batch_size,
int num_kv_splits,
int page_size // must equal KV_TILE (32)
) {
constexpr float sm_scale = 1.0f / 24.0f;
const float kv_scale = *kv_scale_ptr;
constexpr int NUM_HEAD_GROUPS = (NHEAD + HEAD_GROUP - 1) / HEAD_GROUP;
const int q_pos = blockIdx.x;
const int split_id = blockIdx.y;
const int tid = threadIdx.x;
const int wave_id = tid / WAVESIZE;
const int lane_id = tid % WAVESIZE;
const int lane_group = lane_id / 16;
const int lane_col = lane_id % 16;
// Paged address resolution
const int batch_id = batch_map[q_pos];
const int pg_start = paged_kv_indptr[batch_id];
const int pg_end = paged_kv_indptr[batch_id + 1];
const int num_pages = pg_end - pg_start;
const int last_pg_len = kv_last_page_len[batch_id];
const int total_kv_len = (num_pages > 1) ? (num_pages - 1) * page_size + last_pg_len : last_pg_len;
// Page-aligned splits: round up to page boundary
int pages_per_split = (num_pages + num_kv_splits - 1) / num_kv_splits;
int split_page_start = split_id * pages_per_split;
int split_page_end = min(split_page_start + pages_per_split, num_pages);
int split_start = split_page_start * page_size;
int split_end = (split_page_end < num_pages)
? split_page_end * page_size
: total_kv_len;
__shared__ LDS lds;
for (int hg = 0; hg < NUM_HEAD_GROUPS; hg++) {
int h_start = hg * HEAD_GROUP;
int h_count = min(HEAD_GROUP, NHEAD - h_start);
float my_head_max = -1e30f;
float my_head_sum = 0.0f;
mfma_acc_t v_acc[V_MFMAS_PER_WAVE];
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++)
v_acc[m] = {0.0f, 0.0f, 0.0f, 0.0f};
if (split_start >= split_end) {
if (num_kv_splits > 1) {
__hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM
+ split_id * V_DIM;
float* ml = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits
+ split_id;
__hip_bfloat16 zero_bf16 = __float2bfloat16(0.0f);
#pragma unroll
for (int k = 0; k < 4; k++) {
int head = lane_group * 4 + k;
if (head < h_count) {
int abs_h = h_start + head;
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++) {
int v_idx = wave_id * 128 + (m / 4) * 64 + lane_col * 4 + (m % 4);
if (v_idx < V_DIM)
mo[abs_h * num_kv_splits * V_DIM + v_idx] = zero_bf16;
}
}
}
if (wave_id == 0 && lane_group == 0 && lane_col < h_count)
ml[(h_start + lane_col) * num_kv_splits] = -1e30f;
}
__syncthreads();
continue;
}
// First tile prefetch — resolve page address
int cur_page = split_page_start;
int first_phys = kv_indices[pg_start + cur_page];
int first_count = min(KV_TILE, split_end - split_start);
issue_kv_loads(0, kv_pages + (size_t)first_phys * page_size * QK_DIM,
first_count, tid);
// Inline Q BF16→FP8 conversion (overlaps with async KV load)
const __hip_bfloat16* q_head_base = q_bf16 + (size_t)(q_pos * NHEAD + h_start) * QK_DIM;
float my_scale = 0.0f;
if (lane_col < h_count) {
my_scale = kv_scale * sm_scale * LOG2E_F;
}
uint32_t q_cache[QK_MFMAS][8];
#pragma unroll
for (int t = 0; t < QK_MFMAS; t++) {
int elem_start = t * FP8_MFMA_K + lane_group * 32;
if (elem_start + 32 <= QK_DIM && lane_col < h_count) {
const uint32_t* wp = reinterpret_cast<const uint32_t*>(
q_head_base + (size_t)lane_col * QK_DIM + elem_start);
#pragma unroll
for (int j = 0; j < 8; j++) {
uint32_t w0 = wp[j * 2];
uint32_t w1 = wp[j * 2 + 1];
uint32_t pk;
float s1 = 1.0f;
asm("v_cvt_scalef32_pk_fp8_bf16 %0, %1, %2"
: "=v"(pk) : "v"(w0), "v"(s1));
asm("v_cvt_scalef32_pk_fp8_bf16 %0, %1, %2 op_sel:[0,0,1]"
: "+v"(pk) : "v"(w1), "v"(s1));
q_cache[t][j] = pk;
}
} else {
q_cache[t][0]=0; q_cache[t][1]=0; q_cache[t][2]=0; q_cache[t][3]=0;
q_cache[t][4]=0; q_cache[t][5]=0; q_cache[t][6]=0; q_cache[t][7]=0;
}
}
wait_kv_loads();
__syncthreads();
// Tile loop — page-aligned: each tile = one page (page_size=KV_TILE)
int buf = 0;
for (int tile_start = split_start; tile_start < split_end; tile_start += KV_TILE) {
int tile_count = min(KV_TILE, split_end - tile_start);
int next_page = cur_page + 1;
int next_tile_start = tile_start + KV_TILE;
int has_next = (next_tile_start < split_end);
if (__builtin_expect(has_next, 1)) {
int next_phys = kv_indices[pg_start + next_page];
int next_count = min(KV_TILE, split_end - next_tile_start);
issue_kv_loads((buf ^ 1) * LDS_BUF_SIZE,
kv_pages + (size_t)next_phys * page_size * QK_DIM,
next_count, tid);
}
const uint8_t* kv_lds = lds.kv_fp8[buf];
// === QK scoring (IDENTICAL to non-paged) ===
mfma_acc_t acc_lo = {0.0f, 0.0f, 0.0f, 0.0f};
mfma_acc_t acc_hi = {0.0f, 0.0f, 0.0f, 0.0f};
if constexpr (USE_SETPRIO) asm volatile("s_setprio 3" :::);
#pragma unroll
for (int t = 0; t < QK_MFMAS; t++) {
mfma_input_t frag_q = {q_cache[t][0], q_cache[t][1], q_cache[t][2], q_cache[t][3],
q_cache[t][4], q_cache[t][5], q_cache[t][6], q_cache[t][7]};
int kv_byte = t * FP8_MFMA_K + lane_group * 32;
mfma_input_t frag_kv_lo, frag_kv_hi;
if (kv_byte + 32 <= QK_DIM) {
load_frag_b_wide(kv_lds, lane_col * LDS_KV_STRIDE + kv_byte, frag_kv_lo);
load_frag_b_wide(kv_lds, (lane_col + KV_HALF) * LDS_KV_STRIDE + kv_byte, frag_kv_hi);
} else {
frag_kv_lo[0]=0; frag_kv_lo[1]=0; frag_kv_lo[2]=0; frag_kv_lo[3]=0;
frag_kv_lo[4]=0; frag_kv_lo[5]=0; frag_kv_lo[6]=0; frag_kv_lo[7]=0;
frag_kv_hi[0]=0; frag_kv_hi[1]=0; frag_kv_hi[2]=0; frag_kv_hi[3]=0;
frag_kv_hi[4]=0; frag_kv_hi[5]=0; frag_kv_hi[6]=0; frag_kv_hi[7]=0;
}
acc_lo = mfma_f32_16x16x128_fp8(frag_kv_lo, frag_q, acc_lo);
acc_hi = mfma_f32_16x16x128_fp8(frag_kv_hi, frag_q, acc_hi);
}
if constexpr (USE_SETPRIO) asm volatile("s_setprio 0" ::: "memory");
// === V preload (IDENTICAL) ===
uint32_t raw4_cur[8], raw4_next_v[8];
int blo = lane_group * 4;
int bhi = lane_group * 4 + 16;
int v_off0 = wave_id * 128 + lane_col * 4;
int v_off1 = v_off0 + 64;
if constexpr (PRELOAD_V) {
#pragma unroll
for (int i = 0; i < 4; i++) {
raw4_cur[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(blo + i) * LDS_KV_STRIDE + v_off0]);
raw4_cur[i+4] = *reinterpret_cast<const uint32_t*>(&kv_lds[(bhi + i) * LDS_KV_STRIDE + v_off0]);
raw4_next_v[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(blo + i) * LDS_KV_STRIDE + v_off1]);
raw4_next_v[i+4] = *reinterpret_cast<const uint32_t*>(&kv_lds[(bhi + i) * LDS_KV_STRIDE + v_off1]);
}
}
// === Softmax (IDENTICAL) ===
float partial_max = -1e30f;
partial_max = fmaxf(partial_max, fmaxf(acc_lo[0], acc_lo[1]));
partial_max = fmaxf(partial_max, fmaxf(acc_lo[2], acc_lo[3]));
partial_max = fmaxf(partial_max, fmaxf(acc_hi[0], acc_hi[1]));
partial_max = fmaxf(partial_max, fmaxf(acc_hi[2], acc_hi[3]));
partial_max = fmaxf(partial_max, __shfl_xor(partial_max, 16));
partial_max = fmaxf(partial_max, __shfl_xor(partial_max, 32));
float scaled_partial_max = partial_max * my_scale;
float old_max = my_head_max;
float new_max = fmaxf(old_max, scaled_partial_max);
my_head_max = new_max;
float neg_max = -new_max;
if (scaled_partial_max > old_max && old_max > -1e29f) {
float my_rescale = __builtin_amdgcn_exp2f(old_max - new_max);
my_head_sum *= my_rescale;
#pragma unroll
for (int k = 0; k < 4; k++) {
float r = __shfl(my_rescale, lane_group * 4 + k);
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++)
v_acc[m][k] *= r;
}
}
uint32_t a_packed[4];
float tile_sum_partial = 0.0f;
#pragma unroll
for (int j = 0; j < 4; j += 2) {
float es0 = __builtin_amdgcn_exp2f(fmaf(acc_lo[j], my_scale, neg_max));
float es1 = __builtin_amdgcn_exp2f(fmaf(acc_lo[j+1], my_scale, neg_max));
tile_sum_partial += es0 + es1;
a_packed[j / 2] = pack_bf16x2(es0, es1);
}
#pragma unroll
for (int j = 0; j < 4; j += 2) {
float es0 = __builtin_amdgcn_exp2f(fmaf(acc_hi[j], my_scale, neg_max));
float es1 = __builtin_amdgcn_exp2f(fmaf(acc_hi[j+1], my_scale, neg_max));
tile_sum_partial += es0 + es1;
a_packed[2 + j / 2] = pack_bf16x2(es0, es1);
}
tile_sum_partial += __shfl_xor(tile_sum_partial, 16);
tile_sum_partial += __shfl_xor(tile_sum_partial, 32);
my_head_sum += tile_sum_partial;
// === V accumulation (IDENTICAL) ===
union { uint32_t u[4]; mfma_bf16_input_t v; } a_cvt;
a_cvt.u[0] = a_packed[0]; a_cvt.u[1] = a_packed[1];
a_cvt.u[2] = a_packed[2]; a_cvt.u[3] = a_packed[3];
mfma_bf16_input_t a_bf16 = a_cvt.v;
{
if constexpr (!PRELOAD_V) {
#pragma unroll
for (int i = 0; i < 4; i++) {
raw4_cur[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(blo + i) * LDS_KV_STRIDE + v_off0]);
raw4_cur[i+4] = *reinterpret_cast<const uint32_t*>(&kv_lds[(bhi + i) * LDS_KV_STRIDE + v_off0]);
raw4_next_v[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(blo + i) * LDS_KV_STRIDE + v_off1]);
raw4_next_v[i+4] = *reinterpret_cast<const uint32_t*>(&kv_lds[(bhi + i) * LDS_KV_STRIDE + v_off1]);
}
}
if constexpr (USE_SETPRIO) asm volatile("s_setprio 3" :::);
uint32_t bp0[4], bp1[4], bp2[4], bp3[4];
v_convert_bf16_dispatch<DIRECT_CVT>(raw4_cur, bp0, bp1, bp2, bp3);
v_acc[0] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp0), v_acc[0]);
v_acc[1] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp1), v_acc[1]);
v_acc[2] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp2), v_acc[2]);
v_acc[3] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp3), v_acc[3]);
uint32_t bp0b[4], bp1b[4], bp2b[4], bp3b[4];
v_convert_bf16_dispatch<DIRECT_CVT>(raw4_next_v, bp0b, bp1b, bp2b, bp3b);
v_acc[4] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp0b), v_acc[4]);
v_acc[5] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp1b), v_acc[5]);
v_acc[6] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp2b), v_acc[6]);
v_acc[7] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp3b), v_acc[7]);
if constexpr (USE_SETPRIO) asm volatile("s_setprio 0" :::);
}
if (__builtin_expect(has_next, 1)) wait_kv_loads();
__syncthreads();
buf ^= 1;
cur_page++;
}
// === Output write (IDENTICAL to non-paged) ===
if (num_kv_splits == 1) {
#pragma unroll
for (int k = 0; k < 4; k++) {
int abs_h = h_start + lane_group * 4 + k;
float hs = __shfl(my_head_sum, lane_group * 4 + k);
float final_scale = (hs > 0.0f) ? kv_scale / hs : 0.0f;
__hip_bfloat16* o_ptr = output + (size_t)q_pos * NHEAD * V_DIM + abs_h * V_DIM;
#pragma unroll
for (int g = 0; g < 2; g++) {
int base_idx = wave_id * 128 + g * 64 + lane_col * 4;
uint32_t pk0 = pack_bf16x2(v_acc[g*4][k] * final_scale, v_acc[g*4+1][k] * final_scale);
uint32_t pk1 = pack_bf16x2(v_acc[g*4+2][k] * final_scale, v_acc[g*4+3][k] * final_scale);
*reinterpret_cast<uint64_t*>(&o_ptr[base_idx]) = ((uint64_t)pk1 << 32) | (uint64_t)pk0;
}
}
} else {
__hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM
+ split_id * V_DIM;
float* ml = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits
+ split_id;
#pragma unroll
for (int k = 0; k < 4; k++) {
int abs_h = h_start + lane_group * 4 + k;
float hs = __shfl(my_head_sum, lane_group * 4 + k);
float final_scale = (hs > 0.0f) ? kv_scale / hs : 0.0f;
#pragma unroll
for (int g = 0; g < 2; g++) {
int base_idx = wave_id * 128 + g * 64 + lane_col * 4;
uint32_t pk0 = pack_bf16x2(v_acc[g*4][k] * final_scale, v_acc[g*4+1][k] * final_scale);
uint32_t pk1 = pack_bf16x2(v_acc[g*4+2][k] * final_scale, v_acc[g*4+3][k] * final_scale);
*reinterpret_cast<uint64_t*>(&mo[abs_h * num_kv_splits * V_DIM + base_idx]) = ((uint64_t)pk1 << 32) | (uint64_t)pk0;
}
}
if (wave_id == 0 && lane_group == 0) {
int abs_h = h_start + lane_col;
ml[abs_h * num_kv_splits] = (my_head_sum > 0.0f) ?
my_head_max * LN2_F + __logf(my_head_sum) : -1e30f;
}
}
__syncthreads();
}
}
// ============================================================================
// Stage 1 — V2: occ=1, 5-buffer, FP8 K=128 V accumulation
// ============================================================================
// Key differences from bf16 path:
// - __launch_bounds__(256, 1): 512 VGPRs, 160KB LDS
// - 5-buffer LDS (4 retained + 1 prefetch) for 128-token megatile V MFMA
// - FP8 mfma_f32_16x16x128 for V accumulation (4x throughput vs BF16 K=32)
// - Attention weights accumulated across 4 tiles, converted to FP8 for V MFMA
template <int NHEAD>
__global__ void
__launch_bounds__(256, 1)
mla_decode_stage1_v2(
const __hip_bfloat16* __restrict__ q_bf16,
const uint8_t* __restrict__ kv_fp8,
const float* __restrict__ kv_scale_ptr,
__hip_bfloat16* __restrict__ mid_o,
float* __restrict__ mid_lse,
__hip_bfloat16* __restrict__ output,
const int32_t* __restrict__ batch_map,
const int32_t* __restrict__ kv_indptr,
int batch_size,
int num_kv_splits
) {
constexpr float sm_scale = 1.0f / 24.0f;
const float kv_scale = *kv_scale_ptr;
constexpr int NUM_HEAD_GROUPS = (NHEAD + HEAD_GROUP - 1) / HEAD_GROUP;
const int q_pos = blockIdx.x;
const int split_id = blockIdx.y;
const int tid = threadIdx.x;
const int wave_id = tid / WAVESIZE;
const int lane_id = tid % WAVESIZE;
const int lane_group = lane_id / 16;
const int lane_col = lane_id % 16;
const int batch_id = batch_map[q_pos];
const int kv_start = kv_indptr[batch_id];
const int kv_end = kv_indptr[batch_id + 1];
const int total_kv_len = kv_end - kv_start;
int positions_per_split = (total_kv_len + num_kv_splits - 1) / num_kv_splits;
positions_per_split = (positions_per_split + KV_TILE - 1) / KV_TILE * KV_TILE;
int split_start = split_id * positions_per_split;
int split_end = min(split_start + positions_per_split, total_kv_len);
__shared__ LDS_V2 lds;
for (int hg = 0; hg < NUM_HEAD_GROUPS; hg++) {
int h_start = hg * HEAD_GROUP;
int h_count = min(HEAD_GROUP, NHEAD - h_start);
float my_head_max = -1e30f;
float my_head_sum = 0.0f;
mfma_acc_t v_acc[V_MFMAS_PER_WAVE];
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++)
v_acc[m] = {0.0f, 0.0f, 0.0f, 0.0f};
if (split_start >= split_end) {
// Empty split: write zeros (v2 layout)
if (num_kv_splits > 1) {
__hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM + split_id * V_DIM;
float* ml = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits + split_id;
__hip_bfloat16 zero_bf16 = __float2bfloat16(0.0f);
for (int n = 0; n < 4; n++) {
int head = lane_group * 4 + n;
if (head < h_count) {
int abs_h = h_start + head;
for (int vm = 0; vm < V_MFMAS_PER_WAVE; vm++) {
int v_dim = wave_id * 128 + vm * 16 + lane_col;
if (v_dim < V_DIM) mo[abs_h * num_kv_splits * V_DIM + v_dim] = zero_bf16;
}
}
}
if (wave_id == 0 && lane_group == 0 && lane_col < h_count)
ml[(h_start + lane_col) * num_kv_splits] = -1e30f;
}
__syncthreads();
continue;
}
// Inline Q BF16→FP8 conversion
const __hip_bfloat16* q_head_base = q_bf16 + (size_t)(q_pos * NHEAD + h_start) * QK_DIM;
float my_scale = 0.0f;
if (lane_col < h_count) my_scale = kv_scale * sm_scale * LOG2E_F;
uint32_t q_cache[QK_MFMAS][8];
#pragma unroll
for (int t = 0; t < QK_MFMAS; t++) {
int elem_start = t * FP8_MFMA_K + lane_group * 32;
if (elem_start + 32 <= QK_DIM && lane_col < h_count) {
const uint32_t* wp = reinterpret_cast<const uint32_t*>(
q_head_base + (size_t)lane_col * QK_DIM + elem_start);
#pragma unroll
for (int j = 0; j < 8; j++) {
uint32_t w0 = wp[j * 2];
uint32_t w1 = wp[j * 2 + 1];
uint32_t pk;
float s1 = 1.0f;
asm("v_cvt_scalef32_pk_fp8_bf16 %0, %1, %2"
: "=v"(pk) : "v"(w0), "v"(s1));
asm("v_cvt_scalef32_pk_fp8_bf16 %0, %1, %2 op_sel:[0,0,1]"
: "+v"(pk) : "v"(w1), "v"(s1));
q_cache[t][j] = pk;
}
} else {
#pragma unroll
for (int j = 0; j < 8; j++) q_cache[t][j] = 0;
}
}
// Prefetch first tile into buffer 0
int first_count = min(KV_TILE, split_end - split_start);
issue_kv_loads(0, kv_fp8 + (size_t)(kv_start + split_start) * QK_DIM, first_count, tid);
wait_kv_loads();
__syncthreads();
// Tile loop with 4-tile megatile grouping
int buf_idx = 0;
int tile_in_mega = 0;
// Attention weights accumulated across megatile: [tile][8] per lane
float attn_w[4][8];
for (int tile_start = split_start; tile_start < split_end; tile_start += KV_TILE) {
int tile_count = min(KV_TILE, split_end - tile_start);
int next_tile_start = tile_start + KV_TILE;
int has_next = (next_tile_start < split_end);
// Prefetch next tile
int next_buf = (buf_idx + 1) % V2_NUM_BUFS;
if (has_next) {
int next_count = min(KV_TILE, split_end - next_tile_start);
issue_kv_loads(next_buf * LDS_BUF_SIZE,
kv_fp8 + (size_t)(kv_start + next_tile_start) * QK_DIM,
next_count, tid);
}
const uint8_t* kv_lds = lds.kv_fp8[buf_idx];
// === QK scoring (identical to bf16 path) ===
mfma_acc_t acc_lo = {0.0f, 0.0f, 0.0f, 0.0f};
mfma_acc_t acc_hi = {0.0f, 0.0f, 0.0f, 0.0f};
#pragma unroll
for (int t = 0; t < QK_MFMAS; t++) {
mfma_input_t frag_q = {q_cache[t][0], q_cache[t][1], q_cache[t][2], q_cache[t][3],
q_cache[t][4], q_cache[t][5], q_cache[t][6], q_cache[t][7]};
int kv_byte = t * FP8_MFMA_K + lane_group * 32;
mfma_input_t frag_kv_lo, frag_kv_hi;
if (kv_byte + 32 <= QK_DIM) {
load_frag_b_wide(kv_lds, lane_col * LDS_KV_STRIDE + kv_byte, frag_kv_lo);
load_frag_b_wide(kv_lds, (lane_col + KV_HALF) * LDS_KV_STRIDE + kv_byte, frag_kv_hi);
} else {
#pragma unroll
for (int j = 0; j < 8; j++) { frag_kv_lo[j] = 0; frag_kv_hi[j] = 0; }
}
acc_lo = mfma_f32_16x16x128_fp8(frag_kv_lo, frag_q, acc_lo);
acc_hi = mfma_f32_16x16x128_fp8(frag_kv_hi, frag_q, acc_hi);
}
// === Softmax: accumulate into megatile weights ===
float partial_max = -1e30f;
partial_max = fmaxf(partial_max, fmaxf(acc_lo[0], acc_lo[1]));
partial_max = fmaxf(partial_max, fmaxf(acc_lo[2], acc_lo[3]));
partial_max = fmaxf(partial_max, fmaxf(acc_hi[0], acc_hi[1]));
partial_max = fmaxf(partial_max, fmaxf(acc_hi[2], acc_hi[3]));
partial_max = fmaxf(partial_max, __shfl_xor(partial_max, 16));
partial_max = fmaxf(partial_max, __shfl_xor(partial_max, 32));
float scaled_partial_max = partial_max * my_scale;
float old_max = my_head_max;
float new_max = fmaxf(old_max, scaled_partial_max);
my_head_max = new_max;
float neg_max = -new_max;
if (scaled_partial_max > old_max && old_max > -1e29f) {
float my_rescale = __builtin_amdgcn_exp2f(old_max - new_max);
my_head_sum *= my_rescale;
// Rescale previously accumulated megatile weights
for (int pt = 0; pt < tile_in_mega; pt++)
for (int j = 0; j < 8; j++)
attn_w[pt][j] *= my_rescale;
// Rescale V accumulators from previous megatiles
#pragma unroll
for (int k = 0; k < 4; k++) {
float r = __shfl(my_rescale, lane_group * 4 + k);
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++)
v_acc[m][k] *= r;
}
}
// Compute exp weights and store in attn_w
float tile_sum_partial = 0.0f;
#pragma unroll
for (int j = 0; j < 4; j++) {
float es = __builtin_amdgcn_exp2f(fmaf(acc_lo[j], my_scale, neg_max));
tile_sum_partial += es;
attn_w[tile_in_mega][j] = es;
}
#pragma unroll
for (int j = 0; j < 4; j++) {
float es = __builtin_amdgcn_exp2f(fmaf(acc_hi[j], my_scale, neg_max));
tile_sum_partial += es;
attn_w[tile_in_mega][4 + j] = es;
}
tile_sum_partial += __shfl_xor(tile_sum_partial, 16);
tile_sum_partial += __shfl_xor(tile_sum_partial, 32);
my_head_sum += tile_sum_partial;
tile_in_mega++;
// === V MFMA every 4 tiles (or at end of split) ===
int is_last_tile = !has_next;
if (tile_in_mega == V2_MEGATILE || is_last_tile) {
// Zero out unused megatile slots
for (int pt = tile_in_mega; pt < V2_MEGATILE; pt++)
for (int j = 0; j < 8; j++)
attn_w[pt][j] = 0.0f;
// === Interleaved tile V MFMA: no shuffle needed ===
// Key insight: interleave tiles within K=128 so each lane_group's
// native attn_w values align with the MFMA operand layout.
//
// A operand layout (16 heads × 128 tokens):
// g%4 = tile index, g/4 = half (0=tokens 0..15, 1=tokens 16..31)
// frag_a[g] = FP8(attn_w[g%4][(g/4)*4 + 0..3])
// Each lane_group already has its 4 tokens' weights — no cross-lane shuffle.
//
// B operand layout (128 tokens × 16 V dims):
// Same interleaving: frag_b[g] reads from tile g%4's LDS buffer,
// tokens (g/4)*16 + lane_group*4 + 0..3
// Compute uniform FP8 scale across all lane_groups (per head)
float amax = 0.0f;
#pragma unroll
for (int tile = 0; tile < V2_MEGATILE; tile++)
#pragma unroll
for (int j = 0; j < 8; j++)
amax = fmaxf(amax, attn_w[tile][j]);
// Reduce across lane_groups so all use the same scale per head
amax = fmaxf(amax, __shfl_xor(amax, 16));
amax = fmaxf(amax, __shfl_xor(amax, 32));
float att_scale = (amax > 1e-12f) ? 224.0f / amax : 0.0f;
float inv_att_scale = (att_scale > 0.0f) ? 1.0f / att_scale : 0.0f;
// Pack A operand: interleaved tile order
mfma_input_t frag_a;
#pragma unroll
for (int g = 0; g < 8; g++) {
int tile = g & 3; // g % 4
int half = g >> 2; // g / 4
float v0 = attn_w[tile][half*4 + 0] * att_scale;
float v1 = attn_w[tile][half*4 + 1] * att_scale;
float v2 = attn_w[tile][half*4 + 2] * att_scale;
float v3 = attn_w[tile][half*4 + 3] * att_scale;
uint32_t pk;
pk = __builtin_amdgcn_cvt_pk_fp8_f32(v0, v1, 0u, false);
pk = __builtin_amdgcn_cvt_pk_fp8_f32(v2, v3, pk, true);
frag_a[g] = pk;
}
// V MFMA: interleaved B operand from 4 LDS buffers
#pragma unroll
for (int vm = 0; vm < V_MFMAS_PER_WAVE; vm++) {
int v_base = wave_id * 128 + vm * 16;
if (v_base >= V_DIM) break;
int v_col = v_base + lane_col;
mfma_input_t frag_b;
if (v_col < V_DIM) {
#pragma unroll
for (int g = 0; g < 8; g++) {
int tile = g & 3;
int half = g >> 2;
int v2_buf = (buf_idx - tile_in_mega + 1 + tile + V2_NUM_BUFS) % V2_NUM_BUFS;
const uint8_t* v_lds = lds.kv_fp8[v2_buf];
int tok_base = half * 16 + lane_group * 4;
uint32_t b = (uint32_t)v_lds[(tok_base ) * LDS_KV_STRIDE + v_col]
| ((uint32_t)v_lds[(tok_base+1) * LDS_KV_STRIDE + v_col] << 8)
| ((uint32_t)v_lds[(tok_base+2) * LDS_KV_STRIDE + v_col] << 16)
| ((uint32_t)v_lds[(tok_base+3) * LDS_KV_STRIDE + v_col] << 24);
frag_b[g] = b;
}
} else {
#pragma unroll
for (int g = 0; g < 8; g++) frag_b[g] = 0;
}
// V MFMA into zeroed temp (unscaled: scale=0 → compiler emits v_mfma_f32_16x16x128_f8f6f4)
mfma_acc_t v_delta = {0.0f, 0.0f, 0.0f, 0.0f};
v_delta = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
frag_a, frag_b, v_delta, 0, 0, 0, 0, 0, 0);
// Correct per-head scale and accumulate
// v_delta[n] = result for head (lane_group*4+n), need that head's inv_att_scale
#pragma unroll
for (int n = 0; n < 4; n++) {
float inv_s = __shfl(inv_att_scale, lane_group * 4 + n);
v_acc[vm][n] += v_delta[n] * inv_s;
}
}
tile_in_mega = 0;
}
if (has_next) wait_kv_loads();
__syncthreads();
buf_idx = next_buf;
}
// === Output write (v2 layout: v_acc[vm][n] = head lg*4+n, V dim wave*128+vm*16+lc) ===
if (num_kv_splits == 1) {
#pragma unroll
for (int n = 0; n < 4; n++) {
int abs_h = h_start + lane_group * 4 + n;
float hs = __shfl(my_head_sum, lane_group * 4 + n);
float final_scale = (hs > 0.0f) ? kv_scale / hs : 0.0f;
__hip_bfloat16* o_ptr = output + (size_t)q_pos * NHEAD * V_DIM + abs_h * V_DIM;
#pragma unroll
for (int vm = 0; vm < V_MFMAS_PER_WAVE; vm++) {
int v_dim = wave_id * 128 + vm * 16 + lane_col;
if (v_dim < V_DIM)
o_ptr[v_dim] = __float2bfloat16(v_acc[vm][n] * final_scale);
}
}
} else {
__hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM + split_id * V_DIM;
float* ml = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits + split_id;
#pragma unroll
for (int n = 0; n < 4; n++) {
int abs_h = h_start + lane_group * 4 + n;
float hs = __shfl(my_head_sum, lane_group * 4 + n);
float final_scale = (hs > 0.0f) ? kv_scale / hs : 0.0f;
#pragma unroll
for (int vm = 0; vm < V_MFMAS_PER_WAVE; vm++) {
int v_dim = wave_id * 128 + vm * 16 + lane_col;
if (v_dim < V_DIM)
mo[abs_h * num_kv_splits * V_DIM + v_dim] =
__float2bfloat16(v_acc[vm][n] * final_scale);
}
}
if (wave_id == 0 && lane_group == 0) {
int abs_h = h_start + lane_col;
ml[abs_h * num_kv_splits] = (my_head_sum > 0.0f) ?
my_head_max * LN2_F + __logf(my_head_sum) : -1e30f;
}
}
__syncthreads();
}
}
// ============================================================================
// Stage 1 — MXFP4 path (native FP4 MFMA + hardware V conversion)
// ============================================================================
// Uses FP8(Q) × FP4(KV) MFMA with software E8M0 scaling for QK.
// V accumulation: hardware cvt_scalef32_pk_f32_fp4 for FP4→F32, then BF16 MFMA.
// 50% HBM bandwidth reduction vs FP8 path.
template <int NHEAD>
__global__ void
__launch_bounds__(256, 3)
mla_decode_stage1_mxfp4_bf16(
const __hip_bfloat16* __restrict__ q_bf16,
const uint8_t* __restrict__ kv_fp4, // MXFP4 packed data
const uint8_t* __restrict__ kv_fp4_scales, // E8M0 block scales
int scales_stride, // stride between scale rows
const float* __restrict__ kv_scale_ptr, // unused (compat)
__hip_bfloat16* __restrict__ mid_o,
float* __restrict__ mid_lse,
__hip_bfloat16* __restrict__ output,
const int32_t* __restrict__ batch_map,
const int32_t* __restrict__ kv_indptr,
int batch_size,
int num_kv_splits
) {
constexpr float sm_scale = 1.0f / 24.0f;
constexpr int NUM_HEAD_GROUPS = (NHEAD + HEAD_GROUP - 1) / HEAD_GROUP;
const int q_pos = blockIdx.x;
const int split_id = blockIdx.y;
const int tid = threadIdx.x;
const int wave_id = tid / WAVESIZE;
const int lane_id = tid % WAVESIZE;
const int lane_group = lane_id / 16;
const int lane_col = lane_id % 16;
const int batch_id = batch_map[q_pos];
const int kv_start = kv_indptr[batch_id];
const int kv_end = kv_indptr[batch_id + 1];
const int total_kv_len = kv_end - kv_start;
int positions_per_split = (total_kv_len + num_kv_splits - 1) / num_kv_splits;
int split_start = split_id * positions_per_split;
int split_end = min(split_start + positions_per_split, total_kv_len);
__shared__ LDS_FP4 lds;
// Inline per-batch-element Q scale computation
{
const __hip_bfloat16* q_batch = q_bf16 + (size_t)q_pos * NHEAD * QK_DIM;
constexpr int q_elems = NHEAD * QK_DIM;
const uint4* q_vec = reinterpret_cast<const uint4*>(q_batch);
constexpr int n_vec = q_elems / 8;
float local_amax = 0.0f;
for (int i = tid; i < n_vec; i += NTHREADS) {
uint4 data = q_vec[i];
const __hip_bfloat16* bv = reinterpret_cast<const __hip_bfloat16*>(&data);
#pragma unroll
for (int j = 0; j < 8; j++)
local_amax = fmaxf(local_amax, fabsf(__bfloat162float(bv[j])));
}
local_amax = fmaxf(local_amax, __shfl_xor(local_amax, 1));
local_amax = fmaxf(local_amax, __shfl_xor(local_amax, 2));
local_amax = fmaxf(local_amax, __shfl_xor(local_amax, 4));
local_amax = fmaxf(local_amax, __shfl_xor(local_amax, 8));
local_amax = fmaxf(local_amax, __shfl_xor(local_amax, 16));
local_amax = fmaxf(local_amax, __shfl_xor(local_amax, 32));
if (lane_id == 0) lds.wave_amax[wave_id] = local_amax;
__syncthreads();
if (tid == 0) {
float amax = fmaxf(fmaxf(lds.wave_amax[0], lds.wave_amax[1]),
fmaxf(lds.wave_amax[2], lds.wave_amax[3]));
amax = fmaxf(amax, 1e-12f);
float scale_f32 = amax / 448.0f;
__hip_bfloat16 scale_bf16 = __float2bfloat16(scale_f32);
lds.wave_amax[0] = __bfloat162float(scale_bf16);
}
__syncthreads();
}
float q_scale = lds.wave_amax[0];
float q_inv_scale = (q_scale > 0.0f) ? 1.0f / q_scale : 0.0f;
for (int hg = 0; hg < NUM_HEAD_GROUPS; hg++) {
int h_start = hg * HEAD_GROUP;
int h_count = min(HEAD_GROUP, NHEAD - h_start);
float my_head_max = -1e30f;
float my_head_sum = 0.0f;
mfma_acc_t v_acc[V_MFMAS_PER_WAVE];
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++)
v_acc[m] = {0.0f, 0.0f, 0.0f, 0.0f};
if (split_start >= split_end) {
if (num_kv_splits > 1) {
__hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM
+ split_id * V_DIM;
float* ml = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits
+ split_id;
__hip_bfloat16 zero_bf16 = __float2bfloat16(0.0f);
#pragma unroll
for (int k = 0; k < 4; k++) {
int head = lane_group * 4 + k;
if (head < h_count) {
int abs_h = h_start + head;
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++) {
int v_idx = wave_id * 128 + (m / 4) * 64 + lane_col * 4 + (m % 4);
if (v_idx < V_DIM)
mo[abs_h * num_kv_splits * V_DIM + v_idx] = zero_bf16;
}
}
}
if (wave_id == 0 && lane_group == 0 && lane_col < h_count)
ml[(h_start + lane_col) * num_kv_splits] = -1e30f;
}
__syncthreads();
continue;
}
// Issue first FP4 tile load
int first_count = min(KV_TILE, split_end - split_start);
issue_kv_fp4_loads(0, kv_fp4 + (size_t)(kv_start + split_start) * KV_FP4_STRIDE,
first_count, tid);
// Q BF16→FP8 conversion (overlapped with FP4 loads)
const __hip_bfloat16* q_head_base = q_bf16 + (size_t)(q_pos * NHEAD + h_start) * QK_DIM;
float my_scale = 0.0f;
if (lane_col < h_count) {
my_scale = q_scale * sm_scale * LOG2E_F;
}
uint32_t q_cache[QK_MFMAS][8];
#pragma unroll
for (int t = 0; t < QK_MFMAS; t++) {
int elem_start = t * FP8_MFMA_K + lane_group * 32;
if (elem_start + 32 <= QK_DIM && lane_col < h_count) {
const uint32_t* wp = reinterpret_cast<const uint32_t*>(
q_head_base + (size_t)lane_col * QK_DIM + elem_start);
#pragma unroll
for (int j = 0; j < 8; j++) {
uint32_t w0 = wp[j * 2];
uint32_t w1 = wp[j * 2 + 1];
float f0 = __bfloat162float(*reinterpret_cast<const __hip_bfloat16*>(&w0)) * q_inv_scale;
float f1 = __bfloat162float(*(reinterpret_cast<const __hip_bfloat16*>(&w0) + 1)) * q_inv_scale;
float f2 = __bfloat162float(*reinterpret_cast<const __hip_bfloat16*>(&w1)) * q_inv_scale;
float f3 = __bfloat162float(*(reinterpret_cast<const __hip_bfloat16*>(&w1) + 1)) * q_inv_scale;
uint32_t pk = __builtin_amdgcn_cvt_pk_fp8_f32(f0, f1, 0u, false);
pk = __builtin_amdgcn_cvt_pk_fp8_f32(f2, f3, pk, true);
q_cache[t][j] = pk;
}
} else {
q_cache[t][0]=0; q_cache[t][1]=0; q_cache[t][2]=0; q_cache[t][3]=0;
q_cache[t][4]=0; q_cache[t][5]=0; q_cache[t][6]=0; q_cache[t][7]=0;
}
}
// Wait for FP4 loads
wait_kv_loads();
__syncthreads();
// Tile loop
int fp4_buf = 0;
for (int tile_start = split_start; tile_start < split_end; tile_start += KV_TILE) {
int tile_count = min(KV_TILE, split_end - tile_start);
int next_start = tile_start + KV_TILE;
int has_next = (next_start < split_end);
if (__builtin_expect(has_next, 1)) {
int next_count = min(KV_TILE, split_end - next_start);
issue_kv_fp4_loads((fp4_buf ^ 1) * LDS_FP4_BUF_SIZE,
kv_fp4 + (size_t)(kv_start + next_start) * KV_FP4_STRIDE,
next_count, tid);
}
// Cooperative load of E8M0 scales for this tile into LDS cache
{
const int total_scale_bytes = tile_count * KV_SCALE_PER_ROW;
for (int idx = tid; idx < total_scale_bytes; idx += NTHREADS) {
int row = idx / KV_SCALE_PER_ROW;
int col = idx % KV_SCALE_PER_ROW;
int kv_pos = kv_start + tile_start + row;
lds.scale_cache[row * KV_SCALE_PER_ROW + col] =
kv_fp4_scales[kv_pos * scales_stride + col];
}
}
__syncthreads();
const uint8_t* fp4_lds = lds.fp4[fp4_buf];
int tile_lo = min(KV_HALF, tile_count);
int tile_hi = max(0, tile_count - KV_HALF);
// ---- QK MFMA: FP8(Q) × FP4(KV) with hardware per-block E8M0 scaling ----
mfma_acc_t acc_lo_v = {0.0f, 0.0f, 0.0f, 0.0f};
asm volatile("s_setprio 3" :::);
#pragma unroll
for (int t = 0; t < QK_MFMAS; t++) {
mfma_input_t frag_a = {q_cache[t][0], q_cache[t][1], q_cache[t][2], q_cache[t][3],
q_cache[t][4], q_cache[t][5], q_cache[t][6], q_cache[t][7]};
mfma_input_t frag_b;
int kv_fp4_byte = t * FP4_MFMA_BYTES + lane_group * 16;
if (kv_fp4_byte + 16 <= KV_FP4_STRIDE && lane_col < tile_lo) {
load_frag_b_fp4(fp4_lds, lane_col * KV_FP4_STRIDE + kv_fp4_byte, frag_b);
} else {
frag_b[0]=0; frag_b[1]=0; frag_b[2]=0; frag_b[3]=0;
frag_b[4]=0; frag_b[5]=0; frag_b[6]=0; frag_b[7]=0;
}
// Pack 4 E8M0 block scales for hardware per-block scaling
int scale_b = 0;
if (lane_col < tile_lo) {
int base_block = t * 4;
#pragma unroll
for (int g = 0; g < 4; g++) {
int bi = base_block + g;
if (bi < KV_SCALE_PER_ROW) {
uint8_t e = lds.scale_cache[lane_col * KV_SCALE_PER_ROW + bi];
scale_b |= ((int)e << (g * 8));
}
}
}
acc_lo_v = mfma_f32_16x16x128_fp8_fp4_scaled(frag_a, frag_b, acc_lo_v, scale_b);
}
asm volatile("s_setprio 0" ::: "memory");
mfma_acc_t acc_hi_v = {0.0f, 0.0f, 0.0f, 0.0f};
if (__builtin_expect(tile_hi > 0, 1)) {
asm volatile("s_setprio 3" :::);
#pragma unroll
for (int t = 0; t < QK_MFMAS; t++) {
mfma_input_t frag_a = {q_cache[t][0], q_cache[t][1], q_cache[t][2], q_cache[t][3],
q_cache[t][4], q_cache[t][5], q_cache[t][6], q_cache[t][7]};
mfma_input_t frag_b;
int kv_fp4_byte = t * FP4_MFMA_BYTES + lane_group * 16;
if (kv_fp4_byte + 16 <= KV_FP4_STRIDE && lane_col < tile_hi) {
load_frag_b_fp4(fp4_lds, (lane_col + KV_HALF) * KV_FP4_STRIDE + kv_fp4_byte, frag_b);
} else {
frag_b[0]=0; frag_b[1]=0; frag_b[2]=0; frag_b[3]=0;
frag_b[4]=0; frag_b[5]=0; frag_b[6]=0; frag_b[7]=0;
}
// Pack 4 E8M0 block scales for hardware per-block scaling
int scale_b = 0;
if (lane_col < tile_hi) {
int base_block = t * 4;
#pragma unroll
for (int g = 0; g < 4; g++) {
int bi = base_block + g;
if (bi < KV_SCALE_PER_ROW) {
uint8_t e = lds.scale_cache[(KV_HALF + lane_col) * KV_SCALE_PER_ROW + bi];
scale_b |= ((int)e << (g * 8));
}
}
}
acc_hi_v = mfma_f32_16x16x128_fp8_fp4_scaled(frag_a, frag_b, acc_hi_v, scale_b);
}
asm volatile("s_setprio 0" ::: "memory");
}
// ---- Score extraction (identical to FP8 path) ----
float scores[8];
float partial_max = -1e30f;
#pragma unroll
for (int j = 0; j < 8; j++) {
int pos = lane_group * 8 + j;
{
int _src = (lane_col / 4) * 16 + (pos < 16 ? pos : pos - 16);
float s_lo = extract_score_4(acc_lo_v[0], acc_lo_v[1], acc_lo_v[2], acc_lo_v[3], lane_col, _src);
float s_hi = extract_score_4(acc_hi_v[0], acc_hi_v[1], acc_hi_v[2], acc_hi_v[3], lane_col, _src);
scores[j] = (pos < 16) ? s_lo : s_hi;
}
if (pos < tile_count)
partial_max = fmaxf(partial_max, scores[j] * my_scale);
}
partial_max = fmaxf(partial_max, __shfl_xor(partial_max, 16));
partial_max = fmaxf(partial_max, __shfl_xor(partial_max, 32));
// ---- V loading: hardware FP4→F32 conversion + BF16 MFMA ----
// Preload FP4 data and E8M0 scales for V groups 0 and 1
uint32_t fp4_cur[8], fp4_next[8];
float v_scale_cur[8], v_scale_next[8];
{
int v_fp8_base0 = wave_id * 128 + lane_col * 4;
int v_fp8_base1 = wave_id * 128 + 64 + lane_col * 4;
load_v_fp4_data(fp4_lds, v_fp8_base0 / 2, lane_group,
lds.scale_cache, v_fp8_base0 / 32,
fp4_cur, v_scale_cur);
load_v_fp4_data(fp4_lds, v_fp8_base1 / 2, lane_group,
lds.scale_cache, v_fp8_base1 / 32,
fp4_next, v_scale_next);
}
float old_max = my_head_max;
float new_max = fmaxf(old_max, partial_max);
my_head_max = new_max;
// Conditional rescale: skip when max unchanged (saves ~37 VALU per stable tile)
if (partial_max > old_max && old_max > -1e29f) {
float my_rescale = __builtin_amdgcn_exp2f(old_max - new_max);
my_head_sum *= my_rescale;
#pragma unroll
for (int k = 0; k < 4; k++) {
float r = __shfl(my_rescale, lane_group * 4 + k);
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++)
v_acc[m][k] *= r;
}
}
uint32_t a_packed[4];
float tile_sum_partial = 0.0f;
#pragma unroll
for (int j = 0; j < 8; j += 2) {
int pos0 = lane_group * 8 + j;
int pos1 = lane_group * 8 + j + 1;
float es0 = (pos0 < tile_count) ? __builtin_amdgcn_exp2f(scores[j] * my_scale - my_head_max) : 0.0f;
float es1 = (pos1 < tile_count) ? __builtin_amdgcn_exp2f(scores[j+1] * my_scale - my_head_max) : 0.0f;
tile_sum_partial += es0 + es1;
a_packed[j / 2] = pack_bf16x2(es0, es1);
}
tile_sum_partial += __shfl_xor(tile_sum_partial, 16);
tile_sum_partial += __shfl_xor(tile_sum_partial, 32);
my_head_sum += tile_sum_partial;
// Zero-copy reinterpret a_packed as mfma_bf16_input_t via union
union { uint32_t u[4]; mfma_bf16_input_t v; } a_cvt;
a_cvt.u[0] = a_packed[0]; a_cvt.u[1] = a_packed[1];
a_cvt.u[2] = a_packed[2]; a_cvt.u[3] = a_packed[3];
mfma_bf16_input_t a_bf16 = a_cvt.v;
// V MFMAs with hardware FP4→BF16 conversion
{
asm volatile("s_setprio 3" :::);
// Group 0: V dims [wave_id*128 + lane_col*4 .. +3]
{
uint32_t bp0[4], bp1[4];
v_cvt_fp4_hw_pair<0>(fp4_cur, v_scale_cur, bp0, bp1);
mfma_bf16_input_t b0 = assemble_bf16_input(bp0);
mfma_acc_t c0 = {v_acc[0][0], v_acc[0][1], v_acc[0][2], v_acc[0][3]};
mfma_acc_t r0 = mfma_f32_16x16x32_bf16(a_bf16, b0, c0);
v_acc[0][0]=r0[0]; v_acc[0][1]=r0[1]; v_acc[0][2]=r0[2]; v_acc[0][3]=r0[3];
mfma_bf16_input_t b1 = assemble_bf16_input(bp1);
mfma_acc_t c1 = {v_acc[1][0], v_acc[1][1], v_acc[1][2], v_acc[1][3]};
mfma_acc_t r1 = mfma_f32_16x16x32_bf16(a_bf16, b1, c1);
v_acc[1][0]=r1[0]; v_acc[1][1]=r1[1]; v_acc[1][2]=r1[2]; v_acc[1][3]=r1[3];
}
{
uint32_t bp2[4], bp3[4];
v_cvt_fp4_hw_pair<1>(fp4_cur, v_scale_cur, bp2, bp3);
mfma_bf16_input_t b2 = assemble_bf16_input(bp2);
mfma_acc_t c2 = {v_acc[2][0], v_acc[2][1], v_acc[2][2], v_acc[2][3]};
mfma_acc_t r2 = mfma_f32_16x16x32_bf16(a_bf16, b2, c2);
v_acc[2][0]=r2[0]; v_acc[2][1]=r2[1]; v_acc[2][2]=r2[2]; v_acc[2][3]=r2[3];
mfma_bf16_input_t b3 = assemble_bf16_input(bp3);
mfma_acc_t c3 = {v_acc[3][0], v_acc[3][1], v_acc[3][2], v_acc[3][3]};
mfma_acc_t r3 = mfma_f32_16x16x32_bf16(a_bf16, b3, c3);
v_acc[3][0]=r3[0]; v_acc[3][1]=r3[1]; v_acc[3][2]=r3[2]; v_acc[3][3]=r3[3];
}
// Group 1: V dims [wave_id*128 + 64 + lane_col*4 .. +3]
{
uint32_t bp4[4], bp5[4];
v_cvt_fp4_hw_pair<0>(fp4_next, v_scale_next, bp4, bp5);
mfma_bf16_input_t b4 = assemble_bf16_input(bp4);
mfma_acc_t c4 = {v_acc[4][0], v_acc[4][1], v_acc[4][2], v_acc[4][3]};
mfma_acc_t r4 = mfma_f32_16x16x32_bf16(a_bf16, b4, c4);
v_acc[4][0]=r4[0]; v_acc[4][1]=r4[1]; v_acc[4][2]=r4[2]; v_acc[4][3]=r4[3];
mfma_bf16_input_t b5 = assemble_bf16_input(bp5);
mfma_acc_t c5 = {v_acc[5][0], v_acc[5][1], v_acc[5][2], v_acc[5][3]};
mfma_acc_t r5 = mfma_f32_16x16x32_bf16(a_bf16, b5, c5);
v_acc[5][0]=r5[0]; v_acc[5][1]=r5[1]; v_acc[5][2]=r5[2]; v_acc[5][3]=r5[3];
}
{
uint32_t bp6[4], bp7[4];
v_cvt_fp4_hw_pair<1>(fp4_next, v_scale_next, bp6, bp7);
mfma_bf16_input_t b6 = assemble_bf16_input(bp6);
mfma_acc_t c6 = {v_acc[6][0], v_acc[6][1], v_acc[6][2], v_acc[6][3]};
mfma_acc_t r6 = mfma_f32_16x16x32_bf16(a_bf16, b6, c6);
v_acc[6][0]=r6[0]; v_acc[6][1]=r6[1]; v_acc[6][2]=r6[2]; v_acc[6][3]=r6[3];
mfma_bf16_input_t b7 = assemble_bf16_input(bp7);
mfma_acc_t c7 = {v_acc[7][0], v_acc[7][1], v_acc[7][2], v_acc[7][3]};
mfma_acc_t r7 = mfma_f32_16x16x32_bf16(a_bf16, b7, c7);
v_acc[7][0]=r7[0]; v_acc[7][1]=r7[1]; v_acc[7][2]=r7[2]; v_acc[7][3]=r7[3];
}
asm volatile("s_setprio 0" ::: "memory");
}
if (__builtin_expect(has_next, 1)) wait_kv_loads();
__syncthreads();
fp4_buf ^= 1;
}
// Output
if (num_kv_splits == 1) {
#pragma unroll
for (int k = 0; k < 4; k++) {
int head = lane_group * 4 + k;
if (head >= h_count) continue;
int abs_h = h_start + head;
float hs = __shfl(my_head_sum, head);
float final_scale = (hs > 0.0f) ? 1.0f / hs : 0.0f;
__hip_bfloat16* o_ptr = output + (size_t)q_pos * NHEAD * V_DIM + abs_h * V_DIM;
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++) {
int v_idx = wave_id * 128 + (m / 4) * 64 + lane_col * 4 + (m % 4);
if (v_idx < V_DIM)
o_ptr[v_idx] = __float2bfloat16(v_acc[m][k] * final_scale);
}
}
} else {
__hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM
+ split_id * V_DIM;
float* ml = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits
+ split_id;
#pragma unroll
for (int k = 0; k < 4; k++) {
int head = lane_group * 4 + k;
if (head >= h_count) continue;
int abs_h = h_start + head;
float hs = __shfl(my_head_sum, head);
float final_scale = (hs > 0.0f) ? 1.0f / hs : 0.0f;
#pragma unroll
for (int m = 0; m < V_MFMAS_PER_WAVE; m++) {
int v_idx = wave_id * 128 + (m / 4) * 64 + lane_col * 4 + (m % 4);
if (v_idx < V_DIM)
mo[abs_h * num_kv_splits * V_DIM + v_idx] = __float2bfloat16(v_acc[m][k] * final_scale);
}
}
if (wave_id == 0 && lane_group == 0 && lane_col < h_count) {
int abs_h = h_start + lane_col;
ml[abs_h * num_kv_splits] = (my_head_sum > 0.0f) ?
my_head_max * LN2_F + __logf(my_head_sum) : -1e30f;
}
}
__syncthreads();
}
}
// ============================================================================
// Stage 2
// ============================================================================
// Batched stage2: processes ALL heads per block, reducing grid by 16×
// For large-bs + small-splits where block scheduling dominates
template <int NHEAD, int NSPLITS>
__global__ void
__launch_bounds__(256, 4)
mla_decode_stage2_batch(
const __hip_bfloat16* __restrict__ mid_o,
const float* __restrict__ mid_lse,
__hip_bfloat16* __restrict__ o
) {
const int q_pos = blockIdx.x;
const int tid = threadIdx.x;
// 256 threads / 16 heads = 16 threads per head
constexpr int TPH = NTHREADS / NHEAD; // 16
const int head_id = tid / TPH;
const int local_tid = tid % TPH;
// Load LSE and compute weights (all in registers)
// Layout: [q_pos][head][split] — contiguous splits per head
float lse[NSPLITS];
const float* lse_base = mid_lse + (size_t)q_pos * NHEAD * NSPLITS + head_id * NSPLITS;
float gmax = -1e30f;
#pragma unroll
for (int s = 0; s < NSPLITS; s++) {
lse[s] = lse_base[s];
gmax = fmaxf(gmax, lse[s]);
}
float w[NSPLITS];
float total = 0.0f;
#pragma unroll
for (int s = 0; s < NSPLITS; s++) {
w[s] = __expf(lse[s] - gmax);
total += w[s];
}
float inv = (total > 0.0f) ? 1.0f / total : 0.0f;
#pragma unroll
for (int s = 0; s < NSPLITS; s++) w[s] *= inv;
// Layout: [q_pos][head][split][v_dim] — contiguous splits per head
const __hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * NSPLITS * V_DIM + head_id * NSPLITS * V_DIM;
__hip_bfloat16* op = o + (size_t)q_pos * NHEAD * V_DIM + head_id * V_DIM;
// Vectorized: 16 threads per head, each handles pairs of adjacent dims via BF16x2
for (int dp = local_tid; dp < V_DIM / 2; dp += TPH) {
int d_base = dp * 2;
float val0 = 0.0f, val1 = 0.0f;
#pragma unroll
for (int s = 0; s < NSPLITS; s++) {
uint32_t packed = *reinterpret_cast<const uint32_t*>(&mo[s * V_DIM + d_base]);
__hip_bfloat16 lo_bf16, hi_bf16;
lo_bf16 = *reinterpret_cast<const __hip_bfloat16*>(&packed);
hi_bf16 = *(reinterpret_cast<const __hip_bfloat16*>(&packed) + 1);
val0 += w[s] * __bfloat162float(lo_bf16);
val1 += w[s] * __bfloat162float(hi_bf16);
}
*reinterpret_cast<uint32_t*>(&op[d_base]) = pack_bf16x2(val0, val1);
}
}
// Templated stage2: compile-time NSPLITS for small split counts
template <int NHEAD, int NSPLITS>
__global__ void
__launch_bounds__(256, 4)
mla_decode_stage2_t(
const __hip_bfloat16* __restrict__ mid_o,
const float* __restrict__ mid_lse,
__hip_bfloat16* __restrict__ o
) {
const int q_pos = blockIdx.x;
const int head_id = blockIdx.y;
const int tid = threadIdx.x;
// Load LSE values into registers (compile-time unrolled)
// Layout: [q_pos][head][split] — contiguous splits per head
float lse[NSPLITS];
const float* lse_base = mid_lse + (size_t)q_pos * NHEAD * NSPLITS + head_id * NSPLITS;
float gmax = -1e30f;
#pragma unroll
for (int s = 0; s < NSPLITS; s++) {
lse[s] = lse_base[s];
gmax = fmaxf(gmax, lse[s]);
}
// Compute weights (fully in registers)
float w[NSPLITS];
float total = 0.0f;
#pragma unroll
for (int s = 0; s < NSPLITS; s++) {
w[s] = __expf(lse[s] - gmax);
total += w[s];
}
float inv = (total > 0.0f) ? 1.0f / total : 0.0f;
#pragma unroll
for (int s = 0; s < NSPLITS; s++) w[s] *= inv;
// Layout: [q_pos][head][split][v_dim] — contiguous splits per head
const __hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * NSPLITS * V_DIM + head_id * NSPLITS * V_DIM;
__hip_bfloat16* op = o + (size_t)q_pos * NHEAD * V_DIM + head_id * V_DIM;
if constexpr (NSPLITS <= 16) {
// Vectorized: BF16x2 packed reads — better for moderate split counts
int d_base = tid * 2;
if (d_base < V_DIM) {
float val0 = 0.0f, val1 = 0.0f;
#pragma unroll
for (int s = 0; s < NSPLITS; s++) {
uint32_t packed = *reinterpret_cast<const uint32_t*>(&mo[s * V_DIM + d_base]);
__hip_bfloat16 lo_bf16, hi_bf16;
lo_bf16 = *reinterpret_cast<const __hip_bfloat16*>(&packed);
hi_bf16 = *(reinterpret_cast<const __hip_bfloat16*>(&packed) + 1);
val0 += w[s] * __bfloat162float(lo_bf16);
val1 += w[s] * __bfloat162float(hi_bf16);
}
*reinterpret_cast<uint32_t*>(&op[d_base]) = pack_bf16x2(val0, val1);
}
} else {
// Scalar: better for high split counts (fewer registers, better scheduling)
for (int d = tid; d < V_DIM; d += NTHREADS) {
float val = 0.0f;
#pragma unroll
for (int s = 0; s < NSPLITS; s++)
val += w[s] * __bfloat162float(mo[s * V_DIM + d]);
op[d] = __float2bfloat16(val);
}
}
}
// Partially-unrolled stage2 for large split counts (avoids icache thrash)
template <int NHEAD, int NSPLITS>
__global__ void
__launch_bounds__(256, 4)
mla_decode_stage2_pu(
const __hip_bfloat16* __restrict__ mid_o,
const float* __restrict__ mid_lse,
__hip_bfloat16* __restrict__ o
) {
const int q_pos = blockIdx.x;
const int head_id = blockIdx.y;
const int tid = threadIdx.x;
__shared__ float s_w[NSPLITS];
// Load LSE and compute weights cooperatively
// Layout: [q_pos][head][split] — contiguous splits per head
const float* lse_base = mid_lse + (size_t)q_pos * NHEAD * NSPLITS + head_id * NSPLITS;
float my_lse = -1e30f;
if (tid < NSPLITS) my_lse = lse_base[tid];
// Wave reduction for gmax (all values in one wave for NSPLITS ≤ 64)
float gmax = my_lse;
#pragma unroll
for (int offset = 32; offset >= 1; offset >>= 1)
gmax = fmaxf(gmax, __shfl_xor(gmax, offset));
// Broadcast from lane 0 to all threads
gmax = __shfl(gmax, 0);
if (tid < NSPLITS) s_w[tid] = __expf(my_lse - gmax);
__syncthreads();
// Compute total in one thread
float total = 0.0f;
if (tid == 0) {
#pragma unroll 8
for (int i = 0; i < NSPLITS; i++) total += s_w[i];
}
total = __shfl(total, 0); // Broadcast
float inv = (total > 0.0f) ? 1.0f / total : 0.0f;
// Pre-multiply weights
if (tid < NSPLITS) s_w[tid] *= inv;
__syncthreads();
// Layout: [q_pos][head][split][v_dim] — contiguous splits per head
const __hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * NSPLITS * V_DIM + head_id * NSPLITS * V_DIM;
__hip_bfloat16* op = o + (size_t)q_pos * NHEAD * V_DIM + head_id * V_DIM;
// Vectorized: each thread handles 2 adjacent dims via BF16x2
int d_base = tid * 2;
if (d_base < V_DIM) {
float val0 = 0.0f, val1 = 0.0f;
#pragma unroll 8
for (int s = 0; s < NSPLITS; s++) {
uint32_t packed = *reinterpret_cast<const uint32_t*>(&mo[s * V_DIM + d_base]);
__hip_bfloat16 lo_bf16, hi_bf16;
lo_bf16 = *reinterpret_cast<const __hip_bfloat16*>(&packed);
hi_bf16 = *(reinterpret_cast<const __hip_bfloat16*>(&packed) + 1);
val0 += s_w[s] * __bfloat162float(lo_bf16);
val1 += s_w[s] * __bfloat162float(hi_bf16);
}
*reinterpret_cast<uint32_t*>(&op[d_base]) = pack_bf16x2(val0, val1);
}
}
// Runtime stage2: for large split counts (32, 64) where full unrolling hurts
template <int NHEAD>
__global__ void
__launch_bounds__(256, 4)
mla_decode_stage2(
const __hip_bfloat16* __restrict__ mid_o,
const float* __restrict__ mid_lse,
__hip_bfloat16* __restrict__ o,
int num_kv_splits
) {
const int q_pos = blockIdx.x;
const int head_id = blockIdx.y;
const int tid = threadIdx.x;
__shared__ float s_lse[256], s_w[256], s_total;
// Layout: [q_pos][head][split] — contiguous splits per head
const float* lse_base = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits + head_id * num_kv_splits;
if (tid < num_kv_splits) s_lse[tid] = lse_base[tid];
__syncthreads();
float gmax = -1e30f;
for (int s = 0; s < num_kv_splits; s++) gmax = fmaxf(gmax, s_lse[s]);
if (tid < num_kv_splits) s_w[tid] = __expf(s_lse[tid] - gmax);
__syncthreads();
if (tid == 0) {
float t = 0.0f;
for (int i = 0; i < num_kv_splits; i++) t += s_w[i];
s_total = t;
}
__syncthreads();
float inv = (s_total > 0.0f) ? 1.0f / s_total : 0.0f;
// Layout: [q_pos][head][split][v_dim] — contiguous splits per head
const __hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM + head_id * num_kv_splits * V_DIM;
__hip_bfloat16* op = o + (size_t)q_pos * NHEAD * V_DIM + head_id * V_DIM;
// Vectorized: each thread handles 2 adjacent dims via BF16x2
int d_base = tid * 2;
if (d_base < V_DIM) {
float val0 = 0.0f, val1 = 0.0f;
for (int s = 0; s < num_kv_splits; s++) {
float wt = s_w[s] * inv;
uint32_t packed = *reinterpret_cast<const uint32_t*>(&mo[s * V_DIM + d_base]);
__hip_bfloat16 lo_bf16, hi_bf16;
lo_bf16 = *reinterpret_cast<const __hip_bfloat16*>(&packed);
hi_bf16 = *(reinterpret_cast<const __hip_bfloat16*>(&packed) + 1);
val0 += wt * __bfloat162float(lo_bf16);
val1 += wt * __bfloat162float(hi_bf16);
}
*reinterpret_cast<uint32_t*>(&op[d_base]) = pack_bf16x2(val0, val1);
}
}
// ============================================================================
// C-linkage entry points
// ============================================================================
// Dead stubs — keep pybind signatures but avoid instantiating dead GPU templates
void pb_bf16_to_fp8(torch::Tensor, torch::Tensor, torch::Tensor) {}
void pb_mla_fwd_fp8(torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor,
torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor,
int64_t, int64_t, int64_t, int64_t) {}
void pb_mla_fwd_fp8_fused(torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor,
torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor,
torch::Tensor, int64_t, int64_t, int64_t, int64_t) {}
void pb_compute_q_scale(torch::Tensor, torch::Tensor, torch::Tensor) {}
// BF16 path (inline Q conversion, per-tensor Q scale)
void pb_mla_fwd_bf16(
torch::Tensor q_bf16, torch::Tensor kv_fp8,
torch::Tensor kv_scale,
torch::Tensor mid_o, torch::Tensor mid_lse,
torch::Tensor output, torch::Tensor batch_map, torch::Tensor kv_indptr,
int64_t total_q, int64_t nhead, int64_t batch_size,
int64_t num_kv_splits
) {
auto qb = reinterpret_cast<const __hip_bfloat16*>(q_bf16.data_ptr());
auto kf = reinterpret_cast<const uint8_t*>(kv_fp8.data_ptr());
auto ksp = reinterpret_cast<const float*>(kv_scale.data_ptr());
auto op = reinterpret_cast<__hip_bfloat16*>(output.data_ptr());
auto bm = reinterpret_cast<const int32_t*>(batch_map.data_ptr());
auto ki = reinterpret_cast<const int32_t*>(kv_indptr.data_ptr());
int n_splits = static_cast<int>(num_kv_splits);
int tq = static_cast<int>(total_q);
int bs = static_cast<int>(batch_size);
dim3 block(256);
dim3 grid(n_splits, tq);
#define DISPATCH_BF16(NH) \
if (n_splits == 1) { \
hipLaunchKernelGGL((mla_decode_stage1_bf16<NH, true, true, false>), grid, block, 0, 0, \
qb, kf, ksp, nullptr, nullptr, op, \
bm, ki, bs, 1); \
} else { \
auto mo = reinterpret_cast<__hip_bfloat16*>(mid_o.data_ptr()); \
auto ml = reinterpret_cast<float*>(mid_lse.data_ptr()); \
hipLaunchKernelGGL((mla_decode_stage1_bf16<NH, true, true, false>), grid, block, 0, 0, \
qb, kf, ksp, mo, ml, op, \
bm, ki, bs, n_splits); \
dim3 g2(tq, NH); \
hipLaunchKernelGGL((mla_decode_stage2<NH>), g2, block, 0, 0, \
mo, ml, op, n_splits); \
}
if (nhead == 16) { DISPATCH_BF16(16) }
else if (nhead == 32) { DISPATCH_BF16(32) }
#undef DISPATCH_BF16
}
void pb_mla_fwd_mxfp4(torch::Tensor, torch::Tensor, torch::Tensor, int64_t,
torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor,
torch::Tensor, int64_t, int64_t, int64_t, int64_t) {}
// ============================================================================
// Fast registered-state dispatcher (minimizes Python→C++ overhead)
// ============================================================================
struct DispatchState {
uint8_t* q_fp8;
float* q_scales;
__hip_bfloat16* mid_o;
float* mid_lse;
__hip_bfloat16* output;
int32_t* batch_map;
int32_t* kv_indptr;
int total_q;
int nhead;
int batch_size;
int num_kv_splits;
int q_rows; // total_q * nhead
int64_t cached_q_ptr;
int use_occ3; // 1 = use occ=3 double-buffer kernel
int use_bf16_fused; // 1 = use BF16 fused path (skip bf16_to_fp8)
int use_direct_cvt; // 1 = use direct FP8→BF16 V conversion (v_cvt_scalef32_pk_bf16_fp8)
int use_preload_v; // 1 = preload V from LDS before softmax (default), 0 = load after
int use_setprio; // 1 = use s_setprio hints (default), 0 = disable
const uint8_t* cached_kf;
const float* cached_ksp;
// Paged attention fields
int use_paged; // 1 = use paged kernel
int32_t* kv_indices;
int32_t* paged_kv_indptr;
int32_t* kv_last_page_len;
int page_size;
};
static std::unordered_map<int64_t, DispatchState> g_states;
void pb_register_state(
int64_t key,
torch::Tensor q_fp8, torch::Tensor q_scales,
torch::Tensor mid_o, torch::Tensor mid_lse,
torch::Tensor output, torch::Tensor batch_map, torch::Tensor kv_indptr,
int64_t total_q, int64_t nhead, int64_t batch_size, int64_t num_kv_splits,
int64_t use_occ3, int64_t use_bf16_fused, int64_t use_direct_cvt, int64_t use_preload_v,
int64_t use_setprio
) {
DispatchState s;
s.q_fp8 = reinterpret_cast<uint8_t*>(q_fp8.data_ptr());
s.q_scales = reinterpret_cast<float*>(q_scales.data_ptr());
s.mid_o = (num_kv_splits > 1) ? reinterpret_cast<__hip_bfloat16*>(mid_o.data_ptr()) : nullptr;
s.mid_lse = (num_kv_splits > 1) ? reinterpret_cast<float*>(mid_lse.data_ptr()) : nullptr;
s.output = reinterpret_cast<__hip_bfloat16*>(output.data_ptr());
s.batch_map = reinterpret_cast<int32_t*>(batch_map.data_ptr());
s.kv_indptr = reinterpret_cast<int32_t*>(kv_indptr.data_ptr());
s.total_q = static_cast<int>(total_q);
s.nhead = static_cast<int>(nhead);
s.batch_size = static_cast<int>(batch_size);
s.num_kv_splits = static_cast<int>(num_kv_splits);
s.q_rows = static_cast<int>(total_q * nhead);
s.cached_q_ptr = 0;
s.use_occ3 = static_cast<int>(use_occ3);
s.use_bf16_fused = static_cast<int>(use_bf16_fused);
s.use_direct_cvt = static_cast<int>(use_direct_cvt);
s.use_preload_v = static_cast<int>(use_preload_v);
s.use_setprio = static_cast<int>(use_setprio);
s.cached_kf = nullptr;
s.cached_ksp = nullptr;
s.use_paged = 0;
s.kv_indices = nullptr;
s.paged_kv_indptr = nullptr;
s.kv_last_page_len = nullptr;
s.page_size = 0;
g_states[key] = s;
}
void pb_register_state_paged(
int64_t key,
torch::Tensor mid_o, torch::Tensor mid_lse,
torch::Tensor output, torch::Tensor batch_map,
torch::Tensor kv_indices_t, torch::Tensor paged_kv_indptr_t,
torch::Tensor kv_last_page_len_t,
int64_t total_q, int64_t nhead, int64_t batch_size, int64_t num_kv_splits,
int64_t page_size
) {
DispatchState s;
s.q_fp8 = nullptr;
s.q_scales = nullptr;
s.mid_o = (num_kv_splits > 1) ? reinterpret_cast<__hip_bfloat16*>(mid_o.data_ptr()) : nullptr;
s.mid_lse = (num_kv_splits > 1) ? reinterpret_cast<float*>(mid_lse.data_ptr()) : nullptr;
s.output = reinterpret_cast<__hip_bfloat16*>(output.data_ptr());
s.batch_map = reinterpret_cast<int32_t*>(batch_map.data_ptr());
s.kv_indptr = nullptr;
s.total_q = static_cast<int>(total_q);
s.nhead = static_cast<int>(nhead);
s.batch_size = static_cast<int>(batch_size);
s.num_kv_splits = static_cast<int>(num_kv_splits);
s.q_rows = static_cast<int>(total_q * nhead);
s.cached_q_ptr = 0;
s.use_occ3 = 1;
s.use_bf16_fused = 1;
s.use_direct_cvt = 1;
s.use_preload_v = 1;
s.use_setprio = 0;
s.cached_kf = nullptr;
s.cached_ksp = nullptr;
s.use_paged = 1;
s.kv_indices = reinterpret_cast<int32_t*>(kv_indices_t.data_ptr());
s.paged_kv_indptr = reinterpret_cast<int32_t*>(paged_kv_indptr_t.data_ptr());
s.kv_last_page_len = reinterpret_cast<int32_t*>(kv_last_page_len_t.data_ptr());
s.page_size = static_cast<int>(page_size);
g_states[key] = s;
}
// Launch BF16 fused stage1 (inline Q conversion, skip bf16_to_fp8)
// CONFIG packs per-shape template bools: bit0=DIRECT_CVT, bit1=PRELOAD_V, bit2=USE_SETPRIO
template <int NH, int CONFIG>
static void launch_kernels_bf16_inner(DispatchState& s, const __hip_bfloat16* qb,
const uint8_t* kf, const float* ksp, hipSt_QQ__t stm) {
constexpr bool DIRECT_CVT = (CONFIG >> 0) & 1;
constexpr bool PRELOAD_V = (CONFIG >> 1) & 1;
constexpr bool USE_SETPRIO = (CONFIG >> 2) & 1;
dim3 block(256);
dim3 grid(s.total_q, s.num_kv_splits);
if (s.num_kv_splits == 1) {
hipLaunchKernelGGL((mla_decode_stage1_bf16<NH, DIRECT_CVT, PRELOAD_V, USE_SETPRIO>), grid, block, 0, stm,
qb, kf, ksp, nullptr, nullptr, s.output,
s.batch_map, s.kv_indptr, s.batch_size, 1);
} else {
hipLaunchKernelGGL((mla_decode_stage1_bf16<NH, DIRECT_CVT, PRELOAD_V, USE_SETPRIO>), grid, block, 0, stm,
qb, kf, ksp, s.mid_o, s.mid_lse, s.output,
s.batch_map, s.kv_indptr, s.batch_size, s.num_kv_splits);
}
if (s.num_kv_splits > 1) {
dim3 g2(s.total_q, NH);
dim3 gb(s.total_q);
switch (s.num_kv_splits) {
case 2: hipLaunchKernelGGL((mla_decode_stage2_batch<NH,2>), gb, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;
case 3: hipLaunchKernelGGL((mla_decode_stage2_batch<NH,3>), gb, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;
case 4: hipLaunchKernelGGL((mla_decode_stage2_batch<NH,4>), gb, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;
case 6: hipLaunchKernelGGL((mla_decode_stage2_t<NH,6>), g2, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;
case 8: hipLaunchKernelGGL((mla_decode_stage2_t<NH,8>), g2, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;
case 12: hipLaunchKernelGGL((mla_decode_stage2_t<NH,12>), g2, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;
case 16: hipLaunchKernelGGL((mla_decode_stage2_t<NH,16>), g2, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;
case 24: hipLaunchKernelGGL((mla_decode_stage2_t<NH,24>), g2, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;
case 32: hipLaunchKernelGGL((mla_decode_stage2_t<NH,32>), g2, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;
default: hipLaunchKernelGGL((mla_decode_stage2<NH>), g2, block, 0, stm, s.mid_o, s.mid_lse, s.output, s.num_kv_splits); break;
}
}
}
template <int NH>
static void launch_kernels_bf16_impl(DispatchState& s, const __hip_bfloat16* qb,
const uint8_t* kf, const float* ksp, hipSt_QQ__t stm) {
// All shapes use CONFIG=3: DIRECT_CVT=1, PRELOAD_V=1, USE_SETPRIO=0
// Hardcode to eliminate 7 dead template instantiations (saves icache, -4% on c3/c5/c6)
launch_kernels_bf16_inner<NH, 3>(s, qb, kf, ksp, stm);
}
static void launch_kernels_bf16(DispatchState& s, const __hip_bfloat16* qb,
const uint8_t* kf, const float* ksp, hipSt_QQ__t stm) {
if (s.nhead == 32) launch_kernels_bf16_impl<32>(s, qb, kf, ksp, stm);
else launch_kernels_bf16_impl<16>(s, qb, kf, ksp, stm);
}
// V2 kernel launch: occ=1, FP8 V accumulation via K=128 megatile
template <int NH>
static void launch_kernels_v2_impl(DispatchState& s, const __hip_bfloat16* qb,
const uint8_t* kf, const float* ksp, hipSt_QQ__t stm) {
dim3 block(256);
dim3 grid(s.total_q, s.num_kv_splits);
if (s.num_kv_splits == 1) {
hipLaunchKernelGGL((mla_decode_stage1_v2<NH>), grid, block, 0, stm,
qb, kf, ksp, nullptr, nullptr, s.output,
s.batch_map, s.kv_indptr, s.batch_size, 1);
} else {
hipLaunchKernelGGL((mla_decode_stage1_v2<NH>), grid, block, 0, stm,
qb, kf, ksp, s.mid_o, s.mid_lse, s.output,
s.batch_map, s.kv_indptr, s.batch_size, s.num_kv_splits);
}
if (s.num_kv_splits > 1) {
dim3 g2(s.total_q, NH);
dim3 gb(s.total_q);
switch (s.num_kv_splits) {
case 2: hipLaunchKernelGGL((mla_decode_stage2_batch<NH,2>), gb, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;
case 3: hipLaunchKernelGGL((mla_decode_stage2_batch<NH,3>), gb, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;
case 4: hipLaunchKernelGGL((mla_decode_stage2_batch<NH,4>), gb, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;
case 6: hipLaunchKernelGGL((mla_decode_stage2_t<NH,6>), g2, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;
case 8: hipLaunchKernelGGL((mla_decode_stage2_t<NH,8>), g2, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;
case 12: hipLaunchKernelGGL((mla_decode_stage2_t<NH,12>), g2, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;
case 16: hipLaunchKernelGGL((mla_decode_stage2_t<NH,16>), g2, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;
default: hipLaunchKernelGGL((mla_decode_stage2<NH>), g2, block, 0, stm, s.mid_o, s.mid_lse, s.output, s.num_kv_splits); break;
}
}
}
static void launch_kernels_v2(DispatchState& s, const __hip_bfloat16* qb,
const uint8_t* kf, const float* ksp, hipSt_QQ__t stm) {
if (s.nhead == 32) launch_kernels_v2_impl<32>(s, qb, kf, ksp, stm);
else launch_kernels_v2_impl<16>(s, qb, kf, ksp, stm);
}
void pb_fast_dispatch(
int64_t key,
torch::Tensor q_bf16,
torch::Tensor kv_fp8,
torch::Tensor kv_scale
) {
auto& s = g_states[key];
auto qb = reinterpret_cast<const __hip_bfloat16*>(q_bf16.data_ptr());
hipSt_QQ__t stm = c10::cuda::getCurrentCUDASt_QQ_().st_QQ_();
s.cached_kf = reinterpret_cast<const uint8_t*>(kv_fp8.data_ptr());
s.cached_ksp = reinterpret_cast<const float*>(kv_scale.data_ptr());
// BF16 fused path: inline Q conversion, skip bf16_to_fp8 kernel
launch_kernels_bf16(s, qb, s.cached_kf, s.cached_ksp, stm);
}
void pb_fast_dispatch_v2(
int64_t key,
torch::Tensor q_bf16,
torch::Tensor kv_fp8,
torch::Tensor kv_scale
) {
auto& s = g_states[key];
auto qb = reinterpret_cast<const __hip_bfloat16*>(q_bf16.data_ptr());
hipSt_QQ__t stm = c10::cuda::getCurrentCUDASt_QQ_().st_QQ_();
s.cached_kf = reinterpret_cast<const uint8_t*>(kv_fp8.data_ptr());
s.cached_ksp = reinterpret_cast<const float*>(kv_scale.data_ptr());
launch_kernels_v2(s, qb, s.cached_kf, s.cached_ksp, stm);
}
void pb_dispatch_cached(int64_t) {}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("bf16_to_fp8", &pb_bf16_to_fp8);
m.def("mla_fwd_fp8", &pb_mla_fwd_fp8);
m.def("mla_fwd_fp8_fused", &pb_mla_fwd_fp8_fused);
m.def("mla_fwd_bf16", &pb_mla_fwd_bf16);
m.def("mla_fwd_mxfp4", &pb_mla_fwd_mxfp4);
m.def("compute_q_scale", &pb_compute_q_scale);
m.def("register_state", &pb_register_state);
m.def("fast_dispatch", &pb_fast_dispatch);
m.def("fast_dispatch_v2", &pb_fast_dispatch_v2);
m.def("dispatch_cached", &pb_dispatch_cached);
}
""".replace('_QQ_', 'ream')
CPP_SOURCE = ""
# ============================================================================
# Build + load
# ============================================================================
QK_DIM = 576
V_DIM = 512
# FP8 path splits (for shapes where FP8 is used)
_SPLITS = {
(4, 1, 1024, 16): 16, # c0: was 32 — fewer splits = less stage2 overhead, 2 tiles/split
(4, 1, 8192, 16): 32, # c1: 32 splits optimal (v919 sweep + v974 confirmed: 24=24.5µs, 48=34.9µs)
(32, 1, 1024, 16): 8, # Case 4: 256 WGs = 1/CU
(32, 1, 8192, 16): 8, # Case 3: optimal (less stage2 overhead)
(32, 4, 1024, 16): 8, # Case 3: preserve 1024 WGs for occ4
(32, 4, 8192, 16): 8, # Case 5: preserve 1024 WGs for occ4
(64, 1, 1024, 16): 8,
(64, 1, 8192, 16): 8,
(256, 1, 1024, 16): 1,
(256, 1, 8192, 16): 1, # c7: splits=1 eliminates stage2 (66µs vs 75µs)
(128, 1, 8192, 16): 6, # Case 6: 768 WGs = 3/CU (was 512 = 2/CU)
(128, 4, 8192, 16): 2, # Case 7: fewer splits for qs=4 (less stage2 overhead)
# tp=4 (32 heads)
(4, 1, 1024, 32): 16,
(4, 4, 8192, 32): 16,
(32, 1, 8192, 32): 16,
(32, 4, 1024, 32): 8,
(32, 1, 1024, 32): 8,
(32, 4, 8192, 32): 8,
(128, 1, 8192, 32): 6,
(128, 4, 8192, 32): 2,
}
# Per-case kernel selection: 0 = occ4, 1 = occ3 double-buffer, 2 = occ2 triple-buffer
_USE_OCC3 = {
(4, 1, 1024, 16): 2, # c0: occ2 triple-buf (13.8→13.5µs)
(4, 1, 8192, 16): 1, # 256 WGs = 1/CU
(32, 1, 1024, 16): 1, # 256 WGs = 1/CU
(32, 1, 8192, 16): 1, # 512 WGs = 2/CU (splits=16), occ3 (occ4 produced NaN)
(64, 1, 1024, 16): 2, # c4: occ2 triple-buf (28.3→27.6µs)
(64, 1, 8192, 16): 1, # 512 WGs = 2/CU
(256, 1, 1024, 16): 1, # 256 WGs = 1/CU
(256, 1, 8192, 16): 2, # c7: occ2 triple-buf (67.6→66.5µs)
(32, 4, 1024, 16): 0, # 1024 WGs (qs=4), occ4 for full 4/CU utilization
(32, 4, 8192, 16): 0, # 1024 WGs (qs=4), occ4 for full 4/CU utilization
(128, 1, 8192, 16): 1, # 512 WGs (qs=1), occ3 for packed V conversion
(128, 4, 8192, 16): 0, # 2048 WGs (qs=4), occ4 for higher occupancy
# tp=4 (32 heads)
(4, 1, 1024, 32): 1,
(4, 4, 8192, 32): 1,
(32, 1, 8192, 32): 1,
(32, 4, 1024, 32): 0,
(32, 1, 1024, 32): 1,
(32, 4, 8192, 32): 0,
(128, 1, 8192, 32): 1,
(128, 4, 8192, 32): 0,
}
# Per-case V conversion method: 0 = cvt_pk+pack (default), 1 = v_cvt_scalef32_pk_bf16_fp8+v_perm_b32
_USE_DIRECT_CVT = {
(4, 1, 1024, 16): 1, # c0: kv=1024
(4, 1, 8192, 16): 1, # c1: test with SLP
(32, 1, 1024, 16): 1, # c2: kv=1024
(32, 1, 8192, 16): 1, # c3: test with SLP
(64, 1, 1024, 16): 1, # c4: kv=1024
(64, 1, 8192, 16): 1, # c5: test with SLP
(256, 1, 1024, 16): 1, # c6: kv=1024
(256, 1, 8192, 16): 1, # c7: test with SLP
}
# Per-case V preload strategy: 1 = preload before softmax (default), 0 = load after softmax
_USE_PRELOAD_V = {}
# Per-case s_setprio hints: 1 = enable (default), 0 = disable
# Setprio boosts MFMA priority but adds VALU overhead — may hurt small-tile cases
_USE_SETPRIO = {}
_ext = None
_cache = {}
_registered = set()
def _get_ext():
global _ext
if _ext is not None:
return _ext
from torch.utils.cpp_extension import load_inline
_ext = load_inline(
name="mla_decode_v994",
cpp_sources=CPP_SOURCE,
cuda_sources=HIP_SOURCE,
extra_cuda_cflags=[
"-Os", "-ffast-math", "--offload-arch=gfx950", "-std=c++17",
"-D__gfx950__",
"-mllvm", "-amdgpu-early-inline-all=true",
"-mllvm", "-amdgpu-function-calls=false",
"-mllvm", "-amdgpu-max-memory-clause=16",
"-mllvm", "-amdgpu-promote-alloca-to-vector-limit=128",
],
verbose=False,
)
return _ext
# Build at import time (before benchmark timing)
_get_ext()
def _custom_kernel_hip(data: input_t) -> output_t:
"""Our custom HIP kernel — fastest for small batch × kv products."""
ext = _ext
cfg = data[4]
bs = cfg["batch_size"]
nh = cfg["num_heads"]
kvsl = cfg["kv_seq_len"]
qs = cfg["q_seq_len"]
key = bs * 10000000 + qs * 1000000 + kvsl * 100 + nh
if key not in _registered:
total_q = bs * cfg["q_seq_len"]
num_splits = _SPLITS.get((bs, qs, kvsl, nh), _SPLITS.get((bs, 1, kvsl, nh), 4))
dev = data[0].device
o = torch.empty((total_q, nh, V_DIM), dtype=torch.bfloat16, device=dev)
q_fp8 = torch.empty(total_q * nh, QK_DIM, dtype=torch.uint8, device=dev)
q_scales = torch.empty(total_q * nh, dtype=torch.float32, device=dev)
dummy = torch.empty(1, dtype=torch.float32, device=dev)
mid_o = dummy
mid_lse = dummy
if num_splits > 1:
mid_o = torch.empty(total_q * num_splits * nh * V_DIM, dtype=torch.bfloat16, device=dev)
mid_lse = torch.empty(total_q * num_splits * nh, dtype=torch.float32, device=dev)
batch_map = torch.arange(bs, dtype=torch.int32, device=dev).repeat_interleave(qs)
kv_indptr_t = torch.arange(bs + 1, dtype=torch.int32, device=dev) * kvsl
use_occ3 = _USE_OCC3.get((bs, qs, kvsl, nh), _USE_OCC3.get((bs, 1, kvsl, nh), 0))
use_bf16_fused = 1
use_direct_cvt = _USE_DIRECT_CVT.get((bs, qs, kvsl, nh), 0)
use_preload_v = _USE_PRELOAD_V.get((bs, qs, kvsl, nh), 1)
use_setprio = _USE_SETPRIO.get((bs, qs, kvsl, nh), 0)
ext.register_state(key, q_fp8, q_scales, mid_o, mid_lse, o,
batch_map, kv_indptr_t, total_q, nh, bs, num_splits, use_occ3, use_bf16_fused, use_direct_cvt, use_preload_v, use_setprio)
_cache[key] = (o, q_fp8, q_scales, mid_o, mid_lse, batch_map, kv_indptr_t)
_registered.add(key)
o = _cache[key][0]
kv = data[1]["fp8"]
ext.fast_dispatch(key, data[0], kv[0], kv[1])
return o
def _custom_kernel_hip_paged(data: input_t) -> output_t:
"""Our custom HIP kernel with paged attention — for large batch × kv shapes."""
ext = _ext
cfg = data[4]
bs = cfg["batch_size"]
nh = cfg["num_heads"]
kvsl = cfg["kv_seq_len"]
qs = cfg["q_seq_len"]
key = bs * 10000000 + qs * 1000000 + kvsl * 100 + nh + 50 # +50 to avoid collision with non-paged keys
if key not in _paged_registered:
total_q = bs * qs
ps = _PAGED_PAGE_SIZE
num_splits = _PAGED_SPLITS.get((bs, qs, kvsl, nh), 4)
dev = data[0].device
pages_per_seq = (kvsl + ps - 1) // ps
ebs = bs # effective batch size (qs=1)
o = torch.empty((total_q, nh, V_DIM), dtype=torch.bfloat16, device=dev)
dummy = torch.empty(1, dtype=torch.float32, device=dev)
mid_o = dummy
mid_lse = dummy
if num_splits > 1:
mid_o = torch.empty(total_q * num_splits * nh * V_DIM, dtype=torch.bfloat16, device=dev)
mid_lse = torch.empty(total_q * num_splits * nh, dtype=torch.float32, device=dev)
batch_map = torch.arange(bs, dtype=torch.int32, device=dev).repeat_interleave(qs)
# Page table setup (contiguous KV → identity page mapping)
kv_indices = torch.arange(ebs * pages_per_seq, dtype=torch.int32, device=dev)
paged_kv_indptr = torch.arange(ebs + 1, dtype=torch.int32, device=dev) * pages_per_seq
kv_last_page_len = torch.full((ebs,), kvsl % ps if kvsl % ps != 0 else ps, dtype=torch.int32, device=dev)
ext.register_state_paged(key, mid_o, mid_lse, o, batch_map,
kv_indices, paged_kv_indptr, kv_last_page_len,
total_q, nh, bs, num_splits, ps)
_paged_cache[key] = (o, kv_indices, paged_kv_indptr, kv_last_page_len)
_paged_registered.add(key)
kv = data[1]["fp8"]
kv_fp8, kv_scale = kv[0], kv[1]
# kv_fp8 shape: [bs*kvsl, 1, QK_DIM] → reshape to pages: [bs*pages_per_seq, page_size, QK_DIM]
ps = _PAGED_PAGE_SIZE
pages_per_seq = kvsl // ps # kvsl is always divisible by 32
o = _paged_cache[key][0]
kv_paged = kv_fp8.view(bs * pages_per_seq, ps, QK_DIM)
ext.fast_dispatch(key, data[0], kv_paged, kv_scale)
return o
# V2 kernel split tuning: occ=1, FP8 K=128 V MFMA
_V2_SPLITS = {
(32, 1, 8192, 16): 8, # 32*8=256 WGs
(64, 1, 8192, 16): 4, # 64*4=256 WGs
(256, 1, 1024, 16): 1, # 256*1=256 WGs
(256, 1, 8192, 16): 1, # 256*1=256 WGs
}
_v2_cache = {}
_v2_registered = set()
def _custom_kernel_hip_v2(data: input_t) -> output_t:
"""V2 HIP kernel: occ=1, FP8 V accumulation via K=128 megatile."""
ext = _ext
cfg = data[4]
bs = cfg["batch_size"]
nh = cfg["num_heads"]
kvsl = cfg["kv_seq_len"]
qs = cfg["q_seq_len"]
key = bs * 10000000 + qs * 1000000 + kvsl * 100 + nh + 70 # +70 to avoid collision
if key not in _v2_registered:
total_q = bs * qs
num_splits = _V2_SPLITS.get((bs, qs, kvsl, nh), 4)
dev = data[0].device
o = torch.empty((total_q, nh, V_DIM), dtype=torch.bfloat16, device=dev)
q_fp8 = torch.empty(total_q * nh, QK_DIM, dtype=torch.uint8, device=dev)
q_scales = torch.empty(total_q * nh, dtype=torch.float32, device=dev)
dummy = torch.empty(1, dtype=torch.float32, device=dev)
mid_o = dummy
mid_lse = dummy
if num_splits > 1:
mid_o = torch.empty(total_q * num_splits * nh * V_DIM, dtype=torch.bfloat16, device=dev)
mid_lse = torch.empty(total_q * num_splits * nh, dtype=torch.float32, device=dev)
batch_map = torch.arange(bs, dtype=torch.int32, device=dev).repeat_interleave(qs)
kv_indptr_t = torch.arange(bs + 1, dtype=torch.int32, device=dev) * kvsl
# Reuse the same register_state — v2 uses same DispatchState layout
ext.register_state(key, q_fp8, q_scales, mid_o, mid_lse, o,
batch_map, kv_indptr_t, total_q, nh, bs, num_splits, 0, 1, 0, 1, 0)
_v2_cache[key] = (o, q_fp8, q_scales, mid_o, mid_lse, batch_map, kv_indptr_t)
_v2_registered.add(key)
o = _v2_cache[key][0]
kv = data[1]["fp8"]
ext.fast_dispatch_v2(key, data[0], kv[0], kv[1])
return o
# ============================================================================
# AITER path — for large batch × kv shapes where paged attention wins
# ============================================================================
from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
_FP8_DTYPE = aiter_dtypes.fp8
_SM_SCALE = 1.0 / (QK_DIM ** 0.5)
# Per-case AITER tuning: (bs, kvsl) -> (page_size, num_kv_splits, fast_mode)
_AITER_TUNE = {
(4, 1024): (1, 32, False), # c0: HIP handles this
(4, 8192): (1, 32, False), # c1: HIP handles this
(32, 1024): (1, 32, False), # c2: HIP handles this
(32, 8192): (8, 16, False), # c3: ps=8 ns=16 + kv_gran=32
(64, 1024): (2, 1, False), # c4: AITER ps=2 beats HIP
(64, 8192): (8, 2, False), # c5
(128, 1024): (2, 8, False),
(128, 8192): (8, 8, False),
(256, 1024): (2, 1, True), # c6
(256, 8192): (8, 32, True), # c7
}
_aiter_cache = {}
_aiter_q_scale = None
def _aiter_build(cfg_key, bs, qsl, kvsl, nh, kv_indptr, dev):
global _aiter_q_scale
if cfg_key in _aiter_cache:
return _aiter_cache[cfg_key]
if _aiter_q_scale is None:
_aiter_q_scale = torch.ones(1, dtype=torch.float32, device=dev)
ps, ns, fm = _AITER_TUNE.get((bs, kvsl), (1, 32, bs <= 4))
kv_gran = 32 if ps == 8 else max(ps, 16)
ebs = bs * qsl
eff_qo_indptr = torch.arange(ebs + 1, dtype=torch.int32, device=dev)
if ps > 1:
pages_per_seq = (kvsl + ps - 1) // ps
kv_last_page_len = torch.full((ebs,),
kvsl % ps if kvsl % ps != 0 else ps,
dtype=torch.int32, device=dev)
kv_indices = torch.arange(ebs * pages_per_seq, dtype=torch.int32, device=dev)
paged_kv_indptr = torch.arange(ebs + 1, dtype=torch.int32, device=dev) * pages_per_seq
else:
kv_last_page_len = torch.full((ebs,), ps, dtype=torch.int32, device=dev)
kv_indices = torch.arange(int(kv_indptr[-1].item()), dtype=torch.int32, device=dev)
paged_kv_indptr = kv_indptr
info = get_mla_metadata_info_v1(ebs, 1, nh, _FP8_DTYPE, _FP8_DTYPE,
is_sparse=False, fast_mode=fm, num_kv_splits=ns, intra_batch_mode=(not fm))
work = [torch.empty(s, dtype=t, device=dev) for s, t in info]
wm, wi, wis, ri, rfm, rpm = work
get_mla_metadata_v1(eff_qo_indptr, paged_kv_indptr, kv_last_page_len,
nh, 1, True, wm, wis, wi, ri, rfm, rpm,
page_size=ps, kv_granularity=kv_gran, max_seqlen_qo=1, uni_seqlen_qo=1,
fast_mode=fm, max_split_per_batch=ns, intra_batch_mode=(not fm),
dtype_q=_FP8_DTYPE, dtype_kv=_FP8_DTYPE)
e = {
'ps': ps, 'ns': ns, 'fm': fm, 'ibm': not fm,
'eqi': eff_qo_indptr, 'kvi': paged_kv_indptr, 'ki': kv_indices, 'klp': kv_last_page_len,
'meta': {'work_meta_data': wm, 'work_indptr': wi, 'work_info_set': wis,
'reduce_indptr': ri, 'reduce_final_map': rfm, 'reduce_partial_map': rpm},
'q_fp8_buf': torch.empty((ebs, nh, QK_DIM), dtype=_FP8_DTYPE, device=dev),
'o_buf': torch.empty((ebs, nh, V_DIM), dtype=torch.bfloat16, device=dev),
}
_aiter_cache[cfg_key] = e
return e
def _custom_kernel_aiter(data: input_t) -> output_t:
"""AITER paged attention — fastest for large batch × kv shapes."""
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
nh = config["num_heads"]
qsl = config["q_seq_len"]
kvsl = config["kv_seq_len"]
c = _aiter_build((bs, qsl, kvsl, nh), bs, qsl, kvsl, nh, kv_indptr, q.device)
r, s = kv_data["fp8"]
ks = s.view(1) if s.numel() == 1 else s
c['q_fp8_buf'].copy_(q.view(bs * qsl, nh, QK_DIM))
ps = c['ps']
if ps > 1:
kv_4d = r.reshape(bs * qsl * ((kvsl + ps - 1) // ps), ps, 1, QK_DIM)
else:
kv_4d = r.view(bs * kvsl, 1, 1, QK_DIM)
mla_decode_fwd(
c['q_fp8_buf'], kv_4d, c['o_buf'], c['eqi'], c['kvi'], c['ki'], c['klp'],
1, page_size=ps, nhead_kv=1, sm_scale=_SM_SCALE, logit_cap=0.0,
num_kv_splits=c['ns'], q_scale=_aiter_q_scale, kv_scale=ks,
intra_batch_mode=c['ibm'], **c['meta'])
return c['o_buf']
# ============================================================================
# Hybrid dispatch — pick best kernel per shape
# ============================================================================
# Shapes where our custom HIP kernel beats AITER (measured on MI355X):
# c0(4,1024)=12.7 vs 22.1, c1(4,8192)=21.3 vs 23.4,
# c2(32,1024)=19.0 vs 23.9, c4(64,1024)=27.3 vs 28.2
# HIP wins: c0(4,1024)=13.4, c1(4,8192)=22.1, c2(32,1024)=19.5
_USE_HIP = {(4, 1, 1024, 16), (4, 1, 8192, 16), (32, 1, 1024, 16)}
# Shapes where we use paged HIP kernel (page_size=32=KV_TILE for exact alignment)
# Currently empty — HIP kernel not competitive for large shapes vs AITER ASM
_USE_HIP_PAGED = set()
_PAGED_PAGE_SIZE = 32 # Must match KV_TILE
# Shapes where V2 kernel (occ=1, FP8 V MFMA) is used
# Start with large shapes where bf16 V path is MFMA-bound
_USE_HIP_V2 = set() # Disabled: v2 megatile 2x slower than AITER due to LDS scatter overhead
# Paged kernel split tuning
_PAGED_SPLITS = {
(32, 1, 8192, 16): 8, # 32*8=256 WGs
(64, 1, 8192, 16): 4, # 64*4=256 WGs
(256, 1, 1024, 16): 1, # 256*1=256 WGs
(256, 1, 8192, 16): 1, # 256*1=256 WGs
}
_paged_cache = {}
_paged_registered = set()
def custom_kernel(data: input_t) -> output_t:
cfg = data[4]
shape_key = (cfg["batch_size"], cfg["q_seq_len"], cfg["kv_seq_len"], cfg["num_heads"])
if shape_key in _USE_HIP:
return _custom_kernel_hip(data)
elif shape_key in _USE_HIP_V2:
return _custom_kernel_hip_v2(data)
elif shape_key in _USE_HIP_PAGED:
return _custom_kernel_hip_paged(data)
else:
return _custom_kernel_aiter(data)
scrolls · 3953 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 693347.
⋯ 284 unchanged linesuint8_t kv_fp8[3][LDS_BUF_SIZE];};+ // OCC1 5-buffer LDS for V2 kernel: 5 × 18,432 = 92,160 bytes < 163,840 (160KB at occ=1)+ // 4 retained buffers for FP8 V MFMA (K=128 = 4×32 tokens), plus 1 prefetch buffer+ constexpr int V2_NUM_BUFS = 5;+ constexpr int V2_MEGATILE = 4;+ constexpr int V2_MEGATILE_TOKENS = V2_MEGATILE * KV_TILE; // 128++ struct LDS_V2 {+ uint8_t kv_fp8[V2_NUM_BUFS][LDS_BUF_SIZE]; // 92,160 bytes+ };+// FP4 path: double-buffered FP4 loads + scale cache (native FP4 MFMA, no FP8 buffer)constexpr int LDS_FP4_BUF_SIZE = KV_TILE * KV_FP4_STRIDE + (WAVESIZE - 1) * 16; // 10,224⋯ 213 unchanged linesasm volatile("s_waitcnt vmcnt(0)" ::: "memory");}+ // Strided KV tile load: loads tile_count tokens from non-contiguous positions.+ // Token i comes from kv_base[(first_token + i*stride) * QK_DIM] and is packed+ // contiguously in lds.kv_fp8[buf_idx] for MFMA processing. Uses VGPR-mediated+ // global_load → ds_write (not DMA). Same LDS layout as issue_kv_loads — callers+ // don't see the difference. Synchronous: all LDS writes complete before return.+ template <typename LDS_T>+ __device__ __forceinline__ void+ issue_kv_loads_strided(+ LDS_T& lds,+ int buf_idx,+ const uint8_t* __restrict__ kv_base,+ int first_token,+ int stride,+ int tile_count,+ int tid+ ) {+ constexpr int ITEMS_PER_TOKEN = QK_DIM / 16; // 576/16 = 36+ const int total_items = tile_count * ITEMS_PER_TOKEN;++ for (int item = tid; item < total_items; item += 256) {+ int tok = item / ITEMS_PER_TOKEN;+ int sub = item - tok * ITEMS_PER_TOKEN;+ int byte_off = sub * 16;++ int global_token = first_token + tok * stride;+ uint4 val = *reinterpret_cast<const uint4*>(+ kv_base + (size_t)global_token * QK_DIM + byte_off);++ *reinterpret_cast<uint4*>(+ &lds.kv_fp8[buf_idx][tok * LDS_KV_STRIDE + byte_off]) = val;+ }+ asm volatile("s_waitcnt vmcnt(0) lgkmcnt(0)" ::: "memory");+ }+// ============================================================================// COMMON: tile loop body (shared between FP8 and BF16 paths)// ============================================================================⋯ 1045 unchanged lines// KEY OPTIMIZATION: Issue first KV tile load BEFORE Q conversion// buffer_load...lds is async — Q conversion overlaps with it// ============================================================- int first_count = min(KV_TILE, split_end - split_start);- issue_kv_loads(0, kv_fp8 + (size_t)(kv_start + split_start) * QK_DIM,+ int first_tile = split_start;++ int first_count = min(KV_TILE, split_end - first_tile);+ issue_kv_loads(0, kv_fp8 + (size_t)(kv_start + first_tile) * QK_DIM,first_count, tid);// Inline Q BF16→FP8 conversion (no scaling — FP8 e4m3 range ±448 covers typical Q values)⋯ 34 unchanged lineswait_kv_loads();__syncthreads();- // Tile loopint buf = 0;- for (int tile_start = split_start; tile_start < split_end; tile_start += KV_TILE) {+ for (int tile_start = first_tile; tile_start < split_end; tile_start += KV_TILE) {int tile_count = min(KV_TILE, split_end - tile_start);int next_start = tile_start + KV_TILE;int has_next = (next_start < split_end);⋯ 193 unchanged lines}// ============================================================================+ // Stage 1 — BF16 fused PAGED path (page_size=KV_TILE=32 for exact alignment)+ // ============================================================================+ // Identical compute to mla_decode_stage1_bf16 but KV addressed via page tables.+ // Only the address computation per tile changes — all MFMA/softmax/V is shared.++ template <int NHEAD, bool DIRECT_CVT = true, bool PRELOAD_V = true, bool USE_SETPRIO = false>+ __global__ void+ __launch_bounds__(256, 3)+ mla_decode_stage1_bf16_paged(+ const __hip_bfloat16* __restrict__ q_bf16,+ const uint8_t* __restrict__ kv_pages, // [total_pages, page_size, QK_DIM]+ const float* __restrict__ kv_scale_ptr,+ __hip_bfloat16* __restrict__ mid_o,+ float* __restrict__ mid_lse,+ __hip_bfloat16* __restrict__ output,+ const int32_t* __restrict__ batch_map,+ const int32_t* __restrict__ kv_indices, // logical->physical page map+ const int32_t* __restrict__ paged_kv_indptr, // [batch+1] page ranges+ const int32_t* __restrict__ kv_last_page_len, // valid tokens in last page+ int batch_size,+ int num_kv_splits,+ int page_size // must equal KV_TILE (32)+ ) {+ constexpr float sm_scale = 1.0f / 24.0f;+ const float kv_scale = *kv_scale_ptr;+ constexpr int NUM_HEAD_GROUPS = (NHEAD + HEAD_GROUP - 1) / HEAD_GROUP;++ const int q_pos = blockIdx.x;+ const int split_id = blockIdx.y;+ const int tid = threadIdx.x;+ const int wave_id = tid / WAVESIZE;+ const int lane_id = tid % WAVESIZE;+ const int lane_group = lane_id / 16;+ const int lane_col = lane_id % 16;++ // Paged address resolution+ const int batch_id = batch_map[q_pos];+ const int pg_start = paged_kv_indptr[batch_id];+ const int pg_end = paged_kv_indptr[batch_id + 1];+ const int num_pages = pg_end - pg_start;+ const int last_pg_len = kv_last_page_len[batch_id];+ const int total_kv_len = (num_pages > 1) ? (num_pages - 1) * page_size + last_pg_len : last_pg_len;++ // Page-aligned splits: round up to page boundary+ int pages_per_split = (num_pages + num_kv_splits - 1) / num_kv_splits;+ int split_page_start = split_id * pages_per_split;+ int split_page_end = min(split_page_start + pages_per_split, num_pages);++ int split_start = split_page_start * page_size;+ int split_end = (split_page_end < num_pages)+ ? split_page_end * page_size+ : total_kv_len;++ __shared__ LDS lds;++ for (int hg = 0; hg < NUM_HEAD_GROUPS; hg++) {+ int h_start = hg * HEAD_GROUP;+ int h_count = min(HEAD_GROUP, NHEAD - h_start);++ float my_head_max = -1e30f;+ float my_head_sum = 0.0f;++ mfma_acc_t v_acc[V_MFMAS_PER_WAVE];+ #pragma unroll+ for (int m = 0; m < V_MFMAS_PER_WAVE; m++)+ v_acc[m] = {0.0f, 0.0f, 0.0f, 0.0f};++ if (split_start >= split_end) {+ if (num_kv_splits > 1) {+ __hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM+ + split_id * V_DIM;+ float* ml = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits+ + split_id;+ __hip_bfloat16 zero_bf16 = __float2bfloat16(0.0f);+ #pragma unroll+ for (int k = 0; k < 4; k++) {+ int head = lane_group * 4 + k;+ if (head < h_count) {+ int abs_h = h_start + head;+ #pragma unroll+ for (int m = 0; m < V_MFMAS_PER_WAVE; m++) {+ int v_idx = wave_id * 128 + (m / 4) * 64 + lane_col * 4 + (m % 4);+ if (v_idx < V_DIM)+ mo[abs_h * num_kv_splits * V_DIM + v_idx] = zero_bf16;+ }+ }+ }+ if (wave_id == 0 && lane_group == 0 && lane_col < h_count)+ ml[(h_start + lane_col) * num_kv_splits] = -1e30f;+ }+ __syncthreads();+ continue;+ }++ // First tile prefetch — resolve page address+ int cur_page = split_page_start;+ int first_phys = kv_indices[pg_start + cur_page];+ int first_count = min(KV_TILE, split_end - split_start);+ issue_kv_loads(0, kv_pages + (size_t)first_phys * page_size * QK_DIM,+ first_count, tid);++ // Inline Q BF16→FP8 conversion (overlaps with async KV load)+ const __hip_bfloat16* q_head_base = q_bf16 + (size_t)(q_pos * NHEAD + h_start) * QK_DIM;++ float my_scale = 0.0f;+ if (lane_col < h_count) {+ my_scale = kv_scale * sm_scale * LOG2E_F;+ }++ uint32_t q_cache[QK_MFMAS][8];+ #pragma unroll+ for (int t = 0; t < QK_MFMAS; t++) {+ int elem_start = t * FP8_MFMA_K + lane_group * 32;+ if (elem_start + 32 <= QK_DIM && lane_col < h_count) {+ const uint32_t* wp = reinterpret_cast<const uint32_t*>(+ q_head_base + (size_t)lane_col * QK_DIM + elem_start);+ #pragma unroll+ for (int j = 0; j < 8; j++) {+ uint32_t w0 = wp[j * 2];+ uint32_t w1 = wp[j * 2 + 1];+ uint32_t pk;+ float s1 = 1.0f;+ asm("v_cvt_scalef32_pk_fp8_bf16 %0, %1, %2"+ : "=v"(pk) : "v"(w0), "v"(s1));+ asm("v_cvt_scalef32_pk_fp8_bf16 %0, %1, %2 op_sel:[0,0,1]"+ : "+v"(pk) : "v"(w1), "v"(s1));+ q_cache[t][j] = pk;+ }+ } else {+ q_cache[t][0]=0; q_cache[t][1]=0; q_cache[t][2]=0; q_cache[t][3]=0;+ q_cache[t][4]=0; q_cache[t][5]=0; q_cache[t][6]=0; q_cache[t][7]=0;+ }+ }++ wait_kv_loads();+ __syncthreads();++ // Tile loop — page-aligned: each tile = one page (page_size=KV_TILE)+ int buf = 0;+ for (int tile_start = split_start; tile_start < split_end; tile_start += KV_TILE) {+ int tile_count = min(KV_TILE, split_end - tile_start);+ int next_page = cur_page + 1;+ int next_tile_start = tile_start + KV_TILE;+ int has_next = (next_tile_start < split_end);++ if (__builtin_expect(has_next, 1)) {+ int next_phys = kv_indices[pg_start + next_page];+ int next_count = min(KV_TILE, split_end - next_tile_start);+ issue_kv_loads((buf ^ 1) * LDS_BUF_SIZE,+ kv_pages + (size_t)next_phys * page_size * QK_DIM,+ next_count, tid);+ }++ const uint8_t* kv_lds = lds.kv_fp8[buf];++ // === QK scoring (IDENTICAL to non-paged) ===+ mfma_acc_t acc_lo = {0.0f, 0.0f, 0.0f, 0.0f};+ mfma_acc_t acc_hi = {0.0f, 0.0f, 0.0f, 0.0f};++ if constexpr (USE_SETPRIO) asm volatile("s_setprio 3" :::);+ #pragma unroll+ for (int t = 0; t < QK_MFMAS; t++) {+ mfma_input_t frag_q = {q_cache[t][0], q_cache[t][1], q_cache[t][2], q_cache[t][3],+ q_cache[t][4], q_cache[t][5], q_cache[t][6], q_cache[t][7]};+ int kv_byte = t * FP8_MFMA_K + lane_group * 32;++ mfma_input_t frag_kv_lo, frag_kv_hi;+ if (kv_byte + 32 <= QK_DIM) {+ load_frag_b_wide(kv_lds, lane_col * LDS_KV_STRIDE + kv_byte, frag_kv_lo);+ load_frag_b_wide(kv_lds, (lane_col + KV_HALF) * LDS_KV_STRIDE + kv_byte, frag_kv_hi);+ } else {+ frag_kv_lo[0]=0; frag_kv_lo[1]=0; frag_kv_lo[2]=0; frag_kv_lo[3]=0;+ frag_kv_lo[4]=0; frag_kv_lo[5]=0; frag_kv_lo[6]=0; frag_kv_lo[7]=0;+ frag_kv_hi[0]=0; frag_kv_hi[1]=0; frag_kv_hi[2]=0; frag_kv_hi[3]=0;+ frag_kv_hi[4]=0; frag_kv_hi[5]=0; frag_kv_hi[6]=0; frag_kv_hi[7]=0;+ }++ acc_lo = mfma_f32_16x16x128_fp8(frag_kv_lo, frag_q, acc_lo);+ acc_hi = mfma_f32_16x16x128_fp8(frag_kv_hi, frag_q, acc_hi);+ }+ if constexpr (USE_SETPRIO) asm volatile("s_setprio 0" ::: "memory");++ // === V preload (IDENTICAL) ===+ uint32_t raw4_cur[8], raw4_next_v[8];+ int blo = lane_group * 4;+ int bhi = lane_group * 4 + 16;+ int v_off0 = wave_id * 128 + lane_col * 4;+ int v_off1 = v_off0 + 64;+ if constexpr (PRELOAD_V) {+ #pragma unroll+ for (int i = 0; i < 4; i++) {+ raw4_cur[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(blo + i) * LDS_KV_STRIDE + v_off0]);+ raw4_cur[i+4] = *reinterpret_cast<const uint32_t*>(&kv_lds[(bhi + i) * LDS_KV_STRIDE + v_off0]);+ raw4_next_v[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(blo + i) * LDS_KV_STRIDE + v_off1]);+ raw4_next_v[i+4] = *reinterpret_cast<const uint32_t*>(&kv_lds[(bhi + i) * LDS_KV_STRIDE + v_off1]);+ }+ }++ // === Softmax (IDENTICAL) ===+ float partial_max = -1e30f;+ partial_max = fmaxf(partial_max, fmaxf(acc_lo[0], acc_lo[1]));+ partial_max = fmaxf(partial_max, fmaxf(acc_lo[2], acc_lo[3]));+ partial_max = fmaxf(partial_max, fmaxf(acc_hi[0], acc_hi[1]));+ partial_max = fmaxf(partial_max, fmaxf(acc_hi[2], acc_hi[3]));+ partial_max = fmaxf(partial_max, __shfl_xor(partial_max, 16));+ partial_max = fmaxf(partial_max, __shfl_xor(partial_max, 32));++ float scaled_partial_max = partial_max * my_scale;+ float old_max = my_head_max;+ float new_max = fmaxf(old_max, scaled_partial_max);+ my_head_max = new_max;+ float neg_max = -new_max;++ if (scaled_partial_max > old_max && old_max > -1e29f) {+ float my_rescale = __builtin_amdgcn_exp2f(old_max - new_max);+ my_head_sum *= my_rescale;+ #pragma unroll+ for (int k = 0; k < 4; k++) {+ float r = __shfl(my_rescale, lane_group * 4 + k);+ #pragma unroll+ for (int m = 0; m < V_MFMAS_PER_WAVE; m++)+ v_acc[m][k] *= r;+ }+ }++ uint32_t a_packed[4];+ float tile_sum_partial = 0.0f;+ #pragma unroll+ for (int j = 0; j < 4; j += 2) {+ float es0 = __builtin_amdgcn_exp2f(fmaf(acc_lo[j], my_scale, neg_max));+ float es1 = __builtin_amdgcn_exp2f(fmaf(acc_lo[j+1], my_scale, neg_max));+ tile_sum_partial += es0 + es1;+ a_packed[j / 2] = pack_bf16x2(es0, es1);+ }+ #pragma unroll+ for (int j = 0; j < 4; j += 2) {+ float es0 = __builtin_amdgcn_exp2f(fmaf(acc_hi[j], my_scale, neg_max));+ float es1 = __builtin_amdgcn_exp2f(fmaf(acc_hi[j+1], my_scale, neg_max));+ tile_sum_partial += es0 + es1;+ a_packed[2 + j / 2] = pack_bf16x2(es0, es1);+ }++ tile_sum_partial += __shfl_xor(tile_sum_partial, 16);+ tile_sum_partial += __shfl_xor(tile_sum_partial, 32);+ my_head_sum += tile_sum_partial;++ // === V accumulation (IDENTICAL) ===+ union { uint32_t u[4]; mfma_bf16_input_t v; } a_cvt;+ a_cvt.u[0] = a_packed[0]; a_cvt.u[1] = a_packed[1];+ a_cvt.u[2] = a_packed[2]; a_cvt.u[3] = a_packed[3];+ mfma_bf16_input_t a_bf16 = a_cvt.v;++ {+ if constexpr (!PRELOAD_V) {+ #pragma unroll+ for (int i = 0; i < 4; i++) {+ raw4_cur[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(blo + i) * LDS_KV_STRIDE + v_off0]);+ raw4_cur[i+4] = *reinterpret_cast<const uint32_t*>(&kv_lds[(bhi + i) * LDS_KV_STRIDE + v_off0]);+ raw4_next_v[i] = *reinterpret_cast<const uint32_t*>(&kv_lds[(blo + i) * LDS_KV_STRIDE + v_off1]);+ raw4_next_v[i+4] = *reinterpret_cast<const uint32_t*>(&kv_lds[(bhi + i) * LDS_KV_STRIDE + v_off1]);+ }+ }++ if constexpr (USE_SETPRIO) asm volatile("s_setprio 3" :::);+ uint32_t bp0[4], bp1[4], bp2[4], bp3[4];+ v_convert_bf16_dispatch<DIRECT_CVT>(raw4_cur, bp0, bp1, bp2, bp3);+ v_acc[0] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp0), v_acc[0]);+ v_acc[1] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp1), v_acc[1]);+ v_acc[2] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp2), v_acc[2]);+ v_acc[3] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp3), v_acc[3]);++ uint32_t bp0b[4], bp1b[4], bp2b[4], bp3b[4];+ v_convert_bf16_dispatch<DIRECT_CVT>(raw4_next_v, bp0b, bp1b, bp2b, bp3b);+ v_acc[4] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp0b), v_acc[4]);+ v_acc[5] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp1b), v_acc[5]);+ v_acc[6] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp2b), v_acc[6]);+ v_acc[7] = mfma_f32_16x16x32_bf16(a_bf16, assemble_bf16_input(bp3b), v_acc[7]);+ if constexpr (USE_SETPRIO) asm volatile("s_setprio 0" :::);+ }++ if (__builtin_expect(has_next, 1)) wait_kv_loads();+ __syncthreads();+ buf ^= 1;+ cur_page++;+ }++ // === Output write (IDENTICAL to non-paged) ===+ if (num_kv_splits == 1) {+ #pragma unroll+ for (int k = 0; k < 4; k++) {+ int abs_h = h_start + lane_group * 4 + k;+ float hs = __shfl(my_head_sum, lane_group * 4 + k);+ float final_scale = (hs > 0.0f) ? kv_scale / hs : 0.0f;+ __hip_bfloat16* o_ptr = output + (size_t)q_pos * NHEAD * V_DIM + abs_h * V_DIM;+ #pragma unroll+ for (int g = 0; g < 2; g++) {+ int base_idx = wave_id * 128 + g * 64 + lane_col * 4;+ uint32_t pk0 = pack_bf16x2(v_acc[g*4][k] * final_scale, v_acc[g*4+1][k] * final_scale);+ uint32_t pk1 = pack_bf16x2(v_acc[g*4+2][k] * final_scale, v_acc[g*4+3][k] * final_scale);+ *reinterpret_cast<uint64_t*>(&o_ptr[base_idx]) = ((uint64_t)pk1 << 32) | (uint64_t)pk0;+ }+ }+ } else {+ __hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM+ + split_id * V_DIM;+ float* ml = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits+ + split_id;+ #pragma unroll+ for (int k = 0; k < 4; k++) {+ int abs_h = h_start + lane_group * 4 + k;+ float hs = __shfl(my_head_sum, lane_group * 4 + k);+ float final_scale = (hs > 0.0f) ? kv_scale / hs : 0.0f;+ #pragma unroll+ for (int g = 0; g < 2; g++) {+ int base_idx = wave_id * 128 + g * 64 + lane_col * 4;+ uint32_t pk0 = pack_bf16x2(v_acc[g*4][k] * final_scale, v_acc[g*4+1][k] * final_scale);+ uint32_t pk1 = pack_bf16x2(v_acc[g*4+2][k] * final_scale, v_acc[g*4+3][k] * final_scale);+ *reinterpret_cast<uint64_t*>(&mo[abs_h * num_kv_splits * V_DIM + base_idx]) = ((uint64_t)pk1 << 32) | (uint64_t)pk0;+ }+ }+ if (wave_id == 0 && lane_group == 0) {+ int abs_h = h_start + lane_col;+ ml[abs_h * num_kv_splits] = (my_head_sum > 0.0f) ?+ my_head_max * LN2_F + __logf(my_head_sum) : -1e30f;+ }+ }+ __syncthreads();+ }+ }++ // ============================================================================+ // Stage 1 — V2: occ=1, 5-buffer, FP8 K=128 V accumulation+ // ============================================================================+ // Key differences from bf16 path:+ // - __launch_bounds__(256, 1): 512 VGPRs, 160KB LDS+ // - 5-buffer LDS (4 retained + 1 prefetch) for 128-token megatile V MFMA+ // - FP8 mfma_f32_16x16x128 for V accumulation (4x throughput vs BF16 K=32)+ // - Attention weights accumulated across 4 tiles, converted to FP8 for V MFMA++ template <int NHEAD>+ __global__ void+ __launch_bounds__(256, 1)+ mla_decode_stage1_v2(+ const __hip_bfloat16* __restrict__ q_bf16,+ const uint8_t* __restrict__ kv_fp8,+ const float* __restrict__ kv_scale_ptr,+ __hip_bfloat16* __restrict__ mid_o,+ float* __restrict__ mid_lse,+ __hip_bfloat16* __restrict__ output,+ const int32_t* __restrict__ batch_map,+ const int32_t* __restrict__ kv_indptr,+ int batch_size,+ int num_kv_splits+ ) {+ constexpr float sm_scale = 1.0f / 24.0f;+ const float kv_scale = *kv_scale_ptr;+ constexpr int NUM_HEAD_GROUPS = (NHEAD + HEAD_GROUP - 1) / HEAD_GROUP;++ const int q_pos = blockIdx.x;+ const int split_id = blockIdx.y;+ const int tid = threadIdx.x;+ const int wave_id = tid / WAVESIZE;+ const int lane_id = tid % WAVESIZE;+ const int lane_group = lane_id / 16;+ const int lane_col = lane_id % 16;++ const int batch_id = batch_map[q_pos];+ const int kv_start = kv_indptr[batch_id];+ const int kv_end = kv_indptr[batch_id + 1];+ const int total_kv_len = kv_end - kv_start;++ int positions_per_split = (total_kv_len + num_kv_splits - 1) / num_kv_splits;+ positions_per_split = (positions_per_split + KV_TILE - 1) / KV_TILE * KV_TILE;+ int split_start = split_id * positions_per_split;+ int split_end = min(split_start + positions_per_split, total_kv_len);++ __shared__ LDS_V2 lds;++ for (int hg = 0; hg < NUM_HEAD_GROUPS; hg++) {+ int h_start = hg * HEAD_GROUP;+ int h_count = min(HEAD_GROUP, NHEAD - h_start);++ float my_head_max = -1e30f;+ float my_head_sum = 0.0f;++ mfma_acc_t v_acc[V_MFMAS_PER_WAVE];+ #pragma unroll+ for (int m = 0; m < V_MFMAS_PER_WAVE; m++)+ v_acc[m] = {0.0f, 0.0f, 0.0f, 0.0f};++ if (split_start >= split_end) {+ // Empty split: write zeros (v2 layout)+ if (num_kv_splits > 1) {+ __hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM + split_id * V_DIM;+ float* ml = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits + split_id;+ __hip_bfloat16 zero_bf16 = __float2bfloat16(0.0f);+ for (int n = 0; n < 4; n++) {+ int head = lane_group * 4 + n;+ if (head < h_count) {+ int abs_h = h_start + head;+ for (int vm = 0; vm < V_MFMAS_PER_WAVE; vm++) {+ int v_dim = wave_id * 128 + vm * 16 + lane_col;+ if (v_dim < V_DIM) mo[abs_h * num_kv_splits * V_DIM + v_dim] = zero_bf16;+ }+ }+ }+ if (wave_id == 0 && lane_group == 0 && lane_col < h_count)+ ml[(h_start + lane_col) * num_kv_splits] = -1e30f;+ }+ __syncthreads();+ continue;+ }++ // Inline Q BF16→FP8 conversion+ const __hip_bfloat16* q_head_base = q_bf16 + (size_t)(q_pos * NHEAD + h_start) * QK_DIM;+ float my_scale = 0.0f;+ if (lane_col < h_count) my_scale = kv_scale * sm_scale * LOG2E_F;++ uint32_t q_cache[QK_MFMAS][8];+ #pragma unroll+ for (int t = 0; t < QK_MFMAS; t++) {+ int elem_start = t * FP8_MFMA_K + lane_group * 32;+ if (elem_start + 32 <= QK_DIM && lane_col < h_count) {+ const uint32_t* wp = reinterpret_cast<const uint32_t*>(+ q_head_base + (size_t)lane_col * QK_DIM + elem_start);+ #pragma unroll+ for (int j = 0; j < 8; j++) {+ uint32_t w0 = wp[j * 2];+ uint32_t w1 = wp[j * 2 + 1];+ uint32_t pk;+ float s1 = 1.0f;+ asm("v_cvt_scalef32_pk_fp8_bf16 %0, %1, %2"+ : "=v"(pk) : "v"(w0), "v"(s1));+ asm("v_cvt_scalef32_pk_fp8_bf16 %0, %1, %2 op_sel:[0,0,1]"+ : "+v"(pk) : "v"(w1), "v"(s1));+ q_cache[t][j] = pk;+ }+ } else {+ #pragma unroll+ for (int j = 0; j < 8; j++) q_cache[t][j] = 0;+ }+ }++ // Prefetch first tile into buffer 0+ int first_count = min(KV_TILE, split_end - split_start);+ issue_kv_loads(0, kv_fp8 + (size_t)(kv_start + split_start) * QK_DIM, first_count, tid);++ wait_kv_loads();+ __syncthreads();++ // Tile loop with 4-tile megatile grouping+ int buf_idx = 0;+ int tile_in_mega = 0;+ // Attention weights accumulated across megatile: [tile][8] per lane+ float attn_w[4][8];++ for (int tile_start = split_start; tile_start < split_end; tile_start += KV_TILE) {+ int tile_count = min(KV_TILE, split_end - tile_start);+ int next_tile_start = tile_start + KV_TILE;+ int has_next = (next_tile_start < split_end);++ // Prefetch next tile+ int next_buf = (buf_idx + 1) % V2_NUM_BUFS;+ if (has_next) {+ int next_count = min(KV_TILE, split_end - next_tile_start);+ issue_kv_loads(next_buf * LDS_BUF_SIZE,+ kv_fp8 + (size_t)(kv_start + next_tile_start) * QK_DIM,+ next_count, tid);+ }++ const uint8_t* kv_lds = lds.kv_fp8[buf_idx];++ // === QK scoring (identical to bf16 path) ===+ mfma_acc_t acc_lo = {0.0f, 0.0f, 0.0f, 0.0f};+ mfma_acc_t acc_hi = {0.0f, 0.0f, 0.0f, 0.0f};++ #pragma unroll+ for (int t = 0; t < QK_MFMAS; t++) {+ mfma_input_t frag_q = {q_cache[t][0], q_cache[t][1], q_cache[t][2], q_cache[t][3],+ q_cache[t][4], q_cache[t][5], q_cache[t][6], q_cache[t][7]};+ int kv_byte = t * FP8_MFMA_K + lane_group * 32;++ mfma_input_t frag_kv_lo, frag_kv_hi;+ if (kv_byte + 32 <= QK_DIM) {+ load_frag_b_wide(kv_lds, lane_col * LDS_KV_STRIDE + kv_byte, frag_kv_lo);+ load_frag_b_wide(kv_lds, (lane_col + KV_HALF) * LDS_KV_STRIDE + kv_byte, frag_kv_hi);+ } else {+ #pragma unroll+ for (int j = 0; j < 8; j++) { frag_kv_lo[j] = 0; frag_kv_hi[j] = 0; }+ }+ acc_lo = mfma_f32_16x16x128_fp8(frag_kv_lo, frag_q, acc_lo);+ acc_hi = mfma_f32_16x16x128_fp8(frag_kv_hi, frag_q, acc_hi);+ }++ // === Softmax: accumulate into megatile weights ===+ float partial_max = -1e30f;+ partial_max = fmaxf(partial_max, fmaxf(acc_lo[0], acc_lo[1]));+ partial_max = fmaxf(partial_max, fmaxf(acc_lo[2], acc_lo[3]));+ partial_max = fmaxf(partial_max, fmaxf(acc_hi[0], acc_hi[1]));+ partial_max = fmaxf(partial_max, fmaxf(acc_hi[2], acc_hi[3]));+ partial_max = fmaxf(partial_max, __shfl_xor(partial_max, 16));+ partial_max = fmaxf(partial_max, __shfl_xor(partial_max, 32));++ float scaled_partial_max = partial_max * my_scale;+ float old_max = my_head_max;+ float new_max = fmaxf(old_max, scaled_partial_max);+ my_head_max = new_max;+ float neg_max = -new_max;++ if (scaled_partial_max > old_max && old_max > -1e29f) {+ float my_rescale = __builtin_amdgcn_exp2f(old_max - new_max);+ my_head_sum *= my_rescale;+ // Rescale previously accumulated megatile weights+ for (int pt = 0; pt < tile_in_mega; pt++)+ for (int j = 0; j < 8; j++)+ attn_w[pt][j] *= my_rescale;+ // Rescale V accumulators from previous megatiles+ #pragma unroll+ for (int k = 0; k < 4; k++) {+ float r = __shfl(my_rescale, lane_group * 4 + k);+ #pragma unroll+ for (int m = 0; m < V_MFMAS_PER_WAVE; m++)+ v_acc[m][k] *= r;+ }+ }++ // Compute exp weights and store in attn_w+ float tile_sum_partial = 0.0f;+ #pragma unroll+ for (int j = 0; j < 4; j++) {+ float es = __builtin_amdgcn_exp2f(fmaf(acc_lo[j], my_scale, neg_max));+ tile_sum_partial += es;+ attn_w[tile_in_mega][j] = es;+ }+ #pragma unroll+ for (int j = 0; j < 4; j++) {+ float es = __builtin_amdgcn_exp2f(fmaf(acc_hi[j], my_scale, neg_max));+ tile_sum_partial += es;+ attn_w[tile_in_mega][4 + j] = es;+ }+ tile_sum_partial += __shfl_xor(tile_sum_partial, 16);+ tile_sum_partial += __shfl_xor(tile_sum_partial, 32);+ my_head_sum += tile_sum_partial;++ tile_in_mega++;++ // === V MFMA every 4 tiles (or at end of split) ===+ int is_last_tile = !has_next;+ if (tile_in_mega == V2_MEGATILE || is_last_tile) {+ // Zero out unused megatile slots+ for (int pt = tile_in_mega; pt < V2_MEGATILE; pt++)+ for (int j = 0; j < 8; j++)+ attn_w[pt][j] = 0.0f;++ // === Interleaved tile V MFMA: no shuffle needed ===+ // Key insight: interleave tiles within K=128 so each lane_group's+ // native attn_w values align with the MFMA operand layout.+ //+ // A operand layout (16 heads × 128 tokens):+ // g%4 = tile index, g/4 = half (0=tokens 0..15, 1=tokens 16..31)+ // frag_a[g] = FP8(attn_w[g%4][(g/4)*4 + 0..3])+ // Each lane_group already has its 4 tokens' weights — no cross-lane shuffle.+ //+ // B operand layout (128 tokens × 16 V dims):+ // Same interleaving: frag_b[g] reads from tile g%4's LDS buffer,+ // tokens (g/4)*16 + lane_group*4 + 0..3++ // Compute uniform FP8 scale across all lane_groups (per head)+ float amax = 0.0f;+ #pragma unroll+ for (int tile = 0; tile < V2_MEGATILE; tile++)+ #pragma unroll+ for (int j = 0; j < 8; j++)+ amax = fmaxf(amax, attn_w[tile][j]);+ // Reduce across lane_groups so all use the same scale per head+ amax = fmaxf(amax, __shfl_xor(amax, 16));+ amax = fmaxf(amax, __shfl_xor(amax, 32));++ float att_scale = (amax > 1e-12f) ? 224.0f / amax : 0.0f;+ float inv_att_scale = (att_scale > 0.0f) ? 1.0f / att_scale : 0.0f;++ // Pack A operand: interleaved tile order+ mfma_input_t frag_a;+ #pragma unroll+ for (int g = 0; g < 8; g++) {+ int tile = g & 3; // g % 4+ int half = g >> 2; // g / 4+ float v0 = attn_w[tile][half*4 + 0] * att_scale;+ float v1 = attn_w[tile][half*4 + 1] * att_scale;+ float v2 = attn_w[tile][half*4 + 2] * att_scale;+ float v3 = attn_w[tile][half*4 + 3] * att_scale;+ uint32_t pk;+ pk = __builtin_amdgcn_cvt_pk_fp8_f32(v0, v1, 0u, false);+ pk = __builtin_amdgcn_cvt_pk_fp8_f32(v2, v3, pk, true);+ frag_a[g] = pk;+ }++ // V MFMA: interleaved B operand from 4 LDS buffers+ #pragma unroll+ for (int vm = 0; vm < V_MFMAS_PER_WAVE; vm++) {+ int v_base = wave_id * 128 + vm * 16;+ if (v_base >= V_DIM) break;+ int v_col = v_base + lane_col;++ mfma_input_t frag_b;+ if (v_col < V_DIM) {+ #pragma unroll+ for (int g = 0; g < 8; g++) {+ int tile = g & 3;+ int half = g >> 2;+ int v2_buf = (buf_idx - tile_in_mega + 1 + tile + V2_NUM_BUFS) % V2_NUM_BUFS;+ const uint8_t* v_lds = lds.kv_fp8[v2_buf];+ int tok_base = half * 16 + lane_group * 4;+ uint32_t b = (uint32_t)v_lds[(tok_base ) * LDS_KV_STRIDE + v_col]+ | ((uint32_t)v_lds[(tok_base+1) * LDS_KV_STRIDE + v_col] << 8)+ | ((uint32_t)v_lds[(tok_base+2) * LDS_KV_STRIDE + v_col] << 16)+ | ((uint32_t)v_lds[(tok_base+3) * LDS_KV_STRIDE + v_col] << 24);+ frag_b[g] = b;+ }+ } else {+ #pragma unroll+ for (int g = 0; g < 8; g++) frag_b[g] = 0;+ }++ // V MFMA into zeroed temp (unscaled: scale=0 → compiler emits v_mfma_f32_16x16x128_f8f6f4)+ mfma_acc_t v_delta = {0.0f, 0.0f, 0.0f, 0.0f};+ v_delta = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(+ frag_a, frag_b, v_delta, 0, 0, 0, 0, 0, 0);++ // Correct per-head scale and accumulate+ // v_delta[n] = result for head (lane_group*4+n), need that head's inv_att_scale+ #pragma unroll+ for (int n = 0; n < 4; n++) {+ float inv_s = __shfl(inv_att_scale, lane_group * 4 + n);+ v_acc[vm][n] += v_delta[n] * inv_s;+ }+ }++ tile_in_mega = 0;+ }++ if (has_next) wait_kv_loads();+ __syncthreads();+ buf_idx = next_buf;+ }++ // === Output write (v2 layout: v_acc[vm][n] = head lg*4+n, V dim wave*128+vm*16+lc) ===+ if (num_kv_splits == 1) {+ #pragma unroll+ for (int n = 0; n < 4; n++) {+ int abs_h = h_start + lane_group * 4 + n;+ float hs = __shfl(my_head_sum, lane_group * 4 + n);+ float final_scale = (hs > 0.0f) ? kv_scale / hs : 0.0f;+ __hip_bfloat16* o_ptr = output + (size_t)q_pos * NHEAD * V_DIM + abs_h * V_DIM;+ #pragma unroll+ for (int vm = 0; vm < V_MFMAS_PER_WAVE; vm++) {+ int v_dim = wave_id * 128 + vm * 16 + lane_col;+ if (v_dim < V_DIM)+ o_ptr[v_dim] = __float2bfloat16(v_acc[vm][n] * final_scale);+ }+ }+ } else {+ __hip_bfloat16* mo = mid_o + (size_t)q_pos * NHEAD * num_kv_splits * V_DIM + split_id * V_DIM;+ float* ml = mid_lse + (size_t)q_pos * NHEAD * num_kv_splits + split_id;+ #pragma unroll+ for (int n = 0; n < 4; n++) {+ int abs_h = h_start + lane_group * 4 + n;+ float hs = __shfl(my_head_sum, lane_group * 4 + n);+ float final_scale = (hs > 0.0f) ? kv_scale / hs : 0.0f;+ #pragma unroll+ for (int vm = 0; vm < V_MFMAS_PER_WAVE; vm++) {+ int v_dim = wave_id * 128 + vm * 16 + lane_col;+ if (v_dim < V_DIM)+ mo[abs_h * num_kv_splits * V_DIM + v_dim] =+ __float2bfloat16(v_acc[vm][n] * final_scale);+ }+ }+ if (wave_id == 0 && lane_group == 0) {+ int abs_h = h_start + lane_col;+ ml[abs_h * num_kv_splits] = (my_head_sum > 0.0f) ?+ my_head_max * LN2_F + __logf(my_head_sum) : -1e30f;+ }+ }+ __syncthreads();+ }+ }++ // ============================================================================// Stage 1 — MXFP4 path (native FP4 MFMA + hardware V conversion)// ============================================================================// Uses FP8(Q) × FP4(KV) MFMA with software E8M0 scaling for QK.⋯ 785 unchanged linesint use_setprio; // 1 = use s_setprio hints (default), 0 = disableconst uint8_t* cached_kf;const float* cached_ksp;+ // Paged attention fields+ int use_paged; // 1 = use paged kernel+ int32_t* kv_indices;+ int32_t* paged_kv_indptr;+ int32_t* kv_last_page_len;+ int page_size;};static std::unordered_map<int64_t, DispatchState> g_states;⋯ 28 unchanged liness.use_setprio = static_cast<int>(use_setprio);s.cached_kf = nullptr;s.cached_ksp = nullptr;+ s.use_paged = 0;+ s.kv_indices = nullptr;+ s.paged_kv_indptr = nullptr;+ s.kv_last_page_len = nullptr;+ s.page_size = 0;g_states[key] = s;}+ void pb_register_state_paged(+ int64_t key,+ torch::Tensor mid_o, torch::Tensor mid_lse,+ torch::Tensor output, torch::Tensor batch_map,+ torch::Tensor kv_indices_t, torch::Tensor paged_kv_indptr_t,+ torch::Tensor kv_last_page_len_t,+ int64_t total_q, int64_t nhead, int64_t batch_size, int64_t num_kv_splits,+ int64_t page_size+ ) {+ DispatchState s;+ s.q_fp8 = nullptr;+ s.q_scales = nullptr;+ s.mid_o = (num_kv_splits > 1) ? reinterpret_cast<__hip_bfloat16*>(mid_o.data_ptr()) : nullptr;+ s.mid_lse = (num_kv_splits > 1) ? reinterpret_cast<float*>(mid_lse.data_ptr()) : nullptr;+ s.output = reinterpret_cast<__hip_bfloat16*>(output.data_ptr());+ s.batch_map = reinterpret_cast<int32_t*>(batch_map.data_ptr());+ s.kv_indptr = nullptr;+ s.total_q = static_cast<int>(total_q);+ s.nhead = static_cast<int>(nhead);+ s.batch_size = static_cast<int>(batch_size);+ s.num_kv_splits = static_cast<int>(num_kv_splits);+ s.q_rows = static_cast<int>(total_q * nhead);+ s.cached_q_ptr = 0;+ s.use_occ3 = 1;+ s.use_bf16_fused = 1;+ s.use_direct_cvt = 1;+ s.use_preload_v = 1;+ s.use_setprio = 0;+ s.cached_kf = nullptr;+ s.cached_ksp = nullptr;+ s.use_paged = 1;+ s.kv_indices = reinterpret_cast<int32_t*>(kv_indices_t.data_ptr());+ s.paged_kv_indptr = reinterpret_cast<int32_t*>(paged_kv_indptr_t.data_ptr());+ s.kv_last_page_len = reinterpret_cast<int32_t*>(kv_last_page_len_t.data_ptr());+ s.page_size = static_cast<int>(page_size);+ g_states[key] = s;+ }+// Launch BF16 fused stage1 (inline Q conversion, skip bf16_to_fp8)// CONFIG packs per-shape template bools: bit0=DIRECT_CVT, bit1=PRELOAD_V, bit2=USE_SETPRIOtemplate <int NH, int CONFIG>⋯ 47 unchanged lineselse launch_kernels_bf16_impl<16>(s, qb, kf, ksp, stm);}+ // V2 kernel launch: occ=1, FP8 V accumulation via K=128 megatile+ template <int NH>+ static void launch_kernels_v2_impl(DispatchState& s, const __hip_bfloat16* qb,+ const uint8_t* kf, const float* ksp, hipSt_QQ__t stm) {+ dim3 block(256);+ dim3 grid(s.total_q, s.num_kv_splits);++ if (s.num_kv_splits == 1) {+ hipLaunchKernelGGL((mla_decode_stage1_v2<NH>), grid, block, 0, stm,+ qb, kf, ksp, nullptr, nullptr, s.output,+ s.batch_map, s.kv_indptr, s.batch_size, 1);+ } else {+ hipLaunchKernelGGL((mla_decode_stage1_v2<NH>), grid, block, 0, stm,+ qb, kf, ksp, s.mid_o, s.mid_lse, s.output,+ s.batch_map, s.kv_indptr, s.batch_size, s.num_kv_splits);+ }+ if (s.num_kv_splits > 1) {+ dim3 g2(s.total_q, NH);+ dim3 gb(s.total_q);+ switch (s.num_kv_splits) {+ case 2: hipLaunchKernelGGL((mla_decode_stage2_batch<NH,2>), gb, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;+ case 3: hipLaunchKernelGGL((mla_decode_stage2_batch<NH,3>), gb, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;+ case 4: hipLaunchKernelGGL((mla_decode_stage2_batch<NH,4>), gb, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;+ case 6: hipLaunchKernelGGL((mla_decode_stage2_t<NH,6>), g2, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;+ case 8: hipLaunchKernelGGL((mla_decode_stage2_t<NH,8>), g2, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;+ case 12: hipLaunchKernelGGL((mla_decode_stage2_t<NH,12>), g2, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;+ case 16: hipLaunchKernelGGL((mla_decode_stage2_t<NH,16>), g2, block, 0, stm, s.mid_o, s.mid_lse, s.output); break;+ default: hipLaunchKernelGGL((mla_decode_stage2<NH>), g2, block, 0, stm, s.mid_o, s.mid_lse, s.output, s.num_kv_splits); break;+ }+ }+ }++ static void launch_kernels_v2(DispatchState& s, const __hip_bfloat16* qb,+ const uint8_t* kf, const float* ksp, hipSt_QQ__t stm) {+ if (s.nhead == 32) launch_kernels_v2_impl<32>(s, qb, kf, ksp, stm);+ else launch_kernels_v2_impl<16>(s, qb, kf, ksp, stm);+ }+void pb_fast_dispatch(int64_t key,torch::Tensor q_bf16,⋯ 12 unchanged lineslaunch_kernels_bf16(s, qb, s.cached_kf, s.cached_ksp, stm);}+ void pb_fast_dispatch_v2(+ int64_t key,+ torch::Tensor q_bf16,+ torch::Tensor kv_fp8,+ torch::Tensor kv_scale+ ) {+ auto& s = g_states[key];+ auto qb = reinterpret_cast<const __hip_bfloat16*>(q_bf16.data_ptr());++ hipSt_QQ__t stm = c10::cuda::getCurrentCUDASt_QQ_().st_QQ_();++ s.cached_kf = reinterpret_cast<const uint8_t*>(kv_fp8.data_ptr());+ s.cached_ksp = reinterpret_cast<const float*>(kv_scale.data_ptr());++ launch_kernels_v2(s, qb, s.cached_kf, s.cached_ksp, stm);+ }+void pb_dispatch_cached(int64_t) {}PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {⋯ 5 unchanged linesm.def("compute_q_scale", &pb_compute_q_scale);m.def("register_state", &pb_register_state);m.def("fast_dispatch", &pb_fast_dispatch);+ m.def("fast_dispatch_v2", &pb_fast_dispatch_v2);m.def("dispatch_cached", &pb_dispatch_cached);}""".replace('_QQ_', 'ream')⋯ 18 unchanged lines(64, 1, 1024, 16): 8,(64, 1, 8192, 16): 8,(256, 1, 1024, 16): 1,- (256, 1, 8192, 16): 2,+ (256, 1, 8192, 16): 1, # c7: splits=1 eliminates stage2 (66µs vs 75µs)(128, 1, 8192, 16): 6, # Case 6: 768 WGs = 3/CU (was 512 = 2/CU)(128, 4, 8192, 16): 2, # Case 7: fewer splits for qs=4 (less stage2 overhead)# tp=4 (32 heads)⋯ 9 unchanged lines# Per-case kernel selection: 0 = occ4, 1 = occ3 double-buffer, 2 = occ2 triple-buffer_USE_OCC3 = {- (4, 1, 1024, 16): 1, # 128 WGs = 0.5/CU+ (4, 1, 1024, 16): 2, # c0: occ2 triple-buf (13.8→13.5µs)(4, 1, 8192, 16): 1, # 256 WGs = 1/CU(32, 1, 1024, 16): 1, # 256 WGs = 1/CU(32, 1, 8192, 16): 1, # 512 WGs = 2/CU (splits=16), occ3 (occ4 produced NaN)- (64, 1, 1024, 16): 1, # 512 WGs = 2/CU+ (64, 1, 1024, 16): 2, # c4: occ2 triple-buf (28.3→27.6µs)(64, 1, 8192, 16): 1, # 512 WGs = 2/CU(256, 1, 1024, 16): 1, # 256 WGs = 1/CU- (256, 1, 8192, 16): 1, # 768 WGs = 3/CU+ (256, 1, 8192, 16): 2, # c7: occ2 triple-buf (67.6→66.5µs)(32, 4, 1024, 16): 0, # 1024 WGs (qs=4), occ4 for full 4/CU utilization(32, 4, 8192, 16): 0, # 1024 WGs (qs=4), occ4 for full 4/CU utilization(128, 1, 8192, 16): 1, # 512 WGs (qs=1), occ3 for packed V conversion⋯ 39 unchanged linesreturn _extfrom torch.utils.cpp_extension import load_inline_ext = load_inline(- name="mla_decode_v846",+ name="mla_decode_v994",cpp_sources=CPP_SOURCE,cuda_sources=HIP_SOURCE,extra_cuda_cflags=[⋯ 1 unchanged lines"-D__gfx950__","-mllvm", "-amdgpu-early-inline-all=true","-mllvm", "-amdgpu-function-calls=false",- "-mllvm", "-amdgpu-early-ifcvt=true","-mllvm", "-amdgpu-max-memory-clause=16","-mllvm", "-amdgpu-promote-alloca-to-vector-limit=128",],⋯ 43 unchanged lines_cache[key] = (o, q_fp8, q_scales, mid_o, mid_lse, batch_map, kv_indptr_t)_registered.add(key)+ o = _cache[key][0]kv = data[1]["fp8"]ext.fast_dispatch(key, data[0], kv[0], kv[1])- return _cache[key][0]+ return o+ def _custom_kernel_hip_paged(data: input_t) -> output_t:+ """Our custom HIP kernel with paged attention — for large batch × kv shapes."""+ ext = _ext+ cfg = data[4]+ bs = cfg["batch_size"]+ nh = cfg["num_heads"]+ kvsl = cfg["kv_seq_len"]+ qs = cfg["q_seq_len"]+ key = bs * 10000000 + qs * 1000000 + kvsl * 100 + nh + 50 # +50 to avoid collision with non-paged keys++ if key not in _paged_registered:+ total_q = bs * qs+ ps = _PAGED_PAGE_SIZE+ num_splits = _PAGED_SPLITS.get((bs, qs, kvsl, nh), 4)+ dev = data[0].device++ pages_per_seq = (kvsl + ps - 1) // ps+ ebs = bs # effective batch size (qs=1)++ o = torch.empty((total_q, nh, V_DIM), dtype=torch.bfloat16, device=dev)+ dummy = torch.empty(1, dtype=torch.float32, device=dev)+ mid_o = dummy+ mid_lse = dummy+ if num_splits > 1:+ mid_o = torch.empty(total_q * num_splits * nh * V_DIM, dtype=torch.bfloat16, device=dev)+ mid_lse = torch.empty(total_q * num_splits * nh, dtype=torch.float32, device=dev)+ batch_map = torch.arange(bs, dtype=torch.int32, device=dev).repeat_interleave(qs)++ # Page table setup (contiguous KV → identity page mapping)+ kv_indices = torch.arange(ebs * pages_per_seq, dtype=torch.int32, device=dev)+ paged_kv_indptr = torch.arange(ebs + 1, dtype=torch.int32, device=dev) * pages_per_seq+ kv_last_page_len = torch.full((ebs,), kvsl % ps if kvsl % ps != 0 else ps, dtype=torch.int32, device=dev)++ ext.register_state_paged(key, mid_o, mid_lse, o, batch_map,+ kv_indices, paged_kv_indptr, kv_last_page_len,+ total_q, nh, bs, num_splits, ps)+ _paged_cache[key] = (o, kv_indices, paged_kv_indptr, kv_last_page_len)+ _paged_registered.add(key)++ kv = data[1]["fp8"]+ kv_fp8, kv_scale = kv[0], kv[1]+ # kv_fp8 shape: [bs*kvsl, 1, QK_DIM] → reshape to pages: [bs*pages_per_seq, page_size, QK_DIM]+ ps = _PAGED_PAGE_SIZE+ pages_per_seq = kvsl // ps # kvsl is always divisible by 32+ o = _paged_cache[key][0]+ kv_paged = kv_fp8.view(bs * pages_per_seq, ps, QK_DIM)+ ext.fast_dispatch(key, data[0], kv_paged, kv_scale)+ return o+++ # V2 kernel split tuning: occ=1, FP8 K=128 V MFMA+ _V2_SPLITS = {+ (32, 1, 8192, 16): 8, # 32*8=256 WGs+ (64, 1, 8192, 16): 4, # 64*4=256 WGs+ (256, 1, 1024, 16): 1, # 256*1=256 WGs+ (256, 1, 8192, 16): 1, # 256*1=256 WGs+ }++ _v2_cache = {}+ _v2_registered = set()+++ def _custom_kernel_hip_v2(data: input_t) -> output_t:+ """V2 HIP kernel: occ=1, FP8 V accumulation via K=128 megatile."""+ ext = _ext+ cfg = data[4]+ bs = cfg["batch_size"]+ nh = cfg["num_heads"]+ kvsl = cfg["kv_seq_len"]+ qs = cfg["q_seq_len"]+ key = bs * 10000000 + qs * 1000000 + kvsl * 100 + nh + 70 # +70 to avoid collision++ if key not in _v2_registered:+ total_q = bs * qs+ num_splits = _V2_SPLITS.get((bs, qs, kvsl, nh), 4)+ dev = data[0].device++ o = torch.empty((total_q, nh, V_DIM), dtype=torch.bfloat16, device=dev)+ q_fp8 = torch.empty(total_q * nh, QK_DIM, dtype=torch.uint8, device=dev)+ q_scales = torch.empty(total_q * nh, dtype=torch.float32, device=dev)+ dummy = torch.empty(1, dtype=torch.float32, device=dev)+ mid_o = dummy+ mid_lse = dummy+ if num_splits > 1:+ mid_o = torch.empty(total_q * num_splits * nh * V_DIM, dtype=torch.bfloat16, device=dev)+ mid_lse = torch.empty(total_q * num_splits * nh, dtype=torch.float32, device=dev)+ batch_map = torch.arange(bs, dtype=torch.int32, device=dev).repeat_interleave(qs)+ kv_indptr_t = torch.arange(bs + 1, dtype=torch.int32, device=dev) * kvsl++ # Reuse the same register_state — v2 uses same DispatchState layout+ ext.register_state(key, q_fp8, q_scales, mid_o, mid_lse, o,+ batch_map, kv_indptr_t, total_q, nh, bs, num_splits, 0, 1, 0, 1, 0)+ _v2_cache[key] = (o, q_fp8, q_scales, mid_o, mid_lse, batch_map, kv_indptr_t)+ _v2_registered.add(key)++ o = _v2_cache[key][0]+ kv = data[1]["fp8"]+ ext.fast_dispatch_v2(key, data[0], kv[0], kv[1])+ return o++# ============================================================================# AITER path — for large batch × kv shapes where paged attention wins# ============================================================================⋯ 6 unchanged lines# Per-case AITER tuning: (bs, kvsl) -> (page_size, num_kv_splits, fast_mode)_AITER_TUNE = {- (4, 1024): (1, 32, False), # harmonized: all ns=32 fm=False- (4, 8192): (1, 32, False),- (32, 1024): (1, 32, False),- (32, 8192): (8, 32, False), # sweep best: 26.3µs- (64, 1024): (2, 32, False), # harmonized (was ns=1 fm=True)- (64, 8192): (8, 32, False), # harmonized (was ns=2)- (128, 1024): (2, 32, False),- (128, 8192): (8, 32, False),- (256, 1024): (2, 32, False), # harmonized- (256, 8192): (8, 32, False), # sweep: 41.3µs standalone+ (4, 1024): (1, 32, False), # c0: HIP handles this+ (4, 8192): (1, 32, False), # c1: HIP handles this+ (32, 1024): (1, 32, False), # c2: HIP handles this+ (32, 8192): (8, 16, False), # c3: ps=8 ns=16 + kv_gran=32+ (64, 1024): (2, 1, False), # c4: AITER ps=2 beats HIP+ (64, 8192): (8, 2, False), # c5+ (128, 1024): (2, 8, False),+ (128, 8192): (8, 8, False),+ (256, 1024): (2, 1, True), # c6+ (256, 8192): (8, 32, True), # c7}_aiter_cache = {}⋯ 8 unchanged lines_aiter_q_scale = torch.ones(1, dtype=torch.float32, device=dev)ps, ns, fm = _AITER_TUNE.get((bs, kvsl), (1, 32, bs <= 4))- kv_gran = max(ps, 16)+ kv_gran = 32 if ps == 8 else max(ps, 16)ebs = bs * qsleff_qo_indptr = torch.arange(ebs + 1, dtype=torch.int32, device=dev)⋯ 21 unchanged linesdtype_q=_FP8_DTYPE, dtype_kv=_FP8_DTYPE)e = {- 'ps': ps, 'ns': ns, 'fm': fm,+ 'ps': ps, 'ns': ns, 'fm': fm, 'ibm': not fm,'eqi': eff_qo_indptr, 'kvi': paged_kv_indptr, 'ki': kv_indices, 'klp': kv_last_page_len,'meta': {'work_meta_data': wm, 'work_indptr': wi, 'work_info_set': wis,'reduce_indptr': ri, 'reduce_final_map': rfm, 'reduce_partial_map': rpm},⋯ 7 unchanged linesdef _custom_kernel_aiter(data: input_t) -> output_t:"""AITER paged attention — fastest for large batch × kv shapes."""q, kv_data, qo_indptr, kv_indptr, config = data- bs, nh = config["batch_size"], config["num_heads"]- qsl, kvsl = config["q_seq_len"], config["kv_seq_len"]+ bs = config["batch_size"]+ nh = config["num_heads"]+ qsl = config["q_seq_len"]+ kvsl = config["kv_seq_len"]c = _aiter_build((bs, qsl, kvsl, nh), bs, qsl, kvsl, nh, kv_indptr, q.device)r, s = kv_data["fp8"]ks = s.view(1) if s.numel() == 1 else s- ps = c['ps']- q_fp8 = c['q_fp8_buf']- q_fp8.copy_(q.view(bs * qsl, nh, QK_DIM))+ c['q_fp8_buf'].copy_(q.view(bs * qsl, nh, QK_DIM))+ ps = c['ps']if ps > 1:- pages_per_seq = (kvsl + ps - 1) // ps- ebs = bs * qsl- kv_4d = r.view(bs, kvsl, 1, QK_DIM).reshape(ebs * pages_per_seq, ps, 1, QK_DIM)+ kv_4d = r.reshape(bs * qsl * ((kvsl + ps - 1) // ps), ps, 1, QK_DIM)else:kv_4d = r.view(bs * kvsl, 1, 1, QK_DIM)- o = c['o_buf']mla_decode_fwd(- q_fp8, kv_4d, o, c['eqi'], c['kvi'], c['ki'], c['klp'],+ c['q_fp8_buf'], kv_4d, c['o_buf'], c['eqi'], c['kvi'], c['ki'], c['klp'],1, page_size=ps, nhead_kv=1, sm_scale=_SM_SCALE, logit_cap=0.0,num_kv_splits=c['ns'], q_scale=_aiter_q_scale, kv_scale=ks,- intra_batch_mode=(not c['fm']), **c['meta'])- return o+ intra_batch_mode=c['ibm'], **c['meta'])+ return c['o_buf']# ============================================================================⋯ 4 unchanged lines# c0(4,1024)=12.7 vs 22.1, c1(4,8192)=21.3 vs 23.4,# c2(32,1024)=19.0 vs 23.9, c4(64,1024)=27.3 vs 28.2# HIP wins: c0(4,1024)=13.4, c1(4,8192)=22.1, c2(32,1024)=19.5- _USE_HIP = {(4, 1, 1024, 16), (4, 1, 8192, 16), (32, 1, 1024, 16), (64, 1, 1024, 16)}+ _USE_HIP = {(4, 1, 1024, 16), (4, 1, 8192, 16), (32, 1, 1024, 16)}+ # Shapes where we use paged HIP kernel (page_size=32=KV_TILE for exact alignment)+ # Currently empty — HIP kernel not competitive for large shapes vs AITER ASM+ _USE_HIP_PAGED = set()+ _PAGED_PAGE_SIZE = 32 # Must match KV_TILE+ # Shapes where V2 kernel (occ=1, FP8 V MFMA) is used+ # Start with large shapes where bf16 V path is MFMA-bound+ _USE_HIP_V2 = set() # Disabled: v2 megatile 2x slower than AITER due to LDS scatter overhead++ # Paged kernel split tuning+ _PAGED_SPLITS = {+ (32, 1, 8192, 16): 8, # 32*8=256 WGs+ (64, 1, 8192, 16): 4, # 64*4=256 WGs+ (256, 1, 1024, 16): 1, # 256*1=256 WGs+ (256, 1, 8192, 16): 1, # 256*1=256 WGs+ }++ _paged_cache = {}+ _paged_registered = set()++def custom_kernel(data: input_t) -> output_t:cfg = data[4]shape_key = (cfg["batch_size"], cfg["q_seq_len"], cfg["kv_seq_len"], cfg["num_heads"])if shape_key in _USE_HIP:return _custom_kernel_hip(data)+ elif shape_key in _USE_HIP_V2:⋯ diff truncated
scrolls · 1201 diff lines total
Best evidence level for this revision: reported
JSON