submission 589095
div22 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1330 lines, June 9 Researcher Reciprocity License v1.0.
solution_109.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-589095?include=source"interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, mxfp4
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:dadf36afa196c062e24ec08ec2d4dd9df4a56fee22510d4fd048a2be16c315b5
license declaredunknown
license concludedunknown
authorsdiv22
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
const auto k_byte_start = ks_start * 128; // in bf16 elements (128 per k-tile = 64 bytes FP4 = 256 bf16 bytes for 128 elements)shared-memory
extern __shared__ uint8_t smem_raw[];split-k
template<int BM, int BN, int NWARPS, bool SPLITK, bool A_VALID, bool B_VALID,tile-m = 0
constexpr auto full_m = (BM >= 16) ? C_M / BM : 0;tile-n = 64
constexpr auto WAVES_N = 4; // BN=64, A in registers — more N-reuseKernel source
solution_109.py1330 lines
"""
Solution 8: Fused quant+GEMM for PATH A (K=512) shapes.
Based on solution_3 (11.642μs). Single kernel launch for 3 of 6 shapes.
"""
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
from typing import Tuple
import torch
from torch.utils.cpp_extension import load_inline
import uuid
HIP_KERNEL = r"""
#include <hip/hip_runtime.h>
#include <stdint.h>
using int4_v = int __attribute__((ext_vector_type(4)));
using float4_v = float __attribute__((ext_vector_type(4)));
using bf16x2 = __bf16 __attribute__((ext_vector_type(2)));
static constexpr auto FP4_E2M1 = 4;
__device__ float4_v __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
int4_v a, int4_v b, float4_v c,
int cbsz, int blgp, int op_sel_a, int scale_a, int op_sel_b, int scale_b
) __asm("llvm.amdgcn.mfma.scale.f32.16x16x128.f8f6f4.v4i32.v4i32");
__device__ __forceinline__ auto load16(const uint8_t* __restrict__ p) {
return *reinterpret_cast<const int4_v*>(p);
}
__device__ __forceinline__ auto load16_nt(const uint8_t* __restrict__ p) {
return __builtin_nontemporal_load(reinterpret_cast<const int4_v*>(p));
}
__device__ __forceinline__ auto float_to_bf16(float f) {
bf16x2 v;
v[0] = static_cast<__bf16>(f);
auto r = uint16_t{};
__builtin_memcpy(&r, &v, sizeof(r));
return r;
}
template<int C_M, int C_K>
__global__ void __launch_bounds__(128, (C_M * (C_K / 32) <= 256) ? 2 : 4)
mxfp4_quant(
const __bf16* __restrict__ A_bf16,
uint8_t* __restrict__ A_fp4,
uint8_t* __restrict__ A_scale)
{
constexpr auto KS = C_K / 32;
constexpr auto K2 = C_K / 2;
const auto group = blockIdx.x * 128 + threadIdx.x;
const auto row = group / KS;
const auto kg = group % KS;
if (row >= C_M) return;
const auto row_k = (long)row * C_K;
const auto* src = A_bf16 + row_k + kg * 32;
const auto w0 = *reinterpret_cast<const int4_v*>(src);
const auto w1 = *reinterpret_cast<const int4_v*>(src + 8);
const auto w2 = *reinterpret_cast<const int4_v*>(src + 16);
const auto w3 = *reinterpret_cast<const int4_v*>(src + 24);
const auto* p0 = reinterpret_cast<const uint32_t*>(&w0);
const auto* p1 = reinterpret_cast<const uint32_t*>(&w1);
const auto* p2 = reinterpret_cast<const uint32_t*>(&w2);
const auto* p3 = reinterpret_cast<const uint32_t*>(&w3);
auto absMax = 1e-10f;
#define AMAX_PAIR(pair) { \
const auto lo = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair)); \
const auto hi = __builtin_bit_cast(__bf16, static_cast<uint16_t>((pair) >> 16)); \
const auto flo = __builtin_elementwise_abs(static_cast<float>(lo)); \
const auto fhi = __builtin_elementwise_abs(static_cast<float>(hi)); \
absMax = (flo > absMax) ? flo : absMax; \
absMax = (fhi > absMax) ? fhi : absMax; \
}
AMAX_PAIR(p0[0]) AMAX_PAIR(p0[1]) AMAX_PAIR(p0[2]) AMAX_PAIR(p0[3])
AMAX_PAIR(p1[0]) AMAX_PAIR(p1[1]) AMAX_PAIR(p1[2]) AMAX_PAIR(p1[3])
AMAX_PAIR(p2[0]) AMAX_PAIR(p2[1]) AMAX_PAIR(p2[2]) AMAX_PAIR(p2[3])
AMAX_PAIR(p3[0]) AMAX_PAIR(p3[1]) AMAX_PAIR(p3[2]) AMAX_PAIR(p3[3])
#undef AMAX_PAIR
const auto u32 = __builtin_bit_cast(uint32_t, absMax);
const auto amax_exp = ((u32 + 0x200000u) >> 23) & 0xFFu;
const auto inv_exp = (amax_exp >= 2u) ? (amax_exp - 2u) : 0u;
A_scale[row_k / 32 + kg] = static_cast<uint8_t>(inv_exp);
const auto hw_scale = __builtin_bit_cast(float, static_cast<uint32_t>(inv_exp) << 23);
#define CVT(d, pair, sel) \
d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
d, __builtin_bit_cast(bf16x2, (pair)), hw_scale, sel)
auto d0 = 0u, d1 = 0u, d2 = 0u, d3 = 0u;
CVT(d0, p0[0], 0); CVT(d0, p0[1], 1); CVT(d0, p0[2], 2); CVT(d0, p0[3], 3);
CVT(d1, p1[0], 0); CVT(d1, p1[1], 1); CVT(d1, p1[2], 2); CVT(d1, p1[3], 3);
CVT(d2, p2[0], 0); CVT(d2, p2[1], 1); CVT(d2, p2[2], 2); CVT(d2, p2[3], 3);
CVT(d3, p3[0], 0); CVT(d3, p3[1], 1); CVT(d3, p3[2], 2); CVT(d3, p3[3], 3);
#undef CVT
auto* dst = reinterpret_cast<int4_v*>(A_fp4 + row_k / 2 + kg * 16);
*dst = int4_v{static_cast<int>(d0), static_cast<int>(d1),
static_cast<int>(d2), static_cast<int>(d3)};
}
template<bool ALWAYS_VALID>
__device__ __forceinline__ auto load_or_zero(bool rt_valid, const uint8_t* p) {
if constexpr (ALWAYS_VALID) return load16(p);
else return rt_valid ? load16(p) : int4_v{0,0,0,0};
}
template<int WAVES_M, int WAVES_N, int CHUNK_K>
struct LdsLayout {
static constexpr auto A_ROW = CHUNK_K * 64 + 16; // +16B pad per row
static constexpr auto A_TILE = 16 * A_ROW;
static constexpr auto A_SIZE = WAVES_M * A_TILE;
static constexpr auto B_KGRP_STRIDE = 16 * 16; // 16 lrows × 16B = 256B
static constexpr auto B_KT_STRIDE = 5 * B_KGRP_STRIDE; // 5 slots (4+1 pad)
static constexpr auto B_TILE = CHUNK_K * B_KT_STRIDE;
static constexpr auto B_SIZE = WAVES_N * B_TILE;
static constexpr auto BUF_SIZE = A_SIZE + B_SIZE;
static constexpr auto TOTAL_LDS = 2 * BUF_SIZE;
__device__ __forceinline__ static constexpr auto a_off(uint8_t* const buf) { return buf; }
__device__ __forceinline__ static constexpr auto b_off(uint8_t* const buf) { return buf + A_SIZE; }
__device__ __forceinline__ static constexpr auto a_idx(const auto wm, const auto row, const auto k_idx) {
return wm * A_TILE + (row ^ (k_idx & 7)) * A_ROW + k_idx * 16;
}
__device__ __forceinline__ static constexpr auto b_idx(const auto wn, const auto kt, const auto kgrp, const auto lrow) {
return wn * B_TILE + kt * B_KT_STRIDE
+ (kgrp ^ (lrow >> 2)) * B_KGRP_STRIDE + lrow * 16;
}
};
template<int WAVES_M, int WAVES_N, int CHUNK_K, int NWARPS>
struct LoadCounts {
static constexpr auto NTHREADS = NWARPS * 64;
static constexpr auto A_TOTAL = WAVES_M * 16 * CHUNK_K * 4;
static constexpr auto B_TOTAL = WAVES_N * CHUNK_K * 4 * 16;
static constexpr auto A_PER_THREAD = (A_TOTAL + NTHREADS - 1) / NTHREADS;
static constexpr auto B_PER_THREAD = (B_TOTAL + NTHREADS - 1) / NTHREADS;
};
template<int BM, int BN, int NWARPS, bool SPLITK, bool A_VALID, bool B_VALID,
int CKT, int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_KPS,
int C_M, int C_TILE_OFF_X, int C_TILE_OFF_Y>
__global__ void __launch_bounds__(NWARPS * 64, (NWARPS <= 2) ? 4 : 2)
__attribute__((amdgpu_flat_work_group_size(NWARPS * 64, NWARPS * 64)))
mxfp4_gemm(
const uint8_t* __restrict__ A,
const uint8_t* __restrict__ As,
const uint8_t* __restrict__ Bsh,
const uint8_t* __restrict__ Bssh,
float* __restrict__ C_partial,
uint16_t* __restrict__ C_final)
{
static_assert(BN % 16 == 0);
constexpr auto WAVES_M = (BM + 15) / 16;
constexpr auto WAVES_N = BN / 16;
static_assert(WAVES_M * WAVES_N == NWARPS);
constexpr auto M = C_M;
constexpr auto N = C_N;
constexpr auto K = C_K;
constexpr auto scaleN = C_SCALEN;
constexpr auto K2 = K / 2;
constexpr auto KS = K / 32;
constexpr auto bsh_n_stride = (long)(K / 64) * 512;
const auto ks_idx = static_cast<int>(blockIdx.z);
const auto lane = static_cast<int>(threadIdx.x % 64);
const auto wave = static_cast<int>(threadIdx.x / 64);
const auto wave_m = wave / WAVES_N;
const auto wave_n = wave % WAVES_N;
const auto tid = static_cast<int>(threadIdx.x);
const auto tile_m_base = (static_cast<int>(blockIdx.y) + C_TILE_OFF_Y) * BM;
const auto tile_n_base = (static_cast<int>(blockIdx.x) + C_TILE_OFF_X) * BN;
const auto tile_m = tile_m_base + wave_m * 16;
const auto tile_n = tile_n_base + wave_n * 16;
constexpr auto ktiles_per_split = C_KPS;
const auto ks_start = ks_idx * ktiles_per_split;
const auto lrow = lane % 16;
const auto kgrp = lane / 16;
__builtin_assume(lrow >= 0 && lrow < 16);
__builtin_assume(kgrp >= 0 && kgrp < 4);
const auto gm = tile_m + lrow;
const auto gn = tile_n + lrow;
const auto a_rt = A_VALID | (gm < M);
const auto b_rt = B_VALID | (gn < N);
const uint8_t* as_row = nullptr;
if constexpr (A_VALID) {
as_row = As + (long)gm * KS;
} else {
if (a_rt) as_row = As + (long)gm * KS;
}
auto bssh_base = 0;
if constexpr (B_VALID) {
bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);
} else {
if (b_rt) bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);
}
float4_v acc{0.f, 0.f, 0.f, 0.f};
// PATH A: CKT <= 4 — Direct global loads, no LDS
if constexpr (CKT <= 4) {
if (tile_m >= M || tile_n >= N) return;
const auto n_tile = tile_n / 16;
const auto bsh_lane_base = Bsh + (long)n_tile * bsh_n_stride + (long)lrow * 16;
const auto k_half_off = (kgrp & 1) * 256;
const auto k_blk_base = kgrp >> 1;
constexpr auto a_kt_stride = 64L;
constexpr auto bsh_kt_stride = 1024L;
const uint8_t* a_ptr = nullptr;
const uint8_t* bsh_ptr = nullptr;
const uint8_t* bssh_ptr = nullptr;
if constexpr (A_VALID) {
a_ptr = A + (long)gm * K2 + (long)ks_start * a_kt_stride + kgrp * 16;
} else {
if (a_rt) a_ptr = A + (long)gm * K2 + (long)ks_start * a_kt_stride + kgrp * 16;
}
if constexpr (B_VALID) {
bsh_ptr = bsh_lane_base + (long)(ks_start * 2 + k_blk_base) * 512 + k_half_off;
bssh_ptr = Bssh + bssh_base + (ks_start & 1) * 2 + (ks_start >> 1) * 256;
} else {
if (b_rt) {
bsh_ptr = bsh_lane_base + (long)(ks_start * 2 + k_blk_base) * 512 + k_half_off;
bssh_ptr = Bssh + bssh_base + (ks_start & 1) * 2 + (ks_start >> 1) * 256;
}
}
const auto bssh_step0 = (ks_start & 1) ? 254 : 2;
const auto bssh_step1 = 256 - bssh_step0;
// PATH A: A from regular load (small, L1 reuse), B from .cg (large, L1 bypass)
#define DO_MFMA(a_off, b_off, bssh_off, ks_val) \
{ \
const auto av = load_or_zero<A_VALID>(a_rt, a_ptr + (a_off) * a_kt_stride); \
const auto bv = load_or_zero<B_VALID>(b_rt, bsh_ptr + (b_off) * bsh_kt_stride); \
const auto ks = (ks_val) * 4 + kgrp; \
auto sa = 0, sb = 0; \
if constexpr (A_VALID) sa = static_cast<int>(as_row[ks]); \
else sa = (a_rt & (ks < KS)) ? static_cast<int>(as_row[ks]) : 127; \
if constexpr (B_VALID) sb = static_cast<int>(*(bssh_ptr + (bssh_off))); \
else sb = (b_rt & (ks < KS)) ? static_cast<int>(*(bssh_ptr + (bssh_off))) : 127; \
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(av,bv,acc,FP4_E2M1,FP4_E2M1,0,sa,0,sb); \
}
static_assert(CKT > 0);
#pragma unroll
for (auto q = 0; q < (CKT / 4); ++q) {
DO_MFMA(0, 0, 0, ks_start + q*4)
DO_MFMA(1, 1, bssh_step0, ks_start + q*4 + 1)
DO_MFMA(2, 2, bssh_step0 + bssh_step1, ks_start + q*4 + 2)
DO_MFMA(3, 3, bssh_step0 + bssh_step1 + bssh_step0, ks_start + q*4 + 3)
a_ptr += 4 * a_kt_stride;
bsh_ptr += 4 * bsh_kt_stride;
bssh_ptr += 512;
}
if constexpr ((CKT % 4) >= 2) {
DO_MFMA(0, 0, 0, ks_start + (CKT/4)*4)
DO_MFMA(1, 1, bssh_step0, ks_start + (CKT/4)*4 + 1)
a_ptr += 2 * a_kt_stride;
bsh_ptr += 2 * bsh_kt_stride;
bssh_ptr += 256;
}
if constexpr ((CKT % 2) == 1) {
DO_MFMA(0, 0, 0, ks_start + CKT - 1)
}
#undef DO_MFMA
// PATH B: CKT > 4 — LDS double-buffered with register-staged pipelining
} else {
constexpr auto CHUNK_K = 4;
constexpr auto NUM_CHUNKS = CKT / CHUNK_K;
constexpr auto TAIL_KT = CKT % CHUNK_K;
using Lds = LdsLayout<WAVES_M, WAVES_N, CHUNK_K>;
using LC = LoadCounts<WAVES_M, WAVES_N, CHUNK_K, NWARPS>;
extern __shared__ uint8_t smem_raw[];
auto* buf0 = smem_raw;
auto* buf1 = smem_raw + Lds::BUF_SIZE;
const auto tile_valid = (tile_m < M) && (tile_n < N);
int a_row[LC::A_PER_THREAD];
int a_k_idx[LC::A_PER_THREAD];
int a_lds_offs[LC::A_PER_THREAD];
bool a_valid[LC::A_PER_THREAD];
#pragma unroll
for (auto i = 0; i < LC::A_PER_THREAD; ++i) {
auto linear = i * LC::NTHREADS + tid;
if (linear >= LC::A_TOTAL) {
a_valid[i] = false;
a_lds_offs[i] = -1;
} else {
auto k_idx = linear % (CHUNK_K * 4);
auto m_local = (linear / (CHUNK_K * 4)) % 16;
auto wave_m_idx = linear / (16 * CHUNK_K * 4);
a_row[i] = tile_m_base + wave_m_idx * 16 + m_local;
a_k_idx[i] = k_idx;
a_lds_offs[i] = Lds::a_idx(wave_m_idx, m_local, k_idx);
if constexpr (A_VALID) a_valid[i] = tile_valid;
else a_valid[i] = tile_valid && (a_row[i] < M);
}
}
int b_kt_local[LC::B_PER_THREAD]; // kt within chunk (0..CHUNK_K-1)
long b_base_off[LC::B_PER_THREAD]; // global offset without k_blk term
int b_kgrp_half[LC::B_PER_THREAD]; // (b_kgrp / 2) for k_blk calc
int b_lds_kgrp[LC::B_PER_THREAD]; // b_kgrp for LDS index
int b_lds_lrow[LC::B_PER_THREAD]; // b_lrow for LDS index
int b_lds_wn[LC::B_PER_THREAD]; // wave_n_idx for LDS index
int b_lds_offs[LC::B_PER_THREAD];
bool b_valid[LC::B_PER_THREAD];
#pragma unroll
for (auto i = 0; i < LC::B_PER_THREAD; ++i) {
auto linear = i * LC::NTHREADS + tid;
if (linear >= LC::B_TOTAL) {
b_valid[i] = false;
b_lds_offs[i] = -1;
} else {
auto b_lrow = linear % 16;
auto b_kgrp = (linear / 16) % 4;
auto kt = (linear / 64) % CHUNK_K;
auto wave_n_idx = linear / (CHUNK_K * 64);
auto b_tile_n = tile_n_base + wave_n_idx * 16;
auto n_tile = b_tile_n / 16;
auto k_half = (b_kgrp & 1) * 256;
b_kt_local[i] = kt;
b_base_off[i] = (long)n_tile * bsh_n_stride + k_half + (long)b_lrow * 16;
b_kgrp_half[i] = b_kgrp / 2;
b_lds_kgrp[i] = b_kgrp;
b_lds_lrow[i] = b_lrow;
b_lds_wn[i] = wave_n_idx;
b_lds_offs[i] = Lds::b_idx(wave_n_idx, kt, b_kgrp, b_lrow);
if constexpr (B_VALID) b_valid[i] = tile_valid;
else b_valid[i] = tile_valid && (b_tile_n + b_lrow < N);
}
}
int4_v a_regs[LC::A_PER_THREAD];
int4_v b_regs[LC::B_PER_THREAD];
#define ISSUE_LOADS(ks_base) \
{ \
_Pragma("unroll") \
for (auto i = 0; i < LC::A_PER_THREAD; ++i) { \
a_regs[i] = int4_v{0,0,0,0}; \
if (a_valid[i]) { \
auto k_byte = ((ks_base) * 4 + a_k_idx[i]) * 16; \
if constexpr (A_VALID) \
a_regs[i] = load16(A + (long)a_row[i] * K2 + k_byte); \
else if (k_byte + 16 <= K2) \
a_regs[i] = load16(A + (long)a_row[i] * K2 + k_byte); \
} \
} \
_Pragma("unroll") \
for (auto i = 0; i < LC::B_PER_THREAD; ++i) { \
b_regs[i] = int4_v{0,0,0,0}; \
if (b_valid[i]) { \
auto ks = (ks_base) + b_kt_local[i]; \
auto k_blk = ks * 2 + b_kgrp_half[i]; \
auto global_off = b_base_off[i] + (long)k_blk * 512; \
b_regs[i] = load16_nt(Bsh + global_off); \
} \
} \
}
#define STORE_TO_LDS(buf) \
{ \
auto* _sa = Lds::a_off(buf); \
auto* _sb = Lds::b_off(buf); \
_Pragma("unroll") \
for (auto i = 0; i < LC::A_PER_THREAD; ++i) { \
if (a_lds_offs[i] >= 0) \
*reinterpret_cast<int4_v*>(_sa + a_lds_offs[i]) = a_regs[i]; \
} \
_Pragma("unroll") \
for (auto i = 0; i < LC::B_PER_THREAD; ++i) { \
if (b_lds_offs[i] >= 0) \
*reinterpret_cast<int4_v*>(_sb + b_lds_offs[i]) = b_regs[i]; \
} \
}
#define COMPUTE_CHUNK(buf, chunk_ks) \
{ \
/* Compute LDS byte offsets for each k-tile's A and B loads. */ \
/* buf-smem_raw gives the buffer offset (0 or BUF_SIZE). Since the */ \
/* single extern __shared__ allocation starts at LDS offset 0, these */ \
/* integer offsets are the ds_read_b128 address VGPR values directly. */ \
const unsigned _buf_off = static_cast<unsigned>((buf) - smem_raw); \
const unsigned _a_base = _buf_off; \
const unsigned _b_base = _buf_off + static_cast<unsigned>(Lds::A_SIZE); \
const unsigned _addr_a0 = _a_base + static_cast<unsigned>(Lds::a_idx(wave_m, lrow, 0 * 4 + kgrp)); \
const unsigned _addr_b0 = _b_base + static_cast<unsigned>(Lds::b_idx(wave_n, 0, kgrp, lrow)); \
const unsigned _addr_a1 = _a_base + static_cast<unsigned>(Lds::a_idx(wave_m, lrow, 1 * 4 + kgrp)); \
const unsigned _addr_b1 = _b_base + static_cast<unsigned>(Lds::b_idx(wave_n, 1, kgrp, lrow)); \
const unsigned _addr_a2 = _a_base + static_cast<unsigned>(Lds::a_idx(wave_m, lrow, 2 * 4 + kgrp)); \
const unsigned _addr_b2 = _b_base + static_cast<unsigned>(Lds::b_idx(wave_n, 2, kgrp, lrow)); \
const unsigned _addr_a3 = _a_base + static_cast<unsigned>(Lds::a_idx(wave_m, lrow, 3 * 4 + kgrp)); \
const unsigned _addr_b3 = _b_base + static_cast<unsigned>(Lds::b_idx(wave_n, 3, kgrp, lrow)); \
/* Pre-load all 8 scale values */ \
int _sa0, _sb0, _sa1, _sb1, _sa2, _sb2, _sa3, _sb3; \
if constexpr (A_VALID) { \
auto _sa_base = as_row + (chunk_ks) * 4; \
auto _sa_vec = *reinterpret_cast<const int4_v*>(_sa_base); \
auto _sa_shift = kgrp * 8; \
_sa0 = (reinterpret_cast<const uint32_t*>(&_sa_vec)[0] >> _sa_shift) & 0xFF; \
_sa1 = (reinterpret_cast<const uint32_t*>(&_sa_vec)[1] >> _sa_shift) & 0xFF; \
_sa2 = (reinterpret_cast<const uint32_t*>(&_sa_vec)[2] >> _sa_shift) & 0xFF; \
_sa3 = (reinterpret_cast<const uint32_t*>(&_sa_vec)[3] >> _sa_shift) & 0xFF; \
} else { \
auto _ks0 = (chunk_ks) * 4 + kgrp; \
_sa0 = (a_rt & (_ks0 < KS)) ? static_cast<int>(as_row[_ks0]) : 127; \
_sa1 = (a_rt & (_ks0+4 < KS)) ? static_cast<int>(as_row[_ks0+4]) : 127; \
_sa2 = (a_rt & (_ks0+8 < KS)) ? static_cast<int>(as_row[_ks0+8]) : 127; \
_sa3 = (a_rt & (_ks0+12 < KS)) ? static_cast<int>(as_row[_ks0+12]) : 127; \
} \
{ \
auto _ksv0 = (chunk_ks); \
auto* _bp0 = Bssh + bssh_base + (_ksv0 & 1) * 2 + (_ksv0 >> 1) * 256; \
auto* _bp1 = Bssh + bssh_base + ((_ksv0+1) & 1) * 2 + ((_ksv0+1) >> 1) * 256; \
auto* _bp2 = Bssh + bssh_base + ((_ksv0+2) & 1) * 2 + ((_ksv0+2) >> 1) * 256; \
auto* _bp3 = Bssh + bssh_base + ((_ksv0+3) & 1) * 2 + ((_ksv0+3) >> 1) * 256; \
if constexpr (B_VALID) { \
_sb0 = static_cast<int>(*_bp0); \
_sb1 = static_cast<int>(*_bp1); \
_sb2 = static_cast<int>(*_bp2); \
_sb3 = static_cast<int>(*_bp3); \
} else { \
auto _ks0 = (chunk_ks) * 4 + kgrp; \
_sb0 = (b_rt & (_ks0 < KS)) ? static_cast<int>(*_bp0) : 127; \
_sb1 = (b_rt & (_ks0+4 < KS)) ? static_cast<int>(*_bp1) : 127; \
_sb2 = (b_rt & (_ks0+8 < KS)) ? static_cast<int>(*_bp2) : 127; \
_sb3 = (b_rt & (_ks0+12 < KS)) ? static_cast<int>(*_bp3) : 127; \
} \
} \
/* Inline ASM: software-pipelined ds_read_b128 + v_mfma_scale */ \
/* Double-buffered: even set (ae,be) and odd set (ao,bo) */ \
/* Pipeline: */ \
/* Issue loads for kt0,kt1 → wait kt0 → MFMA kt0 */ \
/* Issue loads for kt2 → wait kt1 → MFMA kt1 */ \
/* Issue loads for kt3 → wait kt2 → MFMA kt2 */ \
/* wait kt3 → MFMA kt3 */ \
/* Each MFMA executes for 16 cycles, hiding 2 ds_read_b128 latency */ \
int4_v _ae, _be, _ao, _bo; \
__builtin_amdgcn_sched_barrier(0x020); \
asm volatile( \
"s_setprio 3 \n\t" \
/* --- Issue 4 LDS reads for kt=0 (even) and kt=1 (odd) --- */ \
"ds_read_b128 %[ae], %[aa0] \n\t" \
"ds_read_b128 %[be], %[ab0] \n\t" \
"ds_read_b128 %[ao], %[aa1] \n\t" \
"ds_read_b128 %[bo], %[ab1] \n\t" \
/* Wait for kt=0 data (even set); 2 still in flight (ao,bo) */ \
"s_waitcnt lgkmcnt(2) \n\t" \
/* --- MFMA kt=0: consume even set --- */ \
"v_mfma_scale_f32_16x16x128_f8f6f4 %[acc], %[ae], %[be], %[acc], %[sa0], %[sb0] op_sel_hi:[0,0,0] cbsz:4 blgp:4\n\t" \
/* --- Issue 2 LDS reads for kt=2 (reuse even set) --- */ \
"ds_read_b128 %[ae], %[aa2] \n\t" \
"ds_read_b128 %[be], %[ab2] \n\t" \
/* Wait for kt=1 data (odd set); 2 still in flight (ae,be) */ \
"s_waitcnt lgkmcnt(2) \n\t" \
/* --- MFMA kt=1: consume odd set --- */ \
"v_mfma_scale_f32_16x16x128_f8f6f4 %[acc], %[ao], %[bo], %[acc], %[sa1], %[sb1] op_sel_hi:[0,0,0] cbsz:4 blgp:4\n\t" \
/* --- Issue 2 LDS reads for kt=3 (reuse odd set) --- */ \
"ds_read_b128 %[ao], %[aa3] \n\t" \
"ds_read_b128 %[bo], %[ab3] \n\t" \
/* Wait for kt=2 data (even set); 2 still in flight (ao,bo) */ \
"s_waitcnt lgkmcnt(2) \n\t" \
/* --- MFMA kt=2: consume even set --- */ \
"v_mfma_scale_f32_16x16x128_f8f6f4 %[acc], %[ae], %[be], %[acc], %[sa2], %[sb2] op_sel_hi:[0,0,0] cbsz:4 blgp:4\n\t" \
/* Wait for kt=3 data (odd set); 0 in flight */ \
"s_waitcnt lgkmcnt(0) \n\t" \
/* --- MFMA kt=3: consume odd set --- */ \
"v_mfma_scale_f32_16x16x128_f8f6f4 %[acc], %[ao], %[bo], %[acc], %[sa3], %[sb3] op_sel_hi:[0,0,0] cbsz:4 blgp:4\n\t" \
"s_setprio 0 \n\t" \
: [acc] "+v"(acc), \
[ae] "=&v"(_ae), [be] "=&v"(_be), \
[ao] "=&v"(_ao), [bo] "=&v"(_bo) \
: [aa0] "v"(_addr_a0), [ab0] "v"(_addr_b0), \
[aa1] "v"(_addr_a1), [ab1] "v"(_addr_b1), \
[aa2] "v"(_addr_a2), [ab2] "v"(_addr_b2), \
[aa3] "v"(_addr_a3), [ab3] "v"(_addr_b3), \
[sa0] "v"(_sa0), [sb0] "v"(_sb0), \
[sa1] "v"(_sa1), [sb1] "v"(_sb1), \
[sa2] "v"(_sa2), [sb2] "v"(_sb2), \
[sa3] "v"(_sa3), [sb3] "v"(_sb3) \
); \
__builtin_amdgcn_sched_barrier(0x020); \
}
auto cur_ks = ks_start;
ISSUE_LOADS(cur_ks);
asm volatile("s_waitcnt vmcnt(0) lgkmcnt(0)" ::: "memory");
STORE_TO_LDS(buf0);
__syncthreads();
auto* cur_buf = buf0;
auto* nxt_buf = buf1;
#pragma unroll
for (auto c = 0; c < NUM_CHUNKS - 1; ++c) {
ISSUE_LOADS(cur_ks + CHUNK_K);
COMPUTE_CHUNK(cur_buf, cur_ks);
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
STORE_TO_LDS(nxt_buf);
__syncthreads();
auto* tmp = cur_buf;
cur_buf = nxt_buf;
nxt_buf = tmp;
cur_ks += CHUNK_K;
}
COMPUTE_CHUNK(cur_buf, cur_ks);
if constexpr (TAIL_KT > 0) {
cur_ks += CHUNK_K;
ISSUE_LOADS(cur_ks);
asm volatile("s_waitcnt vmcnt(0) lgkmcnt(0)" ::: "memory");
STORE_TO_LDS(buf0);
__syncthreads();
// Tail compute — only TAIL_KT tiles, not full CHUNK_K
auto* _ca = Lds::a_off(buf0);
auto* _cb = Lds::b_off(buf0);
#pragma unroll
for (auto kt = 0; kt < TAIL_KT; ++kt) {
auto av = *reinterpret_cast<const int4_v*>(
_ca + Lds::a_idx(wave_m, lrow, kt * 4 + kgrp));
auto bv = *reinterpret_cast<const int4_v*>(
_cb + Lds::b_idx(wave_n, kt, kgrp, lrow));
auto ks_val = cur_ks + kt;
auto ks = ks_val * 4 + kgrp;
auto sa = 0, sb = 0;
if constexpr (A_VALID) sa = static_cast<int>(as_row[ks]);
else sa = (a_rt & (ks < KS)) ? static_cast<int>(as_row[ks]) : 127;
auto* bssh_p = Bssh + bssh_base + (ks_val & 1) * 2 + (ks_val >> 1) * 256;
if constexpr (B_VALID) sb = static_cast<int>(*bssh_p);
else sb = (b_rt & (ks < KS)) ? static_cast<int>(*bssh_p) : 127;
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
av, bv, acc, FP4_E2M1, FP4_E2M1, 0, sa, 0, sb);
}
}
#undef ISSUE_LOADS
#undef STORE_TO_LDS
#undef COMPUTE_CHUNK
if (!tile_valid) return;
}
const auto out_col = tile_n + lrow;
const auto out_row_base = tile_m + kgrp * 4;
if constexpr (!B_VALID) { if (out_col >= N) return; }
constexpr auto out_rows_always_valid = A_VALID && (BM >= 16);
if constexpr (SPLITK) {
auto* c_out = C_partial + (long)ks_idx * M * N + (long)out_row_base * N + out_col;
#pragma unroll
for (auto i = 0; i < 4; ++i, c_out += N) {
if constexpr (out_rows_always_valid) *c_out = acc[i];
else if (out_row_base + i < M) *c_out = acc[i];
}
} else {
auto* c_out = C_final + (long)out_row_base * N + out_col;
#pragma unroll
for (auto i = 0; i < 4; ++i, c_out += N) {
if constexpr (out_rows_always_valid) *c_out = float_to_bf16(acc[i]);
else if (out_row_base + i < M) *c_out = float_to_bf16(acc[i]);
}
}
}
template<int C_N, int C_NUM_KSPLIT, int C_M>
__global__ void mxfp4_reduce(
const float* __restrict__ C_partial,
uint16_t* __restrict__ C_out)
{
constexpr auto M = C_M;
constexpr auto N = C_N;
constexpr auto mn_stride = (long)M * N;
const auto col = blockIdx.x * 32 + threadIdx.x;
const auto row = blockIdx.y * 16 + threadIdx.y;
if (row >= M || col >= N) return;
const auto mn = (long)row * N + col;
auto sum = 0.f;
const auto* ptr = C_partial + mn;
#pragma unroll
for (auto k = 0; k < C_NUM_KSPLIT; ++k)
sum += ptr[k * mn_stride];
bf16x2 v;
v[0] = static_cast<__bf16>(sum);
auto r = uint16_t{};
__builtin_memcpy(&r, &v, sizeof(r));
C_out[mn] = r;
}
// Fused Quant+GEMM kernel for PATH A (K=512, small M).
// Single kernel launch: loads bf16 A into LDS, quants to FP4 in registers,
// then iterates N-tiles reusing A from registers. Eliminates quant kernel.
template<int C_M, int C_K, int C_N, int C_SCALEN>
__global__ void __launch_bounds__(((C_M + 15) / 16) * 4 * 64, 2)
mxfp4_fused_quant_gemm(
const __bf16* __restrict__ A_bf16,
const uint8_t* __restrict__ Bsh,
const uint8_t* __restrict__ Bssh,
uint16_t* __restrict__ C_final)
{
constexpr auto KS = C_K / 32;
constexpr auto K2 = C_K / 2;
constexpr auto CKT = C_K / 128; // k-tiles for MFMA
constexpr auto WAVES_M = (C_M + 15) / 16; // 1 for M<=16, 2 for M=32
// BN=64 for fused kernel: A in registers, more N-waves share A for free
constexpr auto WAVES_N = 4;
constexpr auto BN = WAVES_N * 16; // 64
constexpr auto NWARPS = WAVES_M * WAVES_N;
const auto tid = threadIdx.x;
const auto lane = tid % 64;
const auto wave_id = tid / 64;
const auto lrow = lane % 16;
const auto kgrp = lane / 16;
const auto wave_m = wave_id / WAVES_N;
const auto wave_n = wave_id % WAVES_N;
// M-row and N-tile for this wave
const auto tile_m = wave_m * 16;
const auto tile_n = static_cast<int>(blockIdx.x) * BN + wave_n * 16;
const auto my_row = tile_m + lrow; // actual M-row this lane handles
constexpr auto A_BF16_BYTES = C_M * C_K * 2;
extern __shared__ uint8_t smem[];
{
// 256 threads, 16 bytes each = 4096 bytes per iteration
constexpr auto BYTES_PER_ITER = NWARPS * 64 * 16;
constexpr auto NITERS = (A_BF16_BYTES + BYTES_PER_ITER - 1) / BYTES_PER_ITER;
#pragma unroll
for (auto iter = 0; iter < NITERS; ++iter) {
const auto offset = iter * BYTES_PER_ITER + tid * 16;
if (offset + 16 <= A_BF16_BYTES) {
*reinterpret_cast<int4_v*>(smem + offset) =
*reinterpret_cast<const int4_v*>(
reinterpret_cast<const uint8_t*>(A_bf16) + offset);
}
}
}
__syncthreads();
// Each lane handles row=lrow, and quants CKT groups (one per k-tile)
// kgrp selects which 32-element group within the k-tile
int4_v a_fp4_regs[CKT];
int a_scale_regs[CKT];
constexpr auto a_row_always_valid = (C_M >= ((C_M + 15) / 16) * 16);
const auto a_row_valid = a_row_always_valid ? true : (my_row < C_M);
#pragma unroll
for (auto kt = 0; kt < CKT; ++kt) {
const auto kg = kt * 4 + kgrp; // quant group index
if (a_row_valid) {
const auto lds_off = my_row * C_K * 2 + kg * 64;
const auto w0 = *reinterpret_cast<const int4_v*>(smem + lds_off);
const auto w1 = *reinterpret_cast<const int4_v*>(smem + lds_off + 16);
const auto w2 = *reinterpret_cast<const int4_v*>(smem + lds_off + 32);
const auto w3 = *reinterpret_cast<const int4_v*>(smem + lds_off + 48);
const auto* p0 = reinterpret_cast<const uint32_t*>(&w0);
const auto* p1 = reinterpret_cast<const uint32_t*>(&w1);
const auto* p2 = reinterpret_cast<const uint32_t*>(&w2);
const auto* p3 = reinterpret_cast<const uint32_t*>(&w3);
auto absMax = 1e-10f;
#define AMAX_F(pair) { \
const auto lo = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair)); \
const auto hi = __builtin_bit_cast(__bf16, static_cast<uint16_t>((pair) >> 16)); \
const auto flo = __builtin_elementwise_abs(static_cast<float>(lo)); \
const auto fhi = __builtin_elementwise_abs(static_cast<float>(hi)); \
absMax = (flo > absMax) ? flo : absMax; \
absMax = (fhi > absMax) ? fhi : absMax; \
}
AMAX_F(p0[0]) AMAX_F(p0[1]) AMAX_F(p0[2]) AMAX_F(p0[3])
AMAX_F(p1[0]) AMAX_F(p1[1]) AMAX_F(p1[2]) AMAX_F(p1[3])
AMAX_F(p2[0]) AMAX_F(p2[1]) AMAX_F(p2[2]) AMAX_F(p2[3])
AMAX_F(p3[0]) AMAX_F(p3[1]) AMAX_F(p3[2]) AMAX_F(p3[3])
#undef AMAX_F
const auto u32 = __builtin_bit_cast(uint32_t, absMax);
const auto amax_exp = ((u32 + 0x200000u) >> 23) & 0xFFu;
const auto inv_exp = (amax_exp >= 2u) ? (amax_exp - 2u) : 0u;
a_scale_regs[kt] = static_cast<int>(inv_exp);
const auto hw_scale = __builtin_bit_cast(float, static_cast<uint32_t>(inv_exp) << 23);
#define CVT_F(d, pair, sel) \
d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
d, __builtin_bit_cast(bf16x2, (pair)), hw_scale, sel)
auto d0 = 0u, d1 = 0u, d2 = 0u, d3 = 0u;
CVT_F(d0, p0[0], 0); CVT_F(d0, p0[1], 1); CVT_F(d0, p0[2], 2); CVT_F(d0, p0[3], 3);
CVT_F(d1, p1[0], 0); CVT_F(d1, p1[1], 1); CVT_F(d1, p1[2], 2); CVT_F(d1, p1[3], 3);
CVT_F(d2, p2[0], 0); CVT_F(d2, p2[1], 1); CVT_F(d2, p2[2], 2); CVT_F(d2, p2[3], 3);
CVT_F(d3, p3[0], 0); CVT_F(d3, p3[1], 1); CVT_F(d3, p3[2], 2); CVT_F(d3, p3[3], 3);
#undef CVT_F
a_fp4_regs[kt] = int4_v{static_cast<int>(d0), static_cast<int>(d1),
static_cast<int>(d2), static_cast<int>(d3)};
} else {
a_fp4_regs[kt] = int4_v{0, 0, 0, 0};
a_scale_regs[kt] = 127;
}
}
// tile_n < C_N guaranteed by grid dimensions (C_N / BN blocks)
float4_v acc{0.f, 0.f, 0.f, 0.f};
const auto n_tile = tile_n / 16;
const auto bsh_n_stride = (long)(C_K / 64) * 512;
const auto bsh_lane_base = Bsh + (long)n_tile * bsh_n_stride + (long)lrow * 16;
const auto k_half_off = (kgrp & 1) * 256;
const auto k_blk_base = kgrp >> 1;
// B_scale decode
const auto gn = tile_n + lrow;
constexpr auto scaleN = (C_K / 32 + 7) / 8 * 8;
const auto bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64
+ (gn >> 5) * (32 * scaleN);
#pragma unroll
for (auto kt = 0; kt < CKT; ++kt) {
const auto bv = load16(bsh_lane_base + (long)(kt * 2 + k_blk_base) * 512 + k_half_off);
const auto ks_val = kt;
const auto sb = static_cast<int>(*(Bssh + bssh_base
+ (ks_val & 1) * 2 + (ks_val >> 1) * 256));
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a_fp4_regs[kt], bv, acc, FP4_E2M1, FP4_E2M1,
0, a_scale_regs[kt], 0, sb);
}
const auto out_col = tile_n + lrow;
const auto out_row_base = tile_m + kgrp * 4;
if (out_col >= C_N) return;
auto* c_out = C_final + (long)out_row_base * C_N + out_col;
#pragma unroll
for (auto i = 0; i < 4; ++i, c_out += C_N) {
if (out_row_base + i < C_M)
*c_out = float_to_bf16(acc[i]);
}
}
// Fused Quant+GEMM for splitK — each split block quants its K-slice of A.
// Eliminates separate quant kernel launch. Still needs reduce kernel.
template<int C_M, int C_K, int C_N, int C_SCALEN, int C_KPS, int C_NUM_KSPLIT>
__global__ void __launch_bounds__(((C_M + 15) / 16) * 4 * 64, 2)
mxfp4_fused_quant_gemm_splitk(
const __bf16* __restrict__ A_bf16,
const uint8_t* __restrict__ Bsh,
const uint8_t* __restrict__ Bssh,
float* __restrict__ C_partial)
{
constexpr auto CKT = C_KPS; // k-tiles per split
constexpr auto WAVES_M = (C_M + 15) / 16;
constexpr auto WAVES_N = 4; // BN=64, A in registers — more N-reuse
constexpr auto BN = WAVES_N * 16;
constexpr auto NWARPS = WAVES_M * WAVES_N;
constexpr auto K2 = C_K / 2;
const auto tid = threadIdx.x;
const auto lane = tid % 64;
const auto wave_id = tid / 64;
const auto lrow = lane % 16;
const auto kgrp = lane / 16;
const auto wave_m = wave_id / WAVES_N;
const auto wave_n = wave_id % WAVES_N;
const auto ks_idx = static_cast<int>(blockIdx.z);
const auto tile_m = wave_m * 16;
const auto tile_n = static_cast<int>(blockIdx.x) * BN + wave_n * 16;
const auto my_row = tile_m + lrow;
// K-slice for this split: k_start..k_start+CKT*128
const auto ks_start = ks_idx * C_KPS;
const auto k_byte_start = ks_start * 128; // in bf16 elements (128 per k-tile = 64 bytes FP4 = 256 bf16 bytes for 128 elements)
// For M=16, CKT=4: need 16 rows × 512 bf16 = 16KB
constexpr auto A_SLICE_ELEMS = C_M * C_KPS * 128; // bf16 elements in this K-slice
constexpr auto A_SLICE_BYTES = A_SLICE_ELEMS * 2;
extern __shared__ uint8_t smem[];
{
constexpr auto BYTES_PER_ITER = NWARPS * 64 * 16;
constexpr auto NITERS = (A_SLICE_BYTES + BYTES_PER_ITER - 1) / BYTES_PER_ITER;
#pragma unroll
for (auto iter = 0; iter < NITERS; ++iter) {
const auto offset = iter * BYTES_PER_ITER + tid * 16;
if (offset + 16 <= A_SLICE_BYTES) {
// Source: A_bf16 row-major, we need elements [row, k_byte_start..k_byte_start+CKT*128]
// LDS offset = offset within the slice
// Global offset = row * K + k_byte_start + col_within_slice
const auto slice_elem = offset / 2; // bf16 element index within slice
const auto row = slice_elem / (C_KPS * 128);
const auto col = slice_elem % (C_KPS * 128);
const auto global_byte_off = (long)row * C_K * 2 + (long)(k_byte_start + col) * 2;
if (row < C_M) {
*reinterpret_cast<int4_v*>(smem + offset) =
*reinterpret_cast<const int4_v*>(
reinterpret_cast<const uint8_t*>(A_bf16) + global_byte_off);
}
}
}
}
__syncthreads();
int4_v a_fp4_regs[CKT];
int a_scale_regs[CKT];
constexpr auto a_row_always_valid = (C_M >= ((C_M + 15) / 16) * 16);
const auto a_row_valid = a_row_always_valid ? true : (my_row < C_M);
#pragma unroll
for (auto kt = 0; kt < CKT; ++kt) {
const auto kg = kt * 4 + kgrp;
if (a_row_valid) {
// LDS stores the K-slice contiguously: row * (CKT*128*2) + kg * 64
const auto lds_off = my_row * C_KPS * 128 * 2 + kg * 64;
const auto w0 = *reinterpret_cast<const int4_v*>(smem + lds_off);
const auto w1 = *reinterpret_cast<const int4_v*>(smem + lds_off + 16);
const auto w2 = *reinterpret_cast<const int4_v*>(smem + lds_off + 32);
const auto w3 = *reinterpret_cast<const int4_v*>(smem + lds_off + 48);
const auto* p0 = reinterpret_cast<const uint32_t*>(&w0);
const auto* p1 = reinterpret_cast<const uint32_t*>(&w1);
const auto* p2 = reinterpret_cast<const uint32_t*>(&w2);
const auto* p3 = reinterpret_cast<const uint32_t*>(&w3);
auto absMax = 1e-10f;
#define AMAX_S(pair) { \
const auto lo = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair)); \
const auto hi = __builtin_bit_cast(__bf16, static_cast<uint16_t>((pair) >> 16)); \
const auto flo = __builtin_elementwise_abs(static_cast<float>(lo)); \
const auto fhi = __builtin_elementwise_abs(static_cast<float>(hi)); \
absMax = (flo > absMax) ? flo : absMax; \
absMax = (fhi > absMax) ? fhi : absMax; \
}
AMAX_S(p0[0]) AMAX_S(p0[1]) AMAX_S(p0[2]) AMAX_S(p0[3])
AMAX_S(p1[0]) AMAX_S(p1[1]) AMAX_S(p1[2]) AMAX_S(p1[3])
AMAX_S(p2[0]) AMAX_S(p2[1]) AMAX_S(p2[2]) AMAX_S(p2[3])
AMAX_S(p3[0]) AMAX_S(p3[1]) AMAX_S(p3[2]) AMAX_S(p3[3])
#undef AMAX_S
const auto u32 = __builtin_bit_cast(uint32_t, absMax);
const auto amax_exp = ((u32 + 0x200000u) >> 23) & 0xFFu;
const auto inv_exp = (amax_exp >= 2u) ? (amax_exp - 2u) : 0u;
a_scale_regs[kt] = static_cast<int>(inv_exp);
const auto hw_scale = __builtin_bit_cast(float, static_cast<uint32_t>(inv_exp) << 23);
#define CVT_S(d, pair, sel) \
d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
d, __builtin_bit_cast(bf16x2, (pair)), hw_scale, sel)
auto d0 = 0u, d1 = 0u, d2 = 0u, d3 = 0u;
CVT_S(d0, p0[0], 0); CVT_S(d0, p0[1], 1); CVT_S(d0, p0[2], 2); CVT_S(d0, p0[3], 3);
CVT_S(d1, p1[0], 0); CVT_S(d1, p1[1], 1); CVT_S(d1, p1[2], 2); CVT_S(d1, p1[3], 3);
CVT_S(d2, p2[0], 0); CVT_S(d2, p2[1], 1); CVT_S(d2, p2[2], 2); CVT_S(d2, p2[3], 3);
CVT_S(d3, p3[0], 0); CVT_S(d3, p3[1], 1); CVT_S(d3, p3[2], 2); CVT_S(d3, p3[3], 3);
#undef CVT_S
a_fp4_regs[kt] = int4_v{static_cast<int>(d0), static_cast<int>(d1),
static_cast<int>(d2), static_cast<int>(d3)};
} else {
a_fp4_regs[kt] = int4_v{0, 0, 0, 0};
a_scale_regs[kt] = 127;
}
}
// tile_n < C_N guaranteed by grid dimensions (C_N / BN blocks)
float4_v acc{0.f, 0.f, 0.f, 0.f};
const auto n_tile = tile_n / 16;
const auto bsh_n_stride = (long)(C_K / 64) * 512;
const auto bsh_lane_base = Bsh + (long)n_tile * bsh_n_stride + (long)lrow * 16;
const auto k_half_off = (kgrp & 1) * 256;
const auto k_blk_base = kgrp >> 1;
const auto gn = tile_n + lrow;
constexpr auto scaleN = (C_K / 32 + 7) / 8 * 8;
const auto bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64
+ (gn >> 5) * (32 * scaleN);
#pragma unroll
for (auto kt = 0; kt < CKT; ++kt) {
const auto global_kt = ks_start + kt;
const auto bv = load16(bsh_lane_base + (long)(global_kt * 2 + k_blk_base) * 512 + k_half_off);
const auto ks_val = global_kt;
const auto sb = static_cast<int>(*(Bssh + bssh_base
+ (ks_val & 1) * 2 + (ks_val >> 1) * 256));
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a_fp4_regs[kt], bv, acc, FP4_E2M1, FP4_E2M1,
0, a_scale_regs[kt], 0, sb);
}
const auto out_col = tile_n + lrow;
const auto out_row_base = tile_m + kgrp * 4;
if (out_col >= C_N) return;
auto* c_out = C_partial + (long)ks_idx * C_M * C_N + (long)out_row_base * C_N + out_col;
#pragma unroll
for (auto i = 0; i < 4; ++i, c_out += C_N) {
if (out_row_base + i < C_M)
*c_out = acc[i];
}
}
template<int C_M, int C_K>
void launch_quant(
const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale)
{
constexpr auto KS = C_K / 32;
constexpr auto n_groups = C_M * KS;
constexpr dim3 block{128};
constexpr dim3 grid{static_cast<uint32_t>((n_groups + 127) / 128)};
mxfp4_quant<C_M, C_K><<<grid, block>>>(A_bf16, A_fp4, A_scale);
}
template<int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_KPS, int C_NUM_KSPLIT, int C_M>
void launch_gemm_nk(
const uint8_t* A, const uint8_t* As,
const uint8_t* Bsh, const uint8_t* Bssh,
float* C_partial, uint16_t* C_final)
{
constexpr auto do_splitk = C_NUM_KSPLIT > 1;
// M→BM dispatch resolved at compile time
constexpr auto BM = (C_M <= 8) ? 8 : (C_M <= 32) ? 16 : (C_M <= 128) ? 32 : 64;
constexpr auto BN = 32;
constexpr auto NWARPS = ((BM + 15) / 16) * (BN / 16);
static_assert(((BM + 15) / 16) * (BN / 16) == NWARPS);
constexpr auto WAVES_M = (BM + 15) / 16;
constexpr auto WAVES_N = BN / 16;
constexpr auto CHUNK_K = (C_KPS >= 4) ? 4 : C_KPS;
constexpr auto smem_size = (C_KPS > 4)
? LdsLayout<WAVES_M, WAVES_N, CHUNK_K>::TOTAL_LDS
: 0;
constexpr auto full_m = (BM >= 16) ? C_M / BM : 0;
constexpr auto full_n = C_N / BN;
constexpr auto total_m = (C_M + BM - 1) / BM;
constexpr auto total_n = (C_N + BN - 1) / BN;
constexpr auto edge_m = total_m - full_m;
constexpr auto edge_n = total_n - full_n;
constexpr dim3 block{static_cast<uint32_t>(NWARPS * 64)};
// Helper: set smem + launch for a specific (AV, BV, ox, oy) combo
auto sub = [&]<bool AV, bool BV>() {
constexpr auto gx = BV ? full_n : edge_n;
constexpr auto gy = AV ? full_m : edge_m;
constexpr auto ox = BV ? 0 : full_n;
constexpr auto oy = AV ? 0 : full_m;
if constexpr (gx > 0 && gy > 0) {
constexpr auto SK = do_splitk;
if constexpr (smem_size > 0)
(void)hipFuncSetAttribute(
(const void*)mxfp4_gemm<BM,BN,NWARPS,SK,AV,BV,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS,C_M,ox,oy>,
hipFuncAttributeMaxDynamicSharedMemorySize, smem_size);
constexpr dim3 grid{
static_cast<uint32_t>(gx),
static_cast<uint32_t>(gy),
static_cast<uint32_t>(C_NUM_KSPLIT)
};
if constexpr (SK)
mxfp4_gemm<BM,BN,NWARPS,true,AV,BV,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS,C_M,ox,oy>
<<<grid,block,smem_size>>>(A,As,Bsh,Bssh,C_partial,nullptr);
else
mxfp4_gemm<BM,BN,NWARPS,false,AV,BV,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS,C_M,ox,oy>
<<<grid,block,smem_size>>>(A,As,Bsh,Bssh,nullptr,C_final);
}
};
sub.template operator()<true, true >();
sub.template operator()<true, false>();
sub.template operator()<false, true >();
sub.template operator()<false, false>();
if constexpr (do_splitk) {
constexpr dim3 rblock{32, 16};
constexpr dim3 rgrid{
static_cast<uint32_t>((C_N + 31) / 32),
static_cast<uint32_t>((C_M + 15) / 16)
};
mxfp4_reduce<C_N, C_NUM_KSPLIT, C_M><<<rgrid, rblock>>>(C_partial, C_final);
}
}
template<int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_M>
void launch_gemm_nk_7168(
const uint8_t* A, const uint8_t* As,
const uint8_t* Bsh, const uint8_t* Bssh,
float* C_partial, uint16_t* C_final)
{
if constexpr (C_M <= 8)
launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, 8, 7, C_M>(A, As, Bsh, Bssh, C_partial, C_final);
else if constexpr (C_M <= 16)
launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, 4, 14, C_M>(A, As, Bsh, Bssh, C_partial, C_final);
else
launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, 56, 1, C_M>(A, As, Bsh, Bssh, C_partial, C_final);
}
// Fused quant+GEMM launch for PATH A (K=512, non-splitK)
template<int C_M, int C_K, int C_N, int C_SCALEN>
void launch_fused(
const __bf16* A_bf16,
const uint8_t* Bsh, const uint8_t* Bssh,
uint16_t* C_final)
{
constexpr auto WAVES_M = (C_M + 15) / 16;
constexpr auto WAVES_N = 4; // match kernel BN=64
constexpr auto BN = WAVES_N * 16; // 64
constexpr auto NTHREADS = WAVES_M * WAVES_N * 64;
constexpr auto smem_size = C_M * C_K * 2; // A_bf16 in LDS
if constexpr (smem_size > 48 * 1024)
(void)hipFuncSetAttribute(
(const void*)mxfp4_fused_quant_gemm<C_M, C_K, C_N, C_SCALEN>,
hipFuncAttributeMaxDynamicSharedMemorySize, smem_size);
constexpr dim3 grid{static_cast<uint32_t>(C_N / BN)};
constexpr dim3 block{static_cast<uint32_t>(NTHREADS)};
mxfp4_fused_quant_gemm<C_M, C_K, C_N, C_SCALEN>
<<<grid, block, smem_size>>>(A_bf16, Bsh, Bssh, C_final);
}
// Combined quant + GEMM dispatch — fully compile-time
template<int C_M, int C_K, int C_N, int C_SCALEN, int C_TOTAL_KT, int C_KPS, int C_NUM_KSPLIT>
void launch_shape(
const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
const uint8_t* Bsh, const uint8_t* Bssh,
float* C_partial, uint16_t* C_final)
{
// Use fused kernel for PATH A non-splitK shapes where A fits in LDS (<=64KB)
// M=4: 4KB, M=8: 8KB, M=16: 16KB, M=32: 32KB — all fit
// M=64: 64KB — borderline, skip. M=256: 256KB — way too large.
constexpr auto a_bf16_bytes = C_M * C_K * 2;
if constexpr (C_KPS <= 4 && C_NUM_KSPLIT == 1 && a_bf16_bytes <= 32 * 1024) {
launch_fused<C_M, C_K, C_N, C_SCALEN>(A_bf16, Bsh, Bssh, C_final);
return;
}
launch_quant<C_M, C_K>(A_bf16, A_fp4, A_scale);
launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, C_KPS, C_NUM_KSPLIT, C_M>(
A_fp4, A_scale, Bsh, Bssh, C_partial, C_final);
}
template<int C_M, int C_K, int C_N, int C_SCALEN, int C_TOTAL_KT>
void launch_shape_7168(
const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
const uint8_t* Bsh, const uint8_t* Bssh,
float* C_partial, uint16_t* C_final)
{
// For small M where per-split A slice fits in LDS, use fused quant+GEMM splitK
// M<=8: splitK=7 (original), M<=16: splitK=28 (optimized)
constexpr auto C_KPS = (C_M <= 8) ? 8 : (C_M <= 16) ? 2 : 56;
constexpr auto C_NUM_KSPLIT = (C_M <= 8) ? 7 : (C_M <= 16) ? 28 : 1;
constexpr auto a_slice_bytes = C_M * C_KPS * 128 * 2; // bf16 bytes per split
if constexpr (C_NUM_KSPLIT > 1 && a_slice_bytes <= 32 * 1024) {
// Fused: single kernel does quant+GEMM for each K-split
constexpr auto WAVES_M = (C_M + 15) / 16;
constexpr auto WAVES_N = 4; // match splitK kernel BN=64
constexpr auto BN = WAVES_N * 16;
constexpr auto NTHREADS = WAVES_M * WAVES_N * 64;
constexpr auto smem_size = a_slice_bytes;
if constexpr (smem_size > 48 * 1024)
(void)hipFuncSetAttribute(
(const void*)mxfp4_fused_quant_gemm_splitk<C_M, C_K, C_N, C_SCALEN, C_KPS, C_NUM_KSPLIT>,
hipFuncAttributeMaxDynamicSharedMemorySize, smem_size);
constexpr dim3 grid{static_cast<uint32_t>(C_N / BN), 1, static_cast<uint32_t>(C_NUM_KSPLIT)};
constexpr dim3 block{static_cast<uint32_t>(NTHREADS)};
mxfp4_fused_quant_gemm_splitk<C_M, C_K, C_N, C_SCALEN, C_KPS, C_NUM_KSPLIT>
<<<grid, block, smem_size>>>(A_bf16, Bsh, Bssh, C_partial);
// Still need reduce kernel
constexpr dim3 rblock{32, 16};
constexpr dim3 rgrid{
static_cast<uint32_t>((C_N + 31) / 32),
static_cast<uint32_t>((C_M + 15) / 16)
};
mxfp4_reduce<C_N, C_NUM_KSPLIT, C_M><<<rgrid, rblock>>>(C_partial, C_final);
} else {
// Fallback: separate quant + GEMM
launch_quant<C_M, C_K>(A_bf16, A_fp4, A_scale);
launch_gemm_nk_7168<C_N, C_K, C_SCALEN, C_TOTAL_KT, C_M>(
A_fp4, A_scale, Bsh, Bssh, C_partial, C_final);
}
}
// Runtime M dispatch
template<int C_K, int C_N, int C_SCALEN, int C_TOTAL_KT, int C_KPS, int C_NUM_KSPLIT>
void dispatch_shape(
const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
const uint8_t* Bsh, const uint8_t* Bssh,
float* C_partial, uint16_t* C_final, int M)
{
#define CASE_M(val) case val: launch_shape<val,C_K,C_N,C_SCALEN,C_TOTAL_KT,C_KPS,C_NUM_KSPLIT>(A_bf16,A_fp4,A_scale,Bsh,Bssh,C_partial,C_final); return
switch (M) {
CASE_M(4); CASE_M(8); CASE_M(16); CASE_M(32); CASE_M(64); CASE_M(256);
}
#undef CASE_M
}
template<int C_K, int C_N, int C_SCALEN, int C_TOTAL_KT>
void dispatch_shape_7168(
const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
const uint8_t* Bsh, const uint8_t* Bssh,
float* C_partial, uint16_t* C_final, int M)
{
#define CASE_M(val) case val: launch_shape_7168<val,C_K,C_N,C_SCALEN,C_TOTAL_KT>(A_bf16,A_fp4,A_scale,Bsh,Bssh,C_partial,C_final); return
switch (M) {
CASE_M(4); CASE_M(8); CASE_M(16); CASE_M(32); CASE_M(64); CASE_M(256);
}
#undef CASE_M
}
extern "C" void launch_all(
const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
const uint8_t* Bsh, const uint8_t* Bssh,
float* C_partial, uint16_t* C_final,
int M, int N, int K)
{
if (N == 2880 && K == 512)
dispatch_shape<512, 2880, 16, 4, 4, 1>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
else if (N == 2112 && K == 7168)
dispatch_shape_7168<7168, 2112, 224, 56>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
else if (N == 4096 && K == 512)
dispatch_shape<512, 4096, 16, 4, 4, 1>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
else if (N == 7168 && K == 2048)
dispatch_shape<2048, 7168, 64, 16, 16, 1>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
else if (N == 3072 && K == 1536)
dispatch_shape<1536, 3072, 48, 12, 12, 1>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
}
"""
CPP = r"""
#include <torch/extension.h>
#include <c10/core/DeviceGuard.h>
extern "C" void launch_all(const __bf16*, uint8_t*, uint8_t*,
const uint8_t*, const uint8_t*,
float*, uint16_t*, int, int, int);
static int get_num_ksplit(int M, int K) {
auto total_ktiles = K / 128;
if (total_ktiles < 28) return 1;
if (M <= 8) return 7;
if (M <= 16) return 28;
return 1;
}
struct ShapeWorkspace {
at::Tensor A_fp4;
at::Tensor A_scale;
at::Tensor C_partial;
at::Tensor C;
uint8_t* a_fp4_ptr = nullptr;
uint8_t* a_scale_ptr = nullptr;
float* c_partial_ptr = nullptr;
uint16_t* c_final_ptr = nullptr;
int M = 0, N = 0, K = 0;
int num_ksplit = 0;
};
struct BCache {
const uint8_t* bsh_ptr = nullptr;
const uint8_t* bssh_ptr = nullptr;
int64_t bsh_data_ptr = 0;
int64_t bssh_data_ptr = 0;
};
static ShapeWorkspace g_ws[10];
static auto g_ws_count = 0;
static BCache g_bcache;
static auto* find_or_create_ws(int M, int N, int K, int num_ksplit,
const at::TensorOptions& opts) {
for (auto i = 0; i < g_ws_count; ++i) {
if (g_ws[i].M == M && g_ws[i].N == N && g_ws[i].K == K)
return &g_ws[i];
}
auto& ws = g_ws[g_ws_count++];
ws.M = M; ws.N = N; ws.K = K;
ws.num_ksplit = num_ksplit;
auto KS = K / 32;
ws.A_fp4 = at::empty({(int64_t)M, (int64_t)(K / 2)}, opts.dtype(at::kByte));
ws.A_scale = at::empty({(int64_t)M, (int64_t)KS}, opts.dtype(at::kByte));
ws.C = at::empty({(int64_t)M, (int64_t)N}, opts.dtype(at::kBFloat16));
if (num_ksplit > 1)
ws.C_partial = at::empty({(int64_t)num_ksplit, (int64_t)M, (int64_t)N}, opts.dtype(at::kFloat));
ws.a_fp4_ptr = ws.A_fp4.data_ptr<uint8_t>();
ws.a_scale_ptr = ws.A_scale.data_ptr<uint8_t>();
ws.c_partial_ptr = (num_ksplit > 1) ? ws.C_partial.data_ptr<float>() : nullptr;
ws.c_final_ptr = reinterpret_cast<uint16_t*>(ws.C.data_ptr<at::BFloat16>());
return &ws;
}
at::Tensor fwd(const at::Tensor& A,
const at::Tensor& B_q,
const at::Tensor& B_shuffle,
const at::Tensor& B_scale_sh) {
auto guard = at::DeviceGuard(A.device());
const auto M = static_cast<int>(A.size(0));
const auto K = static_cast<int>(A.size(1));
const auto N = static_cast<int>(B_q.size(0));
const __bf16* a_bf16_ptr;
at::Tensor A_bf16;
if (A.scalar_type() == at::kBFloat16 && A.is_contiguous()) {
a_bf16_ptr = reinterpret_cast<const __bf16*>(A.data_ptr<at::BFloat16>());
} else {
A_bf16 = A.to(A.device(), at::kBFloat16, false, false, at::MemoryFormat::Contiguous);
a_bf16_ptr = reinterpret_cast<const __bf16*>(A_bf16.data_ptr<at::BFloat16>());
}
auto bsh_dp = reinterpret_cast<int64_t>(B_shuffle.data_ptr());
auto bssh_dp = reinterpret_cast<int64_t>(B_scale_sh.data_ptr());
if (bsh_dp != g_bcache.bsh_data_ptr || bssh_dp != g_bcache.bssh_data_ptr) {
auto Bsh = B_shuffle.view(at::kByte);
if (!Bsh.is_contiguous()) Bsh = Bsh.contiguous();
auto Bssh = B_scale_sh.view(at::kByte);
if (!Bssh.is_contiguous()) Bssh = Bssh.contiguous();
g_bcache.bsh_ptr = Bsh.data_ptr<uint8_t>();
g_bcache.bssh_ptr = Bssh.data_ptr<uint8_t>();
g_bcache.bsh_data_ptr = bsh_dp;
g_bcache.bssh_data_ptr = bssh_dp;
}
const auto num_ksplit = get_num_ksplit(M, K);
auto* ws = find_or_create_ws(M, N, K, num_ksplit, A.options());
launch_all(
a_bf16_ptr, ws->a_fp4_ptr, ws->a_scale_ptr,
g_bcache.bsh_ptr, g_bcache.bssh_ptr,
ws->c_partial_ptr, ws->c_final_ptr,
M, N, K);
return ws->C;
}
"""
_ext = load_inline(
name=f"g_{uuid.uuid4().hex[:8]}",
cpp_sources=[CPP],
cuda_sources=[HIP_KERNEL],
functions=["fwd"],
with_cuda=True,
extra_cflags=["-O3", "-std=c++20"],
extra_cuda_cflags=[
"-O3",
"--offload-arch=gfx950",
"-ffast-math",
"-ffinite-math-only",
"-munsafe-fp-atomics",
"-std=c++20",
"-mllvm", "-amdgpu-early-inline-all=true",
"-mllvm", "-amdgpu-function-calls=false",
"-mwavefrontsize64",
"-mcumode",
"-mllvm", "--amdgpu-kernarg-preload-count=16",
"-mllvm", "--lsr-drop-solution=1",
"-mllvm", "-amdgpu-coerce-illegal-types=1",
"-fgpu-flush-denormals-to-zero",
"-fno-offload-uniform-block",
"-mllvm", "-amdgpu-loop-prefetch=true",
"-mllvm", "-enable-unroll-and-jam=true",
"-mllvm", "-unroll-threshold=500",
"-mllvm", "-amdgpu-internalize-symbols=true",
],
extra_ldflags=["-lamdhip64"],
)
def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:
"""MXFP4 GEMM v25d_v4_i1: v4_i0 + quant intrinsic dest_sel + LdsLayout + launch_bounds."""
A, _, B_q, B_shuffle, B_scale_sh = data
return _ext.fwd(A.cuda(), B_q, B_shuffle, B_scale_sh)
scrolls · 1330 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 585128.
⋯ 29 unchanged linesreturn *reinterpret_cast<const int4_v*>(p);}- // Non-temporal load: bypass L1, keep in L2.__device__ __forceinline__ auto load16_nt(const uint8_t* __restrict__ p) {return __builtin_nontemporal_load(reinterpret_cast<const int4_v*>(p));}⋯ 6 unchanged linesreturn r;}- // ── Quant kernel — fully templated on M, K ──────────────────────────────────template<int C_M, int C_K>__global__ void __launch_bounds__(128, (C_M * (C_K / 32) <= 256) ? 2 : 4)⋯ 13 unchanged linesconst auto row_k = (long)row * C_K;const auto* src = A_bf16 + row_k + kg * 32;- // 4 wide loads: 64 bytes in 4 × global_load_dwordx4const auto w0 = *reinterpret_cast<const int4_v*>(src);const auto w1 = *reinterpret_cast<const int4_v*>(src + 8);const auto w2 = *reinterpret_cast<const int4_v*>(src + 16);const auto w3 = *reinterpret_cast<const int4_v*>(src + 24);- // Reinterpret as uint32 pairs (zero-cost, same registers)const auto* p0 = reinterpret_cast<const uint32_t*>(&w0);const auto* p1 = reinterpret_cast<const uint32_t*>(&w1);const auto* p2 = reinterpret_cast<const uint32_t*>(&w2);const auto* p3 = reinterpret_cast<const uint32_t*>(&w3);- // absmax across all 32 bf16 values (16 uint32 pairs, each = 2 bf16)auto absMax = 1e-10f;#define AMAX_PAIR(pair) { \const auto lo = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair)); \⋯ 16 unchanged linesconst auto hw_scale = __builtin_bit_cast(float, static_cast<uint32_t>(inv_exp) << 23);- // Pack 16 FP4 pairs into 4 dwords using intrinsic dest_sel (no shift/OR/mask)- // Each __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(old, v2bf16, scale, byte_sel)- // writes 1 byte at position byte_sel in the accumulator dword.#define CVT(d, pair, sel) \d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \d, __builtin_bit_cast(bf16x2, (pair)), hw_scale, sel)⋯ 10 unchanged linesstatic_cast<int>(d2), static_cast<int>(d3)};}- // ── Boundary-checked loads ──────────────────────────────────────────────────template<bool ALWAYS_VALID>__device__ __forceinline__ auto load_or_zero(bool rt_valid, const uint8_t* p) {⋯ 1 unchanged lineselse return rt_valid ? load16(p) : int4_v{0,0,0,0};}- // ── LDS layout for software-pipelined path (CKT > 4) ───────────────────────- //- // Double-buffered. Per buffer:- // A: [WAVES_M][16 rows][CHUNK_K * 4 kgrps * 16 bytes + 16 pad]- // B: [WAVES_N][CHUNK_K][4 kgrps + 1 pad][16 lrows][16 bytes]- //- // Bank conflict strategy:- // A: XOR swizzle on row index — lrow ^ (k_idx & 7)- // B: XOR swizzle on kgrp — kgrp ^ (lrow >> 2)template<int WAVES_M, int WAVES_N, int CHUNK_K>struct LdsLayout {⋯ 1 unchanged linesstatic constexpr auto A_TILE = 16 * A_ROW;static constexpr auto A_SIZE = WAVES_M * A_TILE;- // B: 5 kgrp slots (4 real + 1 pad) to space kgrps 4-banks apartstatic constexpr auto B_KGRP_STRIDE = 16 * 16; // 16 lrows × 16B = 256Bstatic constexpr auto B_KT_STRIDE = 5 * B_KGRP_STRIDE; // 5 slots (4+1 pad)static constexpr auto B_TILE = CHUNK_K * B_KT_STRIDE;⋯ 15 unchanged lines}};- // ── Per-thread load item counts (constexpr) ─────────────────────────────────template<int WAVES_M, int WAVES_N, int CHUNK_K, int NWARPS>struct LoadCounts {⋯ 4 unchanged linesstatic constexpr auto B_PER_THREAD = (B_TOTAL + NTHREADS - 1) / NTHREADS;};- // ── Main GEMM kernel ────────────────────────────────────────────────────────template<int BM, int BN, int NWARPS, bool SPLITK, bool A_VALID, bool B_VALID,int CKT, int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_KPS,int C_M, int C_TILE_OFF_X, int C_TILE_OFF_Y>__global__ void __launch_bounds__(NWARPS * 64, (NWARPS <= 2) ? 4 : 2)+ __attribute__((amdgpu_flat_work_group_size(NWARPS * 64, NWARPS * 64)))mxfp4_gemm(const uint8_t* __restrict__ A,const uint8_t* __restrict__ As,⋯ 32 unchanged linesconst auto lrow = lane % 16;const auto kgrp = lane / 16;+ __builtin_assume(lrow >= 0 && lrow < 16);+ __builtin_assume(kgrp >= 0 && kgrp < 4);const auto gm = tile_m + lrow;const auto gn = tile_n + lrow;⋯ 1 unchanged linesconst auto a_rt = A_VALID | (gm < M);const auto b_rt = B_VALID | (gn < N);- // Scale pointers (loaded from global, 1 byte, L1 cached)const uint8_t* as_row = nullptr;if constexpr (A_VALID) {as_row = As + (long)gm * KS;⋯ 10 unchanged linesfloat4_v acc{0.f, 0.f, 0.f, 0.f};- // ════════════════════════════════════════════════════════════════════════// PATH A: CKT <= 4 — Direct global loads, no LDS- // ════════════════════════════════════════════════════════════════════════if constexpr (CKT <= 4) {if (tile_m >= M || tile_n >= N) return;⋯ 65 unchanged lines}#undef DO_MFMA- // ════════════════════════════════════════════════════════════════════════// PATH B: CKT > 4 — LDS double-buffered with register-staged pipelining- // ════════════════════════════════════════════════════════════════════════} else {constexpr auto CHUNK_K = 4;constexpr auto NUM_CHUNKS = CKT / CHUNK_K;⋯ 8 unchanged linesconst auto tile_valid = (tile_m < M) && (tile_n < N);- // ── Precompute per-item coordinates (hoisted out of hot loop) ────────- // A items: row, k_idx, lds_off — independent of ks_baseint a_row[LC::A_PER_THREAD];int a_k_idx[LC::A_PER_THREAD];int a_lds_offs[LC::A_PER_THREAD];⋯ 17 unchanged lines}}- // B items: global base offset (without ks-dependent k_blk), lds coordsint b_kt_local[LC::B_PER_THREAD]; // kt within chunk (0..CHUNK_K-1)long b_base_off[LC::B_PER_THREAD]; // global offset without k_blk termint b_kgrp_half[LC::B_PER_THREAD]; // (b_kgrp / 2) for k_blk calc⋯ 31 unchanged lines}}- // ── Register buffers for staged loads ────────────────────────────────int4_v a_regs[LC::A_PER_THREAD];int4_v b_regs[LC::B_PER_THREAD];- // ── Macros for inlined issue/store/compute (no lambdas) ──────────────#define ISSUE_LOADS(ks_base) \{ \⋯ 38 unchanged lines#define COMPUTE_CHUNK(buf, chunk_ks) \{ \- auto* _ca = Lds::a_off(buf); \- auto* _cb = Lds::b_off(buf); \- __builtin_amdgcn_sched_barrier(0); \- asm volatile("s_setprio 1" ::: "memory"); \- __builtin_amdgcn_sched_barrier(0); \- _Pragma("unroll") \- for (auto _kt = 0; _kt < CHUNK_K; ++_kt) { \- auto _av = *reinterpret_cast<const int4_v*>( \- _ca + Lds::a_idx(wave_m, lrow, _kt * 4 + kgrp)); \- auto _bv = *reinterpret_cast<const int4_v*>( \- _cb + Lds::b_idx(wave_n, _kt, kgrp, lrow)); \- auto _ks_val = (chunk_ks) + _kt; \- auto _ks = _ks_val * 4 + kgrp; \- auto _sa = 0, _sb = 0; \- if constexpr (A_VALID) _sa = static_cast<int>(as_row[_ks]); \- else _sa = (a_rt & (_ks < KS)) ? static_cast<int>(as_row[_ks]) : 127; \- auto* _bssh_p = Bssh + bssh_base + (_ks_val & 1) * 2 + (_ks_val >> 1) * 256; \- if constexpr (B_VALID) _sb = static_cast<int>(*_bssh_p); \- else _sb = (b_rt & (_ks < KS)) ? static_cast<int>(*_bssh_p) : 127; \- __builtin_amdgcn_sched_group_barrier(0x100, 2, 0); \- __builtin_amdgcn_sched_group_barrier(0x008, 1, 0); \- acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4( \- _av, _bv, acc, FP4_E2M1, FP4_E2M1, 0, _sa, 0, _sb); \+ /* Compute LDS byte offsets for each k-tile's A and B loads. */ \+ /* buf-smem_raw gives the buffer offset (0 or BUF_SIZE). Since the */ \+ /* single extern __shared__ allocation starts at LDS offset 0, these */ \+ /* integer offsets are the ds_read_b128 address VGPR values directly. */ \+ const unsigned _buf_off = static_cast<unsigned>((buf) - smem_raw); \+ const unsigned _a_base = _buf_off; \+ const unsigned _b_base = _buf_off + static_cast<unsigned>(Lds::A_SIZE); \+ const unsigned _addr_a0 = _a_base + static_cast<unsigned>(Lds::a_idx(wave_m, lrow, 0 * 4 + kgrp)); \+ const unsigned _addr_b0 = _b_base + static_cast<unsigned>(Lds::b_idx(wave_n, 0, kgrp, lrow)); \+ const unsigned _addr_a1 = _a_base + static_cast<unsigned>(Lds::a_idx(wave_m, lrow, 1 * 4 + kgrp)); \+ const unsigned _addr_b1 = _b_base + static_cast<unsigned>(Lds::b_idx(wave_n, 1, kgrp, lrow)); \+ const unsigned _addr_a2 = _a_base + static_cast<unsigned>(Lds::a_idx(wave_m, lrow, 2 * 4 + kgrp)); \+ const unsigned _addr_b2 = _b_base + static_cast<unsigned>(Lds::b_idx(wave_n, 2, kgrp, lrow)); \+ const unsigned _addr_a3 = _a_base + static_cast<unsigned>(Lds::a_idx(wave_m, lrow, 3 * 4 + kgrp)); \+ const unsigned _addr_b3 = _b_base + static_cast<unsigned>(Lds::b_idx(wave_n, 3, kgrp, lrow)); \+ /* Pre-load all 8 scale values */ \+ int _sa0, _sb0, _sa1, _sb1, _sa2, _sb2, _sa3, _sb3; \+ if constexpr (A_VALID) { \+ auto _sa_base = as_row + (chunk_ks) * 4; \+ auto _sa_vec = *reinterpret_cast<const int4_v*>(_sa_base); \+ auto _sa_shift = kgrp * 8; \+ _sa0 = (reinterpret_cast<const uint32_t*>(&_sa_vec)[0] >> _sa_shift) & 0xFF; \+ _sa1 = (reinterpret_cast<const uint32_t*>(&_sa_vec)[1] >> _sa_shift) & 0xFF; \+ _sa2 = (reinterpret_cast<const uint32_t*>(&_sa_vec)[2] >> _sa_shift) & 0xFF; \+ _sa3 = (reinterpret_cast<const uint32_t*>(&_sa_vec)[3] >> _sa_shift) & 0xFF; \+ } else { \+ auto _ks0 = (chunk_ks) * 4 + kgrp; \+ _sa0 = (a_rt & (_ks0 < KS)) ? static_cast<int>(as_row[_ks0]) : 127; \+ _sa1 = (a_rt & (_ks0+4 < KS)) ? static_cast<int>(as_row[_ks0+4]) : 127; \+ _sa2 = (a_rt & (_ks0+8 < KS)) ? static_cast<int>(as_row[_ks0+8]) : 127; \+ _sa3 = (a_rt & (_ks0+12 < KS)) ? static_cast<int>(as_row[_ks0+12]) : 127; \} \- __builtin_amdgcn_sched_barrier(0); \- asm volatile("s_setprio 0" ::: "memory"); \- __builtin_amdgcn_sched_barrier(0); \+ { \+ auto _ksv0 = (chunk_ks); \+ auto* _bp0 = Bssh + bssh_base + (_ksv0 & 1) * 2 + (_ksv0 >> 1) * 256; \+ auto* _bp1 = Bssh + bssh_base + ((_ksv0+1) & 1) * 2 + ((_ksv0+1) >> 1) * 256; \+ auto* _bp2 = Bssh + bssh_base + ((_ksv0+2) & 1) * 2 + ((_ksv0+2) >> 1) * 256; \+ auto* _bp3 = Bssh + bssh_base + ((_ksv0+3) & 1) * 2 + ((_ksv0+3) >> 1) * 256; \+ if constexpr (B_VALID) { \+ _sb0 = static_cast<int>(*_bp0); \+ _sb1 = static_cast<int>(*_bp1); \+ _sb2 = static_cast<int>(*_bp2); \+ _sb3 = static_cast<int>(*_bp3); \+ } else { \+ auto _ks0 = (chunk_ks) * 4 + kgrp; \+ _sb0 = (b_rt & (_ks0 < KS)) ? static_cast<int>(*_bp0) : 127; \+ _sb1 = (b_rt & (_ks0+4 < KS)) ? static_cast<int>(*_bp1) : 127; \+ _sb2 = (b_rt & (_ks0+8 < KS)) ? static_cast<int>(*_bp2) : 127; \+ _sb3 = (b_rt & (_ks0+12 < KS)) ? static_cast<int>(*_bp3) : 127; \+ } \+ } \+ /* Inline ASM: software-pipelined ds_read_b128 + v_mfma_scale */ \+ /* Double-buffered: even set (ae,be) and odd set (ao,bo) */ \+ /* Pipeline: */ \+ /* Issue loads for kt0,kt1 → wait kt0 → MFMA kt0 */ \+ /* Issue loads for kt2 → wait kt1 → MFMA kt1 */ \+ /* Issue loads for kt3 → wait kt2 → MFMA kt2 */ \+ /* wait kt3 → MFMA kt3 */ \+ /* Each MFMA executes for 16 cycles, hiding 2 ds_read_b128 latency */ \+ int4_v _ae, _be, _ao, _bo; \+ __builtin_amdgcn_sched_barrier(0x020); \+ asm volatile( \+ "s_setprio 3 \n\t" \+ /* --- Issue 4 LDS reads for kt=0 (even) and kt=1 (odd) --- */ \+ "ds_read_b128 %[ae], %[aa0] \n\t" \+ "ds_read_b128 %[be], %[ab0] \n\t" \+ "ds_read_b128 %[ao], %[aa1] \n\t" \+ "ds_read_b128 %[bo], %[ab1] \n\t" \+ /* Wait for kt=0 data (even set); 2 still in flight (ao,bo) */ \+ "s_waitcnt lgkmcnt(2) \n\t" \+ /* --- MFMA kt=0: consume even set --- */ \+ "v_mfma_scale_f32_16x16x128_f8f6f4 %[acc], %[ae], %[be], %[acc], %[sa0], %[sb0] op_sel_hi:[0,0,0] cbsz:4 blgp:4\n\t" \+ /* --- Issue 2 LDS reads for kt=2 (reuse even set) --- */ \+ "ds_read_b128 %[ae], %[aa2] \n\t" \+ "ds_read_b128 %[be], %[ab2] \n\t" \+ /* Wait for kt=1 data (odd set); 2 still in flight (ae,be) */ \+ "s_waitcnt lgkmcnt(2) \n\t" \+ /* --- MFMA kt=1: consume odd set --- */ \+ "v_mfma_scale_f32_16x16x128_f8f6f4 %[acc], %[ao], %[bo], %[acc], %[sa1], %[sb1] op_sel_hi:[0,0,0] cbsz:4 blgp:4\n\t" \+ /* --- Issue 2 LDS reads for kt=3 (reuse odd set) --- */ \+ "ds_read_b128 %[ao], %[aa3] \n\t" \+ "ds_read_b128 %[bo], %[ab3] \n\t" \+ /* Wait for kt=2 data (even set); 2 still in flight (ao,bo) */ \+ "s_waitcnt lgkmcnt(2) \n\t" \+ /* --- MFMA kt=2: consume even set --- */ \+ "v_mfma_scale_f32_16x16x128_f8f6f4 %[acc], %[ae], %[be], %[acc], %[sa2], %[sb2] op_sel_hi:[0,0,0] cbsz:4 blgp:4\n\t" \+ /* Wait for kt=3 data (odd set); 0 in flight */ \+ "s_waitcnt lgkmcnt(0) \n\t" \+ /* --- MFMA kt=3: consume odd set --- */ \+ "v_mfma_scale_f32_16x16x128_f8f6f4 %[acc], %[ao], %[bo], %[acc], %[sa3], %[sb3] op_sel_hi:[0,0,0] cbsz:4 blgp:4\n\t" \+ "s_setprio 0 \n\t" \+ : [acc] "+v"(acc), \+ [ae] "=&v"(_ae), [be] "=&v"(_be), \+ [ao] "=&v"(_ao), [bo] "=&v"(_bo) \+ : [aa0] "v"(_addr_a0), [ab0] "v"(_addr_b0), \+ [aa1] "v"(_addr_a1), [ab1] "v"(_addr_b1), \+ [aa2] "v"(_addr_a2), [ab2] "v"(_addr_b2), \+ [aa3] "v"(_addr_a3), [ab3] "v"(_addr_b3), \+ [sa0] "v"(_sa0), [sb0] "v"(_sb0), \+ [sa1] "v"(_sa1), [sb1] "v"(_sb1), \+ [sa2] "v"(_sa2), [sb2] "v"(_sb2), \+ [sa3] "v"(_sa3), [sb3] "v"(_sb3) \+ ); \+ __builtin_amdgcn_sched_barrier(0x020); \}- // ── Prologue: load first chunk ───────────────────────────────────────auto cur_ks = ks_start;ISSUE_LOADS(cur_ks);asm volatile("s_waitcnt vmcnt(0) lgkmcnt(0)" ::: "memory");⋯ 3 unchanged linesauto* cur_buf = buf0;auto* nxt_buf = buf1;- // ── Main pipelined loop (unroll 2 to reduce register pressure) ────────#pragma unrollfor (auto c = 0; c < NUM_CHUNKS - 1; ++c) {ISSUE_LOADS(cur_ks + CHUNK_K);⋯ 8 unchanged linescur_ks += CHUNK_K;}- // ── Epilogue: compute last full chunk ────────────────────────────────COMPUTE_CHUNK(cur_buf, cur_ks);- // ── Tail: remaining K-tiles if CKT not divisible by CHUNK_K ─────────if constexpr (TAIL_KT > 0) {cur_ks += CHUNK_K;ISSUE_LOADS(cur_ks);⋯ 30 unchanged linesif (!tile_valid) return;}- // ── Store output (shared by both paths) ─────────────────────────────────const auto out_col = tile_n + lrow;const auto out_row_base = tile_m + kgrp * 4;⋯ 19 unchanged lines}}- // ── Reduce kernel — fully templated ─────────────────────────────────────────template<int C_N, int C_NUM_KSPLIT, int C_M>__global__ void mxfp4_reduce(⋯ 21 unchanged linesC_out[mn] = r;}- // ═══════════════════════════════════════════════════════════════════════════// Fused Quant+GEMM kernel for PATH A (K=512, small M).// Single kernel launch: loads bf16 A into LDS, quants to FP4 in registers,// then iterates N-tiles reusing A from registers. Eliminates quant kernel.- // ═══════════════════════════════════════════════════════════════════════════template<int C_M, int C_K, int C_N, int C_SCALEN>__global__ void __launch_bounds__(((C_M + 15) / 16) * 4 * 64, 2)⋯ 25 unchanged linesconst auto tile_n = static_cast<int>(blockIdx.x) * BN + wave_n * 16;const auto my_row = tile_m + lrow; // actual M-row this lane handles- // ── Phase 1: Cooperative coalesced load A_bf16 -> LDS ──────────────────constexpr auto A_BF16_BYTES = C_M * C_K * 2;extern __shared__ uint8_t smem[];{⋯ 12 unchanged lines}__syncthreads();- // ── Phase 2: Per-lane quant from LDS -> registers ──────────────────────// Each lane handles row=lrow, and quants CKT groups (one per k-tile)// kgrp selects which 32-element group within the k-tileint4_v a_fp4_regs[CKT];⋯ 57 unchanged lines}}- // ── Phase 3: MFMA — A from registers, B from global ───────────────────// tile_n < C_N guaranteed by grid dimensions (C_N / BN blocks)float4_v acc{0.f, 0.f, 0.f, 0.f};⋯ 22 unchanged lines0, a_scale_regs[kt], 0, sb);}- // ── Phase 4: Store bf16 output ────────────────────────────────────────const auto out_col = tile_n + lrow;const auto out_row_base = tile_m + kgrp * 4;if (out_col >= C_N) return;⋯ 6 unchanged lines}}- // ═══════════════════════════════════════════════════════════════════════════// Fused Quant+GEMM for splitK — each split block quants its K-slice of A.// Eliminates separate quant kernel launch. Still needs reduce kernel.- // ═══════════════════════════════════════════════════════════════════════════template<int C_M, int C_K, int C_N, int C_SCALEN, int C_KPS, int C_NUM_KSPLIT>__global__ void __launch_bounds__(((C_M + 15) / 16) * 4 * 64, 2)⋯ 27 unchanged linesconst auto ks_start = ks_idx * C_KPS;const auto k_byte_start = ks_start * 128; // in bf16 elements (128 per k-tile = 64 bytes FP4 = 256 bf16 bytes for 128 elements)- // ── Phase 1: Load A_bf16 K-slice into LDS ──────────────────────────────// For M=16, CKT=4: need 16 rows × 512 bf16 = 16KBconstexpr auto A_SLICE_ELEMS = C_M * C_KPS * 128; // bf16 elements in this K-sliceconstexpr auto A_SLICE_BYTES = A_SLICE_ELEMS * 2;⋯ 22 unchanged lines}__syncthreads();- // ── Phase 2: Quant from LDS to registers ───────────────────────────────int4_v a_fp4_regs[CKT];int a_scale_regs[CKT];constexpr auto a_row_always_valid = (C_M >= ((C_M + 15) / 16) * 16);⋯ 54 unchanged lines}}- // ── Phase 3: MFMA ──────────────────────────────────────────────────────// tile_n < C_N guaranteed by grid dimensions (C_N / BN blocks)float4_v acc{0.f, 0.f, 0.f, 0.f};⋯ 21 unchanged lines0, a_scale_regs[kt], 0, sb);}- // ── Phase 4: Store f32 partial ─────────────────────────────────────────const auto out_col = tile_n + lrow;const auto out_row_base = tile_m + kgrp * 4;if (out_col >= C_N) return;⋯ 6 unchanged lines}}- // ── Launch helpers — fully templated on M ───────────────────────────────────template<int C_M, int C_K>void launch_quant(⋯ 237 unchanged linesauto total_ktiles = K / 128;if (total_ktiles < 28) return 1;if (M <= 8) return 7;- if (M <= 16) return 14;+ if (M <= 16) return 28;return 1;}
scrolls · 445 diff lines total
Best evidence level for this revision: reported
JSON