submission 713996
hq_struggling · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1966 lines, June 9 Researcher Reciprocity License v1.0.
hq_submission_load_inline.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-713996?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:5b04d2aec5785110db2200e01e644d49d8c7cb4e46df4fb4290881723a377b2f
license declaredunknown
license concludedunknown
authorshq_struggling
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
const int k_blocks = K >> 6; // K / 64,也就是每个 row tile 里的 64-fp4 block 数shared-memory
__shared__ uint4 s_aq[2][TM][SPLITS];split-k
__global__ __launch_bounds__(512) void mxfp4_mfma_16x128x128_splitk_kernel(vector-width = uint4
static __device__ __forceinline__ uint4 quantize_32_bf16_to_fp4(warp-specialization
const bool a_producer = tid < TM * SPLITS;Kernel source
hq_submission_load_inline.py1966 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
from task import input_t, output_t
import os
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
os.environ.setdefault("CXX", "clang++")
os.environ.setdefault("MAX_JOBS", "8")
import torch
from torch.utils.cpp_extension import load_inline
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import (
gemm_afp4wfp4_preshuffled_weight_scales,
)
from aiter.utility.fp4_utils import e8m0_shuffle
_THIS_DIR = os.path.dirname(os.path.abspath(__file__))
AITER_INCLUDE_DIR = None
USE_NATIVE_MFMA = os.environ.get("MXFP4_USE_NATIVE_MFMA") == "1"
USE_TRITON_PRESHUFFLE = os.environ.get("MXFP4_USE_TRITON_PRESHUFFLE", "1") == "1"
_module = None
def _resolve_aiter_include_dirs() -> list[str]:
global AITER_INCLUDE_DIR
if AITER_INCLUDE_DIR is None:
for _root in (
os.path.abspath(os.path.join(_THIS_DIR, "..")),
os.path.abspath(os.path.join(_THIS_DIR, "..", "..")),
"/home/runner",
"/home/ubuntu/data/amd_202602",
):
_include_dir = os.path.join(_root, "aiter", "csrc", "include")
if os.path.isdir(_include_dir):
AITER_INCLUDE_DIR = _include_dir
break
if AITER_INCLUDE_DIR is None:
raise FileNotFoundError("Unable to locate aiter/csrc/include for load_inline")
return [AITER_INCLUDE_DIR, os.path.join(AITER_INCLUDE_DIR, "opus")]
# C++ 侧只暴露一个最小入口,真正的实现放在下面的 HIP 源码里。
CPP_SRC = r"""
extern "C" torch::Tensor mxfp4_gemm_kernel(
torch::Tensor A,
torch::Tensor B_shuffle,
torch::Tensor B_scale_sh
);
extern "C" torch::Tensor mxfp4_gemm_prequant(
torch::Tensor A_q,
torch::Tensor A_scale,
torch::Tensor B_shuffle,
torch::Tensor B_scale
);
extern "C" torch::Tensor mxfp4_pack_a_debug(
torch::Tensor A
);
"""
# HIP 源码使用 gfx950 的 scaled MFMA:
# - B 数据继续按 16B 粒度搬运,避免逐 nibble 标量解包
# - A 在 kernel prologue 中按 1x32 动态量化成 MXFP4
# - B scale 直接保持 preshuffled E8M0 布局,在 device 侧按索引读取
# - quant_kernels.cu 里的 scale shuffle 只在“落内存”时需要;这里 A scale
# 直接留在寄存器喂给 MFMA,所以保留同构的索引定义,但不物化成 shuffled tensor
HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include "opus.hpp"
#include <cstdint>
#include <limits>
using namespace opus;
// 基础参数检查,避免把错误类型或非连续张量送进 kernel。
#define CHECK_GPU(x) TORCH_CHECK(x.is_cuda(), #x " must be a GPU tensor")
#define CHECK_CONTIGUOUS(x) TORCH_CHECK(x.is_contiguous(), #x " must be contiguous")
#define CHECK_UINT8(x) TORCH_CHECK(x.scalar_type() == torch::kUInt8, #x " must be uint8")
// B_scale_sh 的物理布局。
static __device__ __forceinline__ int fp4_scale_shuffle_id(int scaleN_pad, int x, int y) {
return (x / 32 * scaleN_pad) * 32 + (y / 8) * 256 + (y % 4) * 64 + (x % 16) * 4 +
(y % 8) / 4 * 2 + (x % 32) / 16;
}
static __device__ __forceinline__ unsigned char load_preshuffled_scale(
const uint8_t* base,
int scaleNPad,
int row,
int group)
{
return base[fp4_scale_shuffle_id(scaleNPad, row, group)];
}
static __device__ __forceinline__ unsigned char mxfp4_scale_byte(float amax) {
// Match torch_dynamic_mxfp4_quant()/fp4_utils.dynamic_mxfp4_quant():
// round amax via (bits + 0x200000) & 0xFF800000, then encode amax * 0.25
// as E8M0.
uint32_t bits = __builtin_bit_cast(uint32_t, amax);
uint32_t rounded = (bits + 0x200000u) & 0xFF800000u;
uint32_t exponent = (rounded >> 23) & 0xFFu;
if (exponent == 0xFFu) {
return static_cast<unsigned char>(0xFFu);
}
return static_cast<unsigned char>(exponent > 2u ? exponent - 2u : 0u);
}
static __device__ __forceinline__ float mxfp4_scale_f32(unsigned char scale_byte) {
if (scale_byte == 0) {
return __builtin_bit_cast(float, 0x00400000u);
}
if (scale_byte == 0xFFu) {
return __builtin_bit_cast(float, 0x7F800001u);
}
return __builtin_bit_cast(float, static_cast<uint32_t>(scale_byte) << 23);
}
static __device__ __forceinline__ unsigned char float_to_fp4_nibble(float x_scaled) {
constexpr uint32_t SIGN_MASK = 0x80000000u;
constexpr float FP4_MAX_NORMAL = 6.0f;
constexpr float FP4_MIN_NORMAL = 1.0f;
constexpr int32_t EXP_BIAS_FP32 = 127;
constexpr int32_t EXP_BIAS_FP4 = 1;
constexpr int32_t MBITS_F32 = 23;
constexpr int32_t MBITS_FP4 = 1;
constexpr uint32_t FP4_SIGN_BIT = 0x8u;
constexpr uint32_t DENORM_MASK_INT =
((EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1) << MBITS_F32;
constexpr float DENORM_MASK_FLOAT = __builtin_bit_cast(float, DENORM_MASK_INT);
constexpr int32_t NORMAL_BIAS =
((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1;
uint32_t bits = __builtin_bit_cast(uint32_t, x_scaled);
uint32_t sign = bits & SIGN_MASK;
bits ^= sign;
float x_abs = __builtin_bit_cast(float, bits);
uint8_t fp4_value = 0;
if (x_abs >= FP4_MAX_NORMAL) {
fp4_value = 0x7u;
} else if (x_abs < FP4_MIN_NORMAL) {
float denorm = x_abs + DENORM_MASK_FLOAT;
uint32_t denorm_bits = __builtin_bit_cast(uint32_t, denorm) - DENORM_MASK_INT;
fp4_value = static_cast<uint8_t>(denorm_bits);
} else {
uint32_t normal_bits = bits;
uint32_t mant_odd = (normal_bits >> (MBITS_F32 - MBITS_FP4)) & 1u;
normal_bits += static_cast<uint32_t>(NORMAL_BIAS);
normal_bits += mant_odd;
normal_bits >>= (MBITS_F32 - MBITS_FP4);
fp4_value = static_cast<uint8_t>(normal_bits);
}
uint8_t sign_lp = static_cast<uint8_t>((sign >> 28) & FP4_SIGN_BIT);
return static_cast<unsigned char>(fp4_value | sign_lp);
}
static __device__ __forceinline__ unsigned char bf16_to_fp4_packed_byte(
const bf16x2_t& src,
float scale_f32)
{
const float quant_scale = 1.0f / scale_f32;
const unsigned char lo = float_to_fp4_nibble(static_cast<float>(src[0]) * quant_scale);
const unsigned char hi = float_to_fp4_nibble(static_cast<float>(src[1]) * quant_scale);
return static_cast<unsigned char>((lo & 0xFu) | ((hi & 0xFu) << 4));
}
static __device__ __forceinline__ int preshuffled_mxfp4_block_offset(
int n,
int k0,
int K,
int split)
{
// B_shuffle 的物理布局
// -------------------------
// 逻辑张量:
// B_shuffle[n, k_byte],形状为 [N, K/2]
//
// 打包事实:
// - 1 个字节存 2 个 fp4 值
// - 16 个字节 = 32 个 fp4 值
// - 64 个 fp4 值 = 32 个字节
//
// tile 顺序:
// [n_block = n / 16][k_block = k / 64][half = 0/1][row_in_16][byte_in_16]
//
// 一个物理 subtile 是 16 行 x 16 字节 = 256 字节:
//
// 第 0 行 -> 16 个 packed 字节
// 第 1 行 -> 16 个 packed 字节
// ...
// 第15 行 -> 16 个 packed 字节
//
// `split` 用来选择当前 K step 里读取哪一个 32-fp4 chunk。
// 下面的 `k0` 以 fp4 元素为单位,所以这里要换算成字节偏移。
const int n_block = n >> 4; // 每个 tile 覆盖 16 行
const int n_in = n & 15;
const int k_blocks = K >> 6; // K / 64,也就是每个 row tile 里的 64-fp4 block 数
const int kb = (k0 >> 6) + (split >> 1); // 当前 K block,按 split 分组修正
const int c = split & 1; // 选择 64-fp4 block 里的左/右半个 tile
const int block = (n_block * k_blocks + kb) * 2 + c;
return (block * 16 + n_in) * 16; // 每个 half tile 是 16 行 x 16 字节
}
static __device__ __forceinline__ const uint8_t* preshuffled_mxfp4_block_ptr(
const uint8_t* base,
int n,
int k0,
int K,
int split)
{
return base + preshuffled_mxfp4_block_offset(n, k0, K, split);
}
static __device__ __forceinline__ const uint8_t* row_major_mxfp4_block_ptr(
const uint8_t* base,
int row,
int k0,
int K,
int split)
{
return base + row * (K >> 1) + (k0 >> 1) + split * 16;
}
static __device__ __forceinline__ void store_bf16x4_scatter(
uint16_t* __restrict__ dst,
int base_row,
int col,
int stride_n,
int row_limit,
const fp32x4_t& acc)
{
#pragma unroll
for (int i = 0; i < 4; ++i) {
const int row = base_row + i;
if (row < row_limit) {
dst[row * stride_n + col] = fp32_to_bf16_rtn_raw(acc[i]);
}
}
}
static __device__ __forceinline__ void store_fp32x4_scatter(
float* __restrict__ dst,
int base_row,
int col,
int stride_n,
int row_limit,
const fp32x4_t& acc)
{
#pragma unroll
for (int i = 0; i < 4; ++i) {
const int row = base_row + i;
if (row < row_limit) {
dst[row * stride_n + col] = acc[i];
}
}
}
static __device__ __forceinline__ fp32x4_t mxfp4_mma_16x16x128(
const i32x8_t& a,
const i32x8_t& b,
const fp32x4_t& c,
int block_sel,
int scale_a,
int scale_b)
{
switch (block_sel) {
case 0:
return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a, b, c, 4, 4, 0, scale_a, 0, scale_b);
case 1:
return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a, b, c, 4, 4, 1, scale_a, 1, scale_b);
case 2:
return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a, b, c, 4, 4, 2, scale_a, 2, scale_b);
default:
return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a, b, c, 4, 4, 3, scale_a, 3, scale_b);
}
}
static __device__ __forceinline__ fp32x16_t mxfp4_mma_32x32x64(
const i32x8_t& a,
const i32x8_t& b,
const fp32x16_t& c,
int block_sel,
int scale_a,
int scale_b)
{
if (block_sel == 0) {
return __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
a, b, c, 4, 4, 0, scale_a, 0, scale_b);
}
return __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
a, b, c, 4, 4, 1, scale_a, 1, scale_b);
}
static __device__ __forceinline__ uint4 quantize_32_bf16_to_fp4(
const bf16_t* src,
unsigned char* scale_out)
{
bf16x8_t v0 = *reinterpret_cast<const bf16x8_t*>(src + 0);
bf16x8_t v1 = *reinterpret_cast<const bf16x8_t*>(src + 8);
bf16x8_t v2 = *reinterpret_cast<const bf16x8_t*>(src + 16);
bf16x8_t v3 = *reinterpret_cast<const bf16x8_t*>(src + 24);
float amax = 1.0e-10f;
#pragma unroll
for (int i = 0; i < 8; ++i) {
float x = static_cast<float>(v0[i]);
x = x < 0.0f ? -x : x;
amax = amax > x ? amax : x;
}
#pragma unroll
for (int i = 0; i < 8; ++i) {
float x = static_cast<float>(v1[i]);
x = x < 0.0f ? -x : x;
amax = amax > x ? amax : x;
}
#pragma unroll
for (int i = 0; i < 8; ++i) {
float x = static_cast<float>(v2[i]);
x = x < 0.0f ? -x : x;
amax = amax > x ? amax : x;
}
#pragma unroll
for (int i = 0; i < 8; ++i) {
float x = static_cast<float>(v3[i]);
x = x < 0.0f ? -x : x;
amax = amax > x ? amax : x;
}
const unsigned char scale_byte = mxfp4_scale_byte(amax);
const float scale_f32 = mxfp4_scale_f32(scale_byte);
union {
uint4 v;
unsigned char b[16];
} out{};
out.b[0] = bf16_to_fp4_packed_byte(bf16x2_t{v0[0], v0[1]}, scale_f32);
out.b[1] = bf16_to_fp4_packed_byte(bf16x2_t{v0[2], v0[3]}, scale_f32);
out.b[2] = bf16_to_fp4_packed_byte(bf16x2_t{v0[4], v0[5]}, scale_f32);
out.b[3] = bf16_to_fp4_packed_byte(bf16x2_t{v0[6], v0[7]}, scale_f32);
out.b[4] = bf16_to_fp4_packed_byte(bf16x2_t{v1[0], v1[1]}, scale_f32);
out.b[5] = bf16_to_fp4_packed_byte(bf16x2_t{v1[2], v1[3]}, scale_f32);
out.b[6] = bf16_to_fp4_packed_byte(bf16x2_t{v1[4], v1[5]}, scale_f32);
out.b[7] = bf16_to_fp4_packed_byte(bf16x2_t{v1[6], v1[7]}, scale_f32);
out.b[8] = bf16_to_fp4_packed_byte(bf16x2_t{v2[0], v2[1]}, scale_f32);
out.b[9] = bf16_to_fp4_packed_byte(bf16x2_t{v2[2], v2[3]}, scale_f32);
out.b[10] = bf16_to_fp4_packed_byte(bf16x2_t{v2[4], v2[5]}, scale_f32);
out.b[11] = bf16_to_fp4_packed_byte(bf16x2_t{v2[6], v2[7]}, scale_f32);
out.b[12] = bf16_to_fp4_packed_byte(bf16x2_t{v3[0], v3[1]}, scale_f32);
out.b[13] = bf16_to_fp4_packed_byte(bf16x2_t{v3[2], v3[3]}, scale_f32);
out.b[14] = bf16_to_fp4_packed_byte(bf16x2_t{v3[4], v3[5]}, scale_f32);
out.b[15] = bf16_to_fp4_packed_byte(bf16x2_t{v3[6], v3[7]}, scale_f32);
*scale_out = scale_byte;
return out.v;
}
__global__ void mxfp4_pack_a_debug_kernel(
const bf16_t* __restrict__ A,
uint8_t* __restrict__ A_q,
int M,
int K)
{
const int group_idx = static_cast<int>(blockIdx.x * blockDim.x + threadIdx.x);
const int groups_per_row = K / 32;
const int total_groups = M * groups_per_row;
if (group_idx >= total_groups) {
return;
}
const int row = group_idx / groups_per_row;
const int group = group_idx % groups_per_row;
unsigned char scale_byte = 127;
union {
uint4 v;
unsigned char b[16];
} packed{};
packed.v = quantize_32_bf16_to_fp4(A + row * K + group * 32, &scale_byte);
const int out_offset = row * (K / 2) + group * 16;
#pragma unroll
for (int i = 0; i < 16; ++i) {
A_q[out_offset + i] = packed.b[i];
}
}
template<int TM, int TN>
__global__ __launch_bounds__(64) void mxfp4_mfma_16x16x128_kernel(
const bf16_t* __restrict__ A,
const uint8_t* __restrict__ B_q,
const uint8_t* __restrict__ B_scale_sh,
uint16_t* __restrict__ C,
int M,
int N,
int K,
int scaleNPad)
{
// 16x16x128 原生路径。
//
// wave -> tile 映射
// 每个 wave 64 个线程
// lane = 0..63
// lane16 = lane % 16
// group4 = lane / 16
//
// group4 选择一个 128-fp4 K step 里的 4 个 32-fp4 切片:
// group4 0 -> k0 + 0..31
// group4 1 -> k0 + 32..63
// group4 2 -> k0 + 64..95
// group4 3 -> k0 + 96..127
//
// 数据流
// global B_shuffle -> 16B 向量加载 -> b_frag(VGPR)-> MFMA
// global A -> BF16->MXFP4 prologue -> a_frag(VGPR)-> MFMA
// acc 一直保留在 FP32 寄存器里,直到最后写回
//
// 这个 kernel 里没有显式的 shared memory staging buffer。
// 这里的 “tile” 只是逻辑访问模式,不是 LDS 缓冲区。
//
// 逻辑张量:
// - A: [M, K] bf16,row-major,在 prologue 里量化成当前 32-fp4 chunk
// - B_q: [N, K/2] packed fp4 字节,已经做过 bpreshuffle,适合 16x16 tile 读取
// - B_scale: [N, K/32] uint8 E8M0 block scale
// - C: [M, N] bf16 输出
constexpr int KSTEP = 128;
const int lane = static_cast<int>(__builtin_amdgcn_workitem_id_x());
const int lane16 = lane & 15;
const int group4 = lane >> 4;
const int m0 = blockIdx.y * TM;
const int n0 = blockIdx.x * TN;
const int row = m0 + lane16;
const int col = n0 + lane16;
const bool a_active = row < M;
const bool b_active = col < N;
const int row_stride_a = K;
// acc 是单个 lane 的 4 个输出所对应的寄存器态 FP32 累加器。
fp32x4_t acc{0.0f, 0.0f, 0.0f, 0.0f};
// 16x16 tile 内部的布局:
// - lane16 选择 [0, 15] 范围内的列
// - group4 选择输出 tile 中的 4 行带
// (group4 = 0..3,每一带分别覆盖 [0..3]、[4..7]、[8..11]、[12..15])
for (int k0 = 0; k0 < K; k0 += KSTEP) {
union {
i32x8_t v;
uint4 lo;
unsigned char bytes[32];
} a_frag{};
union {
i32x8_t v;
uint4 lo;
unsigned char bytes[32];
} b_frag{};
int scale_a = 127;
int scale_b = 127;
if (a_active) {
const bf16_t* a_src = A + row * row_stride_a + k0 + group4 * 32;
unsigned char scale_byte = 127;
a_frag.lo = quantize_32_bf16_to_fp4(a_src, &scale_byte);
scale_a = static_cast<int>(scale_byte);
}
if (b_active) {
// B_q 在 bpreshuffle 之后仍然是逻辑上的 [N, K/2]。
// helper 会把 (column, k0, group4) 映射到当前的 16 行 x 16 字节 subtile。
const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, col, k0, K, group4);
b_frag.lo = *reinterpret_cast<const uint4*>(b_src);
scale_b = static_cast<int>(load_preshuffled_scale(
B_scale_sh, scaleNPad, col, (k0 >> 5) + group4));
}
acc = mxfp4_mma_16x16x128(a_frag.v, b_frag.v, acc, group4, scale_a, scale_b);
}
const int out_row_base = m0 + group4 * 4;
if (b_active) {
// 每个 lane 针对一个列写回 4 个 bf16 输出:
// 行是 out_row_base + [0..3],列是 n0 + lane16。
store_bf16x4_scatter(C, out_row_base, col, N, M, acc);
}
}
template<int TM, int TN>
__global__ __launch_bounds__(64) void mxfp4_mfma_32x32x64_kernel(
const bf16_t* __restrict__ A,
const uint8_t* __restrict__ B_q,
const uint8_t* __restrict__ B_scale_sh,
uint16_t* __restrict__ C,
int M,
int N,
int K,
int scaleNPad)
{
// 32x32x64 原生路径。
//
// wave -> tile 映射
// lane32 = lane % 32
// group = lane / 32 // 0 或 1
//
// group 0 -> 这个 K step 的前 32-fp4 切片
// group 1 -> 这个 K step 的后 32-fp4 切片
//
// 和 16x16 kernel 一样,这里是直接 global load -> VGPR fragment -> MFMA。
// 不会显式构造 shared memory tile。
//
// 逻辑张量:
// - A: [M, K] bf16,row-major,在 prologue 里量化成当前 32-fp4 chunk
// - B_q: [N, K/2] packed fp4 字节,已经 bpreshuffle
// - B_scale: [N, K/32] uint8 E8M0 block scale
// - C: [M, N] bf16 输出
constexpr int KSTEP = 64;
const int lane = static_cast<int>(__builtin_amdgcn_workitem_id_x());
const int lane32 = lane & 31;
const int group = lane >> 5;
const int m0 = blockIdx.y * TM;
const int n0 = blockIdx.x * TN;
const int row = m0 + lane32;
const int col = n0 + lane32;
const bool a_active = row < M;
const bool b_active = col < N;
const int row_stride_a = K;
// acc 打包了最后要散写回去的 4x4 FP32 输出分块。
fp32x16_t acc{0.0f};
for (int k0 = 0; k0 < K; k0 += KSTEP) {
union {
i32x8_t v;
uint4 lo;
unsigned char bytes[32];
} a_frag{};
union {
i32x8_t v;
uint4 lo;
unsigned char bytes[32];
} b_frag{};
int scale_a = 127;
int scale_b = 127;
if (a_active) {
const bf16_t* a_src = A + row * row_stride_a + k0 + group * 32;
unsigned char scale_byte = 127;
a_frag.lo = quantize_32_bf16_to_fp4(a_src, &scale_byte);
scale_a = static_cast<int>(scale_byte);
}
if (b_active) {
// B_q 已经做过 bpreshuffle,所以 helper 会落到当前的 16x16 字节 subtile。
const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, col, k0, K, group);
b_frag.lo = *reinterpret_cast<const uint4*>(b_src);
scale_b = static_cast<int>(load_preshuffled_scale(
B_scale_sh, scaleNPad, col, (k0 >> 5) + group));
}
acc = mxfp4_mma_32x32x64(a_frag.v, b_frag.v, acc, group, scale_a, scale_b);
}
if (b_active) {
// 行布局和 32x32 MFMA kernel 保持一致。
store_bf16x4_scatter(
C, m0 + group * 4 + 0, col, N, M, fp32x4_t{acc[0], acc[1], acc[2], acc[3]});
store_bf16x4_scatter(
C, m0 + group * 4 + 8, col, N, M, fp32x4_t{acc[4], acc[5], acc[6], acc[7]});
store_bf16x4_scatter(
C, m0 + group * 4 + 16, col, N, M, fp32x4_t{acc[8], acc[9], acc[10], acc[11]});
store_bf16x4_scatter(
C, m0 + group * 4 + 24, col, N, M, fp32x4_t{acc[12], acc[13], acc[14], acc[15]});
}
}
__global__ __launch_bounds__(256) void mxfp4_mfma_16x64x128_kernel(
const bf16_t* __restrict__ A,
const uint8_t* __restrict__ B_q,
const uint8_t* __restrict__ B_scale_sh,
uint16_t* __restrict__ C,
int M,
int N,
int K,
int scaleNPad)
{
constexpr int TM = 16;
constexpr int TN = 64;
constexpr int KSTEP = 128;
constexpr int SPLITS = KSTEP / 32;
__shared__ uint4 s_aq[2][TM][SPLITS];
__shared__ unsigned char s_ascale[2][TM][SPLITS];
__shared__ uint4 s_bq[2][TN][SPLITS];
__shared__ unsigned char s_bscale[2][TN][SPLITS];
const int tid = static_cast<int>(threadIdx.x);
const int wave = tid >> 6;
const int lane = tid & 63;
const int lane16 = lane & 15;
const int group4 = lane >> 4;
const int m0 = blockIdx.y * TM;
const int n0 = blockIdx.x * TN;
const int row = m0 + lane16;
const int col = n0 + wave * 16 + lane16;
const bool a_active = row < M;
const bool b_active = col < N;
const int row_stride_a = K;
const int a_q_group = tid / TM;
const int a_q_row = tid - a_q_group * TM;
const bool a_producer = tid < TM * SPLITS;
const int b_split = tid / TN;
const int b_col_local = tid - b_split * TN;
const int b_col = n0 + b_col_local;
fp32x4_t acc{0.0f, 0.0f, 0.0f, 0.0f};
int read_buf = 0;
if (a_producer) {
const int global_row = m0 + a_q_row;
unsigned char scale_byte = 127;
uint4 packed{};
if (global_row < M) {
packed = quantize_32_bf16_to_fp4(
A + global_row * row_stride_a + a_q_group * 32,
&scale_byte);
}
s_aq[read_buf][a_q_row][a_q_group] = packed;
s_ascale[read_buf][a_q_row][a_q_group] = scale_byte;
}
uint4 b_packed{};
unsigned char b_scale = 127;
if (b_col < N) {
const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, 0, K, b_split);
b_packed = *reinterpret_cast<const uint4*>(b_src);
b_scale = load_preshuffled_scale(B_scale_sh, scaleNPad, b_col, b_split);
}
s_bq[read_buf][b_col_local][b_split] = b_packed;
s_bscale[read_buf][b_col_local][b_split] = b_scale;
__syncthreads();
for (int k0 = 0; k0 < K; k0 += KSTEP) {
union {
i32x8_t v;
uint4 lo;
unsigned char bytes[32];
} a_frag{};
union {
i32x8_t v;
uint4 lo;
unsigned char bytes[32];
} b_frag{};
int scale_a = 127;
int scale_b = 127;
if (a_active) {
a_frag.lo = s_aq[read_buf][lane16][group4];
scale_a = static_cast<int>(s_ascale[read_buf][lane16][group4]);
}
if (b_active) {
b_frag.lo = s_bq[read_buf][col - n0][group4];
scale_b = static_cast<int>(s_bscale[read_buf][col - n0][group4]);
}
const int next_k0 = k0 + KSTEP;
uint4 a_next_packed{};
unsigned char a_next_scale = 127;
uint4 b_next_packed{};
unsigned char b_next_scale = 127;
if (next_k0 < K) {
if (a_producer) {
const int global_row = m0 + a_q_row;
if (global_row < M) {
a_next_packed = quantize_32_bf16_to_fp4(
A + global_row * row_stride_a + next_k0 + a_q_group * 32,
&a_next_scale);
}
}
if (b_col < N) {
const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, next_k0, K, b_split);
b_next_packed = *reinterpret_cast<const uint4*>(b_src);
b_next_scale = load_preshuffled_scale(
B_scale_sh, scaleNPad, b_col, (next_k0 >> 5) + b_split);
}
}
acc = mxfp4_mma_16x16x128(a_frag.v, b_frag.v, acc, group4, scale_a, scale_b);
if (next_k0 < K) {
const int write_buf = read_buf ^ 1;
if (a_producer) {
s_aq[write_buf][a_q_row][a_q_group] = a_next_packed;
s_ascale[write_buf][a_q_row][a_q_group] = a_next_scale;
}
s_bq[write_buf][b_col_local][b_split] = b_next_packed;
s_bscale[write_buf][b_col_local][b_split] = b_next_scale;
__syncthreads();
read_buf = write_buf;
}
}
const int out_row_base = m0 + group4 * 4;
if (b_active) {
store_bf16x4_scatter(C, out_row_base, col, N, M, acc);
}
}
__global__ __launch_bounds__(512) void mxfp4_mfma_16x128x128_kernel(
const bf16_t* __restrict__ A,
const uint8_t* __restrict__ B_q,
const uint8_t* __restrict__ B_scale_sh,
uint16_t* __restrict__ C,
int M,
int N,
int K,
int scaleNPad)
{
constexpr int TM = 16;
constexpr int TN = 128;
constexpr int KSTEP = 128;
constexpr int SPLITS = KSTEP / 32;
__shared__ uint4 s_aq[2][TM][SPLITS];
__shared__ unsigned char s_ascale[2][TM][SPLITS];
__shared__ uint4 s_bq[2][TN][SPLITS];
__shared__ unsigned char s_bscale[2][TN][SPLITS];
const int tid = static_cast<int>(threadIdx.x);
const int wave = tid >> 6;
const int lane = tid & 63;
const int lane16 = lane & 15;
const int group4 = lane >> 4;
const int m0 = blockIdx.y * TM;
const int n0 = blockIdx.x * TN;
const int row = m0 + lane16;
const int col = n0 + wave * 16 + lane16;
const bool a_active = row < M;
const bool b_active = col < N;
const int row_stride_a = K;
const int a_q_group = tid / TM;
const int a_q_row = tid - a_q_group * TM;
const bool a_producer = tid < TM * SPLITS;
const int b_split = tid / TN;
const int b_col_local = tid - b_split * TN;
const int b_col = n0 + b_col_local;
fp32x4_t acc{0.0f, 0.0f, 0.0f, 0.0f};
int read_buf = 0;
if (a_producer) {
const int global_row = m0 + a_q_row;
unsigned char scale_byte = 127;
uint4 packed{};
if (global_row < M) {
packed = quantize_32_bf16_to_fp4(
A + global_row * row_stride_a + a_q_group * 32,
&scale_byte);
}
s_aq[read_buf][a_q_row][a_q_group] = packed;
s_ascale[read_buf][a_q_row][a_q_group] = scale_byte;
}
uint4 b_packed{};
unsigned char b_scale = 127;
if (b_col < N) {
const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, 0, K, b_split);
b_packed = *reinterpret_cast<const uint4*>(b_src);
b_scale = load_preshuffled_scale(B_scale_sh, scaleNPad, b_col, b_split);
}
s_bq[read_buf][b_col_local][b_split] = b_packed;
s_bscale[read_buf][b_col_local][b_split] = b_scale;
__syncthreads();
for (int k0 = 0; k0 < K; k0 += KSTEP) {
union {
i32x8_t v;
uint4 lo;
unsigned char bytes[32];
} a_frag{};
union {
i32x8_t v;
uint4 lo;
unsigned char bytes[32];
} b_frag{};
int scale_a = 127;
int scale_b = 127;
if (a_active) {
a_frag.lo = s_aq[read_buf][lane16][group4];
scale_a = static_cast<int>(s_ascale[read_buf][lane16][group4]);
}
if (b_active) {
b_frag.lo = s_bq[read_buf][col - n0][group4];
scale_b = static_cast<int>(s_bscale[read_buf][col - n0][group4]);
}
const int next_k0 = k0 + KSTEP;
uint4 a_next_packed{};
unsigned char a_next_scale = 127;
uint4 b_next_packed{};
unsigned char b_next_scale = 127;
if (next_k0 < K) {
if (a_producer) {
const int global_row = m0 + a_q_row;
if (global_row < M) {
a_next_packed = quantize_32_bf16_to_fp4(
A + global_row * row_stride_a + next_k0 + a_q_group * 32,
&a_next_scale);
}
}
if (b_col < N) {
const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, next_k0, K, b_split);
b_next_packed = *reinterpret_cast<const uint4*>(b_src);
b_next_scale = load_preshuffled_scale(
B_scale_sh, scaleNPad, b_col, (next_k0 >> 5) + b_split);
}
}
acc = mxfp4_mma_16x16x128(a_frag.v, b_frag.v, acc, group4, scale_a, scale_b);
if (next_k0 < K) {
const int write_buf = read_buf ^ 1;
if (a_producer) {
s_aq[write_buf][a_q_row][a_q_group] = a_next_packed;
s_ascale[write_buf][a_q_row][a_q_group] = a_next_scale;
}
s_bq[write_buf][b_col_local][b_split] = b_next_packed;
s_bscale[write_buf][b_col_local][b_split] = b_next_scale;
__syncthreads();
read_buf = write_buf;
}
}
const int out_row_base = m0 + group4 * 4;
if (b_active) {
store_bf16x4_scatter(C, out_row_base, col, N, M, acc);
}
}
__global__ __launch_bounds__(512) void mxfp4_mfma_16x128x128_splitk_kernel(
const bf16_t* __restrict__ A,
const uint8_t* __restrict__ B_q,
const uint8_t* __restrict__ B_scale_sh,
float* __restrict__ partial,
int M,
int N,
int K,
int scaleNPad,
int split_k)
{
constexpr int TM = 16;
constexpr int TN = 128;
constexpr int KSTEP = 128;
constexpr int SPLITS = KSTEP / 32;
__shared__ uint4 s_aq[2][TM][SPLITS];
__shared__ unsigned char s_ascale[2][TM][SPLITS];
__shared__ uint4 s_bq[2][TN][SPLITS];
__shared__ unsigned char s_bscale[2][TN][SPLITS];
const int tid = static_cast<int>(threadIdx.x);
const int wave = tid >> 6;
const int lane = tid & 63;
const int lane16 = lane & 15;
const int group4 = lane >> 4;
const int split_idx = static_cast<int>(blockIdx.z);
const int m0 = blockIdx.y * TM;
const int n0 = blockIdx.x * TN;
const int row = m0 + lane16;
const int col = n0 + wave * 16 + lane16;
const bool a_active = row < M;
const bool b_active = col < N;
const int row_stride_a = K;
const int a_q_group = tid / TM;
const int a_q_row = tid - a_q_group * TM;
const bool a_producer = tid < TM * SPLITS;
const int b_split = tid / TN;
const int b_col_local = tid - b_split * TN;
const int b_col = n0 + b_col_local;
const int k_chunks = K / KSTEP;
const int chunk_begin = (k_chunks * split_idx) / split_k;
const int chunk_end = (k_chunks * (split_idx + 1)) / split_k;
const int k_begin = chunk_begin * KSTEP;
const int k_end = chunk_end * KSTEP;
fp32x4_t acc{0.0f, 0.0f, 0.0f, 0.0f};
if (k_begin >= k_end) {
return;
}
int read_buf = 0;
if (a_producer) {
const int global_row = m0 + a_q_row;
unsigned char scale_byte = 127;
uint4 packed{};
if (global_row < M) {
packed = quantize_32_bf16_to_fp4(
A + global_row * row_stride_a + k_begin + a_q_group * 32,
&scale_byte);
}
s_aq[read_buf][a_q_row][a_q_group] = packed;
s_ascale[read_buf][a_q_row][a_q_group] = scale_byte;
}
uint4 b_packed{};
unsigned char b_scale = 127;
if (b_col < N) {
const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, k_begin, K, b_split);
b_packed = *reinterpret_cast<const uint4*>(b_src);
b_scale = load_preshuffled_scale(B_scale_sh, scaleNPad, b_col, (k_begin >> 5) + b_split);
}
s_bq[read_buf][b_col_local][b_split] = b_packed;
s_bscale[read_buf][b_col_local][b_split] = b_scale;
__syncthreads();
for (int k0 = k_begin; k0 < k_end; k0 += KSTEP) {
union {
i32x8_t v;
uint4 lo;
unsigned char bytes[32];
} a_frag{};
union {
i32x8_t v;
uint4 lo;
unsigned char bytes[32];
} b_frag{};
int scale_a = 127;
int scale_b = 127;
if (a_active) {
a_frag.lo = s_aq[read_buf][lane16][group4];
scale_a = static_cast<int>(s_ascale[read_buf][lane16][group4]);
}
if (b_active) {
b_frag.lo = s_bq[read_buf][col - n0][group4];
scale_b = static_cast<int>(s_bscale[read_buf][col - n0][group4]);
}
const int next_k0 = k0 + KSTEP;
uint4 a_next_packed{};
unsigned char a_next_scale = 127;
uint4 b_next_packed{};
unsigned char b_next_scale = 127;
if (next_k0 < k_end) {
if (a_producer) {
const int global_row = m0 + a_q_row;
if (global_row < M) {
a_next_packed = quantize_32_bf16_to_fp4(
A + global_row * row_stride_a + next_k0 + a_q_group * 32,
&a_next_scale);
}
}
if (b_col < N) {
const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, next_k0, K, b_split);
b_next_packed = *reinterpret_cast<const uint4*>(b_src);
b_next_scale = load_preshuffled_scale(
B_scale_sh, scaleNPad, b_col, (next_k0 >> 5) + b_split);
}
}
acc = mxfp4_mma_16x16x128(a_frag.v, b_frag.v, acc, group4, scale_a, scale_b);
if (next_k0 < k_end) {
const int write_buf = read_buf ^ 1;
if (a_producer) {
s_aq[write_buf][a_q_row][a_q_group] = a_next_packed;
s_ascale[write_buf][a_q_row][a_q_group] = a_next_scale;
}
s_bq[write_buf][b_col_local][b_split] = b_next_packed;
s_bscale[write_buf][b_col_local][b_split] = b_next_scale;
__syncthreads();
read_buf = write_buf;
}
}
const int out_row_base = m0 + group4 * 4;
float* partial_base = partial + split_idx * M * N;
if (b_active) {
store_fp32x4_scatter(partial_base, out_row_base, col, N, M, acc);
}
}
__global__ __launch_bounds__(256) void mxfp4_mfma_32x128x64_kernel(
const bf16_t* __restrict__ A,
const uint8_t* __restrict__ B_q,
const uint8_t* __restrict__ B_scale_sh,
uint16_t* __restrict__ C,
int M,
int N,
int K,
int scaleNPad)
{
constexpr int TM = 32;
constexpr int TN = 128;
constexpr int KSTEP = 64;
constexpr int SPLITS = KSTEP / 32;
__shared__ uint4 s_aq[2][TM][SPLITS];
__shared__ unsigned char s_ascale[2][TM][SPLITS];
__shared__ uint4 s_bq[2][TN][SPLITS];
__shared__ unsigned char s_bscale[2][TN][SPLITS];
const int tid = static_cast<int>(threadIdx.x);
const int wave = tid >> 6;
const int lane = tid & 63;
const int lane32 = lane & 31;
const int group = lane >> 5;
const int m0 = blockIdx.y * TM;
const int n0 = blockIdx.x * TN;
const int row = m0 + lane32;
const int col = n0 + wave * 32 + lane32;
const bool a_active = row < M;
const bool b_active = col < N;
const int row_stride_a = K;
const int a_q_group = tid / TM;
const int a_q_row = tid - a_q_group * TM;
const bool a_producer = tid < TM * SPLITS;
const int b_split = tid / TN;
const int b_col_local = tid - b_split * TN;
const int b_col = n0 + b_col_local;
fp32x16_t acc{0.0f};
int read_buf = 0;
if (a_producer) {
const int global_row = m0 + a_q_row;
unsigned char scale_byte = 127;
uint4 packed{};
if (global_row < M) {
packed = quantize_32_bf16_to_fp4(
A + global_row * row_stride_a + a_q_group * 32,
&scale_byte);
}
s_aq[read_buf][a_q_row][a_q_group] = packed;
s_ascale[read_buf][a_q_row][a_q_group] = scale_byte;
}
uint4 b_packed{};
unsigned char b_scale = 127;
if (b_col < N) {
const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, 0, K, b_split);
b_packed = *reinterpret_cast<const uint4*>(b_src);
b_scale = load_preshuffled_scale(B_scale_sh, scaleNPad, b_col, b_split);
}
s_bq[read_buf][b_col_local][b_split] = b_packed;
s_bscale[read_buf][b_col_local][b_split] = b_scale;
__syncthreads();
for (int k0 = 0; k0 < K; k0 += KSTEP) {
union {
i32x8_t v;
uint4 lo;
unsigned char bytes[32];
} a_frag{};
union {
i32x8_t v;
uint4 lo;
unsigned char bytes[32];
} b_frag{};
int scale_a = 127;
int scale_b = 127;
if (a_active) {
a_frag.lo = s_aq[read_buf][lane32][group];
scale_a = static_cast<int>(s_ascale[read_buf][lane32][group]);
}
if (b_active) {
b_frag.lo = s_bq[read_buf][col - n0][group];
scale_b = static_cast<int>(s_bscale[read_buf][col - n0][group]);
}
const int next_k0 = k0 + KSTEP;
uint4 a_next_packed{};
unsigned char a_next_scale = 127;
uint4 b_next_packed{};
unsigned char b_next_scale = 127;
if (next_k0 < K) {
if (a_producer) {
const int global_row = m0 + a_q_row;
if (global_row < M) {
a_next_packed = quantize_32_bf16_to_fp4(
A + global_row * row_stride_a + next_k0 + a_q_group * 32,
&a_next_scale);
}
}
if (b_col < N) {
const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, next_k0, K, b_split);
b_next_packed = *reinterpret_cast<const uint4*>(b_src);
b_next_scale = load_preshuffled_scale(
B_scale_sh, scaleNPad, b_col, (next_k0 >> 5) + b_split);
}
}
acc = mxfp4_mma_32x32x64(a_frag.v, b_frag.v, acc, group, scale_a, scale_b);
if (next_k0 < K) {
const int write_buf = read_buf ^ 1;
if (a_producer) {
s_aq[write_buf][a_q_row][a_q_group] = a_next_packed;
s_ascale[write_buf][a_q_row][a_q_group] = a_next_scale;
}
s_bq[write_buf][b_col_local][b_split] = b_next_packed;
s_bscale[write_buf][b_col_local][b_split] = b_next_scale;
__syncthreads();
read_buf = write_buf;
}
}
if (b_active) {
store_bf16x4_scatter(
C, m0 + group * 4 + 0, col, N, M, fp32x4_t{acc[0], acc[1], acc[2], acc[3]});
store_bf16x4_scatter(
C, m0 + group * 4 + 8, col, N, M, fp32x4_t{acc[4], acc[5], acc[6], acc[7]});
store_bf16x4_scatter(
C, m0 + group * 4 + 16, col, N, M, fp32x4_t{acc[8], acc[9], acc[10], acc[11]});
store_bf16x4_scatter(
C, m0 + group * 4 + 24, col, N, M, fp32x4_t{acc[12], acc[13], acc[14], acc[15]});
}
}
__global__ __launch_bounds__(384) void mxfp4_mfma_32x192x64_kernel(
const bf16_t* __restrict__ A,
const uint8_t* __restrict__ B_q,
const uint8_t* __restrict__ B_scale_sh,
uint16_t* __restrict__ C,
int M,
int N,
int K,
int scaleNPad)
{
constexpr int TM = 32;
constexpr int TN = 192;
constexpr int KSTEP = 64;
constexpr int SPLITS = KSTEP / 32;
__shared__ uint4 s_aq[2][TM][SPLITS];
__shared__ unsigned char s_ascale[2][TM][SPLITS];
__shared__ uint4 s_bq[2][TN][SPLITS];
__shared__ unsigned char s_bscale[2][TN][SPLITS];
const int tid = static_cast<int>(threadIdx.x);
const int wave = tid >> 6;
const int lane = tid & 63;
const int lane32 = lane & 31;
const int group = lane >> 5;
const int m0 = blockIdx.y * TM;
const int n0 = blockIdx.x * TN;
const int row = m0 + lane32;
const int col = n0 + wave * 32 + lane32;
const bool a_active = row < M;
const bool b_active = col < N;
const int row_stride_a = K;
const int a_q_group = tid / TM;
const int a_q_row = tid - a_q_group * TM;
const bool a_producer = tid < TM * SPLITS;
const int b_split = tid / TN;
const int b_col_local = tid - b_split * TN;
const int b_col = n0 + b_col_local;
fp32x16_t acc{0.0f};
int read_buf = 0;
if (a_producer) {
const int global_row = m0 + a_q_row;
unsigned char scale_byte = 127;
uint4 packed{};
if (global_row < M) {
packed = quantize_32_bf16_to_fp4(
A + global_row * row_stride_a + a_q_group * 32,
&scale_byte);
}
s_aq[read_buf][a_q_row][a_q_group] = packed;
s_ascale[read_buf][a_q_row][a_q_group] = scale_byte;
}
uint4 b_packed{};
unsigned char b_scale = 127;
if (b_col < N) {
const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, 0, K, b_split);
b_packed = *reinterpret_cast<const uint4*>(b_src);
b_scale = load_preshuffled_scale(B_scale_sh, scaleNPad, b_col, b_split);
}
s_bq[read_buf][b_col_local][b_split] = b_packed;
s_bscale[read_buf][b_col_local][b_split] = b_scale;
__syncthreads();
for (int k0 = 0; k0 < K; k0 += KSTEP) {
union {
i32x8_t v;
uint4 lo;
unsigned char bytes[32];
} a_frag{};
union {
i32x8_t v;
uint4 lo;
unsigned char bytes[32];
} b_frag{};
int scale_a = 127;
int scale_b = 127;
if (a_active) {
a_frag.lo = s_aq[read_buf][lane32][group];
scale_a = static_cast<int>(s_ascale[read_buf][lane32][group]);
}
if (b_active) {
b_frag.lo = s_bq[read_buf][col - n0][group];
scale_b = static_cast<int>(s_bscale[read_buf][col - n0][group]);
}
const int next_k0 = k0 + KSTEP;
uint4 a_next_packed{};
unsigned char a_next_scale = 127;
uint4 b_next_packed{};
unsigned char b_next_scale = 127;
if (next_k0 < K) {
if (a_producer) {
const int global_row = m0 + a_q_row;
if (global_row < M) {
a_next_packed = quantize_32_bf16_to_fp4(
A + global_row * row_stride_a + next_k0 + a_q_group * 32,
&a_next_scale);
}
}
if (b_col < N) {
const uint8_t* b_src = preshuffled_mxfp4_block_ptr(B_q, b_col, next_k0, K, b_split);
b_next_packed = *reinterpret_cast<const uint4*>(b_src);
b_next_scale = load_preshuffled_scale(
B_scale_sh, scaleNPad, b_col, (next_k0 >> 5) + b_split);
}
}
acc = mxfp4_mma_32x32x64(a_frag.v, b_frag.v, acc, group, scale_a, scale_b);
if (next_k0 < K) {
const int write_buf = read_buf ^ 1;
if (a_producer) {
s_aq[write_buf][a_q_row][a_q_group] = a_next_packed;
s_ascale[write_buf][a_q_row][a_q_group] = a_next_scale;
}
s_bq[write_buf][b_col_local][b_split] = b_next_packed;
s_bscale[write_buf][b_col_local][b_split] = b_next_scale;
__syncthreads();
read_buf = write_buf;
}
}
if (b_active) {
store_bf16x4_scatter(
C, m0 + group * 4 + 0, col, N, M, fp32x4_t{acc[0], acc[1], acc[2], acc[3]});
store_bf16x4_scatter(
C, m0 + group * 4 + 8, col, N, M, fp32x4_t{acc[4], acc[5], acc[6], acc[7]});
store_bf16x4_scatter(
C, m0 + group * 4 + 16, col, N, M, fp32x4_t{acc[8], acc[9], acc[10], acc[11]});
store_bf16x4_scatter(
C, m0 + group * 4 + 24, col, N, M, fp32x4_t{acc[12], acc[13], acc[14], acc[15]});
}
}
template<int TM, int TN>
__global__ __launch_bounds__(64) void mxfp4_mfma_16x16x128_prequant_kernel(
const uint8_t* __restrict__ A_q,
const uint8_t* __restrict__ A_scale,
const uint8_t* __restrict__ B_q,
const uint8_t* __restrict__ B_scale,
uint16_t* __restrict__ C,
int M,
int N,
int K)
{
constexpr int KSTEP = 128;
const int lane = static_cast<int>(__builtin_amdgcn_workitem_id_x());
const int lane16 = lane & 15;
const int group4 = lane >> 4;
const int m0 = blockIdx.y * TM;
const int n0 = blockIdx.x * TN;
const int row = m0 + lane16;
const int col = n0 + lane16;
const bool a_active = row < M;
const bool b_active = col < N;
const int scale_stride = K / 32;
const int scale_row_a = row * scale_stride;
const int scale_row_b = col * scale_stride;
fp32x4_t acc{0.0f, 0.0f, 0.0f, 0.0f};
for (int k0 = 0; k0 < K; k0 += KSTEP) {
union {
i32x8_t v;
uint4 lo;
unsigned char bytes[32];
} a_frag{};
union {
i32x8_t v;
uint4 lo;
unsigned char bytes[32];
} b_frag{};
int scale_a = 127;
int scale_b = 127;
if (a_active) {
const uint8_t* a_src = row_major_mxfp4_block_ptr(A_q, row, k0, K, group4);
a_frag.lo = *reinterpret_cast<const uint4*>(a_src);
scale_a = static_cast<int>(A_scale[scale_row_a + (k0 >> 5) + group4]);
}
if (b_active) {
const uint8_t* b_src = row_major_mxfp4_block_ptr(B_q, col, k0, K, group4);
b_frag.lo = *reinterpret_cast<const uint4*>(b_src);
scale_b = static_cast<int>(B_scale[scale_row_b + (k0 >> 5) + group4]);
}
acc = mxfp4_mma_16x16x128(a_frag.v, b_frag.v, acc, group4, scale_a, scale_b);
}
const int out_row_base = m0 + group4 * 4;
if (b_active) {
store_bf16x4_scatter(C, out_row_base, col, N, M, acc);
}
}
template<int TM, int TN>
__global__ __launch_bounds__(64) void mxfp4_mfma_32x32x64_prequant_kernel(
const uint8_t* __restrict__ A_q,
const uint8_t* __restrict__ A_scale,
const uint8_t* __restrict__ B_q,
const uint8_t* __restrict__ B_scale,
uint16_t* __restrict__ C,
int M,
int N,
int K)
{
constexpr int KSTEP = 64;
const int lane = static_cast<int>(__builtin_amdgcn_workitem_id_x());
const int lane32 = lane & 31;
const int group = lane >> 5;
const int m0 = blockIdx.y * TM;
const int n0 = blockIdx.x * TN;
const int row = m0 + lane32;
const int col = n0 + lane32;
const bool a_active = row < M;
const bool b_active = col < N;
const int scale_stride = K / 32;
const int scale_row_a = row * scale_stride;
const int scale_row_b = col * scale_stride;
fp32x16_t acc{0.0f};
for (int k0 = 0; k0 < K; k0 += KSTEP) {
union {
i32x8_t v;
uint4 lo;
unsigned char bytes[32];
} a_frag{};
union {
i32x8_t v;
uint4 lo;
unsigned char bytes[32];
} b_frag{};
int scale_a = 127;
int scale_b = 127;
if (a_active) {
const uint8_t* a_src = row_major_mxfp4_block_ptr(A_q, row, k0, K, group);
a_frag.lo = *reinterpret_cast<const uint4*>(a_src);
scale_a = static_cast<int>(A_scale[scale_row_a + (k0 >> 5) + group]);
}
if (b_active) {
const uint8_t* b_src = row_major_mxfp4_block_ptr(B_q, col, k0, K, group);
b_frag.lo = *reinterpret_cast<const uint4*>(b_src);
scale_b = static_cast<int>(B_scale[scale_row_b + (k0 >> 5) + group]);
}
acc = mxfp4_mma_32x32x64(a_frag.v, b_frag.v, acc, group, scale_a, scale_b);
}
if (b_active) {
store_bf16x4_scatter(
C, m0 + group * 4 + 0, col, N, M, fp32x4_t{acc[0], acc[1], acc[2], acc[3]});
store_bf16x4_scatter(
C, m0 + group * 4 + 8, col, N, M, fp32x4_t{acc[4], acc[5], acc[6], acc[7]});
store_bf16x4_scatter(
C, m0 + group * 4 + 16, col, N, M, fp32x4_t{acc[8], acc[9], acc[10], acc[11]});
store_bf16x4_scatter(
C, m0 + group * 4 + 24, col, N, M, fp32x4_t{acc[12], acc[13], acc[14], acc[15]});
}
}
__global__ void reduce_splitk_fp32_to_bf16_kernel(
const float* __restrict__ partial,
uint16_t* __restrict__ C,
int split_k,
int M,
int N)
{
const int idx = static_cast<int>(blockIdx.x * blockDim.x + threadIdx.x);
const int total = M * N;
if (idx >= total) {
return;
}
float sum = 0.0f;
#pragma unroll
for (int s = 0; s < 4; ++s) {
if (s < split_k) {
sum += partial[s * total + idx];
}
}
C[idx] = fp32_to_bf16_rtn_raw(sum);
}
enum class native_kernel_kind_t {
k16x16x128,
k16x64x128,
k16x128x128,
k16x128x128_splitk,
k32x32x64,
k32x128x64,
k32x192x64,
};
struct native_kernel_plan_t {
native_kernel_kind_t kind;
int split_k;
};
constexpr int MI355X_CU_COUNT = 256;
constexpr int MI355X_LATENCY_WAVE_TARGET = MI355X_CU_COUNT * 2;
static inline int ceil_div_int(int x, int y) {
return (x + y - 1) / y;
}
static inline int candidate_waves(
int M,
int N,
int TM,
int TN,
int block_threads,
int split_k = 1)
{
return ceil_div_int(M, TM) * ceil_div_int(N, TN) * (block_threads / 64) * split_k;
}
static inline int candidate_waste_cols(int N, int TN) {
return ceil_div_int(N, TN) * TN - N;
}
static inline int candidate_row_tiles(int M, int TM) {
return ceil_div_int(M, TM);
}
static inline int candidate_stages(int K, int kstep, int split_k = 1) {
const int chunks = K / kstep;
return ceil_div_int(chunks, split_k);
}
static inline int choose_small_m_split_k(int M, int N, int K) {
if (K < 4096 || N < 128) {
return 1;
}
const int base_waves = candidate_waves(M, N, 16, 128, 512);
if (base_waves >= MI355X_LATENCY_WAVE_TARGET) {
return 1;
}
int split_k = 1;
const int max_chunks = K / 128;
while (split_k < 4 && split_k < max_chunks &&
base_waves * split_k < MI355X_LATENCY_WAVE_TARGET) {
split_k <<= 1;
}
return split_k;
}
static inline int64_t score_candidate(
int M,
int N,
int K,
int TM,
int TN,
int block_threads,
int kstep,
int resource_penalty,
int split_k = 1)
{
const int waves = candidate_waves(M, N, TM, TN, block_threads, split_k);
const int stage_count = candidate_stages(K, kstep, split_k);
const int waste_cols = candidate_waste_cols(N, TN);
const int row_tiles = candidate_row_tiles(M, TM);
const int wave_shortfall = waves < MI355X_CU_COUNT ? MI355X_CU_COUNT - waves : 0;
int long_k_block_penalty = 0;
if (row_tiles <= 2 && stage_count >= 24 && block_threads > 64) {
long_k_block_penalty = ((block_threads / 64) - 1) * 256;
}
int64_t score = 0;
score += static_cast<int64_t>(TN) * 8;
score -= static_cast<int64_t>(stage_count) * 64;
score -= static_cast<int64_t>(resource_penalty);
score -= static_cast<int64_t>(waste_cols) * 2;
score -= static_cast<int64_t>(wave_shortfall) * 32;
score -= static_cast<int64_t>(long_k_block_penalty);
score -= static_cast<int64_t>(split_k - 1) * 128;
return score;
}
static inline native_kernel_plan_t choose_native_kernel_plan(int M, int N, int K) {
native_kernel_plan_t best{native_kernel_kind_t::k16x16x128, 1};
int64_t best_score = std::numeric_limits<int64_t>::min();
auto consider = [&](native_kernel_kind_t kind,
int TM,
int TN,
int block_threads,
int kstep,
int resource_penalty,
int split_k = 1) {
const int64_t score =
score_candidate(M, N, K, TM, TN, block_threads, kstep, resource_penalty, split_k);
if (score > best_score) {
best = native_kernel_plan_t{kind, split_k};
best_score = score;
}
};
if (M >= 32) {
consider(native_kernel_kind_t::k32x32x64, 32, 32, 64, 64, 96);
if (N >= 128) {
consider(native_kernel_kind_t::k32x128x64, 32, 128, 256, 64, 512);
}
if (N >= 192) {
consider(native_kernel_kind_t::k32x192x64, 32, 192, 384, 64, 1088);
}
} else {
consider(native_kernel_kind_t::k16x16x128, 16, 16, 64, 128, 96);
if (N >= 64) {
consider(native_kernel_kind_t::k16x64x128, 16, 64, 256, 128, 448);
}
if (N >= 128) {
consider(native_kernel_kind_t::k16x128x128, 16, 128, 512, 128, 1024);
const int split_k = choose_small_m_split_k(M, N, K);
if (split_k > 1) {
consider(native_kernel_kind_t::k16x128x128_splitk, 16, 128, 512, 128, 1024, split_k);
}
}
}
return best;
}
// 检查 HIP launch 是否成功。
static inline void check_hip() {
hipError_t err = hipGetLastError();
TORCH_CHECK(err == hipSuccess, "HIP kernel launch failed: ", hipGetErrorString(err));
}
extern "C" torch::Tensor mxfp4_gemm_kernel(
torch::Tensor A,
torch::Tensor B_shuffle,
torch::Tensor B_scale_sh
) {
CHECK_GPU(A);
CHECK_GPU(B_shuffle);
CHECK_GPU(B_scale_sh);
CHECK_CONTIGUOUS(A);
CHECK_CONTIGUOUS(B_shuffle);
CHECK_CONTIGUOUS(B_scale_sh);
TORCH_CHECK(A.scalar_type() == torch::kBFloat16, "A must be bf16");
CHECK_UINT8(B_shuffle);
CHECK_UINT8(B_scale_sh);
TORCH_CHECK(A.dim() == 2 && B_shuffle.dim() == 2 && B_scale_sh.dim() == 2,
"A, B_shuffle and B_scale_sh must be 2D");
TORCH_CHECK(A.size(1) == B_shuffle.size(1) * 2, "A and B_shuffle must have same K");
TORCH_CHECK(B_scale_sh.size(0) >= B_shuffle.size(0), "B_scale_sh rows must cover N");
TORCH_CHECK(B_scale_sh.size(1) >= A.size(1) / 32, "B_scale_sh cols must cover K/32 groups");
TORCH_CHECK(A.size(1) % 64 == 0, "K must be divisible by 64");
int64_t M = A.size(0);
int64_t N = B_shuffle.size(0);
int64_t K = A.size(1);
auto C = torch::empty({M, N}, A.options().dtype(torch::kBFloat16));
const auto* A_ptr = reinterpret_cast<const bf16_t*>(A.data_ptr());
const auto* B_shuffle_ptr = B_shuffle.data_ptr<uint8_t>();
const auto* B_scale_sh_ptr = B_scale_sh.data_ptr<uint8_t>();
auto* C_ptr = reinterpret_cast<uint16_t*>(C.data_ptr());
const int scaleNPad = static_cast<int>(B_scale_sh.size(1));
const int M_int = static_cast<int>(M);
const int N_int = static_cast<int>(N);
const int K_int = static_cast<int>(K);
const native_kernel_plan_t plan = choose_native_kernel_plan(M_int, N_int, K_int);
switch (plan.kind) {
case native_kernel_kind_t::k32x192x64: {
constexpr int TM = 32;
constexpr int TN = 192;
dim3 block(384, 1, 1);
dim3 grid((N_int + TN - 1) / TN, (M_int + TM - 1) / TM, 1);
hipLaunchKernelGGL(
HIP_KERNEL_NAME(mxfp4_mfma_32x192x64_kernel),
grid,
block,
0,
0,
A_ptr,
B_shuffle_ptr,
B_scale_sh_ptr,
C_ptr,
M_int,
N_int,
K_int,
scaleNPad
);
break;
}
case native_kernel_kind_t::k32x128x64: {
constexpr int TM = 32;
constexpr int TN = 128;
dim3 block(256, 1, 1);
dim3 grid((N_int + TN - 1) / TN, (M_int + TM - 1) / TM, 1);
hipLaunchKernelGGL(
HIP_KERNEL_NAME(mxfp4_mfma_32x128x64_kernel),
grid,
block,
0,
0,
A_ptr,
B_shuffle_ptr,
B_scale_sh_ptr,
C_ptr,
M_int,
N_int,
K_int,
scaleNPad
);
break;
}
case native_kernel_kind_t::k32x32x64: {
constexpr int TM = 32;
constexpr int TN = 32;
dim3 block(64, 1, 1);
dim3 grid((N_int + TN - 1) / TN, (M_int + TM - 1) / TM, 1);
hipLaunchKernelGGL(
HIP_KERNEL_NAME(mxfp4_mfma_32x32x64_kernel<TM, TN>),
grid,
block,
0,
0,
A_ptr,
B_shuffle_ptr,
B_scale_sh_ptr,
C_ptr,
M_int,
N_int,
K_int,
scaleNPad
);
break;
}
case native_kernel_kind_t::k16x128x128_splitk: {
constexpr int TM = 16;
constexpr int TN = 128;
auto partial = torch::empty({plan.split_k, M, N}, A.options().dtype(torch::kFloat32));
auto* partial_ptr = partial.data_ptr<float>();
dim3 block(512, 1, 1);
dim3 grid((N_int + TN - 1) / TN, (M_int + TM - 1) / TM, plan.split_k);
hipLaunchKernelGGL(
HIP_KERNEL_NAME(mxfp4_mfma_16x128x128_splitk_kernel),
grid,
block,
0,
0,
A_ptr,
B_shuffle_ptr,
B_scale_sh_ptr,
partial_ptr,
M_int,
N_int,
K_int,
scaleNPad,
plan.split_k
);
constexpr int REDUCE_THREADS = 256;
const int total = M_int * N_int;
dim3 reduce_block(REDUCE_THREADS, 1, 1);
dim3 reduce_grid((total + REDUCE_THREADS - 1) / REDUCE_THREADS, 1, 1);
hipLaunchKernelGGL(
reduce_splitk_fp32_to_bf16_kernel,
reduce_grid,
reduce_block,
0,
0,
partial_ptr,
C_ptr,
plan.split_k,
M_int,
N_int
);
break;
}
case native_kernel_kind_t::k16x128x128: {
constexpr int TM = 16;
constexpr int TN = 128;
dim3 block(512, 1, 1);
dim3 grid((N_int + TN - 1) / TN, (M_int + TM - 1) / TM, 1);
hipLaunchKernelGGL(
HIP_KERNEL_NAME(mxfp4_mfma_16x128x128_kernel),
grid,
block,
0,
0,
A_ptr,
B_shuffle_ptr,
B_scale_sh_ptr,
C_ptr,
M_int,
N_int,
K_int,
scaleNPad
);
break;
}
case native_kernel_kind_t::k16x64x128: {
constexpr int TM = 16;
constexpr int TN = 64;
dim3 block(256, 1, 1);
dim3 grid((N_int + TN - 1) / TN, (M_int + TM - 1) / TM, 1);
hipLaunchKernelGGL(
HIP_KERNEL_NAME(mxfp4_mfma_16x64x128_kernel),
grid,
block,
0,
0,
A_ptr,
B_shuffle_ptr,
B_scale_sh_ptr,
C_ptr,
M_int,
N_int,
K_int,
scaleNPad
);
break;
}
default: {
constexpr int TM = 16;
constexpr int TN = 16;
dim3 block(64, 1, 1);
dim3 grid((N_int + TN - 1) / TN, (M_int + TM - 1) / TM, 1);
hipLaunchKernelGGL(
HIP_KERNEL_NAME(mxfp4_mfma_16x16x128_kernel<TM, TN>),
grid,
block,
0,
0,
A_ptr,
B_shuffle_ptr,
B_scale_sh_ptr,
C_ptr,
M_int,
N_int,
K_int,
scaleNPad
);
break;
}
}
check_hip();
return C;
}
extern "C" torch::Tensor mxfp4_gemm_prequant(
torch::Tensor A_q,
torch::Tensor A_scale,
torch::Tensor B_shuffle,
torch::Tensor B_scale
) {
CHECK_GPU(A_q);
CHECK_GPU(A_scale);
CHECK_GPU(B_shuffle);
CHECK_GPU(B_scale);
CHECK_CONTIGUOUS(A_q);
CHECK_CONTIGUOUS(A_scale);
CHECK_CONTIGUOUS(B_shuffle);
CHECK_CONTIGUOUS(B_scale);
CHECK_UINT8(A_q);
CHECK_UINT8(A_scale);
CHECK_UINT8(B_shuffle);
CHECK_UINT8(B_scale);
TORCH_CHECK(
A_q.dim() == 2 && A_scale.dim() == 2 && B_shuffle.dim() == 2 && B_scale.dim() == 2,
"A_q, A_scale, B_shuffle and B_scale must be 2D");
TORCH_CHECK(A_q.size(0) == A_scale.size(0), "A_q and A_scale must have same M");
TORCH_CHECK(A_q.size(1) * 2 == B_shuffle.size(1) * 2, "A_q and B_shuffle must have same K");
TORCH_CHECK(A_scale.size(1) * 32 == A_q.size(1) * 2, "A_scale must have K/32 groups");
TORCH_CHECK(B_scale.size(0) == B_shuffle.size(0), "B_shuffle and B_scale must have same N");
TORCH_CHECK(B_scale.size(1) * 32 == A_q.size(1) * 2, "B_scale must have K/32 groups");
TORCH_CHECK((A_q.size(1) * 2) % 64 == 0, "K must be divisible by 64");
int64_t M = A_q.size(0);
int64_t N = B_shuffle.size(0);
int64_t K = A_q.size(1) * 2;
auto C = torch::empty({M, N}, A_q.options().dtype(torch::kBFloat16));
const auto* A_q_ptr = A_q.data_ptr<uint8_t>();
const auto* A_scale_ptr = A_scale.data_ptr<uint8_t>();
const auto* B_shuffle_ptr = B_shuffle.data_ptr<uint8_t>();
const auto* B_scale_ptr = B_scale.data_ptr<uint8_t>();
auto* C_ptr = reinterpret_cast<uint16_t*>(C.data_ptr());
if (M >= 32) {
constexpr int TM = 32;
constexpr int TN = 32;
dim3 block(64, 1, 1);
dim3 grid((N + TN - 1) / TN, (M + TM - 1) / TM, 1);
hipLaunchKernelGGL(
HIP_KERNEL_NAME(mxfp4_mfma_32x32x64_prequant_kernel<TM, TN>),
grid,
block,
0,
0,
A_q_ptr,
A_scale_ptr,
B_shuffle_ptr,
B_scale_ptr,
C_ptr,
static_cast<int>(M),
static_cast<int>(N),
static_cast<int>(K)
);
} else {
constexpr int TM = 16;
constexpr int TN = 16;
dim3 block(64, 1, 1);
dim3 grid((N + TN - 1) / TN, (M + TM - 1) / TM, 1);
hipLaunchKernelGGL(
HIP_KERNEL_NAME(mxfp4_mfma_16x16x128_prequant_kernel<TM, TN>),
grid,
block,
0,
0,
A_q_ptr,
A_scale_ptr,
B_shuffle_ptr,
B_scale_ptr,
C_ptr,
static_cast<int>(M),
static_cast<int>(N),
static_cast<int>(K)
);
}
check_hip();
return C;
}
extern "C" torch::Tensor mxfp4_pack_a_debug(
torch::Tensor A
) {
CHECK_GPU(A);
CHECK_CONTIGUOUS(A);
TORCH_CHECK(A.scalar_type() == torch::kBFloat16, "A must be bf16");
TORCH_CHECK(A.dim() == 2, "A must be 2D");
TORCH_CHECK(A.size(1) % 64 == 0, "K must be divisible by 64");
int64_t M = A.size(0);
int64_t K = A.size(1);
auto A_q = torch::empty({M, K / 2}, A.options().dtype(torch::kUInt8));
const auto* A_ptr = reinterpret_cast<const bf16_t*>(A.data_ptr());
auto* A_q_ptr = A_q.data_ptr<uint8_t>();
const int total_groups = static_cast<int>(M * (K / 32));
constexpr int THREADS = 256;
dim3 block(THREADS, 1, 1);
dim3 grid((total_groups + THREADS - 1) / THREADS, 1, 1);
hipLaunchKernelGGL(
mxfp4_pack_a_debug_kernel,
grid,
block,
0,
0,
A_ptr,
A_q_ptr,
static_cast<int>(M),
static_cast<int>(K)
);
check_hip();
return A_q;
}
"""
def _get_module():
global _module
if _module is None:
_module = load_inline(
name="mxfp4_gemm_release6",
cpp_sources=[CPP_SRC],
cuda_sources=[HIP_SRC],
functions=["mxfp4_gemm_kernel"],
verbose=False,
extra_cflags=["-std=c++20", "-O3"],
extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3"],
extra_include_paths=_resolve_aiter_include_dirs(),
)
return _module
def _view_preshuffled_weight_u8(weight_sh: torch.Tensor) -> torch.Tensor:
rows, cols = weight_sh.shape
assert rows % 16 == 0, "B_shuffle rows must be divisible by 16"
return weight_sh.view(rows // 16, cols * 16)
def _view_preshuffled_scales_u8(scale_sh: torch.Tensor) -> torch.Tensor:
scale_u8 = scale_sh.view(torch.uint8)
rows, cols = scale_u8.shape
assert rows % 32 == 0, "scale rows must be padded to a multiple of 32"
return scale_u8.view(rows // 32, cols * 32)
def _quantize_a_for_preshuffle_gemm(A: torch.Tensor):
A_q_u8, A_scale_u8 = dynamic_mxfp4_quant(A)
A_scale_u8 = A_scale_u8.view(torch.uint8).contiguous()
A_scale_triton = None
if A.shape[0] >= 32:
A_scale_triton = _view_preshuffled_scales_u8(e8m0_shuffle(A_scale_u8))
return A_q_u8, A_scale_u8, A_scale_triton
def custom_kernel(data: input_t) -> output_t:
"""
默认走 aiter 的量化 + ASM GEMM 热路径。
原生 inline MFMA 仅保留为实验分支。
布局说明:
- A 是稠密 bf16 [M, K]
- B_shuffle 是已经做过 bpreshuffle 的 packed fp4 [N, K/2]
- B_scale_sh 保持 shuffle 后的 E8M0 布局
- 默认路径使用参考的 Triton quant + e8m0 shuffle,再交给 aiter ASM GEMM
"""
A, _, _, B_shuffle, B_scale_sh = data
A = A.contiguous()
B_shuffle = B_shuffle.contiguous()
B_scale_sh = B_scale_sh.contiguous()
if not USE_NATIVE_MFMA and USE_TRITON_PRESHUFFLE:
M, K = A.shape
A_q_u8, A_scale_u8, A_scale_triton = _quantize_a_for_preshuffle_gemm(A)
B_triton = _view_preshuffled_weight_u8(B_shuffle.view(torch.uint8))
B_scale_triton = _view_preshuffled_scales_u8(B_scale_sh)
if M < 32:
A_scale_triton = A_scale_u8
return gemm_afp4wfp4_preshuffled_weight_scales(
A_q_u8,
B_triton,
A_scale_triton,
B_scale_triton,
torch.bfloat16,
)
if not USE_NATIVE_MFMA:
A_q_u8, A_scale_u8 = dynamic_mxfp4_quant(A)
A_q = A_q_u8.view(dtypes.fp4x2)
A_scale_sh = e8m0_shuffle(A_scale_u8).view(dtypes.fp8_e8m0)
return aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
return _get_module().mxfp4_gemm_kernel(
A,
B_shuffle.view(torch.uint8),
B_scale_sh.view(torch.uint8),
)
scrolls · 1966 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