submission 544703
div22 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 533 lines, June 9 Researcher Reciprocity License v1.0.
solution_new_21.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-544703?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:0ff26efd1a5091b68bc47ad72cd03fc179ffc4696349d7c4cdb67f7d92932f07
license declaredunknown
license concludedunknown
authorsdiv22
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 GEMM v21 — gfx950 (MI355X) optimized.split-k
3. Improved splitK: only split K=7168 with small M (from v20)tile-m = 0
const int full_m = (BM >= 16) ? M / BM : 0;Kernel source
solution_new_21.py533 lines
"""
MXFP4 GEMM v21 — gfx950 (MI355X) optimized.
Changes from v20:
1. Hardware FP4 conversion via v_cvt_scalef32_pk_fp4_bf16
- Replaces software f32_to_fp4_e2m1 with single instruction
- Fixed scale semantics: pass bit_cast<float>(inv_exp << 23)
where inv_exp = amax_exp - 2, so hardware applies 2^(129-amax_exp)
2. Template GEMM on CKT for constexpr loop unrolling (from v20)
3. Improved splitK: only split K=7168 with small M (from v20)
"""
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
from typing import Tuple
import torch
from torch.utils.cpp_extension import load_inline
import uuid
HIP_KERNEL = r"""
#include <hip/hip_runtime.h>
#include <stdint.h>
using int4_v = int __attribute__((ext_vector_type(4)));
using float4_v = float __attribute__((ext_vector_type(4)));
using bf16x2 = __bf16 __attribute__((ext_vector_type(2)));
static constexpr int FP4_E2M1 = 4;
__device__ float4_v __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
int4_v a, int4_v b, float4_v c,
int cbsz, int blgp, int op_sel_a, int scale_a, int op_sel_b, int scale_b
) __asm("llvm.amdgcn.mfma.scale.f32.16x16x128.f8f6f4.v4i32.v4i32");
__device__ __forceinline__ int4_v load16(const uint8_t* __restrict__ p) {
return *reinterpret_cast<const int4_v*>(p);
}
__device__ __forceinline__ uint16_t float_to_bf16(float f) {
bf16x2 v;
v[0] = static_cast<__bf16>(f);
uint16_t r;
__builtin_memcpy(&r, &v, sizeof(r));
return r;
}
// Hardware BF16x2 -> FP4x2 conversion using v_cvt_scalef32_pk_fp4_bf16
// The instruction extracts biased exponent E from scale float, applies 2^(127-E) compression.
// To get correct compression factor 2^(129 - amax_exp), pass scale with biased_exp = amax_exp - 2.
__device__ __forceinline__ uint8_t hw_bf16x2_to_fp4x2(uint32_t bf16_pair, float scale) {
uint32_t result;
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
: "=v"(result) : "v"(bf16_pair), "v"(scale));
return static_cast<uint8_t>(result & 0xFFu);
}
// Quant kernel: BF16 [M, K] -> FP4x2 [M, K//2] + E8M0 scale [M, K//32]
// Uses hardware v_cvt_scalef32_pk_fp4_bf16 for FP4 conversion
__global__ void __launch_bounds__(128, 4)
mxfp4_quant(
const __bf16* __restrict__ A_bf16,
uint8_t* __restrict__ A_fp4,
uint8_t* __restrict__ A_scale,
int M, int K)
{
const int KS = K / 32;
const int K2 = K / 2;
const int group = blockIdx.x * 128 + threadIdx.x;
const int row = group / KS;
const int kg = group % KS;
if (row >= M) return;
const auto* src = A_bf16 + (long)row * K + kg * 32;
float absMax = 1e-10f;
#pragma unroll
for (int i = 0; i < 32; ++i) {
float v = __builtin_elementwise_abs(static_cast<float>(src[i]));
absMax = (v > absMax) ? v : absMax;
}
uint32_t u32 = __builtin_bit_cast(uint32_t, absMax);
const uint32_t amax_exp = ((u32 + 0x200000u) >> 23) & 0xFFu;
const uint32_t inv_exp = (amax_exp >= 2u) ? (amax_exp - 2u) : 0u;
A_scale[(long)row * KS + kg] = static_cast<uint8_t>(inv_exp);
// Hardware scale: biased_exp = inv_exp = amax_exp - 2
// Instruction applies 2^(127 - inv_exp) = 2^(129 - amax_exp) for compression — correct!
const float hw_scale = __builtin_bit_cast(float, static_cast<uint32_t>(inv_exp) << 23);
const uint32_t* src_u32 = reinterpret_cast<const uint32_t*>(src);
auto* dst = reinterpret_cast<uint8_t*>(A_fp4 + (long)row * K2 + kg * 16);
#pragma unroll
for (int i = 0; i < 16; ++i) {
// Each uint32_t holds a pair of bf16 values (native layout)
dst[i] = hw_bf16x2_to_fp4x2(src_u32[i], hw_scale);
}
}
template<bool ALWAYS_VALID>
__device__ __forceinline__ int4_v load_or_zero(bool rt_valid, const uint8_t* p) {
if constexpr (ALWAYS_VALID) return load16(p);
else return rt_valid ? load16(p) : int4_v{0,0,0,0};
}
// Main GEMM kernel — templated on CKT (ktiles_per_split) for constexpr loop unrolling
// CKT=0 means runtime loop count
template<int BM, int BN, int NWARPS, bool SPLITK, bool A_VALID, bool B_VALID, int CKT = 0>
__global__ void __launch_bounds__(NWARPS * 64, 2)
mxfp4_gemm(
const uint8_t* __restrict__ A,
const uint8_t* __restrict__ As,
const uint8_t* __restrict__ Bsh,
const uint8_t* __restrict__ Bssh,
float* __restrict__ C_partial,
uint16_t* __restrict__ C_final,
int M, int N, int K, int scaleN,
int total_ktiles, int ktiles_per_split,
int tile_off_x, int tile_off_y)
{
static_assert(BN % 16 == 0);
constexpr auto WAVES_M = (BM + 15) / 16;
constexpr auto WAVES_N = BN / 16;
static_assert(WAVES_M * WAVES_N == NWARPS);
const auto ks_idx = blockIdx.z;
const auto lane = threadIdx.x % 64;
const auto wave = threadIdx.x / 64;
const auto wave_m = wave / WAVES_N;
const auto wave_n = wave % WAVES_N;
const auto tile_m = (blockIdx.y + tile_off_y) * BM + wave_m * 16;
const auto tile_n = (blockIdx.x + tile_off_x) * BN + wave_n * 16;
if (tile_m >= M || tile_n >= N) return;
const int ks_start = ks_idx * ktiles_per_split;
const int ks_end = min(ks_start + ktiles_per_split, total_ktiles);
const auto lrow = lane % 16;
const auto kgrp = lane / 16;
const auto gm = tile_m + lrow;
const auto gn = tile_n + lrow;
const auto K2 = K / 2;
const auto KS = K / 32;
const auto a_rt = A_VALID | (gm < M);
const auto b_rt = B_VALID | (gn < N);
const auto n_tile = tile_n / 16;
const auto bsh_lane_base = Bsh + (long)n_tile * (K / 64) * 512 + (long)lrow * 16;
const uint8_t* a_row = nullptr;
const uint8_t* as_row = nullptr;
if constexpr (A_VALID) {
a_row = A + (long)gm * K2;
as_row = As + (long)gm * KS;
} else {
if (a_rt) { a_row = A + (long)gm * K2;
as_row = As + (long)gm * KS; }
}
int bssh_base = 0;
if constexpr (B_VALID) {
bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);
} else {
if (b_rt) bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);
}
const auto k_half_off = (kgrp & 1) * 256;
const auto k_blk_base = kgrp >> 1;
static constexpr long a_kt_stride = 64L;
static constexpr long bsh_kt_stride = 1024L;
const uint8_t* a_ptr = nullptr;
const uint8_t* bsh_ptr = nullptr;
const uint8_t* bssh_ptr = nullptr;
if constexpr (A_VALID) {
a_ptr = a_row + (long)ks_start * a_kt_stride + kgrp * 16;
} else {
if (a_rt) a_ptr = a_row + (long)ks_start * a_kt_stride + kgrp * 16;
}
if constexpr (B_VALID) {
bsh_ptr = bsh_lane_base + (long)(ks_start * 2 + k_blk_base) * 512 + k_half_off;
bssh_ptr = Bssh + bssh_base + (ks_start & 1) * 2 + (ks_start >> 1) * 256;
} else {
if (b_rt) {
bsh_ptr = bsh_lane_base + (long)(ks_start * 2 + k_blk_base) * 512 + k_half_off;
bssh_ptr = Bssh + bssh_base + (ks_start & 1) * 2 + (ks_start >> 1) * 256;
}
}
float4_v acc{0.f, 0.f, 0.f, 0.f};
// bssh step pattern depends on ks_start parity; doesn't change across quads
const auto bssh_step0 = (ks_start & 1) ? 254 : 2;
const auto bssh_step1 = 256 - bssh_step0;
// Macro for one MFMA iteration
#define DO_MFMA(a_off, b_off, bssh_off, ks_val) \
{ \
const auto av = load_or_zero<A_VALID>(a_rt, a_ptr + (a_off) * a_kt_stride); \
const auto bv = load_or_zero<B_VALID>(b_rt, bsh_ptr + (b_off) * bsh_kt_stride); \
const auto ks = (ks_val) * 4 + kgrp; \
int sa, sb; \
if constexpr (A_VALID) sa = static_cast<int>(as_row[ks]); \
else sa = (a_rt & (ks < KS)) ? static_cast<int>(as_row[ks]) : 127; \
if constexpr (B_VALID) sb = static_cast<int>(*(bssh_ptr + (bssh_off))); \
else sb = (b_rt & (ks < KS)) ? static_cast<int>(*(bssh_ptr + (bssh_off))) : 127; \
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(av,bv,acc,FP4_E2M1,FP4_E2M1,0,sa,0,sb); \
}
// 4x unrolled main loop
if constexpr (CKT > 0) {
// Compile-time unrolled
#pragma unroll
for (int q = 0; q < (CKT / 4); ++q) {
DO_MFMA(0, 0, 0, ks_start + q*4)
DO_MFMA(1, 1, bssh_step0, ks_start + q*4 + 1)
DO_MFMA(2, 2, bssh_step0 + bssh_step1, ks_start + q*4 + 2)
DO_MFMA(3, 3, bssh_step0 + bssh_step1 + bssh_step0, ks_start + q*4 + 3)
a_ptr += 4 * a_kt_stride;
bsh_ptr += 4 * bsh_kt_stride;
bssh_ptr += 512;
}
// Compile-time pair (parity same as bssh_step0 since quads advance by 4)
if constexpr ((CKT % 4) >= 2) {
DO_MFMA(0, 0, 0, ks_start + (CKT/4)*4)
DO_MFMA(1, 1, bssh_step0, ks_start + (CKT/4)*4 + 1)
a_ptr += 2 * a_kt_stride;
bsh_ptr += 2 * bsh_kt_stride;
bssh_ptr += 256;
}
// Compile-time single
if constexpr ((CKT % 2) == 1) {
DO_MFMA(0, 0, 0, ks_start + CKT - 1)
}
} else {
// Runtime loop (fallback for unknown shapes)
auto kt = ks_start;
const int ks_count = ks_end - ks_start;
const auto ks_end_quad = ks_start + (ks_count - (ks_count & 3));
const auto ks_end_pair = ks_start + (ks_count - (ks_count & 1));
for (; kt < ks_end_quad; kt += 4) {
DO_MFMA(0, 0, 0, kt)
DO_MFMA(1, 1, bssh_step0, kt + 1)
DO_MFMA(2, 2, bssh_step0 + bssh_step1, kt + 2)
DO_MFMA(3, 3, bssh_step0 + bssh_step1 + bssh_step0, kt + 3)
a_ptr += 4 * a_kt_stride;
bsh_ptr += 4 * bsh_kt_stride;
bssh_ptr += 512;
}
const auto bssh_step0_trail = (kt & 1) ? 254 : 2;
for (; kt < ks_end_pair; kt += 2) {
DO_MFMA(0, 0, 0, kt)
DO_MFMA(1, 1, bssh_step0_trail, kt + 1)
a_ptr += 2 * a_kt_stride;
bsh_ptr += 2 * bsh_kt_stride;
bssh_ptr += 256;
}
if (kt < ks_end) {
DO_MFMA(0, 0, 0, kt)
}
}
#undef DO_MFMA
const auto out_col = tile_n + lrow;
const auto out_row_base = tile_m + kgrp * 4;
if constexpr (!B_VALID) { if (out_col >= N) return; }
constexpr bool out_rows_always_valid = A_VALID && (BM >= 16);
if constexpr (SPLITK) {
auto c_out = C_partial + (long)ks_idx * M * N + (long)out_row_base * N + out_col;
#pragma unroll
for (int i = 0; i < 4; ++i) {
if constexpr (out_rows_always_valid) c_out[i * N] = acc[i];
else if (out_row_base + i < M) c_out[i * N] = acc[i];
}
} else {
auto c_out = C_final + (long)out_row_base * N + out_col;
#pragma unroll
for (int i = 0; i < 4; ++i) {
if constexpr (out_rows_always_valid) c_out[i * N] = float_to_bf16(acc[i]);
else if (out_row_base + i < M) c_out[i * N] = float_to_bf16(acc[i]);
}
}
}
// Reduce kernel
__global__ void mxfp4_reduce(
const float* __restrict__ C_partial,
uint16_t* __restrict__ C_out,
int M, int N, int NUM_KSPLIT)
{
const auto col = blockIdx.x * 32 + threadIdx.x;
const auto row = blockIdx.y * 16 + threadIdx.y;
if (row >= M || col >= N) return;
float sum = 0.f;
const auto mn = (long)row * N + col;
const auto stride = (long)M * N;
for (auto k = 0; k < NUM_KSPLIT; ++k)
sum += C_partial[k * stride + mn];
bf16x2 v;
v[0] = static_cast<__bf16>(sum);
uint16_t r;
__builtin_memcpy(&r, &v, sizeof(r));
C_out[mn] = r;
}
// Launch quant
extern "C" void launch_quant(
const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale, int M, int K)
{
const int KS = K / 32;
const int n_groups = M * KS;
const dim3 block{128};
const dim3 grid{static_cast<uint32_t>((n_groups + 127) / 128)};
mxfp4_quant<<<grid, block>>>(A_bf16, A_fp4, A_scale, M, K);
}
// Templated GEMM launcher
template<int CKT>
void launch_gemm_ckt(
const uint8_t* A, const uint8_t* As,
const uint8_t* Bsh, const uint8_t* Bssh,
float* C_partial, uint16_t* C_final,
int M, int N, int K, int scaleN, int NUM_KSPLIT)
{
const auto total_ktiles = K / 128;
const auto ktiles_per_split = (total_ktiles + NUM_KSPLIT - 1) / NUM_KSPLIT;
const auto do_splitk = NUM_KSPLIT > 1;
auto launch = [&]<int BM, int BN, int NWARPS>() {
static_assert(((BM + 15) / 16) * (BN / 16) == NWARPS);
const int full_m = (BM >= 16) ? M / BM : 0;
const int full_n = N / BN;
const int total_m = (M + BM - 1) / BM;
const int total_n = (N + BN - 1) / BN;
const int edge_m = total_m - full_m;
const int edge_n = total_n - full_n;
const dim3 block{static_cast<uint32_t>(NWARPS * 64)};
auto sub = [&]<bool AV, bool BV>(int gx, int gy, int ox, int oy) {
if (gx <= 0 || gy <= 0) return;
const dim3 grid{
static_cast<uint32_t>(gx),
static_cast<uint32_t>(gy),
static_cast<uint32_t>(NUM_KSPLIT)
};
if (do_splitk)
mxfp4_gemm<BM,BN,NWARPS,true,AV,BV,CKT><<<grid,block>>>(
A,As,Bsh,Bssh,C_partial,nullptr,
M,N,K,scaleN,total_ktiles,ktiles_per_split,ox,oy);
else
mxfp4_gemm<BM,BN,NWARPS,false,AV,BV,CKT><<<grid,block>>>(
A,As,Bsh,Bssh,nullptr,C_final,
M,N,K,scaleN,total_ktiles,ktiles_per_split,ox,oy);
};
sub.template operator()<true, true >(full_n, full_m, 0, 0);
sub.template operator()<true, false>(edge_n, full_m, full_n, 0);
sub.template operator()<false, true >(full_n, edge_m, 0, full_m);
sub.template operator()<false, false>(edge_n, edge_m, full_n, full_m);
if (do_splitk) {
const dim3 rblock{32, 16};
const dim3 rgrid{
static_cast<uint32_t>((N + 31) / 32),
static_cast<uint32_t>((M + 15) / 16)
};
mxfp4_reduce<<<rgrid, rblock>>>(C_partial, C_final, M, N, NUM_KSPLIT);
}
};
if (M <= 8) launch.template operator()< 8, 32, 2>();
else if (M <= 16) launch.template operator()< 16, 32, 2>();
else if (M <= 32) launch.template operator()< 16, 32, 2>();
else if (M <= 64) launch.template operator()< 32, 32, 4>();
else if (M <=128) launch.template operator()< 32, 32, 4>();
else launch.template operator()< 64, 32, 8>();
}
// Main dispatch — routes to CKT-specialized launcher
extern "C" void launch_gemm(
const uint8_t* A, const uint8_t* As,
const uint8_t* Bsh, const uint8_t* Bssh,
float* C_partial, uint16_t* C_final,
int M, int N, int K, int scaleN, int NUM_KSPLIT)
{
const int total_ktiles = K / 128;
const int ktiles_per_split = (total_ktiles + NUM_KSPLIT - 1) / NUM_KSPLIT;
switch (ktiles_per_split) {
case 4: launch_gemm_ckt< 4>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;
case 8: launch_gemm_ckt< 8>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;
case 12: launch_gemm_ckt<12>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;
case 16: launch_gemm_ckt<16>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;
default: launch_gemm_ckt< 0>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;
}
}
"""
CPP = r"""
#include <torch/extension.h>
#include <c10/core/DeviceGuard.h>
extern "C" void launch_quant(const __bf16*, uint8_t*, uint8_t*, int, int);
extern "C" void launch_gemm(const uint8_t*, const uint8_t*, const uint8_t*, const uint8_t*,
float*, uint16_t*, int, int, int, int, int);
static int get_num_ksplit(int M, int K) {
// Only use splitK for K=7168+ (56+ ktiles) with small M
// Matches AITER's tuned configs
int total_ktiles = K / 128;
if (total_ktiles < 28) return 1; // K<=3456: never split
if (M <= 8) return 7; // kps=8 (K=7168)
if (M <= 16) return 14; // kps=4 (K=7168)
return 1;
}
struct Workspace {
at::Tensor A_fp4;
at::Tensor A_scale;
at::Tensor C_partial;
at::Tensor C;
int64_t last_M = -1, last_K = -1, last_N = -1, last_ksplit = -1;
void ensure(int M, int N, int K, int num_ksplit, const at::TensorOptions& opts) {
if (M == last_M && K == last_K && N == last_N && num_ksplit == last_ksplit) return;
int64_t KS = K / 32;
A_fp4 = at::empty({(int64_t)M, (int64_t)(K / 2)}, opts.dtype(at::kByte));
A_scale = at::empty({(int64_t)M, KS}, opts.dtype(at::kByte));
C = at::empty({(int64_t)M, (int64_t)N}, opts.dtype(at::kBFloat16));
if (num_ksplit > 1)
C_partial = at::empty({(int64_t)num_ksplit, (int64_t)M, (int64_t)N}, opts.dtype(at::kFloat));
else
C_partial = at::Tensor();
last_M = M; last_K = K; last_N = N; last_ksplit = num_ksplit;
}
};
static Workspace g_ws;
at::Tensor fwd(const at::Tensor& A,
const at::Tensor& B_q,
const at::Tensor& B_shuffle,
const at::Tensor& B_scale_sh) {
auto guard = at::DeviceGuard(A.device());
at::Tensor A_bf16 = (A.scalar_type() == at::kBFloat16 && A.is_contiguous())
? A : A.to(A.device(), at::kBFloat16, false, false,
at::MemoryFormat::Contiguous);
const int M = A_bf16.size(0);
const int K = A_bf16.size(1);
const int N = B_q.size(0);
const int KS = K / 32;
const int scaleN = ((KS + 7) / 8) * 8;
at::Tensor Bsh = B_shuffle.view(at::kByte);
if (!Bsh.is_contiguous()) Bsh = Bsh.contiguous();
at::Tensor Bssh = B_scale_sh.view(at::kByte);
if (!Bssh.is_contiguous()) Bssh = Bssh.contiguous();
const int num_ksplit = get_num_ksplit(M, K);
g_ws.ensure(M, N, K, num_ksplit, A_bf16.options());
launch_quant(
reinterpret_cast<const __bf16*>(A_bf16.data_ptr<at::BFloat16>()),
g_ws.A_fp4.data_ptr<uint8_t>(),
g_ws.A_scale.data_ptr<uint8_t>(),
M, K);
float* c_partial_ptr = (num_ksplit > 1) ? g_ws.C_partial.data_ptr<float>() : nullptr;
launch_gemm(
g_ws.A_fp4.data_ptr<uint8_t>(), g_ws.A_scale.data_ptr<uint8_t>(),
Bsh.data_ptr<uint8_t>(), Bssh.data_ptr<uint8_t>(),
c_partial_ptr,
reinterpret_cast<uint16_t*>(g_ws.C.data_ptr<at::BFloat16>()),
M, N, K, scaleN, num_ksplit);
return g_ws.C;
}
"""
_ext = load_inline(
name=f"g_{uuid.uuid4().hex[:8]}",
cpp_sources=[CPP],
cuda_sources=[HIP_KERNEL],
functions=["fwd"],
with_cuda=True,
extra_cflags=["-O3", "-std=c++20"],
extra_cuda_cflags=[
"-O3",
"--offload-arch=gfx950",
"-ffast-math",
"-munsafe-fp-atomics",
"-std=c++20",
"-mllvm", "-amdgpu-early-inline-all=true",
"-mllvm", "-amdgpu-function-calls=false",
"-mwavefrontsize64",
"-mcumode",
"-mllvm", "--amdgpu-kernarg-preload-count=16",
"-mllvm", "-enable-post-misched=0",
"-mllvm", "--lsr-drop-solution=1",
"-mllvm", "-amdgpu-coerce-illegal-types=1",
"-fgpu-flush-denormals-to-zero",
"-fno-offload-uniform-block",
],
extra_ldflags=["-lamdhip64"],
)
def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:
"""MXFP4 GEMM v21: hw FP4 conversion (fixed scale) + constexpr loop unrolling."""
A, _, B_q, B_shuffle, B_scale_sh = data
return _ext.fwd(A.cuda(), B_q, B_shuffle, B_scale_sh)
scrolls · 533 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 522107.
"""- FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.- Quant logic follows aiter op_tests/test_gemm_a4w4.py (get_triton_quant(QuantType.per_1x32)).+ MXFP4 GEMM v21 — gfx950 (MI355X) optimized.++ Changes from v20:+ 1. Hardware FP4 conversion via v_cvt_scalef32_pk_fp4_bf16+ - Replaces software f32_to_fp4_e2m1 with single instruction+ - Fixed scale semantics: pass bit_cast<float>(inv_exp << 23)+ where inv_exp = amax_exp - 2, so hardware applies 2^(129-amax_exp)+ 2. Template GEMM on CKT for constexpr loop unrolling (from v20)+ 3. Improved splitK: only split K=7168 with small M (from v20)"""- from task import input_t, output_t+ import os+ os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"+ from typing import Tuple+ import torch+ from torch.utils.cpp_extension import load_inline+ import uuid- def custom_kernel(data: input_t) -> output_t:- """- Reference: MXFP4 per-1x32 quant on A; B_shuffle, B_scale_sh from generate_input.- gemm_a4w4 with bpreshuffle=True.- """- import aiter- from aiter import QuantType, dtypes- A, B, B_q, B_shuffle, B_scale_sh = data- A = A.contiguous()- B = B.contiguous()- m, k = A.shape- n, _ = B.shape+ HIP_KERNEL = r"""+ #include <hip/hip_runtime.h>+ #include <stdint.h>- quant_func = aiter.get_triton_quant(QuantType.per_1x32)- A_q, A_scale_sh = quant_func(A, shuffle=True)- out_gemm = aiter.gemm_a4w4(- A_q,- B_shuffle,- A_scale_sh,- B_scale_sh,- dtype=dtypes.bf16,- bpreshuffle=True,- )- return out_gemm+ using int4_v = int __attribute__((ext_vector_type(4)));+ using float4_v = float __attribute__((ext_vector_type(4)));+ using bf16x2 = __bf16 __attribute__((ext_vector_type(2)));++ static constexpr int FP4_E2M1 = 4;++ __device__ float4_v __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(+ int4_v a, int4_v b, float4_v c,+ int cbsz, int blgp, int op_sel_a, int scale_a, int op_sel_b, int scale_b+ ) __asm("llvm.amdgcn.mfma.scale.f32.16x16x128.f8f6f4.v4i32.v4i32");++ __device__ __forceinline__ int4_v load16(const uint8_t* __restrict__ p) {+ return *reinterpret_cast<const int4_v*>(p);+ }++ __device__ __forceinline__ uint16_t float_to_bf16(float f) {+ bf16x2 v;+ v[0] = static_cast<__bf16>(f);+ uint16_t r;+ __builtin_memcpy(&r, &v, sizeof(r));+ return r;+ }++ // Hardware BF16x2 -> FP4x2 conversion using v_cvt_scalef32_pk_fp4_bf16+ // The instruction extracts biased exponent E from scale float, applies 2^(127-E) compression.+ // To get correct compression factor 2^(129 - amax_exp), pass scale with biased_exp = amax_exp - 2.+ __device__ __forceinline__ uint8_t hw_bf16x2_to_fp4x2(uint32_t bf16_pair, float scale) {+ uint32_t result;+ asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"+ : "=v"(result) : "v"(bf16_pair), "v"(scale));+ return static_cast<uint8_t>(result & 0xFFu);+ }++ // Quant kernel: BF16 [M, K] -> FP4x2 [M, K//2] + E8M0 scale [M, K//32]+ // Uses hardware v_cvt_scalef32_pk_fp4_bf16 for FP4 conversion+ __global__ void __launch_bounds__(128, 4)+ mxfp4_quant(+ const __bf16* __restrict__ A_bf16,+ uint8_t* __restrict__ A_fp4,+ uint8_t* __restrict__ A_scale,+ int M, int K)+ {+ const int KS = K / 32;+ const int K2 = K / 2;+ const int group = blockIdx.x * 128 + threadIdx.x;+ const int row = group / KS;+ const int kg = group % KS;++ if (row >= M) return;++ const auto* src = A_bf16 + (long)row * K + kg * 32;++ float absMax = 1e-10f;+ #pragma unroll+ for (int i = 0; i < 32; ++i) {+ float v = __builtin_elementwise_abs(static_cast<float>(src[i]));+ absMax = (v > absMax) ? v : absMax;+ }++ uint32_t u32 = __builtin_bit_cast(uint32_t, absMax);+ const uint32_t amax_exp = ((u32 + 0x200000u) >> 23) & 0xFFu;+ const uint32_t inv_exp = (amax_exp >= 2u) ? (amax_exp - 2u) : 0u;+ A_scale[(long)row * KS + kg] = static_cast<uint8_t>(inv_exp);++ // Hardware scale: biased_exp = inv_exp = amax_exp - 2+ // Instruction applies 2^(127 - inv_exp) = 2^(129 - amax_exp) for compression — correct!+ const float hw_scale = __builtin_bit_cast(float, static_cast<uint32_t>(inv_exp) << 23);++ const uint32_t* src_u32 = reinterpret_cast<const uint32_t*>(src);+ auto* dst = reinterpret_cast<uint8_t*>(A_fp4 + (long)row * K2 + kg * 16);+ #pragma unroll+ for (int i = 0; i < 16; ++i) {+ // Each uint32_t holds a pair of bf16 values (native layout)+ dst[i] = hw_bf16x2_to_fp4x2(src_u32[i], hw_scale);+ }+ }++ template<bool ALWAYS_VALID>+ __device__ __forceinline__ int4_v load_or_zero(bool rt_valid, const uint8_t* p) {+ if constexpr (ALWAYS_VALID) return load16(p);+ else return rt_valid ? load16(p) : int4_v{0,0,0,0};+ }++ // Main GEMM kernel — templated on CKT (ktiles_per_split) for constexpr loop unrolling+ // CKT=0 means runtime loop count+ template<int BM, int BN, int NWARPS, bool SPLITK, bool A_VALID, bool B_VALID, int CKT = 0>+ __global__ void __launch_bounds__(NWARPS * 64, 2)+ mxfp4_gemm(+ const uint8_t* __restrict__ A,+ const uint8_t* __restrict__ As,+ const uint8_t* __restrict__ Bsh,+ const uint8_t* __restrict__ Bssh,+ float* __restrict__ C_partial,+ uint16_t* __restrict__ C_final,+ int M, int N, int K, int scaleN,+ int total_ktiles, int ktiles_per_split,+ int tile_off_x, int tile_off_y)+ {+ static_assert(BN % 16 == 0);+ constexpr auto WAVES_M = (BM + 15) / 16;+ constexpr auto WAVES_N = BN / 16;+ static_assert(WAVES_M * WAVES_N == NWARPS);++ const auto ks_idx = blockIdx.z;+ const auto lane = threadIdx.x % 64;+ const auto wave = threadIdx.x / 64;+ const auto wave_m = wave / WAVES_N;+ const auto wave_n = wave % WAVES_N;++ const auto tile_m = (blockIdx.y + tile_off_y) * BM + wave_m * 16;+ const auto tile_n = (blockIdx.x + tile_off_x) * BN + wave_n * 16;++ if (tile_m >= M || tile_n >= N) return;++ const int ks_start = ks_idx * ktiles_per_split;+ const int ks_end = min(ks_start + ktiles_per_split, total_ktiles);++ const auto lrow = lane % 16;+ const auto kgrp = lane / 16;++ const auto gm = tile_m + lrow;+ const auto gn = tile_n + lrow;+ const auto K2 = K / 2;+ const auto KS = K / 32;++ const auto a_rt = A_VALID | (gm < M);+ const auto b_rt = B_VALID | (gn < N);++ const auto n_tile = tile_n / 16;+ const auto bsh_lane_base = Bsh + (long)n_tile * (K / 64) * 512 + (long)lrow * 16;++ const uint8_t* a_row = nullptr;+ const uint8_t* as_row = nullptr;+ if constexpr (A_VALID) {+ a_row = A + (long)gm * K2;+ as_row = As + (long)gm * KS;+ } else {+ if (a_rt) { a_row = A + (long)gm * K2;+ as_row = As + (long)gm * KS; }+ }++ int bssh_base = 0;+ if constexpr (B_VALID) {+ bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);+ } else {+ if (b_rt) bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);+ }++ const auto k_half_off = (kgrp & 1) * 256;+ const auto k_blk_base = kgrp >> 1;++ static constexpr long a_kt_stride = 64L;+ static constexpr long bsh_kt_stride = 1024L;++ const uint8_t* a_ptr = nullptr;+ const uint8_t* bsh_ptr = nullptr;+ const uint8_t* bssh_ptr = nullptr;++ if constexpr (A_VALID) {+ a_ptr = a_row + (long)ks_start * a_kt_stride + kgrp * 16;+ } else {+ if (a_rt) a_ptr = a_row + (long)ks_start * a_kt_stride + kgrp * 16;+ }+ if constexpr (B_VALID) {+ bsh_ptr = bsh_lane_base + (long)(ks_start * 2 + k_blk_base) * 512 + k_half_off;+ bssh_ptr = Bssh + bssh_base + (ks_start & 1) * 2 + (ks_start >> 1) * 256;+ } else {+ if (b_rt) {+ bsh_ptr = bsh_lane_base + (long)(ks_start * 2 + k_blk_base) * 512 + k_half_off;+ bssh_ptr = Bssh + bssh_base + (ks_start & 1) * 2 + (ks_start >> 1) * 256;+ }+ }++ float4_v acc{0.f, 0.f, 0.f, 0.f};++ // bssh step pattern depends on ks_start parity; doesn't change across quads+ const auto bssh_step0 = (ks_start & 1) ? 254 : 2;+ const auto bssh_step1 = 256 - bssh_step0;++ // Macro for one MFMA iteration+ #define DO_MFMA(a_off, b_off, bssh_off, ks_val) \+ { \+ const auto av = load_or_zero<A_VALID>(a_rt, a_ptr + (a_off) * a_kt_stride); \+ const auto bv = load_or_zero<B_VALID>(b_rt, bsh_ptr + (b_off) * bsh_kt_stride); \+ const auto ks = (ks_val) * 4 + kgrp; \+ int sa, sb; \+ if constexpr (A_VALID) sa = static_cast<int>(as_row[ks]); \+ else sa = (a_rt & (ks < KS)) ? static_cast<int>(as_row[ks]) : 127; \+ if constexpr (B_VALID) sb = static_cast<int>(*(bssh_ptr + (bssh_off))); \+ else sb = (b_rt & (ks < KS)) ? static_cast<int>(*(bssh_ptr + (bssh_off))) : 127; \+ acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(av,bv,acc,FP4_E2M1,FP4_E2M1,0,sa,0,sb); \+ }++ // 4x unrolled main loop+ if constexpr (CKT > 0) {+ // Compile-time unrolled+ #pragma unroll+ for (int q = 0; q < (CKT / 4); ++q) {+ DO_MFMA(0, 0, 0, ks_start + q*4)+ DO_MFMA(1, 1, bssh_step0, ks_start + q*4 + 1)+ DO_MFMA(2, 2, bssh_step0 + bssh_step1, ks_start + q*4 + 2)+ DO_MFMA(3, 3, bssh_step0 + bssh_step1 + bssh_step0, ks_start + q*4 + 3)+ a_ptr += 4 * a_kt_stride;+ bsh_ptr += 4 * bsh_kt_stride;+ bssh_ptr += 512;+ }+ // Compile-time pair (parity same as bssh_step0 since quads advance by 4)+ if constexpr ((CKT % 4) >= 2) {+ DO_MFMA(0, 0, 0, ks_start + (CKT/4)*4)+ DO_MFMA(1, 1, bssh_step0, ks_start + (CKT/4)*4 + 1)+ a_ptr += 2 * a_kt_stride;+ bsh_ptr += 2 * bsh_kt_stride;+ bssh_ptr += 256;+ }+ // Compile-time single+ if constexpr ((CKT % 2) == 1) {+ DO_MFMA(0, 0, 0, ks_start + CKT - 1)+ }+ } else {+ // Runtime loop (fallback for unknown shapes)+ auto kt = ks_start;+ const int ks_count = ks_end - ks_start;+ const auto ks_end_quad = ks_start + (ks_count - (ks_count & 3));+ const auto ks_end_pair = ks_start + (ks_count - (ks_count & 1));++ for (; kt < ks_end_quad; kt += 4) {+ DO_MFMA(0, 0, 0, kt)+ DO_MFMA(1, 1, bssh_step0, kt + 1)+ DO_MFMA(2, 2, bssh_step0 + bssh_step1, kt + 2)+ DO_MFMA(3, 3, bssh_step0 + bssh_step1 + bssh_step0, kt + 3)+ a_ptr += 4 * a_kt_stride;+ bsh_ptr += 4 * bsh_kt_stride;+ bssh_ptr += 512;+ }+ const auto bssh_step0_trail = (kt & 1) ? 254 : 2;+ for (; kt < ks_end_pair; kt += 2) {+ DO_MFMA(0, 0, 0, kt)+ DO_MFMA(1, 1, bssh_step0_trail, kt + 1)+ a_ptr += 2 * a_kt_stride;+ bsh_ptr += 2 * bsh_kt_stride;+ bssh_ptr += 256;+ }+ if (kt < ks_end) {+ DO_MFMA(0, 0, 0, kt)+ }+ }+ #undef DO_MFMA++ const auto out_col = tile_n + lrow;+ const auto out_row_base = tile_m + kgrp * 4;++ if constexpr (!B_VALID) { if (out_col >= N) return; }++ constexpr bool out_rows_always_valid = A_VALID && (BM >= 16);++ if constexpr (SPLITK) {+ auto c_out = C_partial + (long)ks_idx * M * N + (long)out_row_base * N + out_col;+ #pragma unroll+ for (int i = 0; i < 4; ++i) {+ if constexpr (out_rows_always_valid) c_out[i * N] = acc[i];+ else if (out_row_base + i < M) c_out[i * N] = acc[i];+ }+ } else {+ auto c_out = C_final + (long)out_row_base * N + out_col;+ #pragma unroll+ for (int i = 0; i < 4; ++i) {+ if constexpr (out_rows_always_valid) c_out[i * N] = float_to_bf16(acc[i]);+ else if (out_row_base + i < M) c_out[i * N] = float_to_bf16(acc[i]);+ }+ }+ }++ // Reduce kernel+ __global__ void mxfp4_reduce(+ const float* __restrict__ C_partial,+ uint16_t* __restrict__ C_out,+ int M, int N, int NUM_KSPLIT)+ {+ const auto col = blockIdx.x * 32 + threadIdx.x;+ const auto row = blockIdx.y * 16 + threadIdx.y;+ if (row >= M || col >= N) return;++ float sum = 0.f;+ const auto mn = (long)row * N + col;+ const auto stride = (long)M * N;+ for (auto k = 0; k < NUM_KSPLIT; ++k)+ sum += C_partial[k * stride + mn];++ bf16x2 v;+ v[0] = static_cast<__bf16>(sum);+ uint16_t r;+ __builtin_memcpy(&r, &v, sizeof(r));+ C_out[mn] = r;+ }++ // Launch quant+ extern "C" void launch_quant(+ const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale, int M, int K)+ {+ const int KS = K / 32;+ const int n_groups = M * KS;+ const dim3 block{128};+ const dim3 grid{static_cast<uint32_t>((n_groups + 127) / 128)};+ mxfp4_quant<<<grid, block>>>(A_bf16, A_fp4, A_scale, M, K);+ }++ // Templated GEMM launcher+ template<int CKT>+ void launch_gemm_ckt(+ const uint8_t* A, const uint8_t* As,+ const uint8_t* Bsh, const uint8_t* Bssh,+ float* C_partial, uint16_t* C_final,+ int M, int N, int K, int scaleN, int NUM_KSPLIT)+ {+ const auto total_ktiles = K / 128;+ const auto ktiles_per_split = (total_ktiles + NUM_KSPLIT - 1) / NUM_KSPLIT;+ const auto do_splitk = NUM_KSPLIT > 1;++ auto launch = [&]<int BM, int BN, int NWARPS>() {+ static_assert(((BM + 15) / 16) * (BN / 16) == NWARPS);++ const int full_m = (BM >= 16) ? M / BM : 0;+ const int full_n = N / BN;+ const int total_m = (M + BM - 1) / BM;+ const int total_n = (N + BN - 1) / BN;+ const int edge_m = total_m - full_m;+ const int edge_n = total_n - full_n;++ const dim3 block{static_cast<uint32_t>(NWARPS * 64)};++ auto sub = [&]<bool AV, bool BV>(int gx, int gy, int ox, int oy) {+ if (gx <= 0 || gy <= 0) return;+ const dim3 grid{+ static_cast<uint32_t>(gx),+ static_cast<uint32_t>(gy),+ static_cast<uint32_t>(NUM_KSPLIT)+ };+ if (do_splitk)+ mxfp4_gemm<BM,BN,NWARPS,true,AV,BV,CKT><<<grid,block>>>(+ A,As,Bsh,Bssh,C_partial,nullptr,+ M,N,K,scaleN,total_ktiles,ktiles_per_split,ox,oy);+ else+ mxfp4_gemm<BM,BN,NWARPS,false,AV,BV,CKT><<<grid,block>>>(+ A,As,Bsh,Bssh,nullptr,C_final,+ M,N,K,scaleN,total_ktiles,ktiles_per_split,ox,oy);+ };++ sub.template operator()<true, true >(full_n, full_m, 0, 0);+ sub.template operator()<true, false>(edge_n, full_m, full_n, 0);+ sub.template operator()<false, true >(full_n, edge_m, 0, full_m);+ sub.template operator()<false, false>(edge_n, edge_m, full_n, full_m);++ if (do_splitk) {+ const dim3 rblock{32, 16};+ const dim3 rgrid{+ static_cast<uint32_t>((N + 31) / 32),+ static_cast<uint32_t>((M + 15) / 16)+ };+ mxfp4_reduce<<<rgrid, rblock>>>(C_partial, C_final, M, N, NUM_KSPLIT);+ }+ };++ if (M <= 8) launch.template operator()< 8, 32, 2>();+ else if (M <= 16) launch.template operator()< 16, 32, 2>();+ else if (M <= 32) launch.template operator()< 16, 32, 2>();+ else if (M <= 64) launch.template operator()< 32, 32, 4>();+ else if (M <=128) launch.template operator()< 32, 32, 4>();+ else launch.template operator()< 64, 32, 8>();+ }++ // Main dispatch — routes to CKT-specialized launcher+ extern "C" void launch_gemm(+ const uint8_t* A, const uint8_t* As,+ const uint8_t* Bsh, const uint8_t* Bssh,+ float* C_partial, uint16_t* C_final,+ int M, int N, int K, int scaleN, int NUM_KSPLIT)+ {+ const int total_ktiles = K / 128;+ const int ktiles_per_split = (total_ktiles + NUM_KSPLIT - 1) / NUM_KSPLIT;++ switch (ktiles_per_split) {+ case 4: launch_gemm_ckt< 4>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;+ case 8: launch_gemm_ckt< 8>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;+ case 12: launch_gemm_ckt<12>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;+ case 16: launch_gemm_ckt<16>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;+ default: launch_gemm_ckt< 0>(A,As,Bsh,Bssh,C_partial,C_final,M,N,K,scaleN,NUM_KSPLIT); break;+ }+ }+ """+++ CPP = r"""+ #include <torch/extension.h>+ #include <c10/core/DeviceGuard.h>++ extern "C" void launch_quant(const __bf16*, uint8_t*, uint8_t*, int, int);+ extern "C" void launch_gemm(const uint8_t*, const uint8_t*, const uint8_t*, const uint8_t*,+ float*, uint16_t*, int, int, int, int, int);++ static int get_num_ksplit(int M, int K) {+ // Only use splitK for K=7168+ (56+ ktiles) with small M+ // Matches AITER's tuned configs+ int total_ktiles = K / 128;+ if (total_ktiles < 28) return 1; // K<=3456: never split+ if (M <= 8) return 7; // kps=8 (K=7168)+ if (M <= 16) return 14; // kps=4 (K=7168)+ return 1;+ }++ struct Workspace {+ at::Tensor A_fp4;+ at::Tensor A_scale;+ at::Tensor C_partial;+ at::Tensor C;+ int64_t last_M = -1, last_K = -1, last_N = -1, last_ksplit = -1;++ void ensure(int M, int N, int K, int num_ksplit, const at::TensorOptions& opts) {+ if (M == last_M && K == last_K && N == last_N && num_ksplit == last_ksplit) return;+ int64_t KS = K / 32;+ A_fp4 = at::empty({(int64_t)M, (int64_t)(K / 2)}, opts.dtype(at::kByte));+ A_scale = at::empty({(int64_t)M, KS}, opts.dtype(at::kByte));+ C = at::empty({(int64_t)M, (int64_t)N}, opts.dtype(at::kBFloat16));+ if (num_ksplit > 1)+ C_partial = at::empty({(int64_t)num_ksplit, (int64_t)M, (int64_t)N}, opts.dtype(at::kFloat));+ else+ C_partial = at::Tensor();+ last_M = M; last_K = K; last_N = N; last_ksplit = num_ksplit;+ }+ };++ static Workspace g_ws;++ at::Tensor fwd(const at::Tensor& A,+ const at::Tensor& B_q,+ const at::Tensor& B_shuffle,+ const at::Tensor& B_scale_sh) {+ auto guard = at::DeviceGuard(A.device());++ at::Tensor A_bf16 = (A.scalar_type() == at::kBFloat16 && A.is_contiguous())+ ? A : A.to(A.device(), at::kBFloat16, false, false,+ at::MemoryFormat::Contiguous);++ const int M = A_bf16.size(0);+ const int K = A_bf16.size(1);+ const int N = B_q.size(0);+ const int KS = K / 32;+ const int scaleN = ((KS + 7) / 8) * 8;++ at::Tensor Bsh = B_shuffle.view(at::kByte);+ if (!Bsh.is_contiguous()) Bsh = Bsh.contiguous();+ at::Tensor Bssh = B_scale_sh.view(at::kByte);+ if (!Bssh.is_contiguous()) Bssh = Bssh.contiguous();++ const int num_ksplit = get_num_ksplit(M, K);++ g_ws.ensure(M, N, K, num_ksplit, A_bf16.options());++ launch_quant(+ reinterpret_cast<const __bf16*>(A_bf16.data_ptr<at::BFloat16>()),+ g_ws.A_fp4.data_ptr<uint8_t>(),+ g_ws.A_scale.data_ptr<uint8_t>(),+ M, K);++ float* c_partial_ptr = (num_ksplit > 1) ? g_ws.C_partial.data_ptr<float>() : nullptr;++ launch_gemm(+ g_ws.A_fp4.data_ptr<uint8_t>(), g_ws.A_scale.data_ptr<uint8_t>(),+ Bsh.data_ptr<uint8_t>(), Bssh.data_ptr<uint8_t>(),+ c_partial_ptr,+ reinterpret_cast<uint16_t*>(g_ws.C.data_ptr<at::BFloat16>()),+ M, N, K, scaleN, num_ksplit);++ return g_ws.C;+ }+ """++ _ext = load_inline(+ name=f"g_{uuid.uuid4().hex[:8]}",+ cpp_sources=[CPP],+ cuda_sources=[HIP_KERNEL],+ functions=["fwd"],+ with_cuda=True,+ extra_cflags=["-O3", "-std=c++20"],+ extra_cuda_cflags=[+ "-O3",+ "--offload-arch=gfx950",+ "-ffast-math",+ "-munsafe-fp-atomics",+ "-std=c++20",+ "-mllvm", "-amdgpu-early-inline-all=true",+ "-mllvm", "-amdgpu-function-calls=false",+ "-mwavefrontsize64",+ "-mcumode",+ "-mllvm", "--amdgpu-kernarg-preload-count=16",+ "-mllvm", "-enable-post-misched=0",+ "-mllvm", "--lsr-drop-solution=1",+ "-mllvm", "-amdgpu-coerce-illegal-types=1",+ "-fgpu-flush-denormals-to-zero",+ "-fno-offload-uniform-block",+ ],+ extra_ldflags=["-lamdhip64"],+ )+++ def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:+ """MXFP4 GEMM v21: hw FP4 conversion (fixed scale) + constexpr loop unrolling."""+ A, _, B_q, B_shuffle, B_scale_sh = data+ return _ext.fwd(A.cuda(), B_q, B_shuffle, B_scale_sh)
scrolls · 558 diff lines total
Best evidence level for this revision: reported
JSON