submission 747369
nataliakokoromyti · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 2110 lines, June 9 Researcher Reciprocity License v1.0.
baseline_743863.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-747369?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:95834955d042e0c3d165a1f47d25c8d3ae33b36dd9646a4227cac7ad788db538
license declaredunknown
license concludedunknown
authorsnataliakokoromyti
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ uint8_t klds_buf[1][K_TILE_BYTES];vector-width = uint4
const uint4 lo4 = *reinterpret_cast<const uint4*>(pv);warp-specialization
static constexpr int DUET_PRODUCER_WAVES = 2;Kernel source
baseline_743863.py2110 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
# Trimmed live-path variant of amd_mla_hip_occ2_fused_noblob.py.
# Keeps the active fused IQ2 kernel path and dispatch logic while removing
# disabled variants, unreachable wrappers, and no-blob vestigial scaffolding.
# Original file is preserved unchanged.
import os
from functools import lru_cache
import torch
from torch.utils.cpp_extension import load_inline
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
# Diagnostic: print GPU info at module load (helps tune splits for target GPU)
try:
_gpu_props = torch.cuda.get_device_properties(0)
print(f"[mla] GPU: {_gpu_props.name}, CUs: {_gpu_props.multi_processor_count}", flush=True)
except Exception:
pass
NUM_HEADS = 16
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM
V_HEAD_DIM = KV_LORA_RANK
HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <cstdint>
#include <cmath>
#include <algorithm>
namespace mla {
#define NEG_INF __int_as_float(0xff800000)
static constexpr int NUM_HEADS = 16;
static constexpr int HEADS_PER_WAVE = 4;
static constexpr int D_QK = 576;
static constexpr int D_V = 512;
static constexpr float SM_SCALE = 1.0f / 24.0f;
static constexpr int BLOCK_THREADS = 256;
static constexpr int WAVE_SIZE = 64;
static constexpr int V_ELEMS = D_V / WAVE_SIZE;
static constexpr int TILE_ROWS = 32;
static constexpr int FP8_MAX = 448;
static constexpr int K_COL_BLOCK = 64;
static constexpr int K_NUM_BLOCKS = D_QK / K_COL_BLOCK;
static constexpr int K_NUM_ROWS_PER_SUBBLOCK = 4;
static constexpr int K_NUM_PADDING_DW = 2;
static constexpr int K_NUM_BYTES_PER_ROW = K_COL_BLOCK;
static constexpr int K_NUM_BYTES_PER_SUBBLOCK =
K_NUM_ROWS_PER_SUBBLOCK * K_NUM_BYTES_PER_ROW + K_NUM_PADDING_DW * 4;
static constexpr int K_NUM_BYTES_PER_BLOCK =
K_NUM_BYTES_PER_SUBBLOCK * (TILE_ROWS / K_NUM_ROWS_PER_SUBBLOCK);
static constexpr int K_TILE_BYTES = K_NUM_BYTES_PER_BLOCK * K_NUM_BLOCKS;
static constexpr int V_TILE_BYTES = TILE_ROWS * D_V;
static constexpr int VT_DV_SLICE = 16;
static constexpr int VT_NUM_SLICES = D_V / VT_DV_SLICE;
// Column-major VT layout: 8 bytes per (group, col) contiguous for ds_read_b64
// Padding between groups avoids LDS bank conflicts (128+8=136, not multiple of 128)
static constexpr int VT_GROUP_ROWS = 8;
static constexpr int VT_GROUP_PAD = 8;
static constexpr int VT_GROUP_STRIDE = VT_DV_SLICE * VT_GROUP_ROWS + VT_GROUP_PAD; // 136
static constexpr int VT_NUM_GROUPS = TILE_ROWS / VT_GROUP_ROWS; // 4
static constexpr int VT_SLICE_BYTES = VT_NUM_GROUPS * VT_GROUP_STRIDE; // 544
static constexpr int VT_TILE_BYTES = VT_NUM_SLICES * VT_SLICE_BYTES; // 17408
static constexpr uint32_t BUFFER_RESOURCE_CONFIG = 0x00020000;
using bf16 = hip_bfloat16;
using floatx4 = float __attribute__((ext_vector_type(4)));
using floatx16 = float __attribute__((ext_vector_type(16)));
using intx4 = int __attribute__((ext_vector_type(4)));
using intx8 = int __attribute__((ext_vector_type(8)));
using as3_uint32_ptr = uint32_t __attribute__((address_space(3)))*;
extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
intx4 rsrc,
as3_uint32_ptr lds_ptr,
int size,
int voffset,
int soffset,
int offset,
int aux) __asm("llvm.amdgcn.raw.buffer.load.lds");
struct alignas(16) U128 {
uint32_t x0, x1, x2, x3;
};
union U64Bytes {
uint8_t b[8];
uint64_t u64;
};
struct buffer_resource {
uint64_t ptr;
uint32_t range;
uint32_t config;
};
__device__ __forceinline__ float hw_fp8_to_f32(uint32_t packed) {
return __builtin_amdgcn_cvt_f32_fp8(packed, 0);
}
__device__ __forceinline__ intx4 make_srsrc(const void* ptr, uint32_t range_bytes) {
buffer_resource rsrc = {
reinterpret_cast<uint64_t>(ptr),
range_bytes, BUFFER_RESOURCE_CONFIG};
return *reinterpret_cast<const intx4*>(&rsrc);
}
__device__ __forceinline__ as3_uint32_ptr make_wave_lds_ptr(uintptr_t p) {
uint32_t lane0 = __builtin_amdgcn_readfirstlane(static_cast<uint32_t>(p));
return reinterpret_cast<as3_uint32_ptr>(static_cast<uintptr_t>(lane0));
}
__device__ __forceinline__ U128 ld16u(const uint8_t* p) {
return *reinterpret_cast<const U128*>(p);
}
__device__ __forceinline__ uint64_t ld_u64(const uint8_t* p) {
return *reinterpret_cast<const uint64_t*>(p);
}
__device__ __forceinline__ uint64_t ds_read_tr8_u64(const uint8_t* p)
{
#define __LDS_ADDR __attribute__((address_space(3)))
typedef __attribute__((__vector_size__(2 * sizeof(int)))) int llvm_i32x2_t;
#pragma clang diagnostic push
#pragma clang diagnostic ignored "-Wold-style-cast"
const auto p_lds = (__LDS_ADDR uint8_t*)(const_cast<uint8_t*>(p));
#pragma clang diagnostic pop
auto lds_ptr = reinterpret_cast<__LDS_ADDR llvm_i32x2_t*>(p_lds);
auto bits = __builtin_amdgcn_ds_read_tr8_b64_v2i32(lds_ptr);
return *reinterpret_cast<uint64_t*>(&bits);
#undef __LDS_ADDR
}
__device__ __forceinline__ void st16u(uint8_t* p, U128 v) {
*reinterpret_cast<U128*>(p) = v;
}
__device__ __forceinline__ float to_f(bf16 x) { return static_cast<float>(x); }
__device__ __forceinline__ bf16 to_b(float x) { return bf16(x); }
__device__ __forceinline__ float fast_exp(float x) {
// Inline asm: 2 instructions instead of compiler's 6 (skip range reduction).
// Input x is always in [-54, 0] (attention score diffs bounded by FP8 range),
// so x*log2(e) in [-78, 0], well within v_exp_f32 safe range [-126, 128].
float r;
asm("v_mul_f32 %0, 0x3fb8aa3b, %1\n\t" // r = x * log2(e)
"v_exp_f32 %0, %0" // r = 2^r = exp(x)
: "=v"(r) : "v"(x));
return r;
}
__device__ __forceinline__ int q_off(int qi, int h, int d) {
return (qi * NUM_HEADS + h) * D_QK + d;
}
__device__ __forceinline__ int out_off(int qi, int h, int d) {
return (qi * NUM_HEADS + h) * D_V + d;
}
__device__ __forceinline__ int pml_off(int b, int s, int h, int ns) {
return (b * NUM_HEADS + h) * ns + s;
}
__device__ __forceinline__ int po_off(int b, int s, int h, int d, int ns) {
return ((b * NUM_HEADS + h) * ns + s) * D_V + d;
}
__device__ __forceinline__ int lq8(int h, int d) { return h * D_QK + d; }
__device__ __forceinline__ int kv2_lds_offset(int row, int d)
{
const int block = d / K_COL_BLOCK;
const int d_in_block = d % K_COL_BLOCK;
const int half = row / 16;
const int row16 = row % 16;
const int row_phy = (row16 / 2) * 4 + (row16 % 2);
return block * K_NUM_BYTES_PER_BLOCK +
half * 128 +
(row_phy / 4) * K_NUM_BYTES_PER_SUBBLOCK +
(row_phy % 4) * K_NUM_BYTES_PER_ROW +
d_in_block;
}
__device__ __forceinline__ float wreduce_max(float v) {
v = fmaxf(v, __shfl_down(v, 32));
v = fmaxf(v, __shfl_down(v, 16));
v = fmaxf(v, __shfl_down(v, 8));
v = fmaxf(v, __shfl_down(v, 4));
v = fmaxf(v, __shfl_down(v, 2));
v = fmaxf(v, __shfl_down(v, 1));
return v;
}
__device__ __forceinline__ float wbcast(float v) { return __shfl(v, 0); }
// DPP row_ror butterfly reduction within 16-lane rows.
// row_ror:N rotates right by N within each 16-lane row (wraps around).
// Uses fused v_max_f32_dpp / v_add_f32_dpp to avoid pipeline hazards.
// s_nop between steps ensures the result commits before the next DPP read.
// DPP row_ror butterfly reduction within 16-lane rows.
// Separate v_mov_b32_dpp + regular ALU to avoid RAW hazard.
// s_nop 1 between steps: 1 intervening ALU + 2 nop cycles = 3 cycles,
// sufficient for single-value DPP chain on gfx950.
__device__ __forceinline__ float wave16_max(float v) {
float t;
asm volatile(
"v_mov_b32_dpp %1, %0 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
"v_max_f32 %0, %0, %1\n\t"
"s_nop 1\n\t"
"v_mov_b32_dpp %1, %0 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
"v_max_f32 %0, %0, %1\n\t"
"s_nop 1\n\t"
"v_mov_b32_dpp %1, %0 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
"v_max_f32 %0, %0, %1\n\t"
"s_nop 1\n\t"
"v_mov_b32_dpp %1, %0 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
"v_max_f32 %0, %0, %1"
: "+v"(v), "=&v"(t));
return v;
}
__device__ __forceinline__ float wave16_sum(float v) {
float t;
asm volatile(
"v_mov_b32_dpp %1, %0 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
"v_add_f32 %0, %0, %1\n\t"
"s_nop 1\n\t"
"v_mov_b32_dpp %1, %0 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
"v_add_f32 %0, %0, %1\n\t"
"s_nop 1\n\t"
"v_mov_b32_dpp %1, %0 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
"v_add_f32 %0, %0, %1\n\t"
"s_nop 1\n\t"
"v_mov_b32_dpp %1, %0 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
"v_add_f32 %0, %0, %1"
: "+v"(v), "=&v"(t));
return v;
}
// DPP row_ror butterfly max/sum for 4 values simultaneously.
// 4 intervening ALU ops between DPP read and next use of same register
// provides sufficient latency (4-5 cycles on gfx950) — no s_nop needed.
__device__ __forceinline__ void wave16_max4(float &a, float &b, float &c, float &d) {
float t0, t1, t2, t3;
asm volatile(
"v_mov_b32_dpp %4, %0 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %5, %1 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %6, %2 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %7, %3 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
"v_max_f32 %0, %0, %4\n\t"
"v_max_f32 %1, %1, %5\n\t"
"v_max_f32 %2, %2, %6\n\t"
"v_max_f32 %3, %3, %7\n\t"
"v_mov_b32_dpp %4, %0 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %5, %1 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %6, %2 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %7, %3 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
"v_max_f32 %0, %0, %4\n\t"
"v_max_f32 %1, %1, %5\n\t"
"v_max_f32 %2, %2, %6\n\t"
"v_max_f32 %3, %3, %7\n\t"
"v_mov_b32_dpp %4, %0 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %5, %1 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %6, %2 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %7, %3 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
"v_max_f32 %0, %0, %4\n\t"
"v_max_f32 %1, %1, %5\n\t"
"v_max_f32 %2, %2, %6\n\t"
"v_max_f32 %3, %3, %7\n\t"
"v_mov_b32_dpp %4, %0 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %5, %1 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %6, %2 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %7, %3 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
"v_max_f32 %0, %0, %4\n\t"
"v_max_f32 %1, %1, %5\n\t"
"v_max_f32 %2, %2, %6\n\t"
"v_max_f32 %3, %3, %7"
: "+v"(a), "+v"(b), "+v"(c), "+v"(d),
"=&v"(t0), "=&v"(t1), "=&v"(t2), "=&v"(t3));
}
__device__ __forceinline__ void wave16_sum4(float &a, float &b, float &c, float &d) {
float t0, t1, t2, t3;
asm volatile(
"v_mov_b32_dpp %4, %0 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %5, %1 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %6, %2 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %7, %3 row_ror:8 row_mask:0xf bank_mask:0xf\n\t"
"v_add_f32 %0, %0, %4\n\t"
"v_add_f32 %1, %1, %5\n\t"
"v_add_f32 %2, %2, %6\n\t"
"v_add_f32 %3, %3, %7\n\t"
"v_mov_b32_dpp %4, %0 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %5, %1 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %6, %2 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %7, %3 row_ror:4 row_mask:0xf bank_mask:0xf\n\t"
"v_add_f32 %0, %0, %4\n\t"
"v_add_f32 %1, %1, %5\n\t"
"v_add_f32 %2, %2, %6\n\t"
"v_add_f32 %3, %3, %7\n\t"
"v_mov_b32_dpp %4, %0 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %5, %1 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %6, %2 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %7, %3 row_ror:2 row_mask:0xf bank_mask:0xf\n\t"
"v_add_f32 %0, %0, %4\n\t"
"v_add_f32 %1, %1, %5\n\t"
"v_add_f32 %2, %2, %6\n\t"
"v_add_f32 %3, %3, %7\n\t"
"v_mov_b32_dpp %4, %0 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %5, %1 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %6, %2 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
"v_mov_b32_dpp %7, %3 row_ror:1 row_mask:0xf bank_mask:0xf\n\t"
"v_add_f32 %0, %0, %4\n\t"
"v_add_f32 %1, %1, %5\n\t"
"v_add_f32 %2, %2, %6\n\t"
"v_add_f32 %3, %3, %7"
: "+v"(a), "+v"(b), "+v"(c), "+v"(d),
"=&v"(t0), "=&v"(t1), "=&v"(t2), "=&v"(t3));
}
__device__ __forceinline__ float wave32_max(float v) {
v = fmaxf(v, __shfl_down(v, 16, 32));
v = fmaxf(v, __shfl_down(v, 8, 32));
v = fmaxf(v, __shfl_down(v, 4, 32));
v = fmaxf(v, __shfl_down(v, 2, 32));
v = fmaxf(v, __shfl_down(v, 1, 32));
return __shfl(v, 0, 32);
}
__device__ __forceinline__ float wave32_sum(float v) {
v += __shfl_down(v, 16, 32);
v += __shfl_down(v, 8, 32);
v += __shfl_down(v, 4, 32);
v += __shfl_down(v, 2, 32);
v += __shfl_down(v, 1, 32);
return __shfl(v, 0, 32);
}
struct SM {
float m, l;
};
__device__ __forceinline__ void sm_init(SM& s) {
s.m = NEG_INF;
s.l = 0.f;
}
__device__ __forceinline__ void sm_upd(SM& s, float sc, float& a, float& b) {
float mn = fmaxf(s.m, sc);
a = fast_exp(s.m - mn);
b = fast_exp(sc - mn);
s.l = a * s.l + b;
s.m = mn;
}
__device__ inline void build_fp8_lut(float* lut, float scale, int tid) {
if (tid < 256) {
lut[tid] = hw_fp8_to_f32(static_cast<uint32_t>(tid)) * scale;
}
}
__device__ __forceinline__ uint8_t cvt_fp8_scalar(float x) {
uint32_t w = 0;
x = fminf(fmaxf(x, -static_cast<float>(FP8_MAX)), static_cast<float>(FP8_MAX));
w = __builtin_amdgcn_cvt_pk_fp8_f32(x, 0.0f, w, 0);
return static_cast<uint8_t>(w & 0xffu);
}
__device__ __forceinline__ uint32_t cvt_fp8x4(float a, float b, float c, float d) {
uint32_t out = 0;
out = __builtin_amdgcn_cvt_pk_fp8_f32(a, b, out, 0);
out = __builtin_amdgcn_cvt_pk_fp8_f32(c, d, out, 1);
return out;
}
__device__ inline void stage_q_fp8(
const bf16* __restrict__ q,
uint8_t* q8,
float* q_scales,
int qi,
int wave,
int lane)
{
const int hb = wave * HEADS_PER_WAVE;
float local_max = 0.f;
for (int i = lane; i < HEADS_PER_WAVE * D_QK; i += WAVE_SIZE) {
int h = hb + (i / D_QK);
int d = i % D_QK;
local_max = fmaxf(local_max, fabsf(to_f(q[q_off(qi, h, d)])));
}
float max_abs = wbcast(wreduce_max(local_max));
float q_scale = fmaxf(
max_abs / static_cast<float>(FP8_MAX),
1.0f / static_cast<float>(FP8_MAX));
float inv_q_scale = 1.0f / q_scale;
if (lane == 0) {
q_scales[wave] = q_scale;
}
for (int i = lane * 4; i < HEADS_PER_WAVE * D_QK; i += WAVE_SIZE * 4) {
int remain = HEADS_PER_WAVE * D_QK - i;
if (remain >= 4) {
int h0 = hb + ((i + 0) / D_QK);
int h1 = hb + ((i + 1) / D_QK);
int h2 = hb + ((i + 2) / D_QK);
int h3 = hb + ((i + 3) / D_QK);
int d0 = (i + 0) % D_QK;
int d1 = (i + 1) % D_QK;
int d2 = (i + 2) % D_QK;
int d3 = (i + 3) % D_QK;
uint32_t packed = cvt_fp8x4(
to_f(q[q_off(qi, h0, d0)]) * inv_q_scale,
to_f(q[q_off(qi, h1, d1)]) * inv_q_scale,
to_f(q[q_off(qi, h2, d2)]) * inv_q_scale,
to_f(q[q_off(qi, h3, d3)]) * inv_q_scale);
*reinterpret_cast<uint32_t*>(&q8[hb * D_QK + i]) = packed;
} else {
for (int j = 0; j < remain; ++j) {
int idx = i + j;
int h = hb + (idx / D_QK);
int d = idx % D_QK;
q8[hb * D_QK + idx] = cvt_fp8_scalar(to_f(q[q_off(qi, h, d)]) * inv_q_scale);
}
}
}
}
template <int COL_OFFSET>
__device__ __forceinline__ void direct_load_k_block(
intx4 srsrc,
uintptr_t p_lds_k_warp_base,
int row,
int col_base)
{
constexpr int k_block_idx = COL_OFFSET / K_COL_BLOCK;
constexpr uintptr_t k_lds_block_base =
k_block_idx * K_NUM_BYTES_PER_BLOCK - COL_OFFSET;
const int voffset = row * D_QK + col_base;
llvm_amdgcn_raw_buffer_load_lds(
srsrc,
make_wave_lds_ptr(p_lds_k_warp_base + k_lds_block_base),
4,
voffset,
0,
COL_OFFSET,
0);
}
__device__ inline void stage_k_tile_kv2(
intx4 srsrc,
uint8_t* klds,
int base_token,
int rows,
int wave,
int lane)
{
const int col_base = (lane & 15) * 4;
#pragma unroll
for (int emu = 0; emu < 2; ++emu) {
const int warp_idx = wave + emu * 4;
const int row_base = (lane / 32) * 16 + ((lane / 16) & 1) + warp_idx * 2;
const int row = (row_base < rows) ? (base_token + row_base) : -1;
if (row < 0) continue; // Skip DMA for invalid rows
const uintptr_t p_lds_k_warp_base =
reinterpret_cast<uintptr_t>(klds)
+ warp_idx * K_NUM_BYTES_PER_SUBBLOCK;
direct_load_k_block<0>(srsrc, p_lds_k_warp_base, row, col_base);
direct_load_k_block<64>(srsrc, p_lds_k_warp_base, row, col_base);
direct_load_k_block<128>(srsrc, p_lds_k_warp_base, row, col_base);
direct_load_k_block<192>(srsrc, p_lds_k_warp_base, row, col_base);
direct_load_k_block<256>(srsrc, p_lds_k_warp_base, row, col_base);
direct_load_k_block<320>(srsrc, p_lds_k_warp_base, row, col_base);
direct_load_k_block<384>(srsrc, p_lds_k_warp_base, row, col_base);
direct_load_k_block<448>(srsrc, p_lds_k_warp_base, row, col_base);
direct_load_k_block<512>(srsrc, p_lds_k_warp_base, row, col_base);
}
}
__device__ inline void stage_v_tile_linear(
intx4 srsrc,
uint8_t* vlds,
int base_token,
int rows,
int wave,
int lane)
{
const int seg = lane & 31;
const int which = lane >> 5;
for (int pair = wave; pair < TILE_ROWS / 2; pair += 4) {
const int row0 = pair * 2;
const int local_row = row0 + which;
const int global_row = (local_row < rows) ? (base_token + local_row) : -1;
const int voffset = (global_row >= 0) ? (global_row * D_QK + seg * 16) : 0x80000000;
const uintptr_t p_lds_pair = reinterpret_cast<uintptr_t>(vlds + row0 * D_V);
llvm_amdgcn_raw_buffer_load_lds(
srsrc,
make_wave_lds_ptr(p_lds_pair),
16,
voffset,
0,
0,
0);
}
}
__device__ inline void zero_tail_tiles(uint8_t* klds, uint8_t* vlds, int rows, int tid)
{
if (rows >= TILE_ROWS) {
return;
}
for (int row = rows; row < TILE_ROWS; ++row) {
for (int d = tid; d < D_QK; d += BLOCK_THREADS) {
const int block = d / K_COL_BLOCK;
const int d_in_block = d % K_COL_BLOCK;
const int half = row / 16;
const int row16 = row % 16;
const int row_phy = (row16 / 2) * 4 + (row16 % 2);
const int offs = block * K_NUM_BYTES_PER_BLOCK +
half * 128 +
(row_phy / 4) * K_NUM_BYTES_PER_SUBBLOCK +
(row_phy % 4) * K_NUM_BYTES_PER_ROW +
d_in_block;
klds[offs] = 0;
}
for (int d = tid; d < D_V; d += BLOCK_THREADS) {
vlds[row * D_V + d] = 0;
}
}
}
__device__ __forceinline__ uint64_t pack_q_mfma(
const uint8_t* q8,
int hb,
int lane,
int k_block)
{
int row = lane & 15;
int real_h = hb + (row & 3);
int k_base = k_block + ((lane >> 4) * 8);
return *reinterpret_cast<const uint64_t*>(&q8[lq8(real_h, k_base)]);
}
__device__ __forceinline__ uint64_t load_k_frag_kv2(
const uint8_t* klds,
int lane,
int row_offset,
int k_block)
{
const int row = lane & 15;
const int row_phy = (row / 2) * 4 + (row % 2);
const int col = (lane >> 4) * 8;
const int fixed = (row_offset / 16) * 128
+ (k_block % K_COL_BLOCK)
+ (k_block / K_COL_BLOCK) * K_NUM_BYTES_PER_BLOCK;
const uint8_t* p = klds +
(row_phy / 4) * K_NUM_BYTES_PER_SUBBLOCK +
(row_phy % 4) * K_NUM_BYTES_PER_ROW +
col +
fixed;
return *reinterpret_cast<const uint64_t*>(p);
}
__device__ __forceinline__ floatx4 mfma_fp8_16x16x32(long a, long b, floatx4 c) {
return __builtin_amdgcn_mfma_f32_16x16x32_fp8_fp8(a, b, c, 0, 0, 0);
}
__device__ __forceinline__ floatx16 mfma_fp8_32x32x16(long a, long b, floatx16 c) {
return __builtin_amdgcn_mfma_f32_32x32x16_fp8_fp8(a, b, c, 0, 0, 0);
}
// Scaled MFMA: 16x16x128 with fp8 (E4M3=type 0), 4x more K-reduction per MFMA
// a,b: v8i (32 bytes = 32 fp8 elements per lane)
// scale_a, scale_b: per-block scaling factors (int, passed as sgpr)
__device__ __forceinline__ floatx4 mfma_scale_fp8_16x16x128_noscale(
intx8 a, intx8 b, floatx4 c) {
// opsel=1 bypasses the scale VGPR and uses implicit scale=1 (no scaling)
return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a, b, c, /*Atype=*/0, /*Btype=*/0, /*opsel_a=*/1, 0, /*opsel_b=*/1, 0);
}
__device__ __forceinline__ floatx4 mfma_scale_fp8_16x16x128(
intx8 a, intx8 b, floatx4 c, int scale_a, int scale_b) {
return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a, b, c, /*Atype=*/0, /*Btype=*/0, /*opsel_a=*/0, scale_a, /*opsel_b=*/0, scale_b);
}
template <int TILE>
__device__ inline void attn_direct_kv2_fp8(
const uint8_t* q8,
const float* q_scales,
const float* lut,
float sc,
const uint8_t* __restrict__ kvg,
uint8_t* klds,
uint8_t* vlds,
int hb,
int wave,
int lane,
int rs,
int re,
float* a0,
float* a1,
float* a2,
float* a3,
SM* st)
{
const intx4 srsrc = make_srsrc(kvg, 0xffffffffu);
const float score_scale = q_scales[wave] * sc * SM_SCALE;
for (int tb = rs; tb < re; tb += TILE) {
const int rows = min(TILE, re - tb);
stage_k_tile_kv2(srsrc, klds, tb, rows, wave, lane);
stage_v_tile_linear(srsrc, vlds, tb, rows, wave, lane);
__builtin_amdgcn_s_waitcnt(0);
// Yield briefly while the previous tile DMA drains.
asm volatile("s_sleep 1");
__syncthreads();
zero_tail_tiles(klds, vlds, rows, threadIdx.x);
__syncthreads();
floatx4 acc_lo = {0.f, 0.f, 0.f, 0.f};
floatx4 acc_hi = {0.f, 0.f, 0.f, 0.f};
#pragma unroll
for (int k_block = 0; k_block < D_QK; k_block += 32) {
long a_frag = static_cast<long>(pack_q_mfma(q8, hb, lane, k_block));
long b_lo = static_cast<long>(load_k_frag_kv2(klds, lane, 0, k_block));
long b_hi = static_cast<long>(load_k_frag_kv2(klds, lane, 16, k_block));
acc_lo = mfma_fp8_16x16x32(a_frag, b_lo, acc_lo);
acc_hi = mfma_fp8_16x16x32(a_frag, b_hi, acc_hi);
}
acc_lo.x *= score_scale;
acc_lo.y *= score_scale;
acc_lo.z *= score_scale;
acc_lo.w *= score_scale;
acc_hi.x *= score_scale;
acc_hi.y *= score_scale;
acc_hi.z *= score_scale;
acc_hi.w *= score_scale;
const int rows_lo = rows < 16 ? rows : 16;
for (int r = 0; r < rows_lo; ++r) {
float s0 = __shfl(acc_lo.x, r);
float s1 = __shfl(acc_lo.y, r);
float s2 = __shfl(acc_lo.z, r);
float s3 = __shfl(acc_lo.w, r);
float al[4], bt[4];
sm_upd(st[0], s0, al[0], bt[0]);
sm_upd(st[1], s1, al[1], bt[1]);
sm_upd(st[2], s2, al[2], bt[2]);
sm_upd(st[3], s3, al[3], bt[3]);
#pragma unroll
for (int j = 0; j < V_ELEMS; ++j) {
int vd = lane + j * WAVE_SIZE;
float vv = lut[vlds[r * D_V + vd]];
a0[j] = al[0] * a0[j] + bt[0] * vv;
a1[j] = al[1] * a1[j] + bt[1] * vv;
a2[j] = al[2] * a2[j] + bt[2] * vv;
a3[j] = al[3] * a3[j] + bt[3] * vv;
}
}
for (int r = 16; r < rows; ++r) {
int rr = r - 16;
float s0 = __shfl(acc_hi.x, rr);
float s1 = __shfl(acc_hi.y, rr);
float s2 = __shfl(acc_hi.z, rr);
float s3 = __shfl(acc_hi.w, rr);
float al[4], bt[4];
sm_upd(st[0], s0, al[0], bt[0]);
sm_upd(st[1], s1, al[1], bt[1]);
sm_upd(st[2], s2, al[2], bt[2]);
sm_upd(st[3], s3, al[3], bt[3]);
#pragma unroll
for (int j = 0; j < V_ELEMS; ++j) {
int vd = lane + j * WAVE_SIZE;
float vv = lut[vlds[r * D_V + vd]];
a0[j] = al[0] * a0[j] + bt[0] * vv;
a1[j] = al[1] * a1[j] + bt[1] * vv;
a2[j] = al[2] * a2[j] + bt[2] * vv;
a3[j] = al[3] * a3[j] + bt[3] * vv;
}
}
__syncthreads();
}
}
static constexpr int DUET_BLOCK_THREADS = 512;
static constexpr int DUET_WAVES_PER_BLOCK = DUET_BLOCK_THREADS / WAVE_SIZE;
static constexpr int DUET_D_SLICE = D_V / DUET_WAVES_PER_BLOCK;
static constexpr int DUET_PRODUCER_WAVES = 2;
static constexpr int NONPRODUCER_WAVES = DUET_WAVES_PER_BLOCK - DUET_PRODUCER_WAVES;
static constexpr int SCORE_K_BLOCKS = D_QK / 32;
// Parallel Q staging: all 512 threads (32 per head, wave32 reduction)
__device__ inline void stage_q_fp8_per_head(
const bf16* __restrict__ q,
uint8_t* q8,
float* q_scales,
int qi,
int wave,
int lane)
{
// 8 waves × 2 heads/wave = 16 heads, 32 threads per head
const int h = wave * 2 + (lane >> 5);
const int lt = lane & 31;
// Parallel max reduction: 32 threads each read 18 elements
float local_max = 0.0f;
#pragma unroll 4
for (int d = lt; d < D_QK; d += 32) {
local_max = fmaxf(local_max, fabsf(to_f(q[q_off(qi, h, d)])));
}
float max_abs = wave32_max(local_max);
const float q_scale = fmaxf(max_abs / static_cast<float>(FP8_MAX),
1.0f / static_cast<float>(FP8_MAX));
const float inv_q_scale = 1.0f / q_scale;
if (lt == 0) {
q_scales[h] = q_scale;
}
// Parallel quantize: 4 bytes at a time, 32 threads covering 576 elements
#pragma unroll 5
for (int d = lt * 4; d < D_QK; d += 128) {
const uint32_t packed = cvt_fp8x4(
to_f(q[q_off(qi, h, d + 0)]) * inv_q_scale,
to_f(q[q_off(qi, h, d + 1)]) * inv_q_scale,
to_f(q[q_off(qi, h, d + 2)]) * inv_q_scale,
to_f(q[q_off(qi, h, d + 3)]) * inv_q_scale);
*reinterpret_cast<uint32_t*>(&q8[h * D_QK + d]) = packed;
}
}
// Extract V from K-LDS → VT-LDS during softmax (no LDS contention with QK MFMA)
// Only called by non-producer waves while producers compute softmax
__device__ inline void extract_v_from_klds_softmax(
const uint8_t* __restrict__ klds,
uint8_t* vt,
int rows,
int wave,
int lane)
{
const int local_wave = wave - DUET_PRODUCER_WAVES;
if (local_wave < 0 || local_wave >= NONPRODUCER_WAVES) return;
const int worker_id = local_wave * WAVE_SIZE + lane;
const int num_workers = NONPRODUCER_WAVES * WAVE_SIZE;
// Row-major VT with group padding: vt[slice*544 + group*136 + tok_in_group*16 + d]
// Contiguous 16-byte writes per token (2x uint64_t), padding breaks bank conflicts
for (int wu = worker_id; wu < TILE_ROWS * VT_NUM_SLICES; wu += num_workers) {
const int tok = wu / VT_NUM_SLICES;
const int slice = wu % VT_NUM_SLICES;
const int d_base = slice * VT_DV_SLICE;
const int tok_group = tok >> 3;
const int tok_in_group = tok & 7;
uint8_t* dst = vt + slice * VT_SLICE_BYTES + tok_group * VT_GROUP_STRIDE
+ tok_in_group * VT_DV_SLICE;
if (tok >= rows) {
*reinterpret_cast<uint64_t*>(dst) = 0;
*reinterpret_cast<uint64_t*>(dst + 8) = 0;
} else {
const int src_off = kv2_lds_offset(tok, d_base);
*reinterpret_cast<uint64_t*>(dst) =
*reinterpret_cast<const uint64_t*>(klds + src_off);
*reinterpret_cast<uint64_t*>(dst + 8) =
*reinterpret_cast<const uint64_t*>(klds + src_off + 8);
}
}
}
// Extract V from K-LDS → VT buffer for a batch of 16 VT slices
// slice_start: global VT slice to start from (0 for batch 0, 16 for batch 1)
// Output always written to local positions 0-15 in vt buffer
__device__ inline void extract_v_from_klds_batch(
const uint8_t* __restrict__ klds,
uint8_t* vt,
int rows,
int worker_id,
int num_workers,
int slice_start)
{
constexpr int BATCH_SLICES = 16;
for (int wu = worker_id; wu < TILE_ROWS * BATCH_SLICES; wu += num_workers) {
const int tok = wu / BATCH_SLICES;
const int local_slice = wu % BATCH_SLICES;
const int global_slice = slice_start + local_slice;
const int d_base = global_slice * VT_DV_SLICE;
const int tok_group = tok >> 3;
const int tok_in_group = tok & 7;
uint8_t* dst = vt + local_slice * VT_SLICE_BYTES + tok_group * VT_GROUP_STRIDE
+ tok_in_group * VT_DV_SLICE;
if (tok >= rows) {
*reinterpret_cast<uint64_t*>(dst) = 0;
*reinterpret_cast<uint64_t*>(dst + 8) = 0;
} else {
const int src_off = kv2_lds_offset(tok, d_base);
*reinterpret_cast<uint64_t*>(dst) =
*reinterpret_cast<const uint64_t*>(klds + src_off);
*reinterpret_cast<uint64_t*>(dst + 8) =
*reinterpret_cast<const uint64_t*>(klds + src_off + 8);
}
}
}
__device__ __forceinline__ uint64_t load_v_frag_tr8(
const uint8_t* vt,
int vt_slice,
int lane)
{
// Row-major VT with group padding: 8 byte reads at stride 16
const int col = lane & 15;
const int group = lane >> 4;
const uint8_t* base = vt + vt_slice * VT_SLICE_BYTES + group * VT_GROUP_STRIDE;
uint64_t result;
uint8_t* rb = reinterpret_cast<uint8_t*>(&result);
rb[0] = base[0 * 16 + col];
rb[1] = base[1 * 16 + col];
rb[2] = base[2 * 16 + col];
rb[3] = base[3 * 16 + col];
rb[4] = base[4 * 16 + col];
rb[5] = base[5 * 16 + col];
rb[6] = base[6 * 16 + col];
rb[7] = base[7 * 16 + col];
return result;
}
// Read V directly from klds using dword loads + v_perm_b32 for fast byte packing
// Dword loads enable 4-lane broadcast (1 LDS cycle vs 4 for byte reads)
// Total: 8 ds_read_b32 + ~10 ALU = ~26 cycles (fits in 64-cycle MFMA window)
__device__ __forceinline__ uint64_t load_v_from_klds_bytes(
const uint8_t* klds, int vt_slice, int lane)
{
const int col = lane & 15;
const int group = lane >> 4;
const int d = vt_slice * VT_DV_SLICE + col;
const int group_base = (group >= 2 ? 128 : 0)
+ ((group & 1) ? 4 * K_NUM_BYTES_PER_SUBBLOCK : 0);
const uint8_t* rb = klds + (d >> 6) * K_NUM_BYTES_PER_BLOCK + group_base + (d & 63);
const uint8_t b0 = rb[0];
const uint8_t b1 = rb[K_NUM_BYTES_PER_ROW];
const uint8_t b2 = rb[K_NUM_BYTES_PER_SUBBLOCK];
const uint8_t b3 = rb[K_NUM_BYTES_PER_SUBBLOCK + K_NUM_BYTES_PER_ROW];
const uint8_t b4 = rb[2 * K_NUM_BYTES_PER_SUBBLOCK];
const uint8_t b5 = rb[2 * K_NUM_BYTES_PER_SUBBLOCK + K_NUM_BYTES_PER_ROW];
const uint8_t b6 = rb[3 * K_NUM_BYTES_PER_SUBBLOCK];
const uint8_t b7 = rb[3 * K_NUM_BYTES_PER_SUBBLOCK + K_NUM_BYTES_PER_ROW];
return (uint64_t)b0 | ((uint64_t)b1 << 8) | ((uint64_t)b2 << 16) | ((uint64_t)b3 << 24)
| ((uint64_t)b4 << 32) | ((uint64_t)b5 << 40) | ((uint64_t)b6 << 48) | ((uint64_t)b7 << 56);
}
// Load 128-col K fragment from K-LDS for scaled MFMA 16x16x128
// Returns intx8 (32 bytes) for the B operand of mfma_scale_f32_16x16x128
__device__ __forceinline__ intx8 load_k128_frag(
const uint8_t* klds, int lane, int tok_base, int mfma_blk128)
{
const int _row = lane & 15;
const int grp = lane >> 4;
const int _rphy = (_row / 2) * 4 + (_row % 2);
const int half = tok_base / 16;
const int k_block = mfma_blk128 * 2 + grp / 2;
const int col_in_kblock = (grp & 1) * 32;
const uint8_t* base = klds
+ k_block * K_NUM_BYTES_PER_BLOCK
+ half * 128
+ (_rphy / 4) * K_NUM_BYTES_PER_SUBBLOCK
+ (_rphy % 4) * K_NUM_BYTES_PER_ROW
+ col_in_kblock;
intx8 r;
reinterpret_cast<uint64_t*>(&r)[0] = *reinterpret_cast<const uint64_t*>(base);
reinterpret_cast<uint64_t*>(&r)[1] = *reinterpret_cast<const uint64_t*>(base + 8);
reinterpret_cast<uint64_t*>(&r)[2] = *reinterpret_cast<const uint64_t*>(base + 16);
reinterpret_cast<uint64_t*>(&r)[3] = *reinterpret_cast<const uint64_t*>(base + 24);
return r;
}
// Load V fragment directly from K-LDS using transposed read (no separate vt buffer)
__device__ __forceinline__ uint64_t load_v_from_klds_tr8(
const uint8_t* klds, int vt_slice, int lane)
{
const int lc = lane & 15;
const int lg = lane >> 4;
const int row = lg * 8 + (lc >> 1);
const int d = vt_slice * VT_DV_SLICE + (lc & 1) * 8;
const int block = d >> 6;
const int d_in_block = d & 63;
const int half = row >> 4;
const int row16 = row & 15;
const int row_phy = (row16 >> 1) * 4 + (row16 & 1);
const int off = block * K_NUM_BYTES_PER_BLOCK
+ half * 128
+ (row_phy >> 2) * K_NUM_BYTES_PER_SUBBLOCK
+ (row_phy & 3) * K_NUM_BYTES_PER_ROW
+ d_in_block;
return ds_read_tr8_u64(klds + off);
}
template <int COL_OFFSET>
__device__ __forceinline__ void direct_load_k_block_duet(
intx4 srsrc,
uintptr_t p_lds_k_warp_base,
int row,
int col_base)
{
constexpr int k_block_idx = COL_OFFSET / K_COL_BLOCK;
constexpr uintptr_t k_lds_block_base =
k_block_idx * K_NUM_BYTES_PER_BLOCK - COL_OFFSET;
const int voffset = row * D_QK + col_base;
llvm_amdgcn_raw_buffer_load_lds(
srsrc,
make_wave_lds_ptr(p_lds_k_warp_base + k_lds_block_base),
4,
voffset,
0,
COL_OFFSET,
0);
}
__device__ inline void stage_k_tile_kv2_duet(
intx4 srsrc,
uint8_t* klds,
int base_token,
int rows,
int wave,
int lane)
{
const int col_base = (lane & 15) * 4;
const int warp_idx = wave;
const int row_base = (lane / 32) * 16 + ((lane / 16) & 1) + warp_idx * 2;
const int row = (row_base < rows) ? (base_token + row_base) : -1;
if (row < 0) return;
const uintptr_t p_lds_k_warp_base =
reinterpret_cast<uintptr_t>(klds)
+ warp_idx * K_NUM_BYTES_PER_SUBBLOCK;
direct_load_k_block_duet<0>(srsrc, p_lds_k_warp_base, row, col_base);
direct_load_k_block_duet<64>(srsrc, p_lds_k_warp_base, row, col_base);
direct_load_k_block_duet<128>(srsrc, p_lds_k_warp_base, row, col_base);
direct_load_k_block_duet<192>(srsrc, p_lds_k_warp_base, row, col_base);
direct_load_k_block_duet<256>(srsrc, p_lds_k_warp_base, row, col_base);
direct_load_k_block_duet<320>(srsrc, p_lds_k_warp_base, row, col_base);
direct_load_k_block_duet<384>(srsrc, p_lds_k_warp_base, row, col_base);
direct_load_k_block_duet<448>(srsrc, p_lds_k_warp_base, row, col_base);
direct_load_k_block_duet<512>(srsrc, p_lds_k_warp_base, row, col_base);
}
// Coalesced VMEM K-tile loader: HBM -> registers -> LDS
// Each wave reads one row at a time, all 64 lanes access same row at consecutive 16B chunks.
// D_QK=576, 576/16=36 chunks per row. Lanes 0-35 active, 36-63 idle.
// Ensures coalesced HBM access (all active lanes within one cache line region).
__device__ inline void vmem_load_k_tile(
const uint8_t* __restrict__ kv, uint8_t* __restrict__ klds,
int base_token, int rows, int local_wave, int num_waves, int lane)
{
for (int row = local_wave; row < TILE_ROWS; row += num_waves) {
const bool valid = (row < rows);
const int hbm_base = valid ? ((base_token + row) * D_QK) : 0;
const int d = lane << 4; // lane * 16: lanes 0-35 cover d=0,16,...,560
if (d < D_QK) {
U128 data = {};
if (valid) {
data = *reinterpret_cast<const U128*>(&kv[hbm_base + d]);
}
const int lds_off = kv2_lds_offset(row, d);
// Two 8B LDS writes for 16B total
*reinterpret_cast<uint64_t*>(&klds[lds_off]) =
static_cast<uint64_t>(data.x0) | (static_cast<uint64_t>(data.x1) << 32);
*reinterpret_cast<uint64_t*>(&klds[lds_off + 8]) =
static_cast<uint64_t>(data.x2) | (static_cast<uint64_t>(data.x3) << 32);
}
}
}
__device__ inline void zero_tail_k_tile(uint8_t* klds, int rows, int tid)
{
if (rows >= TILE_ROWS) {
return;
}
for (int row = rows; row < TILE_ROWS; ++row) {
for (int d = tid; d < D_QK; d += DUET_BLOCK_THREADS) {
const int block = d / K_COL_BLOCK;
const int d_in_block = d % K_COL_BLOCK;
const int half = row / 16;
const int row16 = row % 16;
const int row_phy = (row16 / 2) * 4 + (row16 % 2);
const int offs = block * K_NUM_BYTES_PER_BLOCK +
half * 128 +
(row_phy / 4) * K_NUM_BYTES_PER_SUBBLOCK +
(row_phy % 4) * K_NUM_BYTES_PER_ROW +
d_in_block;
klds[offs] = 0;
}
}
}
__device__ __forceinline__ uint64_t pack_q_mfma_full(
const uint8_t* q8,
int lane,
int k_block)
{
const int row = lane & 15;
const int k_base = k_block + ((lane >> 4) * 8);
return *reinterpret_cast<const uint64_t*>(&q8[lq8(row, k_base)]);
}
// ===================== light + inline Q quantization v2 ==================
// Takes bf16 Q, quantizes to FP8 into LDS at kernel start (all waves cooperate).
// Then loads Q FP8 from LDS into registers (same register footprint as mla_s1_light).
// Eliminates the separate Q quantization kernel launch (~3us savings on ranked).
template <bool DIRECT_OUT, bool FUSE_S2>
__global__ __launch_bounds__(DUET_BLOCK_THREADS, 2)
void mla_s1_light_iq2(
const uint16_t* __restrict__ q_bf16,
const uint8_t* __restrict__ kv,
const float* __restrict__ kv_scale_ptr,
const int32_t* __restrict__ qo,
const int32_t* __restrict__ kvi,
float* __restrict__ pm,
float* __restrict__ pl,
bf16* __restrict__ po,
bf16* __restrict__ out,
int bs,
int ns)
{
const float kv_scale = *kv_scale_ptr;
const intx4 srsrc = make_srsrc(kv, 0xffffffffu);
const int bid = blockIdx.x;
const int b = DIRECT_OUT ? bid : (bid / ns);
const int sid = DIRECT_OUT ? 0 : (bid % ns);
if (b >= bs) return;
const int tid = threadIdx.x;
const int wave = tid >> 6;
const int lane = tid & 63;
const int lane_col = lane & 15;
const int lane_group = lane >> 4;
const int row_base = lane_group * 4;
const int qs = qo[b];
const int qe = qo[b + 1];
if (qe - qs != 1) return;
const int qi = qs;
const int kvs = kvi[b];
const int kve = kvi[b + 1];
const int kvlen = kve - kvs;
// Single-buffered K tiles for occ=2 (LDS < 32KB)
__shared__ uint8_t klds_buf[1][K_TILE_BYTES];
__shared__ uint8_t p8[NUM_HEADS * TILE_ROWS];
__shared__ float half_max[2][NUM_HEADS];
__shared__ float half_sum[2][NUM_HEADS];
__shared__ uint8_t q_fp8_lds[NUM_HEADS * D_QK + 128];
float* q_sc_lds = reinterpret_cast<float*>(&q_fp8_lds[NUM_HEADS * D_QK]);
__shared__ int fused_last;
int* __done_count = nullptr;
if constexpr (FUSE_S2) {
__done_count = reinterpret_cast<int*>(pm + bs * NUM_HEADS * ns);
}
float local_max[4] = {NEG_INF, NEG_INF, NEG_INF, NEG_INF};
float local_sum[4] = {0.0f, 0.0f, 0.0f, 0.0f};
const bool producer = wave < DUET_PRODUCER_WAVES;
const int ss = DIRECT_OUT ? kvs : (kvs + (kvlen * sid) / ns);
const int se = DIRECT_OUT ? kve : (kvs + (kvlen * (sid + 1)) / ns);
if (se <= ss) {
if constexpr (DIRECT_OUT) {
for (int idx = tid; idx < NUM_HEADS * D_V; idx += DUET_BLOCK_THREADS) {
out[out_off(qi, idx / D_V, idx % D_V)] = to_b(0.0f);
}
return;
}
if (tid < NUM_HEADS) {
pm[pml_off(b, sid, tid, ns)] = NEG_INF;
pl[pml_off(b, sid, tid, ns)] = 0.0f;
}
for (int idx = tid; idx < NUM_HEADS * D_V; idx += DUET_BLOCK_THREADS) {
po[po_off(b, sid, idx / D_V, idx % D_V, ns)] = to_b(0.0f);
}
if constexpr (!FUSE_S2) return;
// Release: drain all stores before signaling
asm volatile("s_waitcnt vmcnt(0) lgkmcnt(0)" ::: "memory");
if (tid == 0) fused_last = atomicAdd(&__done_count[b], 1);
__syncthreads();
if (fused_last != ns - 1) return;
// Acquire: only last CTA needs full fence to see other CTAs' stores
__threadfence();
goto s2_tail_iq2;
}
{ // scope: all vars here invisible at s2_tail_iq2, so goto is valid
// Start first tile DMA EARLY — overlaps with Q quantization below
const int first_rows = min(TILE_ROWS, se - ss);
stage_k_tile_kv2_duet(srsrc, klds_buf[0], ss, first_rows, wave, lane);
// ---- Inline Q quantization: all 8 waves, 2 heads per wave ----
// Runs concurrently with buffer_load_lds DMA above (different LDS regions)
{
const int hh0 = wave * 2;
#pragma unroll
for (int hoff = 0; hoff < 2; ++hoff) {
const int h = hh0 + hoff;
const int gbase = (qi * NUM_HEADS + h) * D_QK;
float la = 0.0f;
float vals[9];
#pragma unroll
for (int i = 0; i < 9; ++i) {
float fv = __uint_as_float((uint32_t)q_bf16[gbase + lane + i * 64] << 16);
vals[i] = fv;
la = fmaxf(la, fabsf(fv));
}
float amax = wbcast(wreduce_max(la));
amax = fmaxf(amax, 1.0f / static_cast<float>(FP8_MAX));
float inv_s = static_cast<float>(FP8_MAX) / amax;
if (lane == 0) q_sc_lds[h] = amax / static_cast<float>(FP8_MAX);
#pragma unroll
for (int i = 0; i < 4; ++i) {
float a = fminf(fmaxf(vals[i*2] * inv_s, -static_cast<float>(FP8_MAX)), static_cast<float>(FP8_MAX));
float bb = fminf(fmaxf(vals[i*2+1] * inv_s, -static_cast<float>(FP8_MAX)), static_cast<float>(FP8_MAX));
uint32_t pk = 0;
pk = __builtin_amdgcn_cvt_pk_fp8_f32(a, bb, pk, 0);
q_fp8_lds[h * D_QK + lane + (i*2) * 64] = static_cast<uint8_t>(pk & 0xFFu);
q_fp8_lds[h * D_QK + lane + (i*2+1) * 64] = static_cast<uint8_t>((pk >> 8) & 0xFFu);
}
{
float a = fminf(fmaxf(vals[8] * inv_s, -static_cast<float>(FP8_MAX)), static_cast<float>(FP8_MAX));
uint32_t pk = 0;
pk = __builtin_amdgcn_cvt_pk_fp8_f32(a, 0.0f, pk, 0);
q_fp8_lds[h * D_QK + lane + 8 * 64] = static_cast<uint8_t>(pk & 0xFFu);
}
}
}
__syncthreads(); // barrier0: Q fp8 in LDS ready
{
// First tile DMA already in flight from above — no need to start it here
const int slice_base = wave * DUET_D_SLICE;
const uint8_t* q_row = producer ? (q_fp8_lds + (lane & 15) * D_QK) : nullptr;
const int q_grp128 = lane >> 4;
const int q_grp32 = q_grp128 * 8;
const int q_grp128_col = q_grp128 * 32;
const float score_scale_mul = kv_scale * SM_SCALE;
intx8 q0_cached = {};
intx8 q1_cached = {};
intx8 q2_cached = {};
intx8 q3_cached = {};
float scale0_cached = 0.0f;
float scale1_cached = 0.0f;
float scale2_cached = 0.0f;
float scale3_cached = 0.0f;
if (producer) {
q0_cached = *reinterpret_cast<const intx8*>(&q_row[q_grp128_col + 0]);
q1_cached = *reinterpret_cast<const intx8*>(&q_row[q_grp128_col + 128]);
q2_cached = *reinterpret_cast<const intx8*>(&q_row[q_grp128_col + 256]);
q3_cached = *reinterpret_cast<const intx8*>(&q_row[q_grp128_col + 384]);
scale0_cached = q_sc_lds[row_base + 0] * score_scale_mul;
scale1_cached = q_sc_lds[row_base + 1] * score_scale_mul;
scale2_cached = q_sc_lds[row_base + 2] * score_scale_mul;
scale3_cached = q_sc_lds[row_base + 3] * score_scale_mul;
}
floatx4 o_acc[4];
#pragma unroll
for (int i = 0; i < 4; ++i) o_acc[i] = {0.0f, 0.0f, 0.0f, 0.0f};
uint64_t v_pre[4];
for (int tb = ss; tb < se; tb += TILE_ROWS) {
const int rows = min(TILE_ROWS, se - tb);
const bool full_tile = rows == TILE_ROWS;
const bool init_tile = tb == ss;
uint8_t* klds = klds_buf[0];
const int next_tb = tb + TILE_ROWS;
const bool has_next = next_tb < se;
const int next_rows = has_next ? min(TILE_ROWS, se - next_tb) : 0;
__builtin_amdgcn_s_waitcnt(0);
__syncthreads();
floatx4 score = {0.0f, 0.0f, 0.0f, 0.0f};
const int tok_base = wave * 16;
const int rows_half = tok_base < rows ? min(16, rows - tok_base) : 0;
if (producer) {
asm volatile("s_setprio 1");
// Compiler-managed QK scoring: no hardcoded VGPR clobbers.
floatx4 sa = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 sb = {0.0f, 0.0f, 0.0f, 0.0f};
const uint64_t qr0 = *reinterpret_cast<const uint64_t*>(&q_row[512 + q_grp32]);
const uint64_t qr1 = *reinterpret_cast<const uint64_t*>(&q_row[544 + q_grp32]);
intx8 ka = load_k128_frag(klds, lane, tok_base, 0);
intx8 kb_data = load_k128_frag(klds, lane, tok_base, 1);
sa = mfma_scale_fp8_16x16x128_noscale(q0_cached, ka, sa);
sb = mfma_scale_fp8_16x16x128_noscale(q1_cached, kb_data, sb);
intx8 ka2 = load_k128_frag(klds, lane, tok_base, 2);
intx8 kb_data2 = load_k128_frag(klds, lane, tok_base, 3);
sa = mfma_scale_fp8_16x16x128_noscale(q2_cached, ka2, sa);
sb = mfma_scale_fp8_16x16x128_noscale(q3_cached, kb_data2, sb);
#pragma unroll
for (int cb = 0; cb < 4; ++cb) {
v_pre[cb] = load_v_from_klds_tr8(klds,
wave * (DUET_D_SLICE / VT_DV_SLICE) + cb, lane);
}
uint64_t kf0 = load_k_frag_kv2(klds, lane, tok_base, 512);
uint64_t kf1 = load_k_frag_kv2(klds, lane, tok_base, 544);
sa = mfma_fp8_16x16x32(static_cast<long>(qr0), static_cast<long>(kf0), sa);
sb = mfma_fp8_16x16x32(static_cast<long>(qr1), static_cast<long>(kf1), sb);
score.x = sa.x + sb.x; score.y = sa.y + sb.y;
score.z = sa.z + sb.z; score.w = sa.w + sb.w;
asm volatile("s_setprio 0");
{
const float scale0 = scale0_cached;
const float scale1 = scale1_cached;
const float scale2 = scale2_cached;
const float scale3 = scale3_cached;
float s0, s1, s2, s3;
if (full_tile) {
s0 = score[0] * scale0;
s1 = score[1] * scale1;
s2 = score[2] * scale2;
s3 = score[3] * scale3;
} else {
s0 = lane_col < rows_half ? score[0] * scale0 : NEG_INF;
s1 = lane_col < rows_half ? score[1] * scale1 : NEG_INF;
s2 = lane_col < rows_half ? score[2] * scale2 : NEG_INF;
s3 = lane_col < rows_half ? score[3] * scale3 : NEG_INF;
}
wave16_max4(s0, s1, s2, s3);
if (lane_col == 0) {
half_max[wave][row_base + 0] = s0;
half_max[wave][row_base + 1] = s1;
half_max[wave][row_base + 2] = s2;
half_max[wave][row_base + 3] = s3;
}
}
} else {
#pragma unroll
for (int cb = 0; cb < 4; ++cb) {
v_pre[cb] = load_v_from_klds_tr8(klds,
wave * (DUET_D_SLICE / VT_DV_SLICE) + cb, lane);
}
}
if (has_next) {
stage_k_tile_kv2_duet(srsrc, klds_buf[0],
next_tb, next_rows, wave, lane);
}
__syncthreads(); // barrier2: half_max ready, V loaded
// All waves: compute alpha + rescale o_acc (overlap with producer softmax)
// Pre-load all half_max values to allow LDS read pipelining
// (avoids serial read-wait-compute chains)
float hm0 = half_max[0][row_base + 0];
float hm1 = half_max[1][row_base + 0];
float hm2 = half_max[0][row_base + 1];
float hm3 = half_max[1][row_base + 1];
float hm4 = half_max[0][row_base + 2];
float hm5 = half_max[1][row_base + 2];
float hm6 = half_max[0][row_base + 3];
float hm7 = half_max[1][row_base + 3];
// Compiler fence: force all 8 LDS reads to issue before continuing.
// Without this, compiler sinks reads near uses → 4 serial waits (~160 cy).
// With fence: 1 batched wait (~43 cy), saving ~117 cycles per tile.
asm volatile("" : "+v"(hm0), "+v"(hm1), "+v"(hm2), "+v"(hm3),
"+v"(hm4), "+v"(hm5), "+v"(hm6), "+v"(hm7));
float alpha[4];
if (init_tile) {
local_max[0] = fmaxf(hm0, hm1);
local_max[1] = fmaxf(hm2, hm3);
local_max[2] = fmaxf(hm4, hm5);
local_max[3] = fmaxf(hm6, hm7);
alpha[0] = 0.0f;
alpha[1] = 0.0f;
alpha[2] = 0.0f;
alpha[3] = 0.0f;
} else if (full_tile) {
float new_m0 = fmaxf(local_max[0], fmaxf(hm0, hm1));
alpha[0] = fast_exp(local_max[0] - new_m0);
local_max[0] = new_m0;
float new_m1 = fmaxf(local_max[1], fmaxf(hm2, hm3));
alpha[1] = fast_exp(local_max[1] - new_m1);
local_max[1] = new_m1;
float new_m2 = fmaxf(local_max[2], fmaxf(hm4, hm5));
alpha[2] = fast_exp(local_max[2] - new_m2);
local_max[2] = new_m2;
float new_m3 = fmaxf(local_max[3], fmaxf(hm6, hm7));
alpha[3] = fast_exp(local_max[3] - new_m3);
local_max[3] = new_m3;
} else {
float new_m0 = fmaxf(local_max[0], fmaxf(hm0, hm1));
alpha[0] = local_sum[0] > 0.0f ? fast_exp(local_max[0] - new_m0) : 0.0f;
local_max[0] = new_m0;
float new_m1 = fmaxf(local_max[1], fmaxf(hm2, hm3));
alpha[1] = local_sum[1] > 0.0f ? fast_exp(local_max[1] - new_m1) : 0.0f;
local_max[1] = new_m1;
float new_m2 = fmaxf(local_max[2], fmaxf(hm4, hm5));
alpha[2] = local_sum[2] > 0.0f ? fast_exp(local_max[2] - new_m2) : 0.0f;
local_max[2] = new_m2;
float new_m3 = fmaxf(local_max[3], fmaxf(hm6, hm7));
alpha[3] = local_sum[3] > 0.0f ? fast_exp(local_max[3] - new_m3) : 0.0f;
local_max[3] = new_m3;
}
#pragma unroll
for (int i = 0; i < 4; ++i) {
o_acc[i].x *= alpha[0]; o_acc[i].y *= alpha[1];
o_acc[i].z *= alpha[2]; o_acc[i].w *= alpha[3];
}
if (producer) {
const float scale0 = scale0_cached;
const float scale1 = scale1_cached;
const float scale2 = scale2_cached;
const float scale3 = scale3_cached;
if (full_tile) {
const int p_col = tok_base + lane_col;
const float s0 = score[0] * scale0;
const float s1 = score[1] * scale1;
const float s2 = score[2] * scale2;
const float s3 = score[3] * scale3;
float p0 = fast_exp(s0 - local_max[0]);
float p1 = fast_exp(s1 - local_max[1]);
float p2 = fast_exp(s2 - local_max[2]);
float p3 = fast_exp(s3 - local_max[3]);
const uint32_t p_pack = cvt_fp8x4(p0, p1, p2, p3);
p8[(row_base + 0) * TILE_ROWS + p_col] = static_cast<uint8_t>(p_pack & 0xffu);
p8[(row_base + 1) * TILE_ROWS + p_col] = static_cast<uint8_t>((p_pack >> 8) & 0xffu);
p8[(row_base + 2) * TILE_ROWS + p_col] = static_cast<uint8_t>((p_pack >> 16) & 0xffu);
p8[(row_base + 3) * TILE_ROWS + p_col] = static_cast<uint8_t>((p_pack >> 24) & 0xffu);
wave16_sum4(p0, p1, p2, p3);
if (lane_col == 0) {
half_sum[wave][row_base + 0] = p0;
half_sum[wave][row_base + 1] = p1;
half_sum[wave][row_base + 2] = p2;
half_sum[wave][row_base + 3] = p3;
}
} else {
const bool valid = lane_col < rows_half;
const int p_col = tok_base + lane_col;
const float s0 = valid ? score[0] * scale0 : NEG_INF;
const float s1 = valid ? score[1] * scale1 : NEG_INF;
const float s2 = valid ? score[2] * scale2 : NEG_INF;
const float s3 = valid ? score[3] * scale3 : NEG_INF;
float p0 = valid ? fast_exp(s0 - local_max[0]) : 0.0f;
float p1 = valid ? fast_exp(s1 - local_max[1]) : 0.0f;
float p2 = valid ? fast_exp(s2 - local_max[2]) : 0.0f;
float p3 = valid ? fast_exp(s3 - local_max[3]) : 0.0f;
const uint32_t p_pack = cvt_fp8x4(p0, p1, p2, p3);
p8[(row_base + 0) * TILE_ROWS + p_col] = static_cast<uint8_t>(p_pack & 0xffu);
p8[(row_base + 1) * TILE_ROWS + p_col] = static_cast<uint8_t>((p_pack >> 8) & 0xffu);
p8[(row_base + 2) * TILE_ROWS + p_col] = static_cast<uint8_t>((p_pack >> 16) & 0xffu);
p8[(row_base + 3) * TILE_ROWS + p_col] = static_cast<uint8_t>((p_pack >> 24) & 0xffu);
wave16_sum4(p0, p1, p2, p3);
if (lane_col == 0) {
half_sum[wave][row_base + 0] = p0;
half_sum[wave][row_base + 1] = p1;
half_sum[wave][row_base + 2] = p2;
half_sum[wave][row_base + 3] = p3;
}
}
}
__syncthreads(); // barrier3: half_sum + p8 ready
// Update local_sum (needs half_sum from barrier3)
#pragma unroll
for (int h = 0; h < 4; ++h) {
const int row = row_base + h;
local_sum[h] = alpha[h] * local_sum[h]
+ half_sum[0][row] + half_sum[1][row];
}
const uint64_t p_frag =
*reinterpret_cast<const uint64_t*>(
&p8[lane_col * TILE_ROWS + lane_group * 8]);
#pragma unroll
for (int col_block = 0; col_block < 4; ++col_block) {
o_acc[col_block] = mfma_fp8_16x16x32(
static_cast<long>(p_frag),
static_cast<long>(v_pre[col_block]),
o_acc[col_block]);
}
}
if constexpr (!DIRECT_OUT) {
if (wave == 0 && lane_col == 0) {
#pragma unroll
for (int h = 0; h < 4; ++h) {
pm[pml_off(b, sid, row_base + h, ns)] = local_max[h];
pl[pml_off(b, sid, row_base + h, ns)] = local_sum[h];
}
}
}
const float inv0 = local_sum[0] > 0.0f ? (kv_scale / local_sum[0]) : 0.0f;
const float inv1 = local_sum[1] > 0.0f ? (kv_scale / local_sum[1]) : 0.0f;
const float inv2 = local_sum[2] > 0.0f ? (kv_scale / local_sum[2]) : 0.0f;
const float inv3 = local_sum[3] > 0.0f ? (kv_scale / local_sum[3]) : 0.0f;
#pragma unroll
for (int col_block = 0; col_block < 4; ++col_block) {
const int dv = slice_base + col_block * 16 + lane_col;
if constexpr (DIRECT_OUT) {
out[out_off(qi, row_base + 0, dv)] = to_b(o_acc[col_block].x * inv0);
out[out_off(qi, row_base + 1, dv)] = to_b(o_acc[col_block].y * inv1);
out[out_off(qi, row_base + 2, dv)] = to_b(o_acc[col_block].z * inv2);
out[out_off(qi, row_base + 3, dv)] = to_b(o_acc[col_block].w * inv3);
} else {
po[po_off(b, sid, row_base + 0, dv, ns)] = to_b(o_acc[col_block].x * inv0);
po[po_off(b, sid, row_base + 1, dv, ns)] = to_b(o_acc[col_block].y * inv1);
po[po_off(b, sid, row_base + 2, dv, ns)] = to_b(o_acc[col_block].z * inv2);
po[po_off(b, sid, row_base + 3, dv, ns)] = to_b(o_acc[col_block].w * inv3);
}
}
} // end inner scope (original block for main computation)
} // end outer scope — first_rows, slice_base, o_acc etc dead; goto s2_tail_iq2 is valid
if constexpr (DIRECT_OUT) return;
if constexpr (!FUSE_S2) return;
// ---- Fused S2: atomic completion + last-CTA reduction ----
// Release: drain all stores before signaling (cheap — no cache invalidation)
asm volatile("s_waitcnt vmcnt(0) lgkmcnt(0)" ::: "memory");
if (tid == 0) fused_last = atomicAdd(&__done_count[b], 1);
__syncthreads();
if (fused_last != ns - 1) return;
// Acquire: only last CTA needs full fence to see other CTAs' stores
__threadfence();
s2_tail_iq2:
// S2 reduction: 512 threads = 16 heads × 32 threads
// Each thread reduces ns splits for 16 contiguous V dims → 32×16 = 512 = D_V
{
const int s2_h = tid >> 5; // head index 0..15
const int s2_lt = tid & 31; // lane within head group 0..31
const int vd_base = s2_lt * 16; // 16 contiguous V dims per thread
const int bh_ns = (b * NUM_HEADS + s2_h) * ns;
float mm = NEG_INF;
float ll = 0.0f;
// 16 V-dim accumulators — reuse VGPRs freed from S1's o_acc/q_frag
float r0=0.f, r1=0.f, r2=0.f, r3=0.f;
float r4=0.f, r5=0.f, r6=0.f, r7=0.f;
float r8=0.f, r9=0.f, r10=0.f, r11=0.f;
float r12=0.f, r13=0.f, r14=0.f, r15=0.f;
for (int s = 0; s < ns; ++s) {
float m = pm[bh_ns + s];
float l = pl[bh_ns + s];
// Online softmax: merge split s into running accumulator
float mn = fmaxf(mm, m);
float a = fast_exp(mm - mn);
float bw = fast_exp(m - mn) * l;
ll = a * ll + bw;
mm = mn;
// Vectorized 128-bit loads: 2 × uint4 = 16 bf16 values (32 bytes)
const bf16* pv = &po[bh_ns * D_V + s * D_V + vd_base];
const uint4 lo4 = *reinterpret_cast<const uint4*>(pv);
const uint4 hi4 = *reinterpret_cast<const uint4*>(pv + 8);
// Unpack bf16 pairs from packed uint32 and FMA into accumulators
// lo4.x = [bf16_1 | bf16_0], lo4.y = [bf16_3 | bf16_2], etc.
r0 = __builtin_fmaf(a, r0, bw * to_f(reinterpret_cast<const bf16*>(&lo4.x)[0]));
r1 = __builtin_fmaf(a, r1, bw * to_f(reinterpret_cast<const bf16*>(&lo4.x)[1]));
r2 = __builtin_fmaf(a, r2, bw * to_f(reinterpret_cast<const bf16*>(&lo4.y)[0]));
r3 = __builtin_fmaf(a, r3, bw * to_f(reinterpret_cast<const bf16*>(&lo4.y)[1]));
r4 = __builtin_fmaf(a, r4, bw * to_f(reinterpret_cast<const bf16*>(&lo4.z)[0]));
r5 = __builtin_fmaf(a, r5, bw * to_f(reinterpret_cast<const bf16*>(&lo4.z)[1]));
r6 = __builtin_fmaf(a, r6, bw * to_f(reinterpret_cast<const bf16*>(&lo4.w)[0]));
r7 = __builtin_fmaf(a, r7, bw * to_f(reinterpret_cast<const bf16*>(&lo4.w)[1]));
r8 = __builtin_fmaf(a, r8, bw * to_f(reinterpret_cast<const bf16*>(&hi4.x)[0]));
r9 = __builtin_fmaf(a, r9, bw * to_f(reinterpret_cast<const bf16*>(&hi4.x)[1]));
r10 = __builtin_fmaf(a, r10, bw * to_f(reinterpret_cast<const bf16*>(&hi4.y)[0]));
r11 = __builtin_fmaf(a, r11, bw * to_f(reinterpret_cast<const bf16*>(&hi4.y)[1]));
r12 = __builtin_fmaf(a, r12, bw * to_f(reinterpret_cast<const bf16*>(&hi4.z)[0]));
r13 = __builtin_fmaf(a, r13, bw * to_f(reinterpret_cast<const bf16*>(&hi4.z)[1]));
r14 = __builtin_fmaf(a, r14, bw * to_f(reinterpret_cast<const bf16*>(&hi4.w)[0]));
r15 = __builtin_fmaf(a, r15, bw * to_f(reinterpret_cast<const bf16*>(&hi4.w)[1]));
}
// Normalize and write final output
float inv = ll > 0.0f ? (1.0f / ll) : 0.0f;
bf16* odst = &out[out_off(qi, s2_h, vd_base)];
// Vectorized 128-bit stores via uint4
uint4 out_lo, out_hi;
reinterpret_cast<bf16*>(&out_lo.x)[0] = to_b(r0 * inv);
reinterpret_cast<bf16*>(&out_lo.x)[1] = to_b(r1 * inv);
reinterpret_cast<bf16*>(&out_lo.y)[0] = to_b(r2 * inv);
reinterpret_cast<bf16*>(&out_lo.y)[1] = to_b(r3 * inv);
reinterpret_cast<bf16*>(&out_lo.z)[0] = to_b(r4 * inv);
reinterpret_cast<bf16*>(&out_lo.z)[1] = to_b(r5 * inv);
reinterpret_cast<bf16*>(&out_lo.w)[0] = to_b(r6 * inv);
reinterpret_cast<bf16*>(&out_lo.w)[1] = to_b(r7 * inv);
reinterpret_cast<bf16*>(&out_hi.x)[0] = to_b(r8 * inv);
reinterpret_cast<bf16*>(&out_hi.x)[1] = to_b(r9 * inv);
reinterpret_cast<bf16*>(&out_hi.y)[0] = to_b(r10 * inv);
reinterpret_cast<bf16*>(&out_hi.y)[1] = to_b(r11 * inv);
reinterpret_cast<bf16*>(&out_hi.z)[0] = to_b(r12 * inv);
reinterpret_cast<bf16*>(&out_hi.z)[1] = to_b(r13 * inv);
reinterpret_cast<bf16*>(&out_hi.w)[0] = to_b(r14 * inv);
reinterpret_cast<bf16*>(&out_hi.w)[1] = to_b(r15 * inv);
*reinterpret_cast<uint4*>(odst) = out_lo;
*reinterpret_cast<uint4*>(odst + 8) = out_hi;
// Self-reset done_count for next invocation
if (tid == 0) __done_count[b] = 0;
}
}
// Fast S2: 1 head per block, contiguous V reads for vectorized loads.
// Grid: (bs, NUM_HEADS), Block: 64. Best for low-ns tails where launch cost matters.
__global__ __launch_bounds__(64, 16)
void mla_s2_fast(
const float* __restrict__ pm,
const float* __restrict__ pl,
const bf16* __restrict__ po,
bf16* __restrict__ out,
const int32_t* __restrict__ qo,
int bs,
int ns)
{
const int b = blockIdx.x;
const int h = blockIdx.y;
if (b >= bs) return;
const int lane = threadIdx.x;
const int qs = qo[b];
const int qe = qo[b + 1];
if (qe - qs != 1) return;
const int qi = qs;
const int vd_base = lane * V_ELEMS;
float mm = NEG_INF;
float ll = 0.0f;
float r[V_ELEMS];
#pragma unroll
for (int j = 0; j < V_ELEMS; ++j) r[j] = 0.0f;
const int po_bh = ((b * NUM_HEADS + h) * ns) * D_V;
for (int s = 0; s < ns; ++s) {
float m = pm[(b * NUM_HEADS + h) * ns + s];
float l = pl[(b * NUM_HEADS + h) * ns + s];
float mn = fmaxf(mm, m);
float a = fast_exp(mm - mn);
float bw = fast_exp(m - mn) * l;
ll = a * ll + bw;
mm = mn;
const bf16* po_ptr = &po[po_bh + s * D_V + vd_base];
#pragma unroll
for (int j = 0; j < V_ELEMS; ++j) {
r[j] = a * r[j] + bw * to_f(po_ptr[j]);
}
}
float inv = ll > 0.0f ? (1.0f / ll) : 0.0f;
#pragma unroll
for (int j = 0; j < V_ELEMS; ++j) {
out[out_off(qi, h, vd_base + j)] = to_b(r[j] * inv);
}
}
__global__ __launch_bounds__(512, 2)
void mla_s2_vpar2(
const float* __restrict__ pm,
const float* __restrict__ pl,
const bf16* __restrict__ po,
bf16* __restrict__ out,
const int32_t* __restrict__ qo,
int bs,
int ns)
{
const int b = blockIdx.x;
const int h = blockIdx.y;
const int vb = blockIdx.z;
if (b >= bs) return;
const int tid = threadIdx.x;
const int group = tid / 64;
const int lane = tid % 64;
const int num_groups = blockDim.x / 64;
const int qs = qo[b];
if (qo[b + 1] - qs != 1) return;
const int qi = qs;
const int vd_base = vb * 128 + lane * 2;
const int bh_ns = (b * NUM_HEADS + h) * ns;
const int splits_per_group = (ns + num_groups - 1) / num_groups;
const int s_start = min(group * splits_per_group, ns);
const int s_end = min(s_start + splits_per_group, ns);
float mm = NEG_INF;
float ll = 0.0f;
float r0 = 0.0f, r1 = 0.0f;
for (int s = s_start; s < s_end; ++s) {
float m = pm[bh_ns + s];
float l = pl[bh_ns + s];
float mn = fmaxf(mm, m);
float a = fast_exp(mm - mn);
float bw = fast_exp(m - mn) * l;
ll = a * ll + bw;
mm = mn;
uint32_t packed = *reinterpret_cast<const uint32_t*>(
&po[bh_ns * D_V + s * D_V + vd_base]);
r0 = a * r0 + bw * to_f(*reinterpret_cast<const bf16*>(&packed));
r1 = a * r1 + bw * to_f(*(reinterpret_cast<const bf16*>(&packed) + 1));
}
__shared__ float smem[8][64][4];
smem[group][lane][0] = mm;
smem[group][lane][1] = ll;
smem[group][lane][2] = r0;
smem[group][lane][3] = r1;
__syncthreads();
if (group == 0) {
for (int g = 1; g < num_groups; ++g) {
if (g * splits_per_group >= ns) break;
float m = smem[g][lane][0];
float l = smem[g][lane][1];
float mn = fmaxf(mm, m);
float a = fast_exp(mm - mn);
float bv = fast_exp(m - mn);
ll = a * ll + bv * l;
mm = mn;
r0 = a * r0 + bv * smem[g][lane][2];
r1 = a * r1 + bv * smem[g][lane][3];
}
float inv = ll > 0.0f ? (1.0f / ll) : 0.0f;
out[out_off(qi, h, vd_base)] = to_b(r0 * inv);
out[out_off(qi, h, vd_base + 1)] = to_b(r1 * inv);
}
}
__global__ __launch_bounds__(64, 4)
void mla_s2_lean(
const float* __restrict__ pm,
const float* __restrict__ pl,
const bf16* __restrict__ po,
bf16* __restrict__ out,
const int32_t* __restrict__ qo,
int bs,
int ns)
{
const int b = blockIdx.x;
const int h = blockIdx.y;
const int vb = blockIdx.z;
if (b >= bs) return;
const int lane = threadIdx.x;
const int qs = qo[b];
if (qo[b + 1] - qs != 1) return;
const int qi = qs;
const int vd_base = vb * 128 + lane * 2;
const int bh_ns = (b * NUM_HEADS + h) * ns;
float mm = NEG_INF;
float ll = 0.0f;
float r0 = 0.0f, r1 = 0.0f;
for (int s = 0; s < ns; ++s) {
float m = pm[bh_ns + s];
float l = pl[bh_ns + s];
float mn = fmaxf(mm, m);
float a = fast_exp(mm - mn);
float bw = fast_exp(m - mn) * l;
ll = a * ll + bw;
mm = mn;
uint32_t packed = *reinterpret_cast<const uint32_t*>(
&po[bh_ns * D_V + s * D_V + vd_base]);
r0 = a * r0 + bw * to_f(*reinterpret_cast<const bf16*>(&packed));
r1 = a * r1 + bw * to_f(*(reinterpret_cast<const bf16*>(&packed) + 1));
}
float inv = ll > 0.0f ? (1.0f / ll) : 0.0f;
out[out_off(qi, h, vd_base)] = to_b(r0 * inv);
out[out_off(qi, h, vd_base + 1)] = to_b(r1 * inv);
}
} // namespace mla
// Raw C-exported launch functions for ctypes fast path (bypasses pybind11 overhead)
extern "C" {
void launch_iq2_s2_light_raw(
const void* q_bf16,
const void* kv, const void* kv_scale,
const void* qo, const void* kvi,
void* pm, void* pl, void* po,
void* out,
int bs, int ns)
{
if (ns == 1) {
// Single-split decode writes final output directly; no reduction bookkeeping needed.
hipLaunchKernelGGL(
(mla::mla_s1_light_iq2<true, false>),
dim3(bs), dim3(mla::DUET_BLOCK_THREADS), 0, 0,
reinterpret_cast<const uint16_t*>(q_bf16),
reinterpret_cast<const uint8_t*>(kv), reinterpret_cast<const float*>(kv_scale),
reinterpret_cast<const int32_t*>(qo),
reinterpret_cast<const int32_t*>(kvi),
reinterpret_cast<float*>(pm),
reinterpret_cast<float*>(pl),
reinterpret_cast<mla::bf16*>(po),
reinterpret_cast<mla::bf16*>(out),
bs, 1);
return;
}
if (ns > 1 && ns <= 4) {
// For small split counts, a separate 1-wave S2 beats the fused last-CTA tail.
hipLaunchKernelGGL(
(mla::mla_s1_light_iq2<false, false>),
dim3(bs * ns), dim3(mla::DUET_BLOCK_THREADS), 0, 0,
reinterpret_cast<const uint16_t*>(q_bf16),
reinterpret_cast<const uint8_t*>(kv), reinterpret_cast<const float*>(kv_scale),
reinterpret_cast<const int32_t*>(qo),
reinterpret_cast<const int32_t*>(kvi),
reinterpret_cast<float*>(pm),
reinterpret_cast<float*>(pl),
reinterpret_cast<mla::bf16*>(po),
reinterpret_cast<mla::bf16*>(out),
bs, ns);
hipLaunchKernelGGL(
mla::mla_s2_fast,
dim3(bs, mla::NUM_HEADS), dim3(64), 0, 0,
reinterpret_cast<const float*>(pm),
reinterpret_cast<const float*>(pl),
reinterpret_cast<const mla::bf16*>(po),
reinterpret_cast<mla::bf16*>(out),
reinterpret_cast<const int32_t*>(qo),
bs, ns);
return;
}
if (ns > 4) {
// Higher split counts favor a standalone reduction kernel over last-CTA fusion.
hipLaunchKernelGGL(
(mla::mla_s1_light_iq2<false, false>),
dim3(bs * ns), dim3(mla::DUET_BLOCK_THREADS), 0, 0,
reinterpret_cast<const uint16_t*>(q_bf16),
reinterpret_cast<const uint8_t*>(kv), reinterpret_cast<const float*>(kv_scale),
reinterpret_cast<const int32_t*>(qo),
reinterpret_cast<const int32_t*>(kvi),
reinterpret_cast<float*>(pm),
reinterpret_cast<float*>(pl),
reinterpret_cast<mla::bf16*>(po),
reinterpret_cast<mla::bf16*>(out),
bs, ns);
if (bs <= 8) {
const int vpar_block = min(ns, 8) * 64;
hipLaunchKernelGGL(
mla::mla_s2_vpar2,
dim3(bs, mla::NUM_HEADS, mla::D_V / 128), dim3(vpar_block), 0, 0,
reinterpret_cast<const float*>(pm),
reinterpret_cast<const float*>(pl),
reinterpret_cast<const mla::bf16*>(po),
reinterpret_cast<mla::bf16*>(out),
reinterpret_cast<const int32_t*>(qo),
bs, ns);
} else {
hipLaunchKernelGGL(
mla::mla_s2_lean,
dim3(bs, mla::NUM_HEADS, mla::D_V / 128), dim3(64), 0, 0,
reinterpret_cast<const float*>(pm),
reinterpret_cast<const float*>(pl),
reinterpret_cast<const mla::bf16*>(po),
reinterpret_cast<mla::bf16*>(out),
reinterpret_cast<const int32_t*>(qo),
bs, ns);
}
return;
}
// Fallback fused S1+S2 path, with only the done_count tail memset.
const size_t pm_head_floats = static_cast<size_t>(bs) * mla::NUM_HEADS * ns;
hipMemsetAsync(
reinterpret_cast<int*>(reinterpret_cast<float*>(pm) + pm_head_floats),
0,
static_cast<size_t>(bs) * sizeof(int),
0);
hipLaunchKernelGGL(
(mla::mla_s1_light_iq2<false, true>),
dim3(bs * ns), dim3(mla::DUET_BLOCK_THREADS), 0, 0,
reinterpret_cast<const uint16_t*>(q_bf16),
reinterpret_cast<const uint8_t*>(kv), reinterpret_cast<const float*>(kv_scale),
reinterpret_cast<const int32_t*>(qo),
reinterpret_cast<const int32_t*>(kvi),
reinterpret_cast<float*>(pm),
reinterpret_cast<float*>(pl),
reinterpret_cast<mla::bf16*>(po),
reinterpret_cast<mla::bf16*>(out),
bs, ns);
}
} // extern "C"
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
// All dispatch goes through ctypes; pybind exports are unused
}
"""
@lru_cache(maxsize=1)
def _ext():
return load_inline(
name="mla_fp8hw_quacksoftract_nosleep",
cpp_sources="",
cuda_sources=HIP_SRC,
functions=None,
extra_cuda_cflags=[
"-O3", "-ffast-math", "--offload-arch=gfx950",
"-mllvm", "-enable-post-misched=1",
"-mllvm", "--lsr-drop-solution=1",
"-mllvm", "-amdgpu-early-inline-all=true",
"-mllvm", "-amdgpu-function-calls=false",
"-mllvm", "-amdgpu-max-memory-clause=64",
],
with_cuda=True,
verbose=False,
)
def _hip_src_dsread2ns1():
src = HIP_SRC
src = src.replace(
"using bf16 = hip_bfloat16;\n",
"using bf16 = hip_bfloat16;\n"
"using floatx2 = float __attribute__((ext_vector_type(2)));\n",
1,
)
src = src.replace(
"using intx4 = int __attribute__((ext_vector_type(4)));\n",
"using intx4 = int __attribute__((ext_vector_type(4)));\n"
"using intx2 = int __attribute__((ext_vector_type(2)));\n",
1,
)
src = src.replace(
"__device__ __forceinline__ void st16u(uint8_t* p, U128 v) {\n"
" *reinterpret_cast<U128*>(p) = v;\n"
"}\n",
"__device__ __forceinline__ void st16u(uint8_t* p, U128 v) {\n"
" *reinterpret_cast<U128*>(p) = v;\n"
"}\n"
"\n"
"__device__ __forceinline__ floatx2 ds_read2_f32_pair(const float* p)\n"
"{\n"
"#define __LDS_ADDR __attribute__((address_space(3)))\n"
"#pragma clang diagnostic push\n"
"#pragma clang diagnostic ignored \"-Wold-style-cast\"\n"
" const auto p_lds = (__LDS_ADDR const float*)(p);\n"
"#pragma clang diagnostic pop\n"
" intx2 bits;\n"
" const uint32_t addr = static_cast<uint32_t>(reinterpret_cast<uintptr_t>(p_lds));\n"
" asm volatile(\n"
" \"ds_read2_b32 %0, %1 offset0:0 offset1:16\"\n"
" : \"=v\"(bits)\n"
" : \"v\"(addr));\n"
" return *reinterpret_cast<floatx2*>(&bits);\n"
"#undef __LDS_ADDR\n"
"}\n",
1,
)
src = src.replace(
" float hm0 = half_max[0][row_base + 0];\n"
" float hm1 = half_max[1][row_base + 0];\n"
" float hm2 = half_max[0][row_base + 1];\n"
" float hm3 = half_max[1][row_base + 1];\n"
" float hm4 = half_max[0][row_base + 2];\n"
" float hm5 = half_max[1][row_base + 2];\n"
" float hm6 = half_max[0][row_base + 3];\n"
" float hm7 = half_max[1][row_base + 3];\n",
" float hm0, hm1, hm2, hm3, hm4, hm5, hm6, hm7;\n"
" if constexpr (DIRECT_OUT) {\n"
" const floatx2 hm01 = ds_read2_f32_pair(&half_max[0][row_base + 0]);\n"
" const floatx2 hm23 = ds_read2_f32_pair(&half_max[0][row_base + 1]);\n"
" const floatx2 hm45 = ds_read2_f32_pair(&half_max[0][row_base + 2]);\n"
" const floatx2 hm67 = ds_read2_f32_pair(&half_max[0][row_base + 3]);\n"
" hm0 = hm01.x; hm1 = hm01.y;\n"
" hm2 = hm23.x; hm3 = hm23.y;\n"
" hm4 = hm45.x; hm5 = hm45.y;\n"
" hm6 = hm67.x; hm7 = hm67.y;\n"
" } else {\n"
" hm0 = half_max[0][row_base + 0];\n"
" hm1 = half_max[1][row_base + 0];\n"
" hm2 = half_max[0][row_base + 1];\n"
" hm3 = half_max[1][row_base + 1];\n"
" hm4 = half_max[0][row_base + 2];\n"
" hm5 = half_max[1][row_base + 2];\n"
" hm6 = half_max[0][row_base + 3];\n"
" hm7 = half_max[1][row_base + 3];\n"
" }\n",
1,
)
src = src.replace(
" local_sum[h] = alpha[h] * local_sum[h]\n"
" + half_sum[0][row] + half_sum[1][row];\n",
" if constexpr (DIRECT_OUT) {\n"
" const floatx2 hs = ds_read2_f32_pair(&half_sum[0][row]);\n"
" local_sum[h] = alpha[h] * local_sum[h] + hs.x + hs.y;\n"
" } else {\n"
" local_sum[h] = alpha[h] * local_sum[h]\n"
" + half_sum[0][row] + half_sum[1][row];\n"
" }\n",
1,
)
return src
@lru_cache(maxsize=1)
def _ext_dsread2ns1():
return load_inline(
name="mla_fp8hw_quacksoftract_nosleep_dualns1",
cpp_sources="",
cuda_sources=_hip_src_dsread2ns1(),
functions=None,
extra_cuda_cflags=[
"-O3", "-ffast-math", "--offload-arch=gfx950",
"-mllvm", "-enable-post-misched=1",
"-mllvm", "--lsr-drop-solution=1",
"-mllvm", "-amdgpu-early-inline-all=true",
"-mllvm", "-amdgpu-function-calls=false",
"-mllvm", "-amdgpu-max-memory-clause=64",
],
with_cuda=True,
verbose=False,
)
def _get_bufs(batch_size, ns, device):
# pm has extra space at tail for done_count (bs ints) used by fused S2
# Layout: [bs*NUM_HEADS*ns floats] + [bs ints padded as floats].
# The payload is fully overwritten in-kernel; only the tail counter is zeroed by HIP.
pm_extra = batch_size # sizeof(int)==sizeof(float)==4, need bs slots
return (
torch.empty((batch_size * NUM_HEADS * ns + pm_extra,), device=device, dtype=torch.float32),
torch.empty((batch_size, NUM_HEADS, ns), device=device, dtype=torch.float32),
torch.empty(
(batch_size, NUM_HEADS, ns, V_HEAD_DIM),
device=device, dtype=torch.bfloat16),
)
def _get_out(batch_size, device):
return torch.empty(
(batch_size, NUM_HEADS, V_HEAD_DIM),
device=device, dtype=torch.bfloat16)
_SPLIT_TABLE_256CU = {
(4, 1024): 32,
(4, 8192): 64,
(32, 1024): 8,
(32, 8192): 8,
(64, 1024): 4,
(64, 8192): 4,
(256, 1024): 1,
(256, 8192): 2,
}
# MI355X has ~512 CUs — need higher splits to fill them
_SPLIT_TABLE_512CU = {
(4, 1024): 32,
(4, 8192): 64,
(32, 1024): 8,
(32, 8192): 8,
(64, 1024): 4,
(64, 8192): 4,
(256, 1024): 1,
(256, 8192): 2,
}
def _detect_split_table():
try:
props = torch.cuda.get_device_properties(0)
cu_count = props.multi_processor_count
if cu_count >= 400:
return _SPLIT_TABLE_512CU
except Exception:
pass
return _SPLIT_TABLE_256CU
_SPLIT_TABLE = _detect_split_table()
def _pick_splits_duet(bs, kvl):
key = (bs, kvl)
if key in _SPLIT_TABLE:
return _SPLIT_TABLE[key]
if kvl <= 2048:
ns_cu = max(1, 256 // bs)
ns_kv = max(1, kvl // 512)
ns = max(ns_cu, ns_kv)
ns = min(ns, max(1, kvl // 128))
p = 1
while p * 2 <= ns:
p *= 2
return max(1, min(p, 64))
ns_cu = min(max(1, 512 // bs), 32)
ns_cu = min(ns_cu, kvl // 32) if kvl >= 32 else 1
kv_div = 2048 if bs >= 64 else 1024
ns_kv = max(1, kvl // kv_div)
ns = max(ns_cu, ns_kv)
p = 1
while p * 2 <= ns:
p *= 2
return max(1, min(p, 64))
import ctypes as _ct
_cached_ext = None
_raw_iq2_s2 = None
_cached_ext_dsread2ns1 = None
_raw_iq2_s2_dsread2ns1 = None
def _init_raw_dispatch(ext):
"""Load the compiled .so with ctypes for fast kernel dispatch."""
so = ext.__file__
lib = _ct.CDLL(so)
P = _ct.c_void_p
lib.launch_iq2_s2_light_raw.restype = None
lib.launch_iq2_s2_light_raw.argtypes = [
P, P, P, P, P, P, P, P, P, _ct.c_int, _ct.c_int]
return lib.launch_iq2_s2_light_raw
def custom_kernel(data):
global _cached_ext, _raw_iq2_s2
global _cached_ext_dsread2ns1, _raw_iq2_s2_dsread2ns1
q, kv_data, qo_indptr, kv_indptr, config = data
bs = int(config["batch_size"])
kvl = int(config["kv_seq_len"])
ns = _SPLIT_TABLE.get((bs, kvl)) or _pick_splits_duet(bs, kvl)
if _raw_iq2_s2 is None:
_cached_ext = _ext()
_raw_iq2_s2 = _init_raw_dispatch(_cached_ext)
pm, pl, po = _get_bufs(bs, ns, q.device)
out = _get_out(bs, q.device)
q_c = q.contiguous()
kv_t = kv_data["fp8"][0].contiguous()
sc_t = kv_data["fp8"][1].float()
qo_c = qo_indptr.contiguous()
kvi_c = kv_indptr.contiguous()
raw_fn = _raw_iq2_s2
if ns == 1:
if _raw_iq2_s2_dsread2ns1 is None:
_cached_ext_dsread2ns1 = _ext_dsread2ns1()
_raw_iq2_s2_dsread2ns1 = _init_raw_dispatch(_cached_ext_dsread2ns1)
raw_fn = _raw_iq2_s2_dsread2ns1
raw_fn(
q_c.data_ptr(),
kv_t.data_ptr(),
sc_t.data_ptr(),
qo_c.data_ptr(),
kvi_c.data_ptr(),
pm.data_ptr(),
pl.data_ptr(),
po.data_ptr(),
out.data_ptr(),
bs,
ns,
)
return out
scrolls · 2110 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON