submission 649634
XiaomingFun233 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 2695 lines, June 9 Researcher Reciprocity License v1.0.
submission_amd_mxfp4_mm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-649634?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:82b8c28ecb0842294c733c27649106a7ef76fb35c25f9c52a078832cd77b9af7
license declaredunknown
license concludedunknown
authorsXiaomingFun233
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
TORCH_CHECK(a_fp4_u8.dim() == 2 && b_u8.dim() == 2, "A/B FP4 must be 2D");shared-memory
__device__ __forceinline__ int smem_swizzled_dword(int logical_row, int logical_dword) {split-k
_SPLITK_ENV = "MXFP4_MM_LOG2_SPLITK"stages = 3
constexpr int PREFETCH_STAGES = 3;tile-k = 32
constexpr int BLOCK_K = 32;tile-m = 64
constexpr int GEMM_BLOCK_M = 64;tile-n = 64
constexpr int GEMM_BLOCK_N = 64;vector-width = uint4
constexpr int LDS_VEC_BYTES = sizeof(uint4);Kernel source
submission_amd_mxfp4_mm.py2695 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
import os
import threading
from dataclasses import dataclass
from typing import Dict, Tuple
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
_B_LAYOUT_ENV = "MXFP4_MM_B_LAYOUT" # raw | shuffle | auto
_B_SHUFFLE_INNER_ENV = "MXFP4_MM_B_SHUFFLE_INNER_MODE" # 0 | 1
_SPLITK_ENV = "MXFP4_MM_LOG2_SPLITK"
_NATIVE_MFMA_ENV = "MXFP4_MM_NATIVE_MFMA" # never | auto | force (force also enables native build)
_FUSED_A_QUANT_ENV = "MXFP4_MM_FUSED_A_QUANT" # 0 | 1
_DEFAULT_B_LAYOUT = "shuffle"
_DEFAULT_B_SHUFFLE_INNER_MODE = 1 # Matches aiter.ops.shuffle.shuffle_weight() tile flattening.
_DEFAULT_NATIVE_MFMA = "auto"
_DEFAULT_FUSED_A_QUANT = "0"
_HIP_LOCK = threading.Lock()
_HIP_MODULE = None
_HIP_BUILD_ERROR = None
_HIP_NATIVE_BUILD_ENABLED = False
_HIP_BUILD_REPORT_EMITTED = False
_RUNTIME_REPORT_LOCK = threading.Lock()
_RUNTIME_REPORTED_SHAPES: set[Tuple[int, int, int]] = set()
_B_SCALE_RAW_LOCK = threading.Lock()
_B_SCALE_RAW_CACHE: Dict[Tuple[int, ...], torch.Tensor] = {}
_B_SCALE_RAW_CACHE_MAX = 8
_WORKSPACE_LOCK = threading.Lock()
_WORKSPACE_CACHE: Dict[Tuple[int, ...], torch.Tensor] = {}
_WORKSPACE_CACHE_MAX = 16
_STATIC_SPLITK_LOG2: Dict[Tuple[int, int, int], int] = {
(4, 2880, 512): 2,
(16, 2112, 7168): 3,
(32, 4096, 512): 2,
(32, 2880, 512): 2,
(64, 7168, 2048): 1,
(256, 3072, 1536): 0,
}
_RANKED_SHAPE_KEYS = frozenset(_STATIC_SPLITK_LOG2)
@dataclass(frozen=True)
class _LaunchPolicy:
log2_k_split: int
use_native_prequant: int
use_native_fused: int
is_ranked_shape: bool
CPP_WRAPPER = r"""
#include <cstdint>
#include <vector>
std::vector<torch::Tensor> hip_quant_mxfp4(torch::Tensor x);
torch::Tensor hip_gemm_mxfp4_prequant(
torch::Tensor a_fp4_u8,
torch::Tensor b_u8,
torch::Tensor a_scale_u8,
torch::Tensor b_scale_u8,
int64_t layout_mode,
int64_t log2_k_split,
torch::Tensor workspace,
int64_t workspace_stride,
int64_t b_shuffle_inner_mode,
int64_t native_mfma_mode);
torch::Tensor hip_gemm_mxfp4(
torch::Tensor a_bf16,
torch::Tensor b_u8,
torch::Tensor b_scale_u8,
int64_t layout_mode,
int64_t log2_k_split,
torch::Tensor workspace,
int64_t workspace_stride,
int64_t b_shuffle_inner_mode,
int64_t native_mfma_mode);
"""
HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <algorithm>
#include <cmath>
#include <cstdint>
#include <vector>
#ifndef MXFP4_ENABLE_NATIVE_FP4_MFMA
#define MXFP4_ENABLE_NATIVE_FP4_MFMA 0
#endif
namespace {
constexpr int BLOCK_K = 32;
constexpr int NATIVE_MFMA_K_BLOCKS = 4;
constexpr int PAD_M_KERNEL = 64;
constexpr int PAD_M_SCALE = 256;
constexpr int PAD_SCALE_N = 8;
constexpr int GEMM_BLOCK_M = 64;
constexpr int GEMM_BLOCK_N = 64;
constexpr int GEMM_THREADS = 256;
constexpr int WAVE_SIZE = 64;
constexpr int WAVE_TILE_M = 32;
constexpr int WAVE_TILE_N = 32;
constexpr int MFMA_TILE_M = 16;
constexpr int MFMA_TILE_N = 16;
constexpr int PACKED_BLOCK_K = BLOCK_K / 2;
constexpr int SMEM_DWORDS = PACKED_BLOCK_K / static_cast<int>(sizeof(uint32_t));
constexpr int LDS_PAD_DWORDS = 1;
constexpr int SMEM_STRIDE_DWORDS = SMEM_DWORDS + LDS_PAD_DWORDS;
constexpr int SMEM_STRIDE = SMEM_STRIDE_DWORDS * static_cast<int>(sizeof(uint32_t));
constexpr int SMEM_SWIZZLE_MASK = SMEM_DWORDS - 1;
constexpr int PREFETCH_STAGES = 3;
constexpr int LDS_VEC_BYTES = sizeof(uint4);
constexpr int A_TILE_VEC_LOADS = (GEMM_BLOCK_M * PACKED_BLOCK_K) / LDS_VEC_BYTES;
constexpr int B_TILE_VEC_LOADS = (GEMM_BLOCK_N * PACKED_BLOCK_K) / LDS_VEC_BYTES;
__device__ __forceinline__ float bf16_to_f32(uint16_t x) {
return __uint_as_float(static_cast<uint32_t>(x) << 16);
}
__device__ __forceinline__ uint8_t quantize_e2m1(float x_scaled) {
// Match aiter dynamic_mxfp4_quant conversion path.
uint32_t qx = __float_as_uint(x_scaled);
uint32_t s = qx & 0x80000000u;
uint32_t e = (qx >> 23) & 0xFFu;
uint32_t m = qx & 0x7FFFFFu;
if (e < 127u) {
uint32_t adjusted_exponents = 127u - (e + 1u);
uint32_t denorm_m = 0x400000u | (m >> 1);
m = (adjusted_exponents >= 32u) ? 0u : (denorm_m >> adjusted_exponents);
}
e = ((e > 126u) ? e : 126u) - 126u;
uint32_t e2m1_tmp = ((((e << 2) | (m >> 21)) + 1u) >> 1);
if (e2m1_tmp > 0x7u) {
e2m1_tmp = 0x7u;
}
return static_cast<uint8_t>((s >> 28) | e2m1_tmp);
}
__device__ __forceinline__ uint8_t amax_to_e8m0_scale(float amax) {
// Match aiter dynamic_mxfp4_quant rounding to power-of-two scale.
uint32_t bits = __float_as_uint(amax);
bits = (bits + 0x200000u) & 0xFF800000u;
int exp_unbiased = static_cast<int>((bits >> 23) & 0xFF) - 127;
int scale_unbiased = exp_unbiased - 2;
if (scale_unbiased < -127) {
scale_unbiased = -127;
}
if (scale_unbiased > 127) {
scale_unbiased = 127;
}
return static_cast<uint8_t>(scale_unbiased + 127);
}
__device__ __forceinline__ int64_t shuffled_scale_offset(
int64_t row,
int64_t col,
int64_t scale_n_pad) {
int64_t bs_offs_0 = row / 32;
int64_t bs_offs_1 = row % 32;
int64_t bs_offs_2 = bs_offs_1 % 16;
bs_offs_1 = bs_offs_1 / 16;
int64_t bs_offs_3 = col / 8;
int64_t bs_offs_4 = col % 8;
int64_t bs_offs_5 = bs_offs_4 % 4;
bs_offs_4 = bs_offs_4 / 4;
return bs_offs_1 + bs_offs_4 * 2 + bs_offs_2 * 4 + bs_offs_5 * 64 +
bs_offs_3 * 256 + bs_offs_0 * 32 * scale_n_pad;
}
__device__ __forceinline__ int64_t get_b_scale_shuffled_offset(
int64_t n,
int64_t k_blk,
int64_t scale_k_pad) {
int64_t n_outer = n / 32;
int64_t n_inner = n % 32;
int64_t k_outer = k_blk / 8;
int64_t k_inner = k_blk % 8;
int64_t n_16 = n_inner % 16;
int64_t n_2 = n_inner / 16;
int64_t k_4 = k_inner % 4;
int64_t k_2 = k_inner / 4;
return n_2 + (k_2 * 2) + (n_16 * 4) + (k_4 * 64) + (k_outer * 256) +
(n_outer * 32 * scale_k_pad);
}
__device__ __forceinline__ int64_t get_b_fp4_shuffled_offset(
int64_t n,
int64_t k_fp4,
int64_t k_fp4_pad,
int64_t inner_mode) {
int64_t n_blk = n / 16;
int64_t k_blk = k_fp4 / 16;
int64_t n_in = n % 16;
int64_t k_in = k_fp4 % 16;
int64_t blk_stride = k_fp4_pad / 16;
int64_t blk_offset = (n_blk * blk_stride + k_blk) * 256;
int64_t inner_offset = (inner_mode == 0) ? (k_in * 16 + n_in) : (n_in * 16 + k_in);
return blk_offset + inner_offset;
}
__device__ __forceinline__ float e8m0_to_f32_fast(uint8_t e) {
if (e == 0) {
return __uint_as_float(0x00400000u);
}
if (e == 0xFF) {
return __uint_as_float(0x7F800001u);
}
return __uint_as_float(static_cast<uint32_t>(e) << 23);
}
using floatx4 = __attribute__((__vector_size__(4 * sizeof(float)))) float;
using bit16x4 = __attribute__((__vector_size__(4 * sizeof(uint16_t)))) uint16_t;
using bit16x8 = __attribute__((__vector_size__(8 * sizeof(uint16_t)))) uint16_t;
using int32x8 = __attribute__((__vector_size__(8 * sizeof(int32_t)))) int32_t;
struct B16x8 {
bit16x4 xy[2];
};
__device__ __forceinline__ floatx4 gcn_mfma16x16x32_bf16(
const B16x8& a,
const B16x8& b,
const floatx4& c) {
#if defined(__gfx950__)
bit16x8 ta = __builtin_shufflevector(a.xy[0], a.xy[1], 0, 1, 2, 3, 4, 5, 6, 7);
bit16x8 tb = __builtin_shufflevector(b.xy[0], b.xy[1], 0, 1, 2, 3, 4, 5, 6, 7);
return __builtin_amdgcn_mfma_f32_16x16x32_bf16(ta, tb, c, 0, 0, 0);
#else
return c;
#endif
}
__device__ __forceinline__ int32x8 zero_int32x8() {
return {0, 0, 0, 0, 0, 0, 0, 0};
}
template <int ScaleIdxA, int ScaleIdxB>
__device__ __forceinline__ floatx4 gcn_mfma16x16x128_f4_scale(
const int32x8& a,
const int32x8& b,
const floatx4& c,
int packed_scale_a,
int packed_scale_b) {
#if defined(__gfx950__) && MXFP4_ENABLE_NATIVE_FP4_MFMA
// LLVM AMDGPU docs: the last four operands are
// scale_idx_a, scale_values_a, scale_idx_b, scale_values_b
// where scale_values_* is a per-lane VGPR value and scale_idx_* is the
// wave-uniform 2-bit byte selector. The selected byte across all 64 lanes
// forms the 64 scale entries consumed by one scaled MFMA instruction.
// ROCm/clang also requires scale_idx_* to be compile-time immediates.
// For FP4 x FP4, both operand format codes are 4.
static_assert(0 <= ScaleIdxA && ScaleIdxA < 4, "ScaleIdxA must be in [0, 3]");
static_assert(0 <= ScaleIdxB && ScaleIdxB < 4, "ScaleIdxB must be in [0, 3]");
return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a, b, c, 4, 4, ScaleIdxA, packed_scale_a, ScaleIdxB, packed_scale_b);
#else
return c;
#endif
}
__device__ __forceinline__ int pack_e8m0x4(
uint8_t s0,
uint8_t s1,
uint8_t s2,
uint8_t s3) {
return static_cast<int>(
static_cast<uint32_t>(s0) |
(static_cast<uint32_t>(s1) << 8) |
(static_cast<uint32_t>(s2) << 16) |
(static_cast<uint32_t>(s3) << 24));
}
__device__ __forceinline__ int splat_e8m0(uint8_t s) {
uint32_t v = static_cast<uint32_t>(s);
return static_cast<int>(v | (v << 8) | (v << 16) | (v << 24));
}
__device__ __forceinline__ int smem_swizzled_dword(int logical_row, int logical_dword) {
return logical_dword ^ (logical_row & SMEM_SWIZZLE_MASK);
}
__device__ __forceinline__ int smem_swizzled_byte_offset(int logical_row, int logical_byte) {
int logical_dword = logical_byte >> 2;
int byte_in_dword = logical_byte & 3;
int physical_dword = smem_swizzled_dword(logical_row, logical_dword);
return logical_row * SMEM_STRIDE + physical_dword * static_cast<int>(sizeof(uint32_t)) +
byte_in_dword;
}
__device__ __forceinline__ void store_uint4_to_smem_swizzled(
uint8_t* smem_dst,
int logical_row,
const uint4& v) {
uint32_t* row_ptr = reinterpret_cast<uint32_t*>(smem_dst + logical_row * SMEM_STRIDE);
row_ptr[smem_swizzled_dword(logical_row, 0)] = v.x;
row_ptr[smem_swizzled_dword(logical_row, 1)] = v.y;
row_ptr[smem_swizzled_dword(logical_row, 2)] = v.z;
row_ptr[smem_swizzled_dword(logical_row, 3)] = v.w;
}
__device__ __forceinline__ void store_u8_to_smem_swizzled(
uint8_t* smem_dst,
int logical_row,
int logical_byte,
uint8_t value) {
smem_dst[smem_swizzled_byte_offset(logical_row, logical_byte)] = value;
}
__device__ __forceinline__ uint32_t load_u32_from_smem_swizzled(
const uint8_t* smem_src,
int logical_row,
int logical_byte) {
const uint32_t* row_ptr =
reinterpret_cast<const uint32_t*>(smem_src + logical_row * SMEM_STRIDE);
return row_ptr[smem_swizzled_dword(logical_row, logical_byte >> 2)];
}
template <bool WRITE_BF16_OUT>
__device__ __forceinline__ void store_output_value(
float* workspace,
__hip_bfloat16* out,
int64_t workspace_stride,
int64_t n,
int64_t split_idx,
int64_t row,
int64_t col,
float value) {
int64_t out_off = row * n + col;
if constexpr (WRITE_BF16_OUT) {
out[out_off] = static_cast<__hip_bfloat16>(value);
} else {
workspace[split_idx * workspace_stride + out_off] = value;
}
}
template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE>
__device__ __forceinline__ uint4 load_b_tile_vec(
const uint8_t* b_u8,
int tid,
int64_t tile_n,
int64_t n,
int64_t k2,
int64_t k_base) {
uint4 v = {0u, 0u, 0u, 0u};
if constexpr (!SHUFFLED_B) {
int64_t global_col = tile_n + tid;
if (global_col < n) {
v = *reinterpret_cast<const uint4*>(b_u8 + global_col * k2 + k_base);
}
} else if constexpr (SHUFFLE_INNER_MODE == 1) {
int64_t global_col = tile_n + tid;
if (global_col < n) {
int64_t boff = get_b_fp4_shuffled_offset(global_col, k_base, k2, 1);
v = *reinterpret_cast<const uint4*>(b_u8 + boff);
}
} else {
int col_group = tid / PACKED_BLOCK_K;
int k_in = tid % PACKED_BLOCK_K;
int local_col_base = col_group * 16;
int64_t global_col_base = tile_n + local_col_base;
uint8_t* v_bytes = reinterpret_cast<uint8_t*>(&v);
if (global_col_base + 15 < n) {
int64_t boff = get_b_fp4_shuffled_offset(global_col_base, k_base + k_in, k2, 0);
v = *reinterpret_cast<const uint4*>(b_u8 + boff);
} else {
#pragma unroll
for (int c = 0; c < 16; ++c) {
int64_t global_col = global_col_base + c;
if (global_col < n) {
int64_t boff = get_b_fp4_shuffled_offset(global_col, k_base + k_in, k2, 0);
v_bytes[c] = b_u8[boff];
}
}
}
}
return v;
}
template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE>
__device__ __forceinline__ void store_b_tile_vec(
uint8_t* smem_dst,
int tid,
const uint4& v) {
if constexpr (!SHUFFLED_B || SHUFFLE_INNER_MODE == 1) {
store_uint4_to_smem_swizzled(smem_dst, tid, v);
} else {
int col_group = tid / PACKED_BLOCK_K;
int k_in = tid % PACKED_BLOCK_K;
int local_col_base = col_group * 16;
const uint8_t* src = reinterpret_cast<const uint8_t*>(&v);
#pragma unroll
for (int c = 0; c < 16; ++c) {
store_u8_to_smem_swizzled(smem_dst, local_col_base + c, k_in, src[c]);
}
}
}
template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE>
__device__ __forceinline__ void clear_gemm_stage_smem(
uint8_t* smem_a,
uint8_t* smem_b,
uint8_t* smem_a_scale,
uint8_t* smem_b_scale,
int tid) {
const uint4 zero = {0u, 0u, 0u, 0u};
if (tid < A_TILE_VEC_LOADS) {
store_uint4_to_smem_swizzled(smem_a, tid, zero);
smem_a_scale[tid] = 127;
}
if (tid < B_TILE_VEC_LOADS) {
store_b_tile_vec<SHUFFLED_B, SHUFFLE_INNER_MODE>(smem_b, tid, zero);
}
if (tid < GEMM_BLOCK_N) {
smem_b_scale[tid] = 127;
}
}
template <bool SHUFFLED_B>
__device__ __forceinline__ uint8_t load_b_scale_code(
const uint8_t* b_scale_u8,
int64_t global_col,
int64_t kb,
int64_t n,
int64_t k_blocks_pad) {
if (global_col >= n) {
return 127;
}
if constexpr (SHUFFLED_B) {
int64_t soff = get_b_scale_shuffled_offset(global_col, kb, k_blocks_pad);
return b_scale_u8[soff];
} else {
return b_scale_u8[global_col * k_blocks_pad + kb];
}
}
__device__ __forceinline__ uint16_t fp4_e2m1_to_bf16_bits(uint8_t nib) {
uint16_t sign = static_cast<uint16_t>(nib & 0x08u) << 12;
uint16_t mag = static_cast<uint16_t>(nib & 0x07u);
if (mag == 0) {
return sign;
}
uint16_t exponent = static_cast<uint16_t>(126 + (mag >> 1));
uint16_t mantissa = ((mag > 1) && (mag & 0x1u)) ? 0x40u : 0u;
return static_cast<uint16_t>(sign | (exponent << 7) | mantissa);
}
__device__ __forceinline__ B16x8 unpack_fp4x8_to_bf16(uint32_t packed) {
B16x8 reg;
const uint8_t* bytes = reinterpret_cast<const uint8_t*>(&packed);
#pragma unroll
for (int i = 0; i < 8; ++i) {
uint8_t byte = bytes[i >> 1];
uint8_t nib = (i & 1) ? static_cast<uint8_t>(byte >> 4)
: static_cast<uint8_t>(byte & 0x0Fu);
uint16_t bits = fp4_e2m1_to_bf16_bits(nib);
if (i < 4) {
reg.xy[0][i] = bits;
} else {
reg.xy[1][i - 4] = bits;
}
}
return reg;
}
__device__ __forceinline__ void quantize_a_row_block_to_fp4(
const __hip_bfloat16* a_bf16,
int64_t global_row,
int64_t k,
int64_t k_block,
int64_t m,
uint4& out_fp4,
uint8_t& out_scale) {
out_fp4 = {0u, 0u, 0u, 0u};
out_scale = 127;
if (global_row >= m) {
return;
}
const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(a_bf16);
const uint8_t* row_ptr =
a_bytes + (global_row * k + k_block * BLOCK_K) * static_cast<int64_t>(sizeof(__hip_bfloat16));
uint4 chunks[4];
#pragma unroll
for (int vec_idx = 0; vec_idx < 4; ++vec_idx) {
chunks[vec_idx] = reinterpret_cast<const uint4*>(row_ptr)[vec_idx];
}
float amax = 0.0f;
#pragma unroll
for (int vec_idx = 0; vec_idx < 4; ++vec_idx) {
uint32_t words[4] = {chunks[vec_idx].x, chunks[vec_idx].y, chunks[vec_idx].z, chunks[vec_idx].w};
#pragma unroll
for (int word_idx = 0; word_idx < 4; ++word_idx) {
uint32_t word = words[word_idx];
float f0 = bf16_to_f32(static_cast<uint16_t>(word & 0xFFFFu));
float f1 = bf16_to_f32(static_cast<uint16_t>((word >> 16) & 0xFFFFu));
amax = fmaxf(amax, fabsf(f0));
amax = fmaxf(amax, fabsf(f1));
}
}
out_scale = amax_to_e8m0_scale(amax);
float inv_scale = 1.0f / e8m0_to_f32_fast(out_scale);
uint8_t* out_bytes = reinterpret_cast<uint8_t*>(&out_fp4);
#pragma unroll
for (int vec_idx = 0; vec_idx < 4; ++vec_idx) {
uint32_t words[4] = {chunks[vec_idx].x, chunks[vec_idx].y, chunks[vec_idx].z, chunks[vec_idx].w};
uint8_t vals[8];
#pragma unroll
for (int word_idx = 0; word_idx < 4; ++word_idx) {
uint32_t word = words[word_idx];
vals[word_idx * 2 + 0] =
quantize_e2m1(bf16_to_f32(static_cast<uint16_t>(word & 0xFFFFu)) * inv_scale);
vals[word_idx * 2 + 1] = quantize_e2m1(
bf16_to_f32(static_cast<uint16_t>((word >> 16) & 0xFFFFu)) * inv_scale);
}
#pragma unroll
for (int pair_idx = 0; pair_idx < 4; ++pair_idx) {
uint8_t lo = vals[pair_idx * 2 + 0];
uint8_t hi = vals[pair_idx * 2 + 1];
out_bytes[vec_idx * 4 + pair_idx] = static_cast<uint8_t>((hi << 4) | lo);
}
}
}
template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE>
__device__ __forceinline__ void load_gemm_stage_to_smem(
const __hip_bfloat16* a_bf16,
const uint8_t* b_u8,
const uint8_t* b_scale_u8,
uint8_t* smem_a,
uint8_t* smem_b,
uint8_t* smem_a_scale,
uint8_t* smem_b_scale,
int tid,
int64_t tile_m,
int64_t tile_n,
int64_t m,
int64_t n,
int64_t k,
int64_t k2,
int64_t global_kb,
int64_t k_blocks_pad) {
const int64_t a_k_base = global_kb * PACKED_BLOCK_K;
if (tid < A_TILE_VEC_LOADS) {
int64_t global_row = tile_m + tid;
uint4 a_vec = {0u, 0u, 0u, 0u};
uint8_t a_scale = 127;
quantize_a_row_block_to_fp4(a_bf16, global_row, k, global_kb, m, a_vec, a_scale);
store_uint4_to_smem_swizzled(smem_a, tid, a_vec);
smem_a_scale[tid] = a_scale;
}
if (tid < B_TILE_VEC_LOADS) {
uint4 v = load_b_tile_vec<SHUFFLED_B, SHUFFLE_INNER_MODE>(b_u8, tid, tile_n, n, k2, a_k_base);
store_b_tile_vec<SHUFFLED_B, SHUFFLE_INNER_MODE>(smem_b, tid, v);
}
if (tid < GEMM_BLOCK_N) {
int64_t global_col = tile_n + tid;
smem_b_scale[tid] =
load_b_scale_code<SHUFFLED_B>(b_scale_u8, global_col, global_kb, n, k_blocks_pad);
}
}
template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE>
__device__ __forceinline__ void load_prequant_stage_to_smem(
const uint8_t* a_fp4,
const uint8_t* b_u8,
const uint8_t* a_scale_u8,
const uint8_t* b_scale_u8,
uint8_t* smem_a,
uint8_t* smem_b,
uint8_t* smem_a_scale,
uint8_t* smem_b_scale,
int tid,
int64_t tile_m,
int64_t tile_n,
int64_t m,
int64_t n,
int64_t k2,
int64_t global_kb,
int64_t k_blocks_pad) {
const int64_t a_k_base = global_kb * PACKED_BLOCK_K;
if (tid < A_TILE_VEC_LOADS) {
uint4 a_vec = {0u, 0u, 0u, 0u};
uint8_t a_scale = 127;
int64_t global_row = tile_m + tid;
if (global_row < m) {
a_vec = *reinterpret_cast<const uint4*>(a_fp4 + global_row * k2 + a_k_base);
a_scale = a_scale_u8[global_row * k_blocks_pad + global_kb];
}
store_uint4_to_smem_swizzled(smem_a, tid, a_vec);
smem_a_scale[tid] = a_scale;
}
if (tid < B_TILE_VEC_LOADS) {
uint4 v = load_b_tile_vec<SHUFFLED_B, SHUFFLE_INNER_MODE>(b_u8, tid, tile_n, n, k2, a_k_base);
store_b_tile_vec<SHUFFLED_B, SHUFFLE_INNER_MODE>(smem_b, tid, v);
}
if (tid < GEMM_BLOCK_N) {
int64_t global_col = tile_n + tid;
smem_b_scale[tid] =
load_b_scale_code<SHUFFLED_B>(b_scale_u8, global_col, global_kb, n, k_blocks_pad);
}
}
__device__ __forceinline__ void accumulate_bf16_stage(
const uint8_t* smem_a,
const uint8_t* smem_b,
const uint8_t* smem_a_scale,
const uint8_t* smem_b_scale,
int row_group,
int local_a_row0,
int local_a_row1,
int local_b_col0,
int local_b_col1,
int local_row_block0,
int local_row_block1,
floatx4& c_acc_00,
floatx4& c_acc_01,
floatx4& c_acc_10,
floatx4& c_acc_11) {
int k_byte_base = row_group * 4;
uint32_t a_pack_0 = load_u32_from_smem_swizzled(smem_a, local_a_row0, k_byte_base);
uint32_t a_pack_1 = load_u32_from_smem_swizzled(smem_a, local_a_row1, k_byte_base);
uint32_t b_pack_0 = load_u32_from_smem_swizzled(smem_b, local_b_col0, k_byte_base);
uint32_t b_pack_1 = load_u32_from_smem_swizzled(smem_b, local_b_col1, k_byte_base);
B16x8 a_reg_0 = unpack_fp4x8_to_bf16(a_pack_0);
B16x8 a_reg_1 = unpack_fp4x8_to_bf16(a_pack_1);
B16x8 b_reg_0 = unpack_fp4x8_to_bf16(b_pack_0);
B16x8 b_reg_1 = unpack_fp4x8_to_bf16(b_pack_1);
floatx4 temp_00 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 temp_01 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 temp_10 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 temp_11 = {0.0f, 0.0f, 0.0f, 0.0f};
temp_00 = gcn_mfma16x16x32_bf16(a_reg_0, b_reg_0, temp_00);
temp_01 = gcn_mfma16x16x32_bf16(a_reg_0, b_reg_1, temp_01);
temp_10 = gcn_mfma16x16x32_bf16(a_reg_1, b_reg_0, temp_10);
temp_11 = gcn_mfma16x16x32_bf16(a_reg_1, b_reg_1, temp_11);
float a_scales_0[4];
float a_scales_1[4];
#pragma unroll
for (int i = 0; i < 4; ++i) {
a_scales_0[i] = e8m0_to_f32_fast(smem_a_scale[local_row_block0 + i]);
a_scales_1[i] = e8m0_to_f32_fast(smem_a_scale[local_row_block1 + i]);
}
float b_scale_0 = e8m0_to_f32_fast(smem_b_scale[local_b_col0]);
float b_scale_1 = e8m0_to_f32_fast(smem_b_scale[local_b_col1]);
#pragma unroll
for (int i = 0; i < 4; ++i) {
c_acc_00[i] += temp_00[i] * (a_scales_0[i] * b_scale_0);
c_acc_01[i] += temp_01[i] * (a_scales_0[i] * b_scale_1);
c_acc_10[i] += temp_10[i] * (a_scales_1[i] * b_scale_0);
c_acc_11[i] += temp_11[i] * (a_scales_1[i] * b_scale_1);
}
}
__device__ __forceinline__ void accumulate_native_group(
const uint8_t smem_a[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_M * SMEM_STRIDE],
const uint8_t smem_b[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_N * SMEM_STRIDE],
const uint8_t smem_a_scale[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_M],
const uint8_t smem_b_scale[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_N],
int row_group,
int local_a_row0,
int local_a_row1,
int local_b_col0,
int local_b_col1,
floatx4& c_acc_00,
floatx4& c_acc_01,
floatx4& c_acc_10,
floatx4& c_acc_11) {
// Fragment contract for __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4:
// each lane contributes exactly one 32-wide K block, selected by
// `row_group = lane / 16`:
// row_group 0 -> K [0, 32)
// row_group 1 -> K [32, 64)
// row_group 2 -> K [64, 96)
// row_group 3 -> K [96, 128)
// The lane's 16 packed FP4 bytes always live in the low 128 bits
// (arg[0..3]) of the int32x8 fragment. The instruction assembles the full
// 128-wide K dimension across all 64 lanes, so we must not shift the data
// into stage-dependent dword slots inside the lane-local register.
//
// The scale operand follows the same distributed contract: each lane passes
// the E8M0 code for its own 32-wide K block in byte0 of the int VGPR.
// CK/Opus both use OPSEL=0 here, and CK notes that current backends read
// byte0 regardless of OPSEL, so only the low byte is semantically relevant.
const int stage_idx = row_group;
int32x8 a_frag_0 = zero_int32x8();
int32x8 a_frag_1 = zero_int32x8();
int32x8 b_frag_0 = zero_int32x8();
int32x8 b_frag_1 = zero_int32x8();
int32_t* a0_words = reinterpret_cast<int32_t*>(&a_frag_0);
int32_t* a1_words = reinterpret_cast<int32_t*>(&a_frag_1);
int32_t* b0_words = reinterpret_cast<int32_t*>(&b_frag_0);
int32_t* b1_words = reinterpret_cast<int32_t*>(&b_frag_1);
#pragma unroll
for (int d = 0; d < 4; ++d) {
int byte_off = d * 4;
a0_words[d] = static_cast<int32_t>(
load_u32_from_smem_swizzled(smem_a[stage_idx], local_a_row0, byte_off));
a1_words[d] = static_cast<int32_t>(
load_u32_from_smem_swizzled(smem_a[stage_idx], local_a_row1, byte_off));
b0_words[d] = static_cast<int32_t>(
load_u32_from_smem_swizzled(smem_b[stage_idx], local_b_col0, byte_off));
b1_words[d] = static_cast<int32_t>(
load_u32_from_smem_swizzled(smem_b[stage_idx], local_b_col1, byte_off));
}
const int packed_a_scale_0 = static_cast<int>(static_cast<uint32_t>(smem_a_scale[stage_idx][local_a_row0]));
const int packed_a_scale_1 = static_cast<int>(static_cast<uint32_t>(smem_a_scale[stage_idx][local_a_row1]));
const int packed_b_scale_0 = static_cast<int>(static_cast<uint32_t>(smem_b_scale[stage_idx][local_b_col0]));
const int packed_b_scale_1 = static_cast<int>(static_cast<uint32_t>(smem_b_scale[stage_idx][local_b_col1]));
c_acc_00 = gcn_mfma16x16x128_f4_scale<0, 0>(
a_frag_0, b_frag_0, c_acc_00, packed_a_scale_0, packed_b_scale_0);
c_acc_01 = gcn_mfma16x16x128_f4_scale<0, 0>(
a_frag_0, b_frag_1, c_acc_01, packed_a_scale_0, packed_b_scale_1);
c_acc_10 = gcn_mfma16x16x128_f4_scale<0, 0>(
a_frag_1, b_frag_0, c_acc_10, packed_a_scale_1, packed_b_scale_0);
c_acc_11 = gcn_mfma16x16x128_f4_scale<0, 0>(
a_frag_1, b_frag_1, c_acc_11, packed_a_scale_1, packed_b_scale_1);
}
__device__ __forceinline__ void accumulate_tail_stages(
const uint8_t smem_a[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_M * SMEM_STRIDE],
const uint8_t smem_b[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_N * SMEM_STRIDE],
const uint8_t smem_a_scale[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_M],
const uint8_t smem_b_scale[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_N],
int tail_k_blocks,
int row_group,
int local_a_row0,
int local_a_row1,
int local_b_col0,
int local_b_col1,
int local_row_block0,
int local_row_block1,
floatx4& c_acc_00,
floatx4& c_acc_01,
floatx4& c_acc_10,
floatx4& c_acc_11) {
#pragma unroll
for (int stage = 0; stage < NATIVE_MFMA_K_BLOCKS; ++stage) {
if (stage < tail_k_blocks) {
accumulate_bf16_stage(
smem_a[stage],
smem_b[stage],
smem_a_scale[stage],
smem_b_scale[stage],
row_group,
local_a_row0,
local_a_row1,
local_b_col0,
local_b_col1,
local_row_block0,
local_row_block1,
c_acc_00,
c_acc_01,
c_acc_10,
c_acc_11);
}
}
}
__global__ void quant_mxfp4_kernel(
const __hip_bfloat16* x,
uint8_t* out_fp4,
uint8_t* out_scale,
int64_t m,
int64_t m_pad,
int64_t k,
int64_t k_blocks_valid,
int64_t k_blocks_pad) {
int64_t linear = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
int64_t total = m_pad * k_blocks_pad;
if (linear >= total) {
return;
}
int64_t row = linear / k_blocks_pad;
int64_t kb = linear % k_blocks_pad;
int64_t scale_off = shuffled_scale_offset(row, kb, k_blocks_pad);
if (row >= m || kb >= k_blocks_valid) {
out_scale[scale_off] = 127;
if (row < m_pad && kb < k_blocks_valid) {
int64_t out_base = row * (k / 2) + kb * (BLOCK_K / 2);
#pragma unroll
for (int i = 0; i < BLOCK_K / 2; ++i) {
out_fp4[out_base + i] = 0;
}
}
return;
}
int64_t in_base = row * k + kb * BLOCK_K;
float vals[BLOCK_K];
float amax = 0.0f;
const uint8_t* x_bytes = reinterpret_cast<const uint8_t*>(x);
const uint8_t* row_ptr = x_bytes + in_base * static_cast<int64_t>(sizeof(__hip_bfloat16));
#pragma unroll
for (int vec_idx = 0; vec_idx < 4; ++vec_idx) {
uint4 v4 = *reinterpret_cast<const uint4*>(row_ptr + vec_idx * 16);
uint32_t w0 = v4.x;
uint32_t w1 = v4.y;
uint32_t w2 = v4.z;
uint32_t w3 = v4.w;
uint16_t b0 = static_cast<uint16_t>(w0 & 0xFFFFu);
uint16_t b1 = static_cast<uint16_t>((w0 >> 16) & 0xFFFFu);
uint16_t b2 = static_cast<uint16_t>(w1 & 0xFFFFu);
uint16_t b3 = static_cast<uint16_t>((w1 >> 16) & 0xFFFFu);
uint16_t b4 = static_cast<uint16_t>(w2 & 0xFFFFu);
uint16_t b5 = static_cast<uint16_t>((w2 >> 16) & 0xFFFFu);
uint16_t b6 = static_cast<uint16_t>(w3 & 0xFFFFu);
uint16_t b7 = static_cast<uint16_t>((w3 >> 16) & 0xFFFFu);
int base = vec_idx * 8;
float f0 = bf16_to_f32(b0);
float f1 = bf16_to_f32(b1);
float f2 = bf16_to_f32(b2);
float f3 = bf16_to_f32(b3);
float f4 = bf16_to_f32(b4);
float f5 = bf16_to_f32(b5);
float f6 = bf16_to_f32(b6);
float f7 = bf16_to_f32(b7);
vals[base + 0] = f0;
vals[base + 1] = f1;
vals[base + 2] = f2;
vals[base + 3] = f3;
vals[base + 4] = f4;
vals[base + 5] = f5;
vals[base + 6] = f6;
vals[base + 7] = f7;
amax = fmaxf(amax, fabsf(f0));
amax = fmaxf(amax, fabsf(f1));
amax = fmaxf(amax, fabsf(f2));
amax = fmaxf(amax, fabsf(f3));
amax = fmaxf(amax, fabsf(f4));
amax = fmaxf(amax, fabsf(f5));
amax = fmaxf(amax, fabsf(f6));
amax = fmaxf(amax, fabsf(f7));
}
uint8_t scale_code = amax_to_e8m0_scale(amax);
out_scale[scale_off] = scale_code;
float scale = e8m0_to_f32_fast(scale_code);
float inv_scale = 1.0f / scale;
int64_t out_base = row * (k / 2) + kb * (BLOCK_K / 2);
#pragma unroll
for (int i = 0; i < BLOCK_K / 2; ++i) {
uint8_t lo = quantize_e2m1(vals[2 * i] * inv_scale);
uint8_t hi = quantize_e2m1(vals[2 * i + 1] * inv_scale);
out_fp4[out_base + i] = static_cast<uint8_t>((hi << 4) | lo);
}
}
template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE, bool WRITE_BF16_OUT>
__global__ void gemm_mxfp4_prequant_kernel(
const uint8_t* a_fp4,
const uint8_t* b_u8,
const uint8_t* a_scale_u8,
const uint8_t* b_scale_u8,
float* workspace,
__hip_bfloat16* out,
int64_t workspace_stride,
int64_t m,
int64_t n,
int64_t k2,
int64_t k_blocks_valid,
int64_t k_blocks_pad,
int64_t split_k) {
const int tid = static_cast<int>(threadIdx.x);
const int wave_id = tid / WAVE_SIZE;
const int lane = tid & (WAVE_SIZE - 1);
const int lane16 = lane & 15;
const int row_group = lane >> 4; // 0..3
const int wave_row = wave_id >> 1;
const int wave_col = wave_id & 1;
int64_t tile_m = static_cast<int64_t>(blockIdx.y) * GEMM_BLOCK_M;
int64_t tile_n = static_cast<int64_t>(blockIdx.x) * GEMM_BLOCK_N;
const int local_wave_m = wave_row * WAVE_TILE_M;
const int local_wave_n = wave_col * WAVE_TILE_N;
const int local_row_block0 = local_wave_m + row_group * 4;
const int local_row_block1 = local_row_block0 + MFMA_TILE_M;
const int local_a_row0 = local_wave_m + lane16;
const int local_a_row1 = local_a_row0 + MFMA_TILE_M;
const int local_b_col0 = local_wave_n + lane16;
const int local_b_col1 = local_b_col0 + MFMA_TILE_N;
const int64_t col0 = tile_n + local_b_col0;
const int64_t col1 = tile_n + local_b_col1;
const int64_t split_idx = static_cast<int64_t>(blockIdx.z);
floatx4 c_acc_00 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 c_acc_01 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 c_acc_10 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 c_acc_11 = {0.0f, 0.0f, 0.0f, 0.0f};
__shared__ __align__(16) uint8_t smem_A[PREFETCH_STAGES][GEMM_BLOCK_M * SMEM_STRIDE];
__shared__ __align__(16) uint8_t smem_B[PREFETCH_STAGES][GEMM_BLOCK_N * SMEM_STRIDE];
__shared__ uint8_t smem_A_scale[PREFETCH_STAGES][GEMM_BLOCK_M];
__shared__ uint8_t smem_B_scale[PREFETCH_STAGES][GEMM_BLOCK_N];
int64_t kb_per_split = (k_blocks_valid + split_k - 1) / split_k;
int64_t kb_start = static_cast<int64_t>(blockIdx.z) * kb_per_split;
int64_t kb_end = kb_start + kb_per_split;
if (kb_end > k_blocks_valid) {
kb_end = k_blocks_valid;
}
if (kb_start >= kb_end) {
if constexpr (!WRITE_BF16_OUT) {
if (col0 < n) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int64_t row = tile_m + local_row_block0 + i;
if (row < m) {
store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col0, 0.0f);
}
row = tile_m + local_row_block1 + i;
if (row < m) {
store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col0, 0.0f);
}
}
}
if (col1 < n) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int64_t row = tile_m + local_row_block0 + i;
if (row < m) {
store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col1, 0.0f);
}
row = tile_m + local_row_block1 + i;
if (row < m) {
store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col1, 0.0f);
}
}
}
}
return;
}
const int64_t tiles_in_split = kb_end - kb_start;
const int64_t prologue_tiles = tiles_in_split < 2 ? tiles_in_split : 2;
for (int64_t i = 0; i < prologue_tiles; ++i) {
load_prequant_stage_to_smem<SHUFFLED_B, SHUFFLE_INNER_MODE>(
a_fp4,
b_u8,
a_scale_u8,
b_scale_u8,
smem_A[static_cast<int>(i)],
smem_B[static_cast<int>(i)],
smem_A_scale[static_cast<int>(i)],
smem_B_scale[static_cast<int>(i)],
tid,
tile_m,
tile_n,
m,
n,
k2,
kb_start + i,
k_blocks_pad);
}
__syncthreads();
for (int64_t tile_idx = 0; tile_idx < tiles_in_split; ++tile_idx) {
const int curr_buf = static_cast<int>(tile_idx % PREFETCH_STAGES);
const int64_t prefetch_tile_idx = tile_idx + 2;
const bool has_prefetch = prefetch_tile_idx < tiles_in_split;
const int prefetch_buf = static_cast<int>(prefetch_tile_idx % PREFETCH_STAGES);
const int64_t prefetch_kb = kb_start + prefetch_tile_idx;
const int64_t prefetch_a_k_base = prefetch_kb * PACKED_BLOCK_K;
uint4 next_a_v = {0u, 0u, 0u, 0u};
uint4 next_b_v = {0u, 0u, 0u, 0u};
uint8_t next_a_scale = 127;
uint8_t next_b_scale = 127;
if (has_prefetch && tid < A_TILE_VEC_LOADS) {
int64_t global_row = tile_m + tid;
if (global_row < m) {
next_a_v = *reinterpret_cast<const uint4*>(a_fp4 + global_row * k2 + prefetch_a_k_base);
next_a_scale = a_scale_u8[global_row * k_blocks_pad + prefetch_kb];
}
}
if (has_prefetch && tid < B_TILE_VEC_LOADS) {
next_b_v = load_b_tile_vec<SHUFFLED_B, SHUFFLE_INNER_MODE>(
b_u8, tid, tile_n, n, k2, prefetch_a_k_base);
}
if (has_prefetch && tid < GEMM_BLOCK_N) {
int64_t global_col = tile_n + tid;
next_b_scale = load_b_scale_code<SHUFFLED_B>(
b_scale_u8, global_col, prefetch_kb, n, k_blocks_pad);
}
accumulate_bf16_stage(
smem_A[curr_buf],
smem_B[curr_buf],
smem_A_scale[curr_buf],
smem_B_scale[curr_buf],
row_group,
local_a_row0,
local_a_row1,
local_b_col0,
local_b_col1,
local_row_block0,
local_row_block1,
c_acc_00,
c_acc_01,
c_acc_10,
c_acc_11);
if (has_prefetch) {
if (tid < A_TILE_VEC_LOADS) {
store_uint4_to_smem_swizzled(smem_A[prefetch_buf], tid, next_a_v);
smem_A_scale[prefetch_buf][tid] = next_a_scale;
}
if (tid < B_TILE_VEC_LOADS) {
store_b_tile_vec<SHUFFLED_B, SHUFFLE_INNER_MODE>(smem_B[prefetch_buf], tid, next_b_v);
}
if (tid < GEMM_BLOCK_N) {
smem_B_scale[prefetch_buf][tid] = next_b_scale;
}
}
__syncthreads();
}
if (col0 < n) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int64_t row = tile_m + local_row_block0 + i;
if (row < m) {
store_output_value<WRITE_BF16_OUT>(
workspace, out, workspace_stride, n, split_idx, row, col0, c_acc_00[i]);
}
row = tile_m + local_row_block1 + i;
if (row < m) {
store_output_value<WRITE_BF16_OUT>(
workspace, out, workspace_stride, n, split_idx, row, col0, c_acc_10[i]);
}
}
}
if (col1 < n) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int64_t row = tile_m + local_row_block0 + i;
if (row < m) {
store_output_value<WRITE_BF16_OUT>(
workspace, out, workspace_stride, n, split_idx, row, col1, c_acc_01[i]);
}
row = tile_m + local_row_block1 + i;
if (row < m) {
store_output_value<WRITE_BF16_OUT>(
workspace, out, workspace_stride, n, split_idx, row, col1, c_acc_11[i]);
}
}
}
}
template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE, bool WRITE_BF16_OUT>
__global__ void gemm_mxfp4_kernel(
const __hip_bfloat16* a_bf16,
const uint8_t* b_u8,
const uint8_t* b_scale_u8,
float* workspace,
__hip_bfloat16* out,
int64_t workspace_stride,
int64_t m,
int64_t n,
int64_t k2,
int64_t k_blocks_valid,
int64_t k_blocks_pad,
int64_t split_k) {
const int tid = static_cast<int>(threadIdx.x);
const int wave_id = tid / WAVE_SIZE;
const int lane = tid & (WAVE_SIZE - 1);
const int lane16 = lane & 15;
const int row_group = lane >> 4; // 0..3
const int wave_row = wave_id >> 1;
const int wave_col = wave_id & 1;
int64_t tile_m = static_cast<int64_t>(blockIdx.y) * GEMM_BLOCK_M;
int64_t tile_n = static_cast<int64_t>(blockIdx.x) * GEMM_BLOCK_N;
const int local_wave_m = wave_row * WAVE_TILE_M;
const int local_wave_n = wave_col * WAVE_TILE_N;
const int local_row_block0 = local_wave_m + row_group * 4;
const int local_row_block1 = local_row_block0 + MFMA_TILE_M;
const int local_a_row0 = local_wave_m + lane16;
const int local_a_row1 = local_a_row0 + MFMA_TILE_M;
const int local_b_col0 = local_wave_n + lane16;
const int local_b_col1 = local_b_col0 + MFMA_TILE_N;
const int64_t col0 = tile_n + local_b_col0;
const int64_t col1 = tile_n + local_b_col1;
const int64_t split_idx = static_cast<int64_t>(blockIdx.z);
floatx4 c_acc_00 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 c_acc_01 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 c_acc_10 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 c_acc_11 = {0.0f, 0.0f, 0.0f, 0.0f};
__shared__ __align__(16) uint8_t smem_A[PREFETCH_STAGES][GEMM_BLOCK_M * SMEM_STRIDE];
__shared__ __align__(16) uint8_t smem_B[PREFETCH_STAGES][GEMM_BLOCK_N * SMEM_STRIDE];
__shared__ uint8_t smem_A_scale[PREFETCH_STAGES][GEMM_BLOCK_M];
__shared__ uint8_t smem_B_scale[PREFETCH_STAGES][GEMM_BLOCK_N];
int64_t kb_per_split = (k_blocks_valid + split_k - 1) / split_k;
int64_t kb_start = static_cast<int64_t>(blockIdx.z) * kb_per_split;
int64_t kb_end = kb_start + kb_per_split;
if (kb_end > k_blocks_valid) {
kb_end = k_blocks_valid;
}
if (kb_start >= kb_end) {
if constexpr (!WRITE_BF16_OUT) {
if (col0 < n) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int64_t row = tile_m + local_row_block0 + i;
if (row < m) {
store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col0, 0.0f);
}
row = tile_m + local_row_block1 + i;
if (row < m) {
store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col0, 0.0f);
}
}
}
if (col1 < n) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int64_t row = tile_m + local_row_block0 + i;
if (row < m) {
store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col1, 0.0f);
}
row = tile_m + local_row_block1 + i;
if (row < m) {
store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col1, 0.0f);
}
}
}
}
return;
}
const int64_t tiles_in_split = kb_end - kb_start;
const int64_t k = k2 * 2;
const int64_t prologue_tiles = tiles_in_split < 2 ? tiles_in_split : 2;
for (int64_t i = 0; i < prologue_tiles; ++i) {
load_gemm_stage_to_smem<SHUFFLED_B, SHUFFLE_INNER_MODE>(
a_bf16,
b_u8,
b_scale_u8,
smem_A[static_cast<int>(i)],
smem_B[static_cast<int>(i)],
smem_A_scale[static_cast<int>(i)],
smem_B_scale[static_cast<int>(i)],
tid,
tile_m,
tile_n,
m,
n,
k,
k2,
kb_start + i,
k_blocks_pad);
}
__syncthreads();
for (int64_t tile_idx = 0; tile_idx < tiles_in_split; ++tile_idx) {
const int curr_buf = static_cast<int>(tile_idx % PREFETCH_STAGES);
const int64_t prefetch_tile_idx = tile_idx + 2;
const bool has_prefetch = prefetch_tile_idx < tiles_in_split;
const int prefetch_buf = static_cast<int>(prefetch_tile_idx % PREFETCH_STAGES);
const int64_t prefetch_kb = kb_start + prefetch_tile_idx;
const int64_t prefetch_a_k_base = prefetch_kb * PACKED_BLOCK_K;
uint4 next_a_v = {0u, 0u, 0u, 0u};
uint4 next_b_v = {0u, 0u, 0u, 0u};
uint8_t next_a_scale = 127;
uint8_t next_b_scale = 127;
if (has_prefetch && tid < A_TILE_VEC_LOADS) {
int64_t global_row = tile_m + tid;
quantize_a_row_block_to_fp4(a_bf16, global_row, k, prefetch_kb, m, next_a_v, next_a_scale);
}
if (has_prefetch && tid < B_TILE_VEC_LOADS) {
next_b_v = load_b_tile_vec<SHUFFLED_B, SHUFFLE_INNER_MODE>(
b_u8, tid, tile_n, n, k2, prefetch_a_k_base);
}
if (has_prefetch && tid < GEMM_BLOCK_N) {
int64_t global_col = tile_n + tid;
next_b_scale = load_b_scale_code<SHUFFLED_B>(
b_scale_u8, global_col, prefetch_kb, n, k_blocks_pad);
}
accumulate_bf16_stage(
smem_A[curr_buf],
smem_B[curr_buf],
smem_A_scale[curr_buf],
smem_B_scale[curr_buf],
row_group,
local_a_row0,
local_a_row1,
local_b_col0,
local_b_col1,
local_row_block0,
local_row_block1,
c_acc_00,
c_acc_01,
c_acc_10,
c_acc_11);
if (has_prefetch) {
if (tid < A_TILE_VEC_LOADS) {
store_uint4_to_smem_swizzled(smem_A[prefetch_buf], tid, next_a_v);
smem_A_scale[prefetch_buf][tid] = next_a_scale;
}
if (tid < B_TILE_VEC_LOADS) {
store_b_tile_vec<SHUFFLED_B, SHUFFLE_INNER_MODE>(smem_B[prefetch_buf], tid, next_b_v);
}
if (tid < GEMM_BLOCK_N) {
smem_B_scale[prefetch_buf][tid] = next_b_scale;
}
}
__syncthreads();
}
if (col0 < n) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int64_t row = tile_m + local_row_block0 + i;
if (row < m) {
store_output_value<WRITE_BF16_OUT>(
workspace, out, workspace_stride, n, split_idx, row, col0, c_acc_00[i]);
}
row = tile_m + local_row_block1 + i;
if (row < m) {
store_output_value<WRITE_BF16_OUT>(
workspace, out, workspace_stride, n, split_idx, row, col0, c_acc_10[i]);
}
}
}
if (col1 < n) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int64_t row = tile_m + local_row_block0 + i;
if (row < m) {
store_output_value<WRITE_BF16_OUT>(
workspace, out, workspace_stride, n, split_idx, row, col1, c_acc_01[i]);
}
row = tile_m + local_row_block1 + i;
if (row < m) {
store_output_value<WRITE_BF16_OUT>(
workspace, out, workspace_stride, n, split_idx, row, col1, c_acc_11[i]);
}
}
}
}
template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE, bool WRITE_BF16_OUT>
__global__ void gemm_mxfp4_prequant_native_kernel(
const uint8_t* a_fp4,
const uint8_t* b_u8,
const uint8_t* a_scale_u8,
const uint8_t* b_scale_u8,
float* workspace,
__hip_bfloat16* out,
int64_t workspace_stride,
int64_t m,
int64_t n,
int64_t k2,
int64_t k_blocks_valid,
int64_t k_blocks_pad,
int64_t split_k) {
const int tid = static_cast<int>(threadIdx.x);
const int wave_id = tid / WAVE_SIZE;
const int lane = tid & (WAVE_SIZE - 1);
const int lane16 = lane & 15;
const int row_group = lane >> 4;
const int wave_row = wave_id >> 1;
const int wave_col = wave_id & 1;
int64_t tile_m = static_cast<int64_t>(blockIdx.y) * GEMM_BLOCK_M;
int64_t tile_n = static_cast<int64_t>(blockIdx.x) * GEMM_BLOCK_N;
const int local_wave_m = wave_row * WAVE_TILE_M;
const int local_wave_n = wave_col * WAVE_TILE_N;
const int local_row_block0 = local_wave_m + row_group * 4;
const int local_row_block1 = local_row_block0 + MFMA_TILE_M;
const int local_a_row0 = local_wave_m + lane16;
const int local_a_row1 = local_a_row0 + MFMA_TILE_M;
const int local_b_col0 = local_wave_n + lane16;
const int local_b_col1 = local_b_col0 + MFMA_TILE_N;
const int64_t col0 = tile_n + local_b_col0;
const int64_t col1 = tile_n + local_b_col1;
const int64_t split_idx = static_cast<int64_t>(blockIdx.z);
floatx4 c_acc_00 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 c_acc_01 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 c_acc_10 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 c_acc_11 = {0.0f, 0.0f, 0.0f, 0.0f};
__shared__ __align__(16) uint8_t smem_A[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_M * SMEM_STRIDE];
__shared__ __align__(16) uint8_t smem_B[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_N * SMEM_STRIDE];
__shared__ uint8_t smem_A_scale[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_M];
__shared__ uint8_t smem_B_scale[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_N];
int64_t kb_per_split = (k_blocks_valid + split_k - 1) / split_k;
int64_t kb_start = static_cast<int64_t>(blockIdx.z) * kb_per_split;
int64_t kb_end = kb_start + kb_per_split;
if (kb_end > k_blocks_valid) {
kb_end = k_blocks_valid;
}
if (kb_start >= kb_end) {
if constexpr (!WRITE_BF16_OUT) {
if (col0 < n) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int64_t row = tile_m + local_row_block0 + i;
if (row < m) {
store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col0, 0.0f);
}
row = tile_m + local_row_block1 + i;
if (row < m) {
store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col0, 0.0f);
}
}
}
if (col1 < n) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int64_t row = tile_m + local_row_block0 + i;
if (row < m) {
store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col1, 0.0f);
}
row = tile_m + local_row_block1 + i;
if (row < m) {
store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col1, 0.0f);
}
}
}
}
return;
}
const int64_t full_kb_end =
kb_start + ((kb_end - kb_start) / NATIVE_MFMA_K_BLOCKS) * NATIVE_MFMA_K_BLOCKS;
for (int64_t kb_group = kb_start; kb_group < full_kb_end; kb_group += NATIVE_MFMA_K_BLOCKS) {
#pragma unroll
for (int stage = 0; stage < NATIVE_MFMA_K_BLOCKS; ++stage) {
load_prequant_stage_to_smem<SHUFFLED_B, SHUFFLE_INNER_MODE>(
a_fp4,
b_u8,
a_scale_u8,
b_scale_u8,
smem_A[stage],
smem_B[stage],
smem_A_scale[stage],
smem_B_scale[stage],
tid,
tile_m,
tile_n,
m,
n,
k2,
kb_group + stage,
k_blocks_pad);
}
__syncthreads();
#if defined(__gfx950__) && MXFP4_ENABLE_NATIVE_FP4_MFMA
accumulate_native_group(
smem_A,
smem_B,
smem_A_scale,
smem_B_scale,
row_group,
local_a_row0,
local_a_row1,
local_b_col0,
local_b_col1,
c_acc_00,
c_acc_01,
c_acc_10,
c_acc_11);
#else
accumulate_tail_stages(
smem_A,
smem_B,
smem_A_scale,
smem_B_scale,
NATIVE_MFMA_K_BLOCKS,
row_group,
local_a_row0,
local_a_row1,
local_b_col0,
local_b_col1,
local_row_block0,
local_row_block1,
c_acc_00,
c_acc_01,
c_acc_10,
c_acc_11);
#endif
__syncthreads();
}
const int tail_k_blocks = static_cast<int>(kb_end - full_kb_end);
if (tail_k_blocks > 0) {
#pragma unroll
for (int stage = 0; stage < NATIVE_MFMA_K_BLOCKS; ++stage) {
if (stage < tail_k_blocks) {
load_prequant_stage_to_smem<SHUFFLED_B, SHUFFLE_INNER_MODE>(
a_fp4,
b_u8,
a_scale_u8,
b_scale_u8,
smem_A[stage],
smem_B[stage],
smem_A_scale[stage],
smem_B_scale[stage],
tid,
tile_m,
tile_n,
m,
n,
k2,
full_kb_end + stage,
k_blocks_pad);
}
}
__syncthreads();
accumulate_tail_stages(
smem_A,
smem_B,
smem_A_scale,
smem_B_scale,
tail_k_blocks,
row_group,
local_a_row0,
local_a_row1,
local_b_col0,
local_b_col1,
local_row_block0,
local_row_block1,
c_acc_00,
c_acc_01,
c_acc_10,
c_acc_11);
}
if (col0 < n) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int64_t row = tile_m + local_row_block0 + i;
if (row < m) {
store_output_value<WRITE_BF16_OUT>(
workspace, out, workspace_stride, n, split_idx, row, col0, c_acc_00[i]);
}
row = tile_m + local_row_block1 + i;
if (row < m) {
store_output_value<WRITE_BF16_OUT>(
workspace, out, workspace_stride, n, split_idx, row, col0, c_acc_10[i]);
}
}
}
if (col1 < n) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int64_t row = tile_m + local_row_block0 + i;
if (row < m) {
store_output_value<WRITE_BF16_OUT>(
workspace, out, workspace_stride, n, split_idx, row, col1, c_acc_01[i]);
}
row = tile_m + local_row_block1 + i;
if (row < m) {
store_output_value<WRITE_BF16_OUT>(
workspace, out, workspace_stride, n, split_idx, row, col1, c_acc_11[i]);
}
}
}
}
template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE, bool WRITE_BF16_OUT>
__global__ void gemm_mxfp4_native_kernel(
const __hip_bfloat16* a_bf16,
const uint8_t* b_u8,
const uint8_t* b_scale_u8,
float* workspace,
__hip_bfloat16* out,
int64_t workspace_stride,
int64_t m,
int64_t n,
int64_t k2,
int64_t k_blocks_valid,
int64_t k_blocks_pad,
int64_t split_k) {
const int tid = static_cast<int>(threadIdx.x);
const int wave_id = tid / WAVE_SIZE;
const int lane = tid & (WAVE_SIZE - 1);
const int lane16 = lane & 15;
const int row_group = lane >> 4;
const int wave_row = wave_id >> 1;
const int wave_col = wave_id & 1;
int64_t tile_m = static_cast<int64_t>(blockIdx.y) * GEMM_BLOCK_M;
int64_t tile_n = static_cast<int64_t>(blockIdx.x) * GEMM_BLOCK_N;
const int local_wave_m = wave_row * WAVE_TILE_M;
const int local_wave_n = wave_col * WAVE_TILE_N;
const int local_row_block0 = local_wave_m + row_group * 4;
const int local_row_block1 = local_row_block0 + MFMA_TILE_M;
const int local_a_row0 = local_wave_m + lane16;
const int local_a_row1 = local_a_row0 + MFMA_TILE_M;
const int local_b_col0 = local_wave_n + lane16;
const int local_b_col1 = local_b_col0 + MFMA_TILE_N;
const int64_t col0 = tile_n + local_b_col0;
const int64_t col1 = tile_n + local_b_col1;
const int64_t split_idx = static_cast<int64_t>(blockIdx.z);
floatx4 c_acc_00 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 c_acc_01 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 c_acc_10 = {0.0f, 0.0f, 0.0f, 0.0f};
floatx4 c_acc_11 = {0.0f, 0.0f, 0.0f, 0.0f};
__shared__ __align__(16) uint8_t smem_A[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_M * SMEM_STRIDE];
__shared__ __align__(16) uint8_t smem_B[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_N * SMEM_STRIDE];
__shared__ uint8_t smem_A_scale[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_M];
__shared__ uint8_t smem_B_scale[NATIVE_MFMA_K_BLOCKS][GEMM_BLOCK_N];
int64_t kb_per_split = (k_blocks_valid + split_k - 1) / split_k;
int64_t kb_start = static_cast<int64_t>(blockIdx.z) * kb_per_split;
int64_t kb_end = kb_start + kb_per_split;
if (kb_end > k_blocks_valid) {
kb_end = k_blocks_valid;
}
if (kb_start >= kb_end) {
if constexpr (!WRITE_BF16_OUT) {
if (col0 < n) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int64_t row = tile_m + local_row_block0 + i;
if (row < m) {
store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col0, 0.0f);
}
row = tile_m + local_row_block1 + i;
if (row < m) {
store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col0, 0.0f);
}
}
}
if (col1 < n) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int64_t row = tile_m + local_row_block0 + i;
if (row < m) {
store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col1, 0.0f);
}
row = tile_m + local_row_block1 + i;
if (row < m) {
store_output_value<false>(workspace, out, workspace_stride, n, split_idx, row, col1, 0.0f);
}
}
}
}
return;
}
const int64_t k = k2 * 2;
const int64_t full_kb_end =
kb_start + ((kb_end - kb_start) / NATIVE_MFMA_K_BLOCKS) * NATIVE_MFMA_K_BLOCKS;
for (int64_t kb_group = kb_start; kb_group < full_kb_end; kb_group += NATIVE_MFMA_K_BLOCKS) {
#pragma unroll
for (int stage = 0; stage < NATIVE_MFMA_K_BLOCKS; ++stage) {
load_gemm_stage_to_smem<SHUFFLED_B, SHUFFLE_INNER_MODE>(
a_bf16,
b_u8,
b_scale_u8,
smem_A[stage],
smem_B[stage],
smem_A_scale[stage],
smem_B_scale[stage],
tid,
tile_m,
tile_n,
m,
n,
k,
k2,
kb_group + stage,
k_blocks_pad);
}
__syncthreads();
#if defined(__gfx950__) && MXFP4_ENABLE_NATIVE_FP4_MFMA
accumulate_native_group(
smem_A,
smem_B,
smem_A_scale,
smem_B_scale,
row_group,
local_a_row0,
local_a_row1,
local_b_col0,
local_b_col1,
c_acc_00,
c_acc_01,
c_acc_10,
c_acc_11);
#else
accumulate_tail_stages(
smem_A,
smem_B,
smem_A_scale,
smem_B_scale,
NATIVE_MFMA_K_BLOCKS,
row_group,
local_a_row0,
local_a_row1,
local_b_col0,
local_b_col1,
local_row_block0,
local_row_block1,
c_acc_00,
c_acc_01,
c_acc_10,
c_acc_11);
#endif
__syncthreads();
}
const int tail_k_blocks = static_cast<int>(kb_end - full_kb_end);
if (tail_k_blocks > 0) {
#pragma unroll
for (int stage = 0; stage < NATIVE_MFMA_K_BLOCKS; ++stage) {
if (stage < tail_k_blocks) {
load_gemm_stage_to_smem<SHUFFLED_B, SHUFFLE_INNER_MODE>(
a_bf16,
b_u8,
b_scale_u8,
smem_A[stage],
smem_B[stage],
smem_A_scale[stage],
smem_B_scale[stage],
tid,
tile_m,
tile_n,
m,
n,
k,
k2,
full_kb_end + stage,
k_blocks_pad);
}
}
__syncthreads();
accumulate_tail_stages(
smem_A,
smem_B,
smem_A_scale,
smem_B_scale,
tail_k_blocks,
row_group,
local_a_row0,
local_a_row1,
local_b_col0,
local_b_col1,
local_row_block0,
local_row_block1,
c_acc_00,
c_acc_01,
c_acc_10,
c_acc_11);
}
if (col0 < n) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int64_t row = tile_m + local_row_block0 + i;
if (row < m) {
store_output_value<WRITE_BF16_OUT>(
workspace, out, workspace_stride, n, split_idx, row, col0, c_acc_00[i]);
}
row = tile_m + local_row_block1 + i;
if (row < m) {
store_output_value<WRITE_BF16_OUT>(
workspace, out, workspace_stride, n, split_idx, row, col0, c_acc_10[i]);
}
}
}
if (col1 < n) {
#pragma unroll
for (int i = 0; i < 4; ++i) {
int64_t row = tile_m + local_row_block0 + i;
if (row < m) {
store_output_value<WRITE_BF16_OUT>(
workspace, out, workspace_stride, n, split_idx, row, col1, c_acc_01[i]);
}
row = tile_m + local_row_block1 + i;
if (row < m) {
store_output_value<WRITE_BF16_OUT>(
workspace, out, workspace_stride, n, split_idx, row, col1, c_acc_11[i]);
}
}
}
}
__global__ void reduce_splitk_kernel(
const float* workspace,
__hip_bfloat16* out,
int64_t workspace_stride,
int64_t m,
int64_t n,
int64_t split_k) {
int64_t row = static_cast<int64_t>(blockIdx.y) * blockDim.y + threadIdx.y;
int64_t col = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
if (row >= m || col >= n) {
return;
}
float sum = 0.0f;
int64_t base = row * n + col;
for (int64_t s = 0; s < split_k; ++s) {
sum += workspace[s * workspace_stride + base];
}
out[base] = static_cast<__hip_bfloat16>(sum);
}
inline void check_hip_error(const char* where) {
hipError_t err = hipGetLastError();
TORCH_CHECK(err == hipSuccess, where, " failed: ", hipGetErrorString(err));
}
template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE, bool WRITE_BF16_OUT>
inline void launch_gemm_mxfp4_variant(
dim3 grid,
dim3 block,
const __hip_bfloat16* a_bf16,
const uint8_t* b_u8,
const uint8_t* b_scale_u8,
float* workspace,
__hip_bfloat16* out,
int64_t workspace_stride,
int64_t m,
int64_t n,
int64_t k2,
int64_t k_blocks_valid,
int64_t k_blocks_pad,
int64_t split_k) {
hipLaunchKernelGGL(
HIP_KERNEL_NAME(gemm_mxfp4_kernel<SHUFFLED_B, SHUFFLE_INNER_MODE, WRITE_BF16_OUT>),
grid,
block,
0,
0,
a_bf16,
b_u8,
b_scale_u8,
workspace,
out,
workspace_stride,
m,
n,
k2,
k_blocks_valid,
k_blocks_pad,
split_k);
}
template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE, bool WRITE_BF16_OUT>
inline void launch_gemm_mxfp4_prequant_variant(
dim3 grid,
dim3 block,
const uint8_t* a_fp4,
const uint8_t* b_u8,
const uint8_t* a_scale_u8,
const uint8_t* b_scale_u8,
float* workspace,
__hip_bfloat16* out,
int64_t workspace_stride,
int64_t m,
int64_t n,
int64_t k2,
int64_t k_blocks_valid,
int64_t k_blocks_pad,
int64_t split_k) {
hipLaunchKernelGGL(
HIP_KERNEL_NAME(gemm_mxfp4_prequant_kernel<SHUFFLED_B, SHUFFLE_INNER_MODE, WRITE_BF16_OUT>),
grid,
block,
0,
0,
a_fp4,
b_u8,
a_scale_u8,
b_scale_u8,
workspace,
out,
workspace_stride,
m,
n,
k2,
k_blocks_valid,
k_blocks_pad,
split_k);
}
template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE, bool WRITE_BF16_OUT>
inline void launch_gemm_mxfp4_prequant_native_variant(
dim3 grid,
dim3 block,
const uint8_t* a_fp4,
const uint8_t* b_u8,
const uint8_t* a_scale_u8,
const uint8_t* b_scale_u8,
float* workspace,
__hip_bfloat16* out,
int64_t workspace_stride,
int64_t m,
int64_t n,
int64_t k2,
int64_t k_blocks_valid,
int64_t k_blocks_pad,
int64_t split_k) {
hipLaunchKernelGGL(
HIP_KERNEL_NAME(gemm_mxfp4_prequant_native_kernel<SHUFFLED_B, SHUFFLE_INNER_MODE, WRITE_BF16_OUT>),
grid,
block,
0,
0,
a_fp4,
b_u8,
a_scale_u8,
b_scale_u8,
workspace,
out,
workspace_stride,
m,
n,
k2,
k_blocks_valid,
k_blocks_pad,
split_k);
}
template <bool SHUFFLED_B, int SHUFFLE_INNER_MODE, bool WRITE_BF16_OUT>
inline void launch_gemm_mxfp4_native_variant(
dim3 grid,
dim3 block,
const __hip_bfloat16* a_bf16,
const uint8_t* b_u8,
const uint8_t* b_scale_u8,
float* workspace,
__hip_bfloat16* out,
int64_t workspace_stride,
int64_t m,
int64_t n,
int64_t k2,
int64_t k_blocks_valid,
int64_t k_blocks_pad,
int64_t split_k) {
hipLaunchKernelGGL(
HIP_KERNEL_NAME(gemm_mxfp4_native_kernel<SHUFFLED_B, SHUFFLE_INNER_MODE, WRITE_BF16_OUT>),
grid,
block,
0,
0,
a_bf16,
b_u8,
b_scale_u8,
workspace,
out,
workspace_stride,
m,
n,
k2,
k_blocks_valid,
k_blocks_pad,
split_k);
}
} // namespace
std::vector<torch::Tensor> hip_quant_mxfp4(torch::Tensor x) {
TORCH_CHECK(x.is_cuda(), "x must be CUDA tensor");
TORCH_CHECK(x.scalar_type() == torch::kBFloat16, "x must be bfloat16");
TORCH_CHECK(x.dim() == 2, "x must be 2D [M, K]");
auto x_contig = x.contiguous();
int64_t m = x_contig.size(0);
int64_t k = x_contig.size(1);
TORCH_CHECK(k % 64 == 0, "k must be divisible by 64");
int64_t k_blocks_valid = k / BLOCK_K;
int64_t k_blocks_pad = ((k_blocks_valid + (PAD_SCALE_N - 1)) / PAD_SCALE_N) * PAD_SCALE_N;
int64_t m_pad_kernel = ((m + (PAD_M_KERNEL - 1)) / PAD_M_KERNEL) * PAD_M_KERNEL;
int64_t m_pad_scale = ((m + (PAD_M_SCALE - 1)) / PAD_M_SCALE) * PAD_M_SCALE;
auto u8_opts = x_contig.options().dtype(torch::kUInt8);
auto out_fp4 = torch::empty({m_pad_kernel, k / 2}, u8_opts);
auto out_scale = torch::full({m_pad_scale, k_blocks_pad}, 127, u8_opts);
int64_t total = m_pad_kernel * k_blocks_pad;
int threads = 256;
int blocks = static_cast<int>((total + threads - 1) / threads);
if (blocks > 0) {
hipLaunchKernelGGL(
quant_mxfp4_kernel,
dim3(blocks),
dim3(threads),
0,
0,
reinterpret_cast<const __hip_bfloat16*>(x_contig.data_ptr()),
reinterpret_cast<uint8_t*>(out_fp4.data_ptr()),
reinterpret_cast<uint8_t*>(out_scale.data_ptr()),
m,
m_pad_kernel,
k,
k_blocks_valid,
k_blocks_pad);
check_hip_error("quant_mxfp4_kernel");
}
return {out_fp4, out_scale};
}
torch::Tensor hip_gemm_mxfp4_prequant(
torch::Tensor a_fp4_u8,
torch::Tensor b_u8,
torch::Tensor a_scale_u8,
torch::Tensor b_scale_u8,
int64_t layout_mode,
int64_t log2_k_split,
torch::Tensor workspace,
int64_t workspace_stride,
int64_t b_shuffle_inner_mode,
int64_t native_mfma_mode) {
TORCH_CHECK(a_fp4_u8.is_cuda(), "a_fp4_u8 must be CUDA");
TORCH_CHECK(b_u8.is_cuda(), "b_u8 must be CUDA");
TORCH_CHECK(a_scale_u8.is_cuda(), "a_scale_u8 must be CUDA");
TORCH_CHECK(b_scale_u8.is_cuda(), "b_scale_u8 must be CUDA");
TORCH_CHECK(a_fp4_u8.scalar_type() == torch::kUInt8, "a_fp4_u8 must be uint8");
TORCH_CHECK(b_u8.scalar_type() == torch::kUInt8, "b_u8 must be uint8");
TORCH_CHECK(a_scale_u8.scalar_type() == torch::kUInt8, "a_scale_u8 must be uint8");
TORCH_CHECK(b_scale_u8.scalar_type() == torch::kUInt8, "b_scale_u8 must be uint8");
TORCH_CHECK(a_fp4_u8.dim() == 2 && b_u8.dim() == 2, "A/B FP4 must be 2D");
TORCH_CHECK(a_scale_u8.dim() == 2 && b_scale_u8.dim() == 2, "A/B scale must be 2D");
TORCH_CHECK(layout_mode == 0 || layout_mode == 1, "layout_mode must be 0 or 1");
TORCH_CHECK(layout_mode == 0 || b_shuffle_inner_mode == 0 || b_shuffle_inner_mode == 1,
"b_shuffle_inner_mode must be 0 or 1 for shuffled B");
TORCH_CHECK(native_mfma_mode == 0 || native_mfma_mode == 1,
"native_mfma_mode must be 0 or 1");
auto a_fp4 = a_fp4_u8.contiguous();
auto b_fp4 = b_u8.contiguous();
auto a_scale = a_scale_u8.contiguous();
auto b_scale = b_scale_u8.contiguous();
int64_t m = a_fp4.size(0);
int64_t n = b_fp4.size(0);
int64_t k2 = a_fp4.size(1);
TORCH_CHECK(b_fp4.size(1) == k2, "A/B K/2 mismatch");
int64_t k = k2 * 2;
TORCH_CHECK(k % BLOCK_K == 0, "K must be divisible by 32");
int64_t k_blocks_valid = k / BLOCK_K;
int64_t k_blocks_pad = a_scale.size(1);
TORCH_CHECK(b_scale.size(1) == k_blocks_pad, "A/B scale padded K-block mismatch");
TORCH_CHECK(a_scale.size(0) >= m, "a_scale rows must cover m");
TORCH_CHECK(b_scale.size(0) >= n, "b_scale rows must cover n");
int64_t split_k = 1;
if (log2_k_split > 0) {
split_k = static_cast<int64_t>(1) << log2_k_split;
}
if (split_k < 1) {
split_k = 1;
}
const bool direct_output = (split_k == 1);
int64_t mn = m * n;
auto ws = workspace;
if (!direct_output) {
if (workspace_stride <= 0) {
workspace_stride = mn;
}
TORCH_CHECK(workspace_stride >= mn, "workspace_stride must be >= m*n");
auto ws_opts = a_fp4.options().dtype(torch::kFloat);
int64_t need = split_k * workspace_stride;
if (!ws.defined() || !ws.is_cuda() || ws.scalar_type() != torch::kFloat || ws.numel() < need) {
ws = torch::empty({split_k, workspace_stride}, ws_opts);
} else {
ws = ws.contiguous().view({split_k, workspace_stride});
}
}
auto out = torch::empty({m, n}, a_fp4.options().dtype(torch::kBFloat16));
dim3 block(GEMM_THREADS);
dim3 grid(
static_cast<unsigned int>((n + GEMM_BLOCK_N - 1) / GEMM_BLOCK_N),
static_cast<unsigned int>((m + GEMM_BLOCK_M - 1) / GEMM_BLOCK_M),
static_cast<unsigned int>(split_k));
const uint8_t* a_ptr = reinterpret_cast<const uint8_t*>(a_fp4.data_ptr());
const uint8_t* b_ptr = reinterpret_cast<const uint8_t*>(b_fp4.data_ptr());
const uint8_t* a_scale_ptr = reinterpret_cast<const uint8_t*>(a_scale.data_ptr());
const uint8_t* b_scale_ptr = reinterpret_cast<const uint8_t*>(b_scale.data_ptr());
float* ws_ptr = direct_output ? nullptr : reinterpret_cast<float*>(ws.data_ptr());
auto* out_ptr = reinterpret_cast<__hip_bfloat16*>(out.data_ptr());
constexpr bool native_build_enabled = static_cast<bool>(MXFP4_ENABLE_NATIVE_FP4_MFMA);
const bool use_native_mfma =
native_build_enabled && (native_mfma_mode != 0) && (k_blocks_valid >= NATIVE_MFMA_K_BLOCKS);
if (direct_output) {
if (layout_mode == 0) {
if (use_native_mfma) {
launch_gemm_mxfp4_prequant_native_variant<false, 0, true>(
grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
} else {
launch_gemm_mxfp4_prequant_variant<false, 0, true>(
grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
}
} else if (b_shuffle_inner_mode == 0) {
if (use_native_mfma) {
launch_gemm_mxfp4_prequant_native_variant<true, 0, true>(
grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
} else {
launch_gemm_mxfp4_prequant_variant<true, 0, true>(
grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
}
} else {
if (use_native_mfma) {
launch_gemm_mxfp4_prequant_native_variant<true, 1, true>(
grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
} else {
launch_gemm_mxfp4_prequant_variant<true, 1, true>(
grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
}
}
} else {
if (layout_mode == 0) {
if (use_native_mfma) {
launch_gemm_mxfp4_prequant_native_variant<false, 0, false>(
grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
} else {
launch_gemm_mxfp4_prequant_variant<false, 0, false>(
grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
}
} else if (b_shuffle_inner_mode == 0) {
if (use_native_mfma) {
launch_gemm_mxfp4_prequant_native_variant<true, 0, false>(
grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
} else {
launch_gemm_mxfp4_prequant_variant<true, 0, false>(
grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
}
} else {
if (use_native_mfma) {
launch_gemm_mxfp4_prequant_native_variant<true, 1, false>(
grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
} else {
launch_gemm_mxfp4_prequant_variant<true, 1, false>(
grid, block, a_ptr, b_ptr, a_scale_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
}
}
}
check_hip_error(use_native_mfma ? "gemm_mxfp4_prequant_native_kernel"
: "gemm_mxfp4_prequant_kernel");
if (!direct_output) {
dim3 rblock(16, 16);
dim3 rgrid(
static_cast<unsigned int>((n + 15) / 16),
static_cast<unsigned int>((m + 15) / 16));
hipLaunchKernelGGL(
reduce_splitk_kernel,
rgrid,
rblock,
0,
0,
reinterpret_cast<const float*>(ws.data_ptr()),
reinterpret_cast<__hip_bfloat16*>(out.data_ptr()),
workspace_stride,
m,
n,
split_k);
check_hip_error("reduce_splitk_kernel");
}
return out;
}
torch::Tensor hip_gemm_mxfp4(
torch::Tensor a_bf16,
torch::Tensor b_u8,
torch::Tensor b_scale_u8,
int64_t layout_mode,
int64_t log2_k_split,
torch::Tensor workspace,
int64_t workspace_stride,
int64_t b_shuffle_inner_mode,
int64_t native_mfma_mode) {
TORCH_CHECK(a_bf16.is_cuda(), "a_bf16 must be CUDA");
TORCH_CHECK(b_u8.is_cuda(), "b_u8 must be CUDA");
TORCH_CHECK(b_scale_u8.is_cuda(), "b_scale_u8 must be CUDA");
TORCH_CHECK(a_bf16.scalar_type() == torch::kBFloat16, "a_bf16 must be bfloat16");
TORCH_CHECK(b_u8.scalar_type() == torch::kUInt8, "b_u8 must be uint8");
TORCH_CHECK(b_scale_u8.scalar_type() == torch::kUInt8, "b_scale_u8 must be uint8");
TORCH_CHECK(a_bf16.dim() == 2 && b_u8.dim() == 2, "A BF16 and B FP4 must be 2D");
TORCH_CHECK(b_scale_u8.dim() == 2, "B scale must be 2D");
TORCH_CHECK(layout_mode == 0 || layout_mode == 1, "layout_mode must be 0 or 1");
TORCH_CHECK(layout_mode == 0 || b_shuffle_inner_mode == 0 || b_shuffle_inner_mode == 1,
"b_shuffle_inner_mode must be 0 or 1 for shuffled B");
auto a = a_bf16.contiguous();
auto b_fp4 = b_u8.contiguous();
auto b_scale = b_scale_u8.contiguous();
int64_t m = a.size(0);
int64_t n = b_fp4.size(0);
int64_t k = a.size(1);
int64_t k2 = b_fp4.size(1);
TORCH_CHECK(k == k2 * 2, "A/B K mismatch");
TORCH_CHECK(k % BLOCK_K == 0, "K must be divisible by 32");
int64_t k_blocks_valid = k / BLOCK_K;
int64_t k_blocks_pad = b_scale.size(1);
TORCH_CHECK(k_blocks_pad >= k_blocks_valid, "B scale padded K-block mismatch");
TORCH_CHECK(b_scale.size(0) >= n, "b_scale rows must cover n");
TORCH_CHECK(native_mfma_mode == 0 || native_mfma_mode == 1,
"native_mfma_mode must be 0 or 1");
int64_t split_k = 1;
if (log2_k_split > 0) {
split_k = static_cast<int64_t>(1) << log2_k_split;
}
if (split_k < 1) {
split_k = 1;
}
const bool direct_output = (split_k == 1);
int64_t mn = m * n;
auto ws = workspace;
if (!direct_output) {
if (workspace_stride <= 0) {
workspace_stride = mn;
}
TORCH_CHECK(workspace_stride >= mn, "workspace_stride must be >= m*n");
auto ws_opts = a.options().dtype(torch::kFloat);
int64_t need = split_k * workspace_stride;
if (!ws.defined() || !ws.is_cuda() || ws.scalar_type() != torch::kFloat || ws.numel() < need) {
ws = torch::empty({split_k, workspace_stride}, ws_opts);
} else {
ws = ws.contiguous().view({split_k, workspace_stride});
}
}
auto out = torch::empty({m, n}, a.options().dtype(torch::kBFloat16));
dim3 block(GEMM_THREADS);
dim3 grid(
static_cast<unsigned int>((n + GEMM_BLOCK_N - 1) / GEMM_BLOCK_N),
static_cast<unsigned int>((m + GEMM_BLOCK_M - 1) / GEMM_BLOCK_M),
static_cast<unsigned int>(split_k));
const __hip_bfloat16* a_ptr = reinterpret_cast<const __hip_bfloat16*>(a.data_ptr());
const uint8_t* b_ptr = reinterpret_cast<const uint8_t*>(b_fp4.data_ptr());
const uint8_t* b_scale_ptr = reinterpret_cast<const uint8_t*>(b_scale.data_ptr());
float* ws_ptr = direct_output ? nullptr : reinterpret_cast<float*>(ws.data_ptr());
auto* out_ptr = reinterpret_cast<__hip_bfloat16*>(out.data_ptr());
constexpr bool native_build_enabled = static_cast<bool>(MXFP4_ENABLE_NATIVE_FP4_MFMA);
const bool use_native_mfma =
native_build_enabled && (native_mfma_mode != 0) && (k_blocks_valid >= NATIVE_MFMA_K_BLOCKS);
if (direct_output) {
if (layout_mode == 0) {
if (use_native_mfma) {
launch_gemm_mxfp4_native_variant<false, 0, true>(
grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
} else {
launch_gemm_mxfp4_variant<false, 0, true>(
grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
}
} else if (b_shuffle_inner_mode == 0) {
if (use_native_mfma) {
launch_gemm_mxfp4_native_variant<true, 0, true>(
grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
} else {
launch_gemm_mxfp4_variant<true, 0, true>(
grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
}
} else {
if (use_native_mfma) {
launch_gemm_mxfp4_native_variant<true, 1, true>(
grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
} else {
launch_gemm_mxfp4_variant<true, 1, true>(
grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
}
}
} else {
if (layout_mode == 0) {
if (use_native_mfma) {
launch_gemm_mxfp4_native_variant<false, 0, false>(
grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
} else {
launch_gemm_mxfp4_variant<false, 0, false>(
grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
}
} else if (b_shuffle_inner_mode == 0) {
if (use_native_mfma) {
launch_gemm_mxfp4_native_variant<true, 0, false>(
grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
} else {
launch_gemm_mxfp4_variant<true, 0, false>(
grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
}
} else {
if (use_native_mfma) {
launch_gemm_mxfp4_native_variant<true, 1, false>(
grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
} else {
launch_gemm_mxfp4_variant<true, 1, false>(
grid, block, a_ptr, b_ptr, b_scale_ptr, ws_ptr, out_ptr,
workspace_stride, m, n, k2, k_blocks_valid, k_blocks_pad, split_k);
}
}
}
check_hip_error(use_native_mfma ? "gemm_mxfp4_native_kernel" : "gemm_mxfp4_kernel");
if (!direct_output) {
dim3 rblock(16, 16);
dim3 rgrid(
static_cast<unsigned int>((n + 15) / 16),
static_cast<unsigned int>((m + 15) / 16));
hipLaunchKernelGGL(
reduce_splitk_kernel,
rgrid,
rblock,
0,
0,
reinterpret_cast<const float*>(ws.data_ptr()),
reinterpret_cast<__hip_bfloat16*>(out.data_ptr()),
workspace_stride,
m,
n,
split_k);
check_hip_error("reduce_splitk_kernel");
}
return out;
}
"""
def _sanitize_b_layout(mode: str) -> str:
mode = (mode or _DEFAULT_B_LAYOUT).strip().lower()
if mode in {"raw", "shuffle", "auto"}:
return mode
return _DEFAULT_B_LAYOUT
def _get_b_layout() -> str:
mode = _sanitize_b_layout(os.getenv(_B_LAYOUT_ENV, _DEFAULT_B_LAYOUT))
if mode == "auto":
return "shuffle"
return mode
def _get_fused_a_quant() -> bool:
raw = os.getenv(_FUSED_A_QUANT_ENV, _DEFAULT_FUSED_A_QUANT)
return str(raw).strip().lower() in {"1", "true", "yes", "on"}
def _sanitize_native_mfma(mode: str) -> str:
mode = (mode or _DEFAULT_NATIVE_MFMA).strip().lower()
if mode in {"never", "auto", "force"}:
return mode
return _DEFAULT_NATIVE_MFMA
def _get_native_mfma_mode() -> str:
return _sanitize_native_mfma(os.getenv(_NATIVE_MFMA_ENV, _DEFAULT_NATIVE_MFMA))
def _sanitize_b_shuffle_inner_mode(mode: str | None) -> int:
if mode is None:
return _DEFAULT_B_SHUFFLE_INNER_MODE
try:
value = int(str(mode).strip())
except ValueError:
return _DEFAULT_B_SHUFFLE_INNER_MODE
return value if value in {0, 1} else _DEFAULT_B_SHUFFLE_INNER_MODE
def _get_b_shuffle_inner_mode() -> int:
return _sanitize_b_shuffle_inner_mode(os.getenv(_B_SHUFFLE_INNER_ENV))
def _get_splitk_override() -> int | None:
raw = os.getenv(_SPLITK_ENV)
if raw is None:
return None
try:
return max(0, int(raw))
except ValueError:
return None
def _e8m0_unshuffle(scale_sh: torch.Tensor) -> torch.Tensor:
if scale_sh.ndim != 2:
raise RuntimeError(f"scale_sh must be 2D, got {tuple(scale_sh.shape)}")
sm, sn = scale_sh.shape
if sm % 32 != 0 or sn % 8 != 0:
raise RuntimeError(f"scale_sh shape must be divisible by (32,8), got {(sm, sn)}")
s = scale_sh.view(torch.uint8)
s = s.view(sm // 32, sn // 8, 4, 16, 2, 2)
s = s.permute(0, 5, 3, 1, 4, 2).contiguous()
s = s.view(sm, sn)
return s.view(scale_sh.dtype)
def _get_b_scale_raw_cached(b_scale_sh: torch.Tensor) -> torch.Tensor:
dev = int(b_scale_sh.device.index) if b_scale_sh.device.index is not None else -1
key = (
int(b_scale_sh.data_ptr()),
int(b_scale_sh.shape[0]),
int(b_scale_sh.shape[1]),
dev,
)
cached = _B_SCALE_RAW_CACHE.get(key)
if cached is not None:
return cached
with _B_SCALE_RAW_LOCK:
cached = _B_SCALE_RAW_CACHE.get(key)
if cached is not None:
return cached
raw = _e8m0_unshuffle(b_scale_sh).contiguous()
if len(_B_SCALE_RAW_CACHE) >= _B_SCALE_RAW_CACHE_MAX:
_B_SCALE_RAW_CACHE.pop(next(iter(_B_SCALE_RAW_CACHE)))
_B_SCALE_RAW_CACHE[key] = raw
return raw
def _get_workspace(device: torch.device, m: int, n: int, split_k: int) -> torch.Tensor:
dev = int(device.index) if device.index is not None else -1
key = (dev, int(m), int(n), int(split_k))
cached = _WORKSPACE_CACHE.get(key)
need = split_k * m * n
if cached is not None and cached.numel() >= need:
return cached
with _WORKSPACE_LOCK:
cached = _WORKSPACE_CACHE.get(key)
if cached is not None and cached.numel() >= need:
return cached
ws = torch.empty((split_k, m * n), dtype=torch.float32, device=device)
if len(_WORKSPACE_CACHE) >= _WORKSPACE_CACHE_MAX:
_WORKSPACE_CACHE.pop(next(iter(_WORKSPACE_CACHE)))
_WORKSPACE_CACHE[key] = ws
return ws
def _pick_splitk_log2(m: int, n: int, k: int) -> int:
override = _get_splitk_override()
if override is not None:
return override
key = (int(m), int(n), int(k))
if key in _STATIC_SPLITK_LOG2:
return _STATIC_SPLITK_LOG2[key]
if m <= 16 and k >= 4096:
return 3
if m <= 32:
return 2
if m <= 64 and k >= 2048:
return 1
return 0
def _resolve_launch_policy(m: int, n: int, k: int, native_mfma_mode: str) -> _LaunchPolicy:
shape_key = (int(m), int(n), int(k))
is_ranked_shape = shape_key in _RANKED_SHAPE_KEYS
use_native_prequant = 0
use_native_fused = 0
if native_mfma_mode == "force":
use_native_prequant = 1
use_native_fused = 1
elif native_mfma_mode == "auto" and is_ranked_shape:
# Ranked benchmark shapes now have a correctness-validated native
# prequant path on gfx950. Keep fused-A native guarded behind `force`
# until that variant gets its own correctness sign-off.
use_native_prequant = 1
return _LaunchPolicy(
log2_k_split=_pick_splitk_log2(*shape_key),
use_native_prequant=use_native_prequant,
use_native_fused=use_native_fused,
is_ranked_shape=is_ranked_shape,
)
def _get_hip_module():
global _HIP_MODULE, _HIP_BUILD_ERROR, _HIP_NATIVE_BUILD_ENABLED, _HIP_BUILD_REPORT_EMITTED
if _HIP_MODULE is not None:
return _HIP_MODULE
if _HIP_BUILD_ERROR is not None:
raise RuntimeError(f"HIP inline build failed previously: {_HIP_BUILD_ERROR}")
with _HIP_LOCK:
if _HIP_MODULE is not None:
return _HIP_MODULE
if _HIP_BUILD_ERROR is not None:
raise RuntimeError(f"HIP inline build failed previously: {_HIP_BUILD_ERROR}")
native_mode = _get_native_mfma_mode()
build_attempts = [0]
if native_mode != "never":
build_attempts = [1, 0]
os.environ.setdefault("CXX", "clang++")
os.environ.setdefault("GPU_ARCHS", "gfx950")
os.environ.setdefault("AITER_GPU_ARCHS", "gfx950")
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
os.environ.setdefault("TORCH_HIP_ARCH_LIST", "gfx950")
last_error = None
for enable_native_fp4_mfma in build_attempts:
try:
_HIP_MODULE = load_inline(
name=f"mxfp4_mm_inline_quant_gemm_v13_nat{enable_native_fp4_mfma}",
cpp_sources=[CPP_WRAPPER],
cuda_sources=[HIP_SRC],
functions=["hip_quant_mxfp4", "hip_gemm_mxfp4_prequant", "hip_gemm_mxfp4"],
verbose=False,
extra_cuda_cflags=[
"-O3",
"-std=c++20",
"--offload-arch=gfx950",
f"-DMXFP4_ENABLE_NATIVE_FP4_MFMA={enable_native_fp4_mfma}",
],
)
_HIP_NATIVE_BUILD_ENABLED = bool(enable_native_fp4_mfma)
_HIP_BUILD_ERROR = None
if not _HIP_BUILD_REPORT_EMITTED:
print(
f"[mxfp4_mm] hip_inline_build requested={native_mode} "
f"built=nat{1 if _HIP_NATIVE_BUILD_ENABLED else 0}",
flush=True,
)
_HIP_BUILD_REPORT_EMITTED = True
break
except Exception as e: # pragma: no cover - runtime dependent
last_error = e
_HIP_MODULE = None
if enable_native_fp4_mfma == 1 and native_mode != "force":
if not _HIP_BUILD_REPORT_EMITTED:
print(
"[mxfp4_mm] hip_inline_build requested="
f"{native_mode} native_build_failed -> retry nat0",
flush=True,
)
_HIP_BUILD_REPORT_EMITTED = True
continue
_HIP_BUILD_ERROR = e
raise RuntimeError(f"HIP inline build failed: {e}") from e
if _HIP_MODULE is None:
_HIP_BUILD_ERROR = last_error
raise RuntimeError(f"HIP inline build failed: {last_error}")
return _HIP_MODULE
def _quant_triton_mxfp4(x: torch.Tensor, shuffle: bool = True):
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
if shuffle:
bs_e8m0 = e8m0_shuffle(bs_e8m0)
return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)
def _detect_b_shuffle_inner_mode(
module,
a_bf16: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
) -> int:
del module, a_bf16, b_shuffle, b_scale_sh
# Keep one fixed B_shuffle inner layout in the steady-state path. The
# pre-shuffled MXFP4 tensor produced by aiter.ops.shuffle.shuffle_weight()
# flattens 16x16 byte tiles in N-major inner order, which matches mode 1.
return _get_b_shuffle_inner_mode()
def _run_hip_full_pipeline(
a_bf16: torch.Tensor,
b_q: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
) -> torch.Tensor:
module = _get_hip_module()
b_layout = _get_b_layout()
native_mfma_mode = _get_native_mfma_mode()
if not _HIP_NATIVE_BUILD_ENABLED and native_mfma_mode != "never":
native_mfma_mode = "never"
m = int(a_bf16.shape[0])
k = int(a_bf16.shape[1])
n = int(b_q.shape[0])
launch_policy = _resolve_launch_policy(m, n, k, native_mfma_mode)
shape_key = (m, n, k)
with _RUNTIME_REPORT_LOCK:
if shape_key not in _RUNTIME_REPORTED_SHAPES:
print(
"[mxfp4_mm] launch "
f"shape={shape_key} requested={_get_native_mfma_mode()} "
f"effective={native_mfma_mode} build=nat{1 if _HIP_NATIVE_BUILD_ENABLED else 0} "
f"prequant_native={launch_policy.use_native_prequant} "
f"fused_native={launch_policy.use_native_fused} "
f"splitk_log2={launch_policy.log2_k_split}",
flush=True,
)
_RUNTIME_REPORTED_SHAPES.add(shape_key)
if b_layout == "shuffle":
b_u8 = b_shuffle.view(torch.uint8).contiguous()
b_scale_u8 = b_scale_sh.view(torch.uint8).contiguous()
layout_mode = 1
b_inner_mode = _detect_b_shuffle_inner_mode(module, a_bf16, b_shuffle, b_scale_sh)
else:
b_u8 = b_q.view(torch.uint8).contiguous()
b_scale_u8 = _get_b_scale_raw_cached(b_scale_sh).view(torch.uint8).contiguous()
layout_mode = 0
b_inner_mode = 0
log2_k_split = launch_policy.log2_k_split
split_k = 1 << log2_k_split
if split_k > 1:
ws = _get_workspace(a_bf16.device, m, n, split_k)
workspace_stride = int(m * n)
else:
ws = torch.empty((0,), dtype=torch.float32, device=a_bf16.device)
workspace_stride = 0
if _get_fused_a_quant():
return module.hip_gemm_mxfp4(
a_bf16,
b_u8,
b_scale_u8,
int(layout_mode),
int(log2_k_split),
ws,
int(workspace_stride),
int(b_inner_mode),
int(launch_policy.use_native_fused),
)
a_q, a_scale_raw = _quant_triton_mxfp4(a_bf16, shuffle=False)
a_fp4_u8 = a_q.view(torch.uint8).contiguous()
a_scale_u8 = a_scale_raw.view(torch.uint8).contiguous()
return module.hip_gemm_mxfp4_prequant(
a_fp4_u8,
b_u8,
a_scale_u8,
b_scale_u8,
int(layout_mode),
int(log2_k_split),
ws,
int(workspace_stride),
int(b_inner_mode),
int(launch_policy.use_native_prequant),
)
def custom_kernel(data: input_t) -> output_t:
"""
Default HIP path:
bf16 A -> Triton per-1x32 quant -> HIP split-k GEMM -> bf16 C.
Default validation mode now forces the prequant native FP4 scale-MFMA
path on gfx950-capable builds, so nat1 compile/runtime issues are not
hidden by the nat0 fallback.
Experimental fused HIP path:
set MXFP4_MM_FUSED_A_QUANT=1 to quantize A inside the GEMM prolog.
Native path controls:
MXFP4_MM_NATIVE_MFMA=auto uses native only for the ranked benchmark shapes.
MXFP4_MM_NATIVE_MFMA=force tries native on every eligible shape, including fused A experiments.
MXFP4_MM_NATIVE_MFMA=never keeps the BF16-unpack fallback.
MXFP4_MM_B_SHUFFLE_INNER_MODE lets debug runs override the fixed shuffled-B inner layout.
"""
A, B, B_q, B_shuffle, B_scale_sh = data
del B
A = A.contiguous()
B_q = B_q.contiguous()
B_shuffle = B_shuffle.contiguous()
B_scale_sh = B_scale_sh.contiguous()
return _run_hip_full_pipeline(
a_bf16=A,
b_q=B_q,
b_shuffle=B_shuffle,
b_scale_sh=B_scale_sh,
)
scrolls · 2695 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