submission 710593
kida023_89704 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 2104 lines, June 9 Researcher Reciprocity License v1.0.
submission_fused.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-710593?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:c99eb1d613262a9dec96724c7d7b84fa4bfd964346fbd3748735eff6ddca308f
license declaredunknown
license concludedunknown
authorskida023_89704
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
static_assert(sizeof(KernelArgs) == 384, "Unexpected FP4 kernarg size");num-warps = 1
num_warps = 1shared-memory
__shared__ __align__(16) uint8_t s_a[2][TILE_M * (TILE_K / 2)];split-k
splitk_enabled: intstages = 1
num_stages = 1tile-k = 256
constexpr int TILE_K = 256;tile-m = 16
static_assert(TILE_M == 16 || TILE_M == 32, "Only TILE_M=16/32 is supported");tile-n = 128
constexpr int TILE_N = 128;Kernel source
submission_fused.py2104 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
from __future__ import annotations
import csv
import os
from dataclasses import dataclass
from functools import lru_cache
from pathlib import Path
os.environ.setdefault("CXX", "clang++")
os.environ.setdefault("HSA_XNACK", "0")
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
_ARCH = "gfx950"
_F4GEMM_SUBDIR = "f4gemm"
_CUSTOM_SMALLM_MAX_M = 32
_BENCHMARK_QUANT_WORKSPACE_MAX = 4
_BENCHMARK_OUT_CACHE_MAX = 4
_CUSTOM_N_TILE = 128
_FAST_FUSED_ENABLE = True
_SUPPORTED_SHAPES = {
(4, 2880, 512),
(8, 2112, 7168),
(16, 2112, 7168),
(16, 3072, 1536),
(32, 2880, 512),
(32, 4096, 512),
(64, 3072, 1536),
(64, 7168, 2048),
(256, 2880, 512),
(256, 3072, 1536),
}
_SMALLK_BENCHMARK_SHAPES = {
(4, 2880, 512),
(32, 2880, 512),
(32, 4096, 512),
}
_COMBINED_HIP_QUANT_GEMM_SHAPES = {
(4, 2880, 512),
(16, 2112, 7168),
(32, 2880, 512),
(32, 4096, 512),
(64, 7168, 2048),
(256, 3072, 1536),
}
# Optional per-shape override:
# (m, n, k): (kernel_name, co_name_with_subdir, log2_k_split)
_KERNEL_OVERRIDE_BY_SHAPE: dict[tuple[int, int, int], tuple[str, str, int]] = {
(4, 2880, 512): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
"f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co",
0,
),
(8, 2112, 7168): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
"f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co",
0,
),
(16, 2112, 7168): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
"f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co",
0,
),
(16, 3072, 1536): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x768E",
"f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x768.co",
0,
),
(32, 2880, 512): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
"f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co",
0,
),
(32, 4096, 512): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
"f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co",
0,
),
(64, 3072, 1536): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
"f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co",
0,
),
(64, 7168, 2048): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
"f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co",
0,
),
(256, 2880, 512): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
"f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co",
0,
),
(256, 3072, 1536): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
"f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co",
0,
),
}
@dataclass(frozen=True)
class _KernelCfg:
tile_m: int
tile_n: int
splitk_enabled: int
bpreshuffle: int
kernel_name: str
co_name: str
@dataclass(frozen=True)
class _TunedKernel:
cu_num: int
m: int
n: int
k: int
split_k: int
kernel_name: str
@dataclass
class _QuantWorkspaceEntry:
device_type: str
device_index: int
m: int
n: int
q_rows: int
q: torch.Tensor
scale_sh: torch.Tensor
@dataclass
class _OutCacheEntry:
device_type: str
device_index: int
dtype: torch.dtype
padded_m: int
n: int
out: torch.Tensor
_B_SHUFFLE_CACHE: dict[tuple[int, int, int, int, int], tuple[torch.Tensor, torch.Tensor, int]] = {}
_BENCHMARK_QUANT_WORKSPACES: list[_QuantWorkspaceEntry] = []
_BENCHMARK_OUT_CACHE: list[_OutCacheEntry] = []
def _is_benchmark_fastpath_shape(shape: tuple[int, int, int]) -> bool:
return shape in _SUPPORTED_SHAPES
def _is_smallk_benchmark_shape(shape: tuple[int, int, int]) -> bool:
return shape in _SMALLK_BENCHMARK_SHAPES
def _load_inline_with_trace(tag: str, **kwargs):
print(f"[submission-ext] tag={tag} stage=load-start", flush=True)
module = load_inline(**kwargs)
print(f"[submission-ext] tag={tag} stage=load-done", flush=True)
return module
_FAST_FUSED_CPP_SRC = r"""
#include <torch/extension.h>
void launch_fast_quant_mxfp4(
torch::Tensor a_bf16,
torch::Tensor q_out,
torch::Tensor scale_sh_out,
int real_m,
int k);
void launch_fast_quant_and_f4gemm(
torch::Tensor a_bf16,
torch::Tensor q_out,
torch::Tensor scale_sh_out,
torch::Tensor b_shuffle,
torch::Tensor b_scale_sh,
torch::Tensor out,
std::string co_path,
std::string kernel_name,
int tile_m,
int tile_n,
int log2_k_split,
int real_m,
int k);
void launch_fast_fused_mxfp4_gemm(
torch::Tensor a_bf16,
torch::Tensor b_shuffle_u8,
torch::Tensor b_scale_sh_u8,
torch::Tensor out,
int padded_n,
int real_n);
"""
_FAST_FUSED_CUDA_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/hip_bfloat16.h>
#include <c10/hip/HIPFunctions.h>
#include <cmath>
#include <cstdint>
#include <mutex>
#include <stdexcept>
#include <string>
#include <unordered_map>
namespace {
using i32x8_t = int __attribute__((ext_vector_type(8)));
using fp32x4_t = float __attribute__((ext_vector_type(4)));
struct p3 {
uint32_t x;
uint32_t y;
uint32_t z;
};
struct p2 {
uint32_t x;
uint32_t y;
};
struct __attribute__((packed)) KernelArgs {
void* ptr_D;
p2 _p0;
void* ptr_C;
p2 _p1;
void* ptr_A;
p2 _p2;
void* ptr_B;
p2 _p3;
float alpha;
p3 _p4;
float beta;
p3 _p5;
uint32_t stride_D0;
p3 _p6;
uint32_t stride_D1;
p3 _p7;
uint32_t stride_C0;
p3 _p8;
uint32_t stride_C1;
p3 _p9;
uint32_t stride_A0;
p3 _p10;
uint32_t stride_A1;
p3 _p11;
uint32_t stride_B0;
p3 _p12;
uint32_t stride_B1;
p3 _p13;
uint32_t Mdim;
p3 _p14;
uint32_t Ndim;
p3 _p15;
uint32_t Kdim;
p3 _p16;
void* ptr_ScaleA;
p2 _p17;
void* ptr_ScaleB;
p2 _p18;
uint32_t stride_ScaleA0;
p3 _p19;
uint32_t stride_ScaleA1;
p3 _p20;
uint32_t stride_ScaleB0;
p3 _p21;
uint32_t stride_ScaleB1;
p3 _p22;
int32_t log2_k_split;
p3 _p23;
};
static_assert(sizeof(KernelArgs) == 384, "Unexpected FP4 kernarg size");
struct CachedKernel {
hipModule_t module = nullptr;
hipFunction_t func = nullptr;
};
std::unordered_map<std::string, CachedKernel>& fast_kernel_cache() {
static std::unordered_map<std::string, CachedKernel> cache;
return cache;
}
std::mutex& fast_kernel_cache_mutex() {
static std::mutex mu;
return mu;
}
void hip_check(hipError_t err, const char* call_name) {
if (err == hipSuccess) {
return;
}
throw std::runtime_error(std::string(call_name) + " failed: " + hipGetErrorString(err));
}
CachedKernel& get_fast_kernel(const std::string& co_path, const std::string& kernel_name) {
const std::string key = co_path + "|" + kernel_name;
std::lock_guard<std::mutex> guard(fast_kernel_cache_mutex());
auto& cache = fast_kernel_cache();
auto it = cache.find(key);
if (it != cache.end()) {
return it->second;
}
CachedKernel entry;
hip_check(hipModuleLoad(&entry.module, co_path.c_str()), "hipModuleLoad");
hip_check(hipModuleGetFunction(&entry.func, entry.module, kernel_name.c_str()), "hipModuleGetFunction");
auto [new_it, _inserted] = cache.emplace(key, entry);
return new_it->second;
}
void pick_quant_launch_config(int real_m, int k, int* threads, int* blocks) {
const int total_tasks = real_m * (k / 32);
if (total_tasks <= 64) {
*threads = 64;
*blocks = 1;
return;
}
if (total_tasks <= 1024) {
*threads = 128;
*blocks = max(1, min(8, (total_tasks + *threads - 1) / *threads));
return;
}
*threads = 256;
*blocks = max(1, min(120, (total_tasks + *threads - 1) / *threads));
}
__device__ __forceinline__ float decode_e8m0(uint8_t x) {
if (x == 0) return __uint_as_float(0x00400000u);
if (x == 0xFF) return __uint_as_float(0x7F800001u);
return __uint_as_float(static_cast<uint32_t>(x) << 23);
}
__device__ __forceinline__ uint8_t amax_to_scale_e8m0(float amax) {
const uint32_t rounded_bits = (__float_as_uint(amax) + 0x200000u) & 0xFF800000u;
int exponent = static_cast<int>((rounded_bits >> 23) & 0xFFu) - 127;
exponent -= 2;
exponent = max(-127, min(127, exponent));
return static_cast<uint8_t>(exponent + 127);
}
__device__ __forceinline__ uint8_t f32_to_fp4_bits(float x) {
constexpr int MBITS_F32 = 23;
constexpr int EBITS_F32 = 8;
constexpr int ebits = 2;
constexpr int mbits = 1;
constexpr uint8_t max_int = static_cast<uint8_t>((1 << (ebits + mbits)) - 1);
constexpr uint8_t sign_mask = static_cast<uint8_t>(1 << (ebits + mbits));
constexpr int exp_bias = (1 << (ebits - 1)) - 1;
constexpr int f32_exp_bias = (1 << (EBITS_F32 - 1)) - 1;
constexpr int magic_adder = (1 << (MBITS_F32 - mbits - 1)) - 1;
constexpr float max_normal = 6.0f;
constexpr float min_normal = 1.0f;
constexpr int denorm_exp = (f32_exp_bias - exp_bias) + (MBITS_F32 - mbits) + 1;
constexpr int denorm_mask_int = denorm_exp << MBITS_F32;
const float denorm_mask_float = __uint_as_float(static_cast<uint32_t>(denorm_mask_int));
const uint32_t raw = __float_as_uint(x);
const uint32_t sign = raw & 0x80000000u;
const uint32_t mag_bits = raw ^ sign;
const float mag = __uint_as_float(mag_bits);
const bool saturate_mask = mag >= max_normal;
const bool denormal_mask = (!saturate_mask) && (mag < min_normal);
const bool normal_mask = (!saturate_mask) && (!denormal_mask);
uint8_t out = max_int;
if (denormal_mask) {
int denormal_x = __float_as_int(mag + denorm_mask_float);
denormal_x -= denorm_mask_int;
out = static_cast<uint8_t>(denormal_x);
} else if (normal_mask) {
int normal_x = static_cast<int>(mag_bits);
const int mant_odd = (normal_x >> (MBITS_F32 - mbits)) & 1;
const int val_to_add = ((exp_bias - f32_exp_bias) << MBITS_F32) + magic_adder;
normal_x += val_to_add;
normal_x += mant_odd;
normal_x >>= (MBITS_F32 - mbits);
out = static_cast<uint8_t>(normal_x);
}
const uint8_t sign_lp = static_cast<uint8_t>((sign >> (MBITS_F32 + EBITS_F32 - mbits - ebits)) & sign_mask);
return static_cast<uint8_t>(out | sign_lp);
}
template <int ScaleASel, int ScaleBSel>
__device__ __forceinline__ fp32x4_t mfma_scale_16x16x128(
i32x8_t a_reg, i32x8_t b_reg, fp32x4_t c_reg, int scale_a_word, int scale_b_word) {
#if defined(__gfx950__)
return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a_reg, b_reg, c_reg, 4, 4, ScaleASel, scale_a_word, ScaleBSel, scale_b_word);
#else
return c_reg;
#endif
}
__device__ __forceinline__ void sched_barrier() {
#if defined(__gfx950__)
__builtin_amdgcn_sched_barrier(0);
#endif
}
__device__ __forceinline__ float reduce_max4(float x) {
x = fmaxf(x, __shfl_xor(x, 1, 4));
x = fmaxf(x, __shfl_xor(x, 2, 4));
return x;
}
__device__ __forceinline__ int swizzle_xor16(int row, int col, int k_blocks16) {
return col ^ ((row % k_blocks16) * 16);
}
__device__ __forceinline__ i32x8_t pack_i64x4_to_i32x8(uint64_t x0, uint64_t x1, uint64_t x2, uint64_t x3) {
struct Pack {
uint64_t u64[4];
};
const Pack p{{x0, x1, x2, x3}};
return __builtin_bit_cast(i32x8_t, p);
}
template <int TILE_M>
__device__ __forceinline__ void quantize_a_tile_to_stage(
const hip_bfloat16* __restrict__ a_bf16, int m, int a_stride, int tile_m_base, int tile_k_base, int tx,
uint8_t* __restrict__ s_a_stage, uint8_t* __restrict__ s_scale_stage) {
constexpr int TILE_K = 256;
constexpr int SCALE_GROUPS = TILE_K / 32;
constexpr int K_BLOCKS16 = TILE_K / 32;
constexpr int BLOCK_THREADS = 256;
constexpr int TASKS_PER_ROW = TILE_K / 8;
constexpr int TASKS_TOTAL = TILE_M * TASKS_PER_ROW;
static_assert(TASKS_TOTAL % BLOCK_THREADS == 0);
constexpr int TASKS_PER_THREAD = TASKS_TOTAL / BLOCK_THREADS;
#pragma unroll
for (int i = 0; i < TASKS_PER_THREAD; ++i) {
const int task_id = i * BLOCK_THREADS + tx;
const int row_local = task_id / TASKS_PER_ROW;
const int col_task = task_id % TASKS_PER_ROW;
const int block32_idx = col_task / 4;
const int in_block_task = col_task % 4;
const bool is_scale_leader = (in_block_task == 0);
const int src_row = tile_m_base + row_local;
float vals[8];
float max_abs = 0.0f;
#pragma unroll
for (int j = 0; j < 8; ++j) {
float v = 0.0f;
if (src_row < m) {
const int src_col = tile_k_base + col_task * 8 + j;
v = static_cast<float>(a_bf16[src_row * a_stride + src_col]);
}
vals[j] = v;
max_abs = fmaxf(max_abs, fabsf(v));
}
const float group_amax = reduce_max4(max_abs);
const uint8_t scale_byte = (src_row < m) ? amax_to_scale_e8m0(group_amax) : static_cast<uint8_t>(127);
const float inv_scale = (src_row < m && scale_byte != 0) ? (1.0f / decode_e8m0(scale_byte)) : 0.0f;
uint32_t packed = 0;
#pragma unroll
for (int j = 0; j < 8; ++j) {
const uint8_t nibble = f32_to_fp4_bits(vals[j] * inv_scale);
packed |= static_cast<uint32_t>(nibble) << (j * 4);
}
const int col_local_bytes = col_task * 4;
const int col_swz_bytes = swizzle_xor16(row_local, col_local_bytes, K_BLOCKS16);
reinterpret_cast<uint32_t*>(s_a_stage + row_local * (TILE_K / 2) + col_swz_bytes)[0] = packed;
if (is_scale_leader) {
s_scale_stage[row_local * SCALE_GROUPS + block32_idx] = scale_byte;
}
}
}
template <int TILE_M>
__global__ __launch_bounds__(256) void fast_fused_mxfp4_mfma_kernel(
const hip_bfloat16* __restrict__ a_bf16,
const uint8_t* __restrict__ b_shuffle,
const uint8_t* __restrict__ b_scale_sh,
hip_bfloat16* __restrict__ out,
int m,
int n_padded,
int real_n,
int k,
int a_stride,
int out_stride) {
constexpr int TILE_K = 256;
constexpr int TILE_N = 128;
constexpr int SCALE_BLOCK_K = 32;
constexpr int SCALE_GROUPS = TILE_K / SCALE_BLOCK_K;
constexpr int K_BLOCKS16 = TILE_K / 32;
constexpr int ROW_GROUPS = TILE_M / 16;
constexpr int ACCUMULATORS_PER_WAVE = ROW_GROUPS * 2;
static_assert(TILE_M == 16 || TILE_M == 32, "Only TILE_M=16/32 is supported");
__shared__ __align__(16) uint8_t s_a[2][TILE_M * (TILE_K / 2)];
__shared__ __align__(4) uint8_t s_scale[2][TILE_M * SCALE_GROUPS];
__shared__ __align__(16) hip_bfloat16 s_out[TILE_M][TILE_N];
const int tx = static_cast<int>(threadIdx.x);
const int lane_id = tx & 63;
const int wave_id = tx >> 6;
const int lane_div_16 = lane_id >> 4;
const int lane_mod_16 = lane_id & 15;
const int tile_m_base = static_cast<int>(blockIdx.y) * TILE_M;
const int tile_n_base = static_cast<int>(blockIdx.x) * TILE_N;
fp32x4_t acc[ACCUMULATORS_PER_WAVE];
#pragma unroll
for (int i = 0; i < ACCUMULATORS_PER_WAVE; ++i) {
acc[i] = fp32x4_t{0.0f, 0.0f, 0.0f, 0.0f};
}
const int k0_blocks = k / 128;
const int scale_k_tiles = k / 256;
const int num_k_tiles = k / TILE_K;
const int row_a_lds = lane_mod_16;
const int col_offset_base_bytes = lane_div_16 * 16;
quantize_a_tile_to_stage<TILE_M>(a_bf16, m, a_stride, tile_m_base, 0, tx, s_a[0], s_scale[0]);
__syncthreads();
for (int tile_idx = 0; tile_idx < num_k_tiles; ++tile_idx) {
const int curr = tile_idx & 1;
const int next = curr ^ 1;
const int tile_k_base = tile_idx * TILE_K;
if (tile_idx + 1 < num_k_tiles) {
quantize_a_tile_to_stage<TILE_M>(
a_bf16, m, a_stride, tile_m_base, (tile_idx + 1) * TILE_K, tx, s_a[next], s_scale[next]);
}
const int row0 = lane_mod_16;
const int blk0 = lane_div_16;
const int blk1 = lane_div_16 + 4;
const uint32_t a_scale_row0_blk0 = static_cast<uint32_t>(s_scale[curr][row0 * SCALE_GROUPS + blk0]);
const uint32_t a_scale_row0_blk1 = static_cast<uint32_t>(s_scale[curr][row0 * SCALE_GROUPS + blk1]);
uint32_t a_scale_word = a_scale_row0_blk0 | (a_scale_row0_blk1 << 16);
if constexpr (TILE_M == 32) {
const int row1 = row0 + 16;
a_scale_word |= static_cast<uint32_t>(s_scale[curr][row1 * SCALE_GROUPS + blk0]) << 8;
a_scale_word |= static_cast<uint32_t>(s_scale[curr][row1 * SCALE_GROUPS + blk1]) << 24;
} else {
a_scale_word |= 0x7F00u | 0x7F000000u;
}
const int n_pack = (tile_n_base + wave_id * 32) / 32;
const int k_pack = tile_k_base / 256;
uint32_t b_scale_word = 0x7F7F7F7F;
if (n_pack < (real_n / 32)) {
const int b_scale_word_idx = (((n_pack * scale_k_tiles + k_pack) * 4 + lane_div_16) * 16 + lane_mod_16);
b_scale_word = reinterpret_cast<const uint32_t*>(b_scale_sh)[b_scale_word_idx];
}
#pragma unroll
for (int k_half = 0; k_half < 2; ++k_half) {
const int col_base_bytes = col_offset_base_bytes + k_half * 64;
const int a_row0_col = swizzle_xor16(row_a_lds, col_base_bytes, K_BLOCKS16);
const uint64_t a00 = reinterpret_cast<const uint64_t*>(s_a[curr] + row_a_lds * (TILE_K / 2) + a_row0_col)[0];
const uint64_t a01 = reinterpret_cast<const uint64_t*>(s_a[curr] + row_a_lds * (TILE_K / 2) + a_row0_col)[1];
const i32x8_t a_vec0 = pack_i64x4_to_i32x8(a00, a01, 0ull, 0ull);
i32x8_t a_vec1 = i32x8_t{};
if constexpr (TILE_M == 32) {
const int a_row1_col = swizzle_xor16(row_a_lds + 16, col_base_bytes, K_BLOCKS16);
const uint64_t a10 = reinterpret_cast<const uint64_t*>(s_a[curr] + (row_a_lds + 16) * (TILE_K / 2) + a_row1_col)[0];
const uint64_t a11 = reinterpret_cast<const uint64_t*>(s_a[curr] + (row_a_lds + 16) * (TILE_K / 2) + a_row1_col)[1];
a_vec1 = pack_i64x4_to_i32x8(a10, a11, 0ull, 0ull);
}
const int n_blk0 = (tile_n_base + wave_id * 32 + lane_mod_16) / 16;
const int n_blk1 = (tile_n_base + wave_id * 32 + 16 + lane_mod_16) / 16;
const int n_intra = lane_mod_16;
const int k0 = tile_k_base / 128 + k_half;
uint64_t b00 = 0;
uint64_t b01 = 0;
if (n_blk0 < (real_n / 16)) {
const int b_idx0 = ((((n_blk0 * k0_blocks + k0) * 4 + lane_div_16) * 16 + n_intra) * 16);
b00 = reinterpret_cast<const uint64_t*>(b_shuffle + b_idx0)[0];
b01 = reinterpret_cast<const uint64_t*>(b_shuffle + b_idx0)[1];
}
const i32x8_t b_vec0 = pack_i64x4_to_i32x8(b00, b01, 0ull, 0ull);
uint64_t b10 = 0;
uint64_t b11 = 0;
if (n_blk1 < (real_n / 16)) {
const int b_idx1 = ((((n_blk1 * k0_blocks + k0) * 4 + lane_div_16) * 16 + n_intra) * 16);
b10 = reinterpret_cast<const uint64_t*>(b_shuffle + b_idx1)[0];
b11 = reinterpret_cast<const uint64_t*>(b_shuffle + b_idx1)[1];
}
const i32x8_t b_vec1 = pack_i64x4_to_i32x8(b10, b11, 0ull, 0ull);
sched_barrier();
if (k_half == 0) {
acc[0] = mfma_scale_16x16x128<0, 0>(
a_vec0, b_vec0, acc[0], static_cast<int>(a_scale_word), static_cast<int>(b_scale_word));
acc[1] = mfma_scale_16x16x128<0, 1>(
a_vec0, b_vec1, acc[1], static_cast<int>(a_scale_word), static_cast<int>(b_scale_word));
if constexpr (TILE_M == 32) {
acc[2] = mfma_scale_16x16x128<1, 0>(
a_vec1, b_vec0, acc[2], static_cast<int>(a_scale_word), static_cast<int>(b_scale_word));
acc[3] = mfma_scale_16x16x128<1, 1>(
a_vec1, b_vec1, acc[3], static_cast<int>(a_scale_word), static_cast<int>(b_scale_word));
}
} else {
acc[0] = mfma_scale_16x16x128<2, 2>(
a_vec0, b_vec0, acc[0], static_cast<int>(a_scale_word), static_cast<int>(b_scale_word));
acc[1] = mfma_scale_16x16x128<2, 3>(
a_vec0, b_vec1, acc[1], static_cast<int>(a_scale_word), static_cast<int>(b_scale_word));
if constexpr (TILE_M == 32) {
acc[2] = mfma_scale_16x16x128<3, 2>(
a_vec1, b_vec0, acc[2], static_cast<int>(a_scale_word), static_cast<int>(b_scale_word));
acc[3] = mfma_scale_16x16x128<3, 3>(
a_vec1, b_vec1, acc[3], static_cast<int>(a_scale_word), static_cast<int>(b_scale_word));
}
}
sched_barrier();
}
if (tile_idx + 1 < num_k_tiles) {
__syncthreads();
}
}
const int col_base = tile_n_base + wave_id * 32 + lane_mod_16;
const int col_local = wave_id * 32 + lane_mod_16;
#pragma unroll
for (int row_group = 0; row_group < ROW_GROUPS; ++row_group) {
#pragma unroll
for (int ii = 0; ii < 4; ++ii) {
const int row_in_tile = row_group * 16 + lane_div_16 * 4 + ii;
s_out[row_in_tile][col_local] = static_cast<hip_bfloat16>(acc[row_group * 2 + 0][ii]);
s_out[row_in_tile][col_local + 16] = static_cast<hip_bfloat16>(acc[row_group * 2 + 1][ii]);
}
}
__syncthreads();
#pragma unroll
for (int row_group = 0; row_group < ROW_GROUPS; ++row_group) {
#pragma unroll
for (int ii = 0; ii < 4; ++ii) {
const int row_in_tile = row_group * 16 + lane_div_16 * 4 + ii;
const int row = tile_m_base + row_in_tile;
if (row >= m) continue;
out[static_cast<int64_t>(row) * out_stride + col_base] = s_out[row_in_tile][col_local];
out[static_cast<int64_t>(row) * out_stride + col_base + 16] = s_out[row_in_tile][col_local + 16];
}
}
}
} // namespace
void launch_fast_fused_mxfp4_gemm(
torch::Tensor a_bf16,
torch::Tensor b_shuffle_u8,
torch::Tensor b_scale_sh_u8,
torch::Tensor out,
int padded_n,
int real_n) {
const int m = static_cast<int>(a_bf16.size(0));
const int k = static_cast<int>(a_bf16.size(1));
const bool use_m16_kernel = m <= 16;
const int tile_m = use_m16_kernel ? 16 : 32;
const dim3 grid(static_cast<unsigned int>(padded_n / 128), static_cast<unsigned int>((m + tile_m - 1) / tile_m), 1);
const dim3 block(256, 1, 1);
if (use_m16_kernel) {
hipLaunchKernelGGL(
HIP_KERNEL_NAME((fast_fused_mxfp4_mfma_kernel<16>)),
grid, block, 0, 0,
reinterpret_cast<const hip_bfloat16*>(a_bf16.data_ptr()),
reinterpret_cast<const uint8_t*>(b_shuffle_u8.data_ptr()),
reinterpret_cast<const uint8_t*>(b_scale_sh_u8.data_ptr()),
reinterpret_cast<hip_bfloat16*>(out.data_ptr()),
m, padded_n, real_n, k,
static_cast<int>(a_bf16.stride(0)), static_cast<int>(out.stride(0)));
} else {
hipLaunchKernelGGL(
HIP_KERNEL_NAME((fast_fused_mxfp4_mfma_kernel<32>)),
grid, block, 0, 0,
reinterpret_cast<const hip_bfloat16*>(a_bf16.data_ptr()),
reinterpret_cast<const uint8_t*>(b_shuffle_u8.data_ptr()),
reinterpret_cast<const uint8_t*>(b_scale_sh_u8.data_ptr()),
reinterpret_cast<hip_bfloat16*>(out.data_ptr()),
m, padded_n, real_n, k,
static_cast<int>(a_bf16.stride(0)), static_cast<int>(out.stride(0)));
}
const hipError_t err = hipGetLastError();
TORCH_CHECK(err == hipSuccess, "fast_fused_mxfp4_mfma_kernel launch failed: ", hipGetErrorString(err));
}
__global__ __launch_bounds__(256) void fast_quant_mxfp4_kernel(
const hip_bfloat16* __restrict__ a_bf16,
uint8_t* __restrict__ q_out,
uint8_t* __restrict__ scale_sh_out,
int real_m,
int k,
int a_stride,
int q_stride,
int scale_stride,
int scale_n_valid,
int scale_n_pad) {
const int block32_per_row = k / 32;
const int total_tasks = real_m * block32_per_row;
const int tid = static_cast<int>(blockIdx.x) * static_cast<int>(blockDim.x) + static_cast<int>(threadIdx.x);
const int stride = static_cast<int>(blockDim.x) * static_cast<int>(gridDim.x);
for (int task = tid; task < total_tasks; task += stride) {
const int row = task / block32_per_row;
const int blk = task % block32_per_row;
const int k_base = blk * 32;
float vals[32];
float max_abs = 0.0f;
#pragma unroll
for (int i = 0; i < 32; ++i) {
const float v = static_cast<float>(a_bf16[static_cast<int64_t>(row) * a_stride + k_base + i]);
vals[i] = v;
max_abs = fmaxf(max_abs, fabsf(v));
}
const uint8_t scale_byte = amax_to_scale_e8m0(max_abs);
const float inv_scale = scale_byte != 0 ? (1.0f / decode_e8m0(scale_byte)) : 0.0f;
uint32_t packed[4] = {0u, 0u, 0u, 0u};
#pragma unroll
for (int pack_idx = 0; pack_idx < 4; ++pack_idx) {
uint32_t out = 0u;
#pragma unroll
for (int j = 0; j < 8; ++j) {
const uint8_t nibble = f32_to_fp4_bits(vals[pack_idx * 8 + j] * inv_scale);
out |= static_cast<uint32_t>(nibble) << (j * 4);
}
packed[pack_idx] = out;
}
uint32_t* q_row_ptr = reinterpret_cast<uint32_t*>(q_out + static_cast<int64_t>(row) * q_stride + blk * 16);
q_row_ptr[0] = packed[0];
q_row_ptr[1] = packed[1];
q_row_ptr[2] = packed[2];
q_row_ptr[3] = packed[3];
if (blk < scale_n_valid) {
const int bs_offs_0 = row / 32;
const int bs_offs_1 = (row % 32) / 16;
const int bs_offs_2 = row % 16;
const int bs_offs_3 = blk / 8;
const int bs_offs_4 = (blk % 8) / 4;
const int bs_offs_5 = blk % 4;
const int flat = (
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)
);
scale_sh_out[flat] = scale_byte;
}
}
}
void launch_fast_quant_mxfp4(
torch::Tensor a_bf16,
torch::Tensor q_out,
torch::Tensor scale_sh_out,
int real_m,
int k) {
TORCH_CHECK(a_bf16.is_cuda(), "a_bf16 must be a HIP tensor");
TORCH_CHECK(q_out.is_cuda(), "q_out must be a HIP tensor");
TORCH_CHECK(scale_sh_out.is_cuda(), "scale_sh_out must be a HIP tensor");
const int scale_n_valid = (k + 31) / 32;
int threads = 256;
int blocks = 1;
pick_quant_launch_config(real_m, k, &threads, &blocks);
hipLaunchKernelGGL(
HIP_KERNEL_NAME(fast_quant_mxfp4_kernel),
dim3(static_cast<unsigned int>(blocks), 1, 1),
dim3(static_cast<unsigned int>(threads), 1, 1),
0,
0,
reinterpret_cast<const hip_bfloat16*>(a_bf16.data_ptr()),
reinterpret_cast<uint8_t*>(q_out.data_ptr()),
reinterpret_cast<uint8_t*>(scale_sh_out.data_ptr()),
real_m,
k,
static_cast<int>(a_bf16.stride(0)),
static_cast<int>(q_out.stride(0)),
static_cast<int>(scale_sh_out.stride(0)),
scale_n_valid,
static_cast<int>(scale_sh_out.stride(0)));
const hipError_t err = hipGetLastError();
TORCH_CHECK(err == hipSuccess, "fast_quant_mxfp4_kernel launch failed: ", hipGetErrorString(err));
}
void launch_fast_quant_and_f4gemm(
torch::Tensor a_bf16,
torch::Tensor q_out,
torch::Tensor scale_sh_out,
torch::Tensor b_shuffle,
torch::Tensor b_scale_sh,
torch::Tensor out,
std::string co_path,
std::string kernel_name,
int tile_m,
int tile_n,
int log2_k_split,
int real_m,
int k) {
TORCH_CHECK(a_bf16.is_cuda(), "a_bf16 must be a HIP tensor");
TORCH_CHECK(q_out.is_cuda(), "q_out must be a HIP tensor");
TORCH_CHECK(scale_sh_out.is_cuda(), "scale_sh_out must be a HIP tensor");
TORCH_CHECK(b_shuffle.is_cuda(), "b_shuffle must be a HIP tensor");
TORCH_CHECK(b_scale_sh.is_cuda(), "b_scale_sh must be a HIP tensor");
TORCH_CHECK(out.is_cuda(), "out must be a HIP tensor");
const int scale_n_valid = (k + 31) / 32;
int threads = 256;
int blocks = 1;
pick_quant_launch_config(real_m, k, &threads, &blocks);
hipLaunchKernelGGL(
HIP_KERNEL_NAME(fast_quant_mxfp4_kernel),
dim3(static_cast<unsigned int>(blocks), 1, 1),
dim3(static_cast<unsigned int>(threads), 1, 1),
0,
0,
reinterpret_cast<const hip_bfloat16*>(a_bf16.data_ptr()),
reinterpret_cast<uint8_t*>(q_out.data_ptr()),
reinterpret_cast<uint8_t*>(scale_sh_out.data_ptr()),
real_m,
k,
static_cast<int>(a_bf16.stride(0)),
static_cast<int>(q_out.stride(0)),
static_cast<int>(scale_sh_out.stride(0)),
scale_n_valid,
static_cast<int>(scale_sh_out.stride(0)));
hip_check(hipGetLastError(), "fast_quant_mxfp4_kernel launch");
auto& kernel = get_fast_kernel(co_path, kernel_name);
KernelArgs args{};
const int m = static_cast<int>(out.size(0));
const int n = static_cast<int>(out.size(1));
args.ptr_D = out.data_ptr();
args.ptr_C = nullptr;
args.ptr_A = q_out.data_ptr();
args.ptr_B = b_shuffle.data_ptr();
args.alpha = 1.0f;
args.beta = 0.0f;
args.stride_C0 = static_cast<uint32_t>(out.stride(0));
args.stride_A0 = static_cast<uint32_t>(q_out.stride(0) * 2);
args.stride_B0 = static_cast<uint32_t>(b_shuffle.stride(0) * 2);
args.Mdim = static_cast<uint32_t>(m);
args.Ndim = static_cast<uint32_t>(n);
args.Kdim = static_cast<uint32_t>(k);
args.ptr_ScaleA = scale_sh_out.data_ptr();
args.ptr_ScaleB = b_scale_sh.data_ptr();
args.stride_ScaleA0 = static_cast<uint32_t>(scale_sh_out.stride(0));
args.stride_ScaleB0 = static_cast<uint32_t>(b_scale_sh.stride(0));
args.log2_k_split = 0;
int gdz = 1;
if (log2_k_split > 0) {
args.log2_k_split = log2_k_split;
const int split_k = 1 << args.log2_k_split;
TORCH_CHECK(k % split_k == 0, "K must be divisible by split-K factor");
if (split_k > 1) {
hip_check(
hipMemsetAsync(out.data_ptr(), 0, out.numel() * out.element_size(), nullptr),
"hipMemsetAsync");
}
const int k_per_tg = ((k / split_k + 255) / 256) * 256;
gdz = (k + k_per_tg - 1) / k_per_tg;
}
size_t arg_size = sizeof(args);
void* config[] = {
HIP_LAUNCH_PARAM_BUFFER_POINTER,
&args,
HIP_LAUNCH_PARAM_BUFFER_SIZE,
&arg_size,
HIP_LAUNCH_PARAM_END,
};
hip_check(
hipModuleLaunchKernel(
kernel.func,
static_cast<unsigned int>((n + tile_n - 1) / tile_n),
static_cast<unsigned int>((m + tile_m - 1) / tile_m),
static_cast<unsigned int>(gdz),
256,
1,
1,
0,
nullptr,
nullptr,
reinterpret_cast<void**>(config)),
"hipModuleLaunchKernel");
}
"""
_CPP_SRC = r"""
#include <torch/extension.h>
#include <c10/hip/HIPFunctions.h>
#include <hip/hip_runtime.h>
#include <cstdint>
#include <mutex>
#include <stdexcept>
#include <string>
#include <unordered_map>
namespace {
struct p3 {
uint32_t x;
uint32_t y;
uint32_t z;
};
struct p2 {
uint32_t x;
uint32_t y;
};
struct __attribute__((packed)) KernelArgs {
void* ptr_D;
p2 _p0;
void* ptr_C;
p2 _p1;
void* ptr_A;
p2 _p2;
void* ptr_B;
p2 _p3;
float alpha;
p3 _p4;
float beta;
p3 _p5;
uint32_t stride_D0;
p3 _p6;
uint32_t stride_D1;
p3 _p7;
uint32_t stride_C0;
p3 _p8;
uint32_t stride_C1;
p3 _p9;
uint32_t stride_A0;
p3 _p10;
uint32_t stride_A1;
p3 _p11;
uint32_t stride_B0;
p3 _p12;
uint32_t stride_B1;
p3 _p13;
uint32_t Mdim;
p3 _p14;
uint32_t Ndim;
p3 _p15;
uint32_t Kdim;
p3 _p16;
void* ptr_ScaleA;
p2 _p17;
void* ptr_ScaleB;
p2 _p18;
uint32_t stride_ScaleA0;
p3 _p19;
uint32_t stride_ScaleA1;
p3 _p20;
uint32_t stride_ScaleB0;
p3 _p21;
uint32_t stride_ScaleB1;
p3 _p22;
int32_t log2_k_split;
p3 _p23;
};
static_assert(sizeof(KernelArgs) == 384, "Unexpected FP4 kernarg size");
struct CachedKernel {
hipModule_t module = nullptr;
hipFunction_t func = nullptr;
};
std::unordered_map<std::string, CachedKernel>& kernel_cache() {
static std::unordered_map<std::string, CachedKernel> cache;
return cache;
}
std::mutex& kernel_cache_mutex() {
static std::mutex mu;
return mu;
}
void hip_check(hipError_t err, const char* call_name) {
if (err == hipSuccess) {
return;
}
throw std::runtime_error(std::string(call_name) + " failed: " + hipGetErrorString(err));
}
CachedKernel& get_kernel(const std::string& co_path, const std::string& kernel_name) {
const std::string key = co_path + "|" + kernel_name;
std::lock_guard<std::mutex> guard(kernel_cache_mutex());
auto& cache = kernel_cache();
auto it = cache.find(key);
if (it != cache.end()) {
return it->second;
}
CachedKernel entry;
hip_check(hipModuleLoad(&entry.module, co_path.c_str()), "hipModuleLoad");
hip_check(hipModuleGetFunction(&entry.func, entry.module, kernel_name.c_str()), "hipModuleGetFunction");
auto [new_it, _inserted] = cache.emplace(key, entry);
return new_it->second;
}
} // namespace
void launch_f4gemm(
torch::Tensor a_q,
torch::Tensor b_shuffle,
torch::Tensor a_scale_sh,
torch::Tensor b_scale_sh,
torch::Tensor out,
std::string co_path,
std::string kernel_name,
int tile_m,
int tile_n,
int log2_k_split) {
TORCH_CHECK(a_q.is_cuda(), "a_q must be a HIP tensor");
TORCH_CHECK(b_shuffle.is_cuda(), "b_shuffle must be a HIP tensor");
TORCH_CHECK(a_scale_sh.is_cuda(), "a_scale_sh must be a HIP tensor");
TORCH_CHECK(b_scale_sh.is_cuda(), "b_scale_sh must be a HIP tensor");
TORCH_CHECK(out.is_cuda(), "out must be a HIP tensor");
auto& kernel = get_kernel(co_path, kernel_name);
KernelArgs args{};
const int m = static_cast<int>(out.size(0));
const int n = static_cast<int>(out.size(1));
const int k = static_cast<int>(a_q.size(1) * 2);
// Match aiter's asm_gemm_a4w4 host-side argument population exactly.
args.ptr_D = out.data_ptr();
args.ptr_C = nullptr;
args.ptr_A = a_q.data_ptr();
args.ptr_B = b_shuffle.data_ptr();
args.alpha = 1.0f;
args.beta = 0.0f;
args.stride_C0 = static_cast<uint32_t>(out.stride(0));
args.stride_A0 = static_cast<uint32_t>(a_q.stride(0) * 2);
args.stride_B0 = static_cast<uint32_t>(b_shuffle.stride(0) * 2);
args.Mdim = static_cast<uint32_t>(m);
args.Ndim = static_cast<uint32_t>(n);
args.Kdim = static_cast<uint32_t>(k);
args.ptr_ScaleA = a_scale_sh.data_ptr();
args.ptr_ScaleB = b_scale_sh.data_ptr();
args.stride_ScaleA0 = static_cast<uint32_t>(a_scale_sh.stride(0));
args.stride_ScaleB0 = static_cast<uint32_t>(b_scale_sh.stride(0));
args.log2_k_split = 0;
int gdz = 1;
if (log2_k_split > 0) {
args.log2_k_split = log2_k_split;
const int split_k = 1 << args.log2_k_split;
TORCH_CHECK(k % split_k == 0, "K must be divisible by split-K factor");
if (split_k > 1) {
hip_check(
hipMemsetAsync(out.data_ptr(), 0, out.numel() * out.element_size(), nullptr),
"hipMemsetAsync");
}
const int k_per_tg = ((k / split_k + 255) / 256) * 256;
gdz = (k + k_per_tg - 1) / k_per_tg;
}
size_t arg_size = sizeof(args);
void* config[] = {
HIP_LAUNCH_PARAM_BUFFER_POINTER,
&args,
HIP_LAUNCH_PARAM_BUFFER_SIZE,
&arg_size,
HIP_LAUNCH_PARAM_END,
};
hip_check(
hipModuleLaunchKernel(
kernel.func,
static_cast<unsigned int>((n + tile_n - 1) / tile_n),
static_cast<unsigned int>((m + tile_m - 1) / tile_m),
static_cast<unsigned int>(gdz),
256,
1,
1,
0,
nullptr,
nullptr,
reinterpret_cast<void**>(config)),
"hipModuleLaunchKernel");
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("launch_f4gemm", &launch_f4gemm, "Launch gfx950 FP4 GEMM code object");
}
"""
@lru_cache(maxsize=1)
def _launcher():
return _load_inline_with_trace(
"launcher",
name="mxfp4_gfx950_launcher_ext",
cpp_sources=[_CPP_SRC],
functions=None,
extra_cflags=["-std=c++17", "-I/opt/rocm/include"],
extra_ldflags=["-lamdhip64", "-L/opt/rocm/lib"],
with_cuda=False,
verbose=False,
)
@lru_cache(maxsize=1)
def _fast_fused_kernel():
return _load_inline_with_trace(
"fast-fused",
name="mxfp4_gfx950_fast_fused_ext",
cpp_sources=[_FAST_FUSED_CPP_SRC],
cuda_sources=[_FAST_FUSED_CUDA_SRC],
functions=[
"launch_fast_quant_mxfp4",
"launch_fast_quant_and_f4gemm",
"launch_fast_fused_mxfp4_gemm",
],
extra_cflags=["-std=c++20", "-I/opt/rocm/include"],
extra_cuda_cflags=_fused_cuda_cflags(),
with_cuda=True,
no_implicit_headers=True,
verbose=False,
)
@lru_cache(maxsize=1)
def _aiter_hsa_root() -> Path:
env_root = os.environ.get("AITER_ASM_DIR")
if env_root:
root = Path(env_root)
if root.exists():
return root
import aiter # type: ignore
pkg_file = Path(aiter.__file__).resolve()
candidates = [
pkg_file.parents[1] / "hsa",
pkg_file.parents[0] / "hsa",
Path("/home/runner/aiter/hsa"),
]
for root in candidates:
if (root / _ARCH / _F4GEMM_SUBDIR).exists():
return root
raise RuntimeError("Unable to locate aiter HSA directory for gfx950 FP4 kernels")
def _fused_cuda_cflags() -> list[str]:
return [
"-O3",
"--offload-arch=gfx950",
"-std=c++20",
"-U__HIP_NO_HALF_OPERATORS__",
"-U__HIP_NO_HALF_CONVERSIONS__",
]
@lru_cache(maxsize=1)
def _load_f4gemm_cfgs() -> tuple[_KernelCfg, ...]:
csv_path = _aiter_hsa_root() / _ARCH / _F4GEMM_SUBDIR / "f4gemm_bf16_per1x32Fp4.csv"
cfgs: list[_KernelCfg] = []
with csv_path.open("r", encoding="utf-8") as f:
for row in csv.DictReader(f):
cfgs.append(
_KernelCfg(
tile_m=int(row["tile_M"]),
tile_n=int(row["tile_N"]),
splitk_enabled=int(row["splitK"]),
bpreshuffle=int(row["bpreshuffle"]),
kernel_name=row["knl_name"],
co_name=f"{_F4GEMM_SUBDIR}/{row['co_name']}",
)
)
return tuple(cfgs)
@lru_cache(maxsize=1)
def _aiter_config_root() -> Path:
import aiter # type: ignore
pkg_file = Path(aiter.__file__).resolve()
candidates = [
pkg_file.parents[0] / "configs",
pkg_file.parents[1] / "aiter" / "configs",
Path("/home/runner/aiter/aiter/configs"),
]
for root in candidates:
if (root / "a4w4_blockscale_tuned_gemm.csv").exists():
return root
raise RuntimeError("Unable to locate aiter tuned GEMM config directory")
@lru_cache(maxsize=1)
def _load_tuned_a4w4_cfgs() -> tuple[_TunedKernel, ...]:
csv_path = _aiter_config_root() / "a4w4_blockscale_tuned_gemm.csv"
cfgs: list[_TunedKernel] = []
with csv_path.open("r", encoding="utf-8") as f:
for row in csv.DictReader(f):
kernel_name = row["kernelName"]
if not kernel_name.startswith("_ZN"):
continue
cfgs.append(
_TunedKernel(
cu_num=int(row["cu_num"]),
m=int(row["M"]),
n=int(row["N"]),
k=int(row["K"]),
split_k=int(row["splitK"]),
kernel_name=kernel_name,
)
)
return tuple(cfgs)
def _quant_mxfp4(x: torch.Tensor, *, shuffle_scale: bool) -> tuple[torch.Tensor, torch.Tensor]:
from aiter import dtypes # type: ignore
from aiter.ops.triton.quant import dynamic_mxfp4_quant # type: ignore
from aiter.utility.fp4_utils import e8m0_shuffle # type: ignore
x_fp4, scale = dynamic_mxfp4_quant(x)
if shuffle_scale:
scale = e8m0_shuffle(scale)
return x_fp4.view(dtypes.fp4x2), scale.view(dtypes.fp8_e8m0)
def _as_u8_view(x: torch.Tensor) -> torch.Tensor:
return x.view(torch.uint8) if x.dtype != torch.uint8 else x
def _pad_rows_u8(x: torch.Tensor, rows: int, fill: int) -> torch.Tensor:
if int(x.shape[0]) >= rows:
return x.contiguous()
padded = torch.full((rows, *x.shape[1:]), fill, dtype=torch.uint8, device=x.device)
padded[: int(x.shape[0])].copy_(x.contiguous())
return padded
def _prepare_b_fast_inputs(
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
*,
n: int,
) -> tuple[torch.Tensor, torch.Tensor, int]:
cache_key = (
int(b_shuffle.data_ptr()),
int(b_scale_sh.data_ptr()),
n,
int(b_shuffle.shape[1]),
int(b_scale_sh.shape[1]),
)
cached = _B_SHUFFLE_CACHE.get(cache_key)
if cached is not None:
return cached
b_shuffle_u8_src = _as_u8_view(b_shuffle).contiguous()
b_scale_sh_u8_src = _as_u8_view(b_scale_sh).contiguous()
# b_shuffle is N-major and must be padded to the fast-path tile width.
# b_scale_sh is already in preshuffled microscale layout and may have a
# larger leading dimension than N because aiter pads it independently.
n_padded = max(
((n + _CUSTOM_N_TILE - 1) // _CUSTOM_N_TILE) * _CUSTOM_N_TILE,
int(b_shuffle_u8_src.shape[0]),
)
b_shuffle_u8 = _pad_rows_u8(b_shuffle_u8_src, n_padded, 0)
b_scale_sh_u8 = b_scale_sh_u8_src
result = (b_shuffle_u8, b_scale_sh_u8, n_padded)
_B_SHUFFLE_CACHE[cache_key] = result
return result
def _pad_a_quant_inputs(
a_q: torch.Tensor,
a_scale_sh: torch.Tensor,
*,
padded_m: int,
) -> tuple[torch.Tensor, torch.Tensor]:
if int(a_q.shape[0]) >= padded_m:
return a_q.contiguous(), a_scale_sh.contiguous()
a_q_u8 = _as_u8_view(a_q).contiguous()
a_q_pad = torch.zeros((padded_m, int(a_q_u8.shape[1])), dtype=torch.uint8, device=a_q.device)
a_q_pad[: int(a_q.shape[0])].copy_(a_q_u8)
if int(a_scale_sh.shape[0]) >= padded_m:
return a_q_pad.view(a_q.dtype), a_scale_sh.contiguous()
a_scale_u8 = _as_u8_view(a_scale_sh).contiguous()
a_scale_pad = torch.full(
(padded_m, int(a_scale_u8.shape[1])),
127,
dtype=torch.uint8,
device=a_scale_sh.device,
)
a_scale_pad[: int(a_scale_sh.shape[0])].copy_(a_scale_u8)
return a_q_pad.view(a_q.dtype), a_scale_pad.view(a_scale_sh.dtype)
@lru_cache(maxsize=1)
def _direct_shuffled_quant_kernel():
import triton
import triton.language as tl
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op # type: ignore
@triton.heuristics(
{
"EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0
and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0,
}
)
@triton.jit
def _kernel(
x_ptr,
x_fp4_ptr,
scale_sh_ptr,
stride_x_m_in,
stride_x_n_in,
stride_x_fp4_m_in,
stride_x_fp4_n_in,
M,
N,
NUM_ITER: tl.constexpr,
NUM_STAGES: tl.constexpr,
SCALE_N_VALID: tl.constexpr,
SCALE_N_PAD: tl.constexpr,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
EVEN_M_N: tl.constexpr,
SCALING_MODE: tl.constexpr,
):
pid_m = tl.program_id(0)
start_n = tl.program_id(1) * NUM_ITER
stride_x_m = tl.cast(stride_x_m_in, tl.int64)
stride_x_n = tl.cast(stride_x_n_in, tl.int64)
stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
num_quant_blocks: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
for pid_n in tl.range(start_n, start_n + NUM_ITER, num_stages=NUM_STAGES):
x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
if EVEN_M_N:
x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
else:
x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(tl.float32)
out_tensor, bs_e8m0 = _mxfp4_quant_op(
x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
)
out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
out_offs = (
out_offs_m[:, None] * stride_x_fp4_m
+ out_offs_n[None, :] * stride_x_fp4_n
)
if EVEN_M_N:
tl.store(x_fp4_ptr + out_offs, out_tensor)
else:
out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)
bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
bs_offs_n = pid_n * num_quant_blocks + tl.arange(0, num_quant_blocks)
bs_offs_0 = bs_offs_m[:, None] // 32
bs_offs_1 = (bs_offs_m[:, None] % 32) // 16
bs_offs_2 = bs_offs_m[:, None] % 16
bs_offs_3 = bs_offs_n[None, :] // 8
bs_offs_4 = (bs_offs_n[None, :] % 8) // 4
bs_offs_5 = bs_offs_n[None, :] % 4
bs_flat_offs = (
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)
)
bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < SCALE_N_VALID)[None, :]
tl.store(scale_sh_ptr + bs_flat_offs, bs_e8m0, mask=bs_mask)
return _kernel
def _quant_mxfp4_direct_shuffled(
x: torch.Tensor,
*,
x_fp4_out: torch.Tensor | None = None,
scale_sh_out: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
import triton
from aiter import dtypes # type: ignore
m, n = map(int, x.shape)
scale_n_valid = (n + 31) // 32
scale_n_pad = ((scale_n_valid + 7) // 8) * 8
scale_m_pad = ((m + 255) // 256) * 256
if x_fp4_out is None:
x_fp4 = torch.empty((m, n // 2), dtype=torch.uint8, device=x.device)
else:
x_fp4 = x_fp4_out
if int(x_fp4.shape[0]) != m:
x_fp4.zero_()
if scale_sh_out is None:
scale_sh = torch.full((scale_m_pad, scale_n_pad), 127, dtype=torch.uint8, device=x.device)
else:
scale_sh = scale_sh_out
scale_sh.fill_(127)
if m <= _CUSTOM_SMALLM_MAX_M:
num_iter = 1
block_size_m = triton.next_power_of_2(m)
block_size_n = 32
num_warps = 1
num_stages = 1
else:
num_iter = 4
block_size_m = 64
block_size_n = 64
num_warps = 4
num_stages = 2
if n <= 16384:
block_size_m = 32
block_size_n = 128
if n <= 1024:
num_iter = 1
num_stages = 1
num_warps = 4
block_size_n = min(256, triton.next_power_of_2(n))
block_size_n = max(32, block_size_n)
block_size_m = min(8, triton.next_power_of_2(m))
grid = (
triton.cdiv(m, block_size_m),
triton.cdiv(n, block_size_n * num_iter),
)
_direct_shuffled_quant_kernel()[grid](
x,
x_fp4,
scale_sh,
*x.stride(),
*x_fp4.stride(),
M=m,
N=n,
NUM_ITER=num_iter,
NUM_STAGES=num_stages,
SCALE_N_VALID=scale_n_valid,
SCALE_N_PAD=scale_n_pad,
BLOCK_SIZE_M=block_size_m,
BLOCK_SIZE_N=block_size_n,
MXFP4_QUANT_BLOCK_SIZE=32,
SCALING_MODE=0,
num_warps=num_warps,
num_stages=1,
waves_per_eu=0,
)
return x_fp4.view(dtypes.fp4x2), scale_sh.view(dtypes.fp8_e8m0)
def _quant_mxfp4_direct_shuffled_k512_smallm(
x: torch.Tensor,
*,
x_fp4_out: torch.Tensor | None = None,
scale_sh_out: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
import triton
from aiter import dtypes # type: ignore
m, n = map(int, x.shape)
if n != 512 or m > _CUSTOM_SMALLM_MAX_M:
raise RuntimeError(f"k=512 small-shape direct quant does not support shape {(m, n)}")
scale_n_pad = 16
scale_m_pad = ((m + 255) // 256) * 256
if x_fp4_out is None:
x_fp4 = torch.empty((m, n // 2), dtype=torch.uint8, device=x.device)
else:
x_fp4 = x_fp4_out
if int(x_fp4.shape[0]) != m:
x_fp4.zero_()
if scale_sh_out is None:
scale_sh = torch.full((scale_m_pad, scale_n_pad), 127, dtype=torch.uint8, device=x.device)
else:
scale_sh = scale_sh_out
scale_sh.fill_(127)
block_size_m = min(8, triton.next_power_of_2(m))
num_warps = 4
block_size_n = 256
num_iter = 2
grid = (
triton.cdiv(m, block_size_m),
triton.cdiv(n, block_size_n * num_iter),
)
_direct_shuffled_quant_kernel()[grid](
x,
x_fp4,
scale_sh,
*x.stride(),
*x_fp4.stride(),
M=m,
N=n,
NUM_ITER=num_iter,
NUM_STAGES=1,
SCALE_N_VALID=16,
SCALE_N_PAD=16,
BLOCK_SIZE_M=block_size_m,
BLOCK_SIZE_N=block_size_n,
MXFP4_QUANT_BLOCK_SIZE=32,
SCALING_MODE=0,
num_warps=num_warps,
num_stages=1,
waves_per_eu=0,
)
return x_fp4.view(dtypes.fp4x2), scale_sh.view(dtypes.fp8_e8m0)
def _quant_mxfp4_direct_shuffled_largek_smallm(
x: torch.Tensor,
*,
x_fp4_out: torch.Tensor | None = None,
scale_sh_out: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
from aiter import dtypes # type: ignore
m, n = map(int, x.shape)
if (m, n) != (16, 7168):
raise RuntimeError(f"large-k small-m direct quant does not support shape {(m, n)}")
scale_n_valid = (n + 31) // 32
scale_n_pad = ((scale_n_valid + 7) // 8) * 8
scale_m_pad = ((m + 255) // 256) * 256
if x_fp4_out is None:
x_fp4 = torch.empty((m, n // 2), dtype=torch.uint8, device=x.device)
else:
x_fp4 = x_fp4_out
if int(x_fp4.shape[0]) != m:
x_fp4.zero_()
if scale_sh_out is None:
scale_sh = torch.full((scale_m_pad, scale_n_pad), 127, dtype=torch.uint8, device=x.device)
else:
scale_sh = scale_sh_out
scale_sh.fill_(127)
grid = (
1,
(n + 511) // 512,
)
_direct_shuffled_quant_kernel()[grid](
x,
x_fp4,
scale_sh,
*x.stride(),
*x_fp4.stride(),
M=m,
N=n,
NUM_ITER=2,
NUM_STAGES=2,
SCALE_N_VALID=scale_n_valid,
SCALE_N_PAD=scale_n_pad,
BLOCK_SIZE_M=16,
BLOCK_SIZE_N=256,
MXFP4_QUANT_BLOCK_SIZE=32,
SCALING_MODE=0,
num_warps=4,
num_stages=1,
waves_per_eu=0,
)
return x_fp4.view(dtypes.fp4x2), scale_sh.view(dtypes.fp8_e8m0)
def _quant_mxfp4_hip_direct_shuffled(
x: torch.Tensor,
*,
x_fp4_out: torch.Tensor | None = None,
scale_sh_out: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
from aiter import dtypes # type: ignore
m, n = map(int, x.shape)
scale_n_valid = (n + 31) // 32
scale_n_pad = ((scale_n_valid + 7) // 8) * 8
scale_m_pad = ((m + 255) // 256) * 256
if x_fp4_out is None:
x_fp4 = torch.empty((m, n // 2), dtype=torch.uint8, device=x.device)
else:
x_fp4 = x_fp4_out
if int(x_fp4.shape[0]) != m:
x_fp4.zero_()
if scale_sh_out is None:
scale_sh = torch.full((scale_m_pad, scale_n_pad), 127, dtype=torch.uint8, device=x.device)
else:
scale_sh = scale_sh_out
scale_sh.fill_(127)
_fast_fused_kernel().launch_fast_quant_mxfp4(x, x_fp4, scale_sh, int(m), int(n))
return x_fp4.view(dtypes.fp4x2), scale_sh.view(dtypes.fp8_e8m0)
def _quant_mxfp4_aiter_hip(
x: torch.Tensor,
*,
x_fp4_out: torch.Tensor | None = None,
scale_sh_out: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
from aiter import dtypes # type: ignore
from aiter.ops.quant import dynamic_per_group_scaled_quant_fp4 # type: ignore
m, n = map(int, x.shape)
scale_n_valid = (n + 31) // 32
scale_n_pad = ((scale_n_valid + 7) // 8) * 8
scale_m_pad = ((m + 255) // 256) * 256
if x_fp4_out is None:
x_fp4 = torch.empty((m, n // 2), dtype=torch.uint8, device=x.device)
else:
x_fp4 = x_fp4_out
if int(x_fp4.shape[0]) != m:
x_fp4.zero_()
if scale_sh_out is None:
scale_sh = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=x.device)
else:
scale_sh = scale_sh_out
dynamic_per_group_scaled_quant_fp4(
x_fp4.view(dtypes.fp4x2),
x,
scale_sh.view(dtypes.fp8_e8m0),
32,
shuffle_scale=True,
)
return x_fp4.view(dtypes.fp4x2), scale_sh.view(dtypes.fp8_e8m0)
def _device_index(device: torch.device) -> int:
return -1 if device.index is None else int(device.index)
def _lookup_benchmark_quant_workspace(
*,
device: torch.device,
m: int,
n: int,
q_rows: int,
) -> tuple[torch.Tensor, torch.Tensor] | None:
device_type = device.type
device_index = _device_index(device)
for idx, entry in enumerate(_BENCHMARK_QUANT_WORKSPACES):
if (
entry.device_type != device_type
or entry.device_index != device_index
or entry.m != m
or entry.n != n
or entry.q_rows != q_rows
):
continue
if idx:
_BENCHMARK_QUANT_WORKSPACES.insert(0, _BENCHMARK_QUANT_WORKSPACES.pop(idx))
entry = _BENCHMARK_QUANT_WORKSPACES[0]
return entry.q, entry.scale_sh
return None
def _get_benchmark_quant_workspace(
*,
device: torch.device,
m: int,
n: int,
q_rows: int | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
q_rows = m if q_rows is None else q_rows
cached = _lookup_benchmark_quant_workspace(device=device, m=m, n=n, q_rows=q_rows)
if cached is not None:
return cached
scale_n_valid = (n + 31) // 32
scale_n_pad = ((scale_n_valid + 7) // 8) * 8
scale_m_pad = ((max(m, q_rows) + 255) // 256) * 256
q = torch.empty((q_rows, n // 2), dtype=torch.uint8, device=device)
scale_sh = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)
_BENCHMARK_QUANT_WORKSPACES.insert(
0,
_QuantWorkspaceEntry(
device_type=device.type,
device_index=_device_index(device),
m=m,
n=n,
q_rows=q_rows,
q=q,
scale_sh=scale_sh,
),
)
del _BENCHMARK_QUANT_WORKSPACES[_BENCHMARK_QUANT_WORKSPACE_MAX:]
return q, scale_sh
def _compute_benchmark_quant(
x: torch.Tensor,
*,
shape: tuple[int, int, int],
q_out: torch.Tensor,
scale_sh_out: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
if shape in {(64, 7168, 2048), (256, 3072, 1536)}:
return _quant_mxfp4_hip_direct_shuffled(
x,
x_fp4_out=q_out,
scale_sh_out=scale_sh_out,
)
if shape == (16, 2112, 7168):
return _quant_mxfp4_direct_shuffled(
x,
x_fp4_out=q_out,
scale_sh_out=scale_sh_out,
)
if _is_smallk_benchmark_shape(shape):
return _quant_mxfp4_direct_shuffled_k512_smallm(
x,
x_fp4_out=q_out,
scale_sh_out=scale_sh_out,
)
return _quant_mxfp4_direct_shuffled(
x,
x_fp4_out=q_out,
scale_sh_out=scale_sh_out,
)
def _lookup_benchmark_out_cache(
*,
device: torch.device,
dtype: torch.dtype,
padded_m: int,
n: int,
) -> torch.Tensor | None:
device_type = device.type
device_index = _device_index(device)
for idx, entry in enumerate(_BENCHMARK_OUT_CACHE):
if (
entry.device_type != device_type
or entry.device_index != device_index
or entry.dtype != dtype
or entry.padded_m != padded_m
or entry.n != n
):
continue
if idx:
_BENCHMARK_OUT_CACHE.insert(0, _BENCHMARK_OUT_CACHE.pop(idx))
entry = _BENCHMARK_OUT_CACHE[0]
return entry.out
return None
def _get_cached_benchmark_out(
*,
device: torch.device,
dtype: torch.dtype,
padded_m: int,
n: int,
) -> torch.Tensor:
cached = _lookup_benchmark_out_cache(
device=device,
dtype=dtype,
padded_m=padded_m,
n=n,
)
if cached is not None:
return cached
out = torch.empty((padded_m, n), dtype=dtype, device=device)
_BENCHMARK_OUT_CACHE.insert(
0,
_OutCacheEntry(
device_type=device.type,
device_index=_device_index(device),
dtype=dtype,
padded_m=padded_m,
n=n,
out=out,
),
)
del _BENCHMARK_OUT_CACHE[_BENCHMARK_OUT_CACHE_MAX:]
return out
def _maybe_contiguous(x: torch.Tensor) -> torch.Tensor:
return x if x.is_contiguous() else x.contiguous()
def _use_fast_fused_path(m: int, n: int, k: int) -> bool:
if (m, n, k) == (16, 2112, 7168):
return False
return _FAST_FUSED_ENABLE and m <= 32 and k >= 1536 and k % 256 == 0 and n > 0
@lru_cache(maxsize=1)
def _aiter_cu_num() -> int:
from aiter.jit.utils.chip_info import get_cu_num # type: ignore
return int(get_cu_num())
def _find_kernel_cfg(kernel_name: str) -> _KernelCfg:
for cfg in _load_f4gemm_cfgs():
if cfg.kernel_name == kernel_name:
return cfg
raise RuntimeError(f"Kernel {kernel_name} not found in FP4 config table")
def _lookup_tuned_kernel(m: int, n: int, k: int, padded_m: int, cu_num: int) -> tuple[_KernelCfg, int] | None:
for candidate_m in (m, padded_m):
for tuned in _load_tuned_a4w4_cfgs():
if (tuned.cu_num, tuned.m, tuned.n, tuned.k) != (cu_num, candidate_m, n, k):
continue
return _find_kernel_cfg(tuned.kernel_name), tuned.split_k
return None
def _pick_kernel(m: int, n: int, k: int, num_cu: int) -> tuple[_KernelCfg, int]:
override = _KERNEL_OVERRIDE_BY_SHAPE.get((m, n, k))
if override is not None:
kernel_name, co_name, log2_k_split = override
for cfg in _load_f4gemm_cfgs():
if cfg.kernel_name == kernel_name and cfg.co_name == co_name:
return cfg, log2_k_split
raise RuntimeError(f"Override kernel not found in FP4 config table: {override}")
padded_m = ((m + 31) // 32) * 32
tuned = _lookup_tuned_kernel(m, n, k, padded_m, num_cu)
if tuned is not None:
return tuned
empty_cu = num_cu
best_round = 1 << 30
best_eff = 1.0
best_cfg: _KernelCfg | None = None
for cfg in _load_f4gemm_cfgs():
if cfg.bpreshuffle != 1:
continue
if cfg.tile_m == 128 and cfg.tile_n == 512 and n % cfg.tile_n != 0:
continue
tg_num_m = (padded_m + cfg.tile_m - 1) // cfg.tile_m
tg_num_n = (n + cfg.tile_n - 1) // cfg.tile_n
tg_num = tg_num_m * tg_num_n
local_round = (tg_num + num_cu - 1) // num_cu
local_eff = (cfg.tile_m * cfg.tile_n) / (cfg.tile_m + cfg.tile_n)
is_earlier_round = local_round < best_round
is_same_round = local_round == best_round
has_sufficient_empty_cu = empty_cu > (local_round * num_cu - tg_num)
has_better_efficiency = local_eff > best_eff
if is_earlier_round or (is_same_round and (has_sufficient_empty_cu or has_better_efficiency)):
best_round = local_round
empty_cu = local_round * num_cu - tg_num
best_eff = local_eff
best_cfg = cfg
if best_cfg is None:
raise RuntimeError(f"No preshuffled FP4 kernel found for shape {(m, n, k)}")
return best_cfg, 0
def _launch_f4gemm(
*,
a_q: torch.Tensor,
b_shuffle: torch.Tensor,
a_scale_sh: torch.Tensor,
b_scale_sh: torch.Tensor,
out: torch.Tensor,
kernel_cfg: _KernelCfg,
log2_k_split: int,
) -> None:
co_path = _aiter_hsa_root() / _ARCH / kernel_cfg.co_name
if not co_path.exists():
raise RuntimeError(f"Missing code object: {co_path}")
_launcher().launch_f4gemm(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
out,
str(co_path),
kernel_cfg.kernel_name,
int(kernel_cfg.tile_m),
int(kernel_cfg.tile_n),
int(log2_k_split),
)
def _launch_f4gemm_official(
*,
a_q: torch.Tensor,
b_shuffle: torch.Tensor,
a_scale_sh: torch.Tensor,
b_scale_sh: torch.Tensor,
out: torch.Tensor,
kernel_cfg: _KernelCfg,
log2_k_split: int,
) -> None:
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm # type: ignore
gemm_a4w4_asm(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
out,
kernel_cfg.kernel_name,
None,
1.0,
0.0,
True,
log2_k_split,
)
def _launch_fused_fast(
*,
a: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
out: torch.Tensor,
n: int,
) -> None:
a_contig = a if a.is_contiguous() else a.contiguous()
b_shuffle_u8, b_scale_sh_u8, n_padded = _prepare_b_fast_inputs(b_shuffle, b_scale_sh, n=n)
_fast_fused_kernel().launch_fast_fused_mxfp4_gemm(
a_contig,
b_shuffle_u8,
b_scale_sh_u8,
out,
int(n_padded),
int(n),
)
def _launch_fast_quant_and_f4gemm(
*,
a: torch.Tensor,
a_q: torch.Tensor,
a_scale_sh: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
out: torch.Tensor,
kernel_cfg: _KernelCfg,
log2_k_split: int,
) -> None:
co_path = _aiter_hsa_root() / _ARCH / kernel_cfg.co_name
if not co_path.exists():
raise RuntimeError(f"Missing code object: {co_path}")
_fast_fused_kernel().launch_fast_quant_and_f4gemm(
a if a.is_contiguous() else a.contiguous(),
_as_u8_view(a_q).contiguous(),
_as_u8_view(a_scale_sh).contiguous(),
b_shuffle,
b_scale_sh,
out,
str(co_path),
kernel_cfg.kernel_name,
int(kernel_cfg.tile_m),
int(kernel_cfg.tile_n),
int(log2_k_split),
int(a.shape[0]),
int(a.shape[1]),
)
def custom_kernel(data: input_t) -> output_t:
a, _b, _b_q, b_shuffle, b_scale_sh = data
a = _maybe_contiguous(a)
m, k = map(int, a.shape)
n = int(b_shuffle.shape[0])
shape = (m, n, k)
padded_m = ((m + 31) // 32) * 32
if _is_benchmark_fastpath_shape(shape):
q_rows = padded_m if shape in _COMBINED_HIP_QUANT_GEMM_SHAPES else m
a_q_ws, a_scale_sh_ws = _get_benchmark_quant_workspace(
device=a.device,
m=m,
n=k,
q_rows=q_rows,
)
if shape in _COMBINED_HIP_QUANT_GEMM_SHAPES:
a_q, a_scale_sh = a_q_ws, a_scale_sh_ws
else:
a_q, a_scale_sh = _compute_benchmark_quant(
a,
shape=shape,
q_out=a_q_ws,
scale_sh_out=a_scale_sh_ws,
)
out = _get_cached_benchmark_out(
device=a.device,
dtype=torch.bfloat16,
padded_m=((m + 31) // 32) * 32,
n=n,
)
else:
a_q, a_scale_sh = _quant_mxfp4(a, shuffle_scale=True)
out = torch.empty((padded_m, n), dtype=torch.bfloat16, device=a.device)
kernel_cfg, log2_k_split = _pick_kernel(m, n, k, _aiter_cu_num())
if shape == (16, 2112, 7168) and int(a_q.shape[0]) < padded_m:
a_q, a_scale_sh = _pad_a_quant_inputs(a_q, a_scale_sh, padded_m=padded_m)
if shape in _COMBINED_HIP_QUANT_GEMM_SHAPES:
_launch_fast_quant_and_f4gemm(
a=a,
a_q=a_q,
a_scale_sh=a_scale_sh,
b_shuffle=b_shuffle,
b_scale_sh=b_scale_sh,
out=out,
kernel_cfg=kernel_cfg,
log2_k_split=log2_k_split,
)
return out[:m]
use_custom_launcher = (
_is_benchmark_fastpath_shape(shape)
and int(a_q.shape[0]) == int(out.shape[0])
and shape != (16, 2112, 7168)
)
if use_custom_launcher:
_launch_f4gemm(
a_q=a_q,
b_shuffle=b_shuffle,
a_scale_sh=a_scale_sh,
b_scale_sh=b_scale_sh,
out=out,
kernel_cfg=kernel_cfg,
log2_k_split=log2_k_split,
)
else:
_launch_f4gemm_official(
a_q=a_q,
b_shuffle=b_shuffle,
a_scale_sh=a_scale_sh,
b_scale_sh=b_scale_sh,
out=out,
kernel_cfg=kernel_cfg,
log2_k_split=log2_k_split,
)
return out[:m]
scrolls · 2104 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