submission 517676
_knarf_04 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 602 lines, June 9 Researcher Reciprocity License v1.0.
submission-hip.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-517676?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:780edc2b295c095b485effd4b73169795cfac10ce5bcee1e64a4da44e0f226fa
license declaredunknown
license concludedunknown
authors_knarf_04
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 GEMM — v11: v10 + pre-allocated buffers.shared-memory
__shared__ uint8_t A_lds[2][BM * BK_HALF_PAD];tile-m = 16
if (M <= 16) { WM_val = 1; BM = 16; BN = 64; }tile-n = 128
constexpr int BN = 128;vector-width = int4
int4 v0, v1, v2, v3;Kernel source
submission-hip.py602 lines
"""
MXFP4 GEMM — v11: v10 + pre-allocated buffers.
- v10 LDS fixes (bank conflict padding, coalesced load, scale loading)
- Pre-allocate A_q, A_scale, ws in Python — eliminates torch::empty from hot path
"""
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
os.environ["CXX"] = "clang++"
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CPP_WRAPPER = """
void mxfp4_fused_gemm(
torch::Tensor A,
torch::Tensor B_shuf,
torch::Tensor B_sc_shuf,
torch::Tensor C,
torch::Tensor A_q,
torch::Tensor A_scale,
torch::Tensor ws,
int M, int N, int K
);
"""
CUDA_SRC = r"""
#include <hip/hip_runtime.h>
#include <torch/extension.h>
#include <stdint.h>
#define WARP_SIZE 64
static __device__ __forceinline__ float bf16_to_f32(uint16_t x) {
uint32_t u = (uint32_t)x << 16;
float f;
__builtin_memcpy(&f, &u, 4);
return f;
}
static __device__ __forceinline__ uint16_t f32_to_bf16(float x) {
uint32_t u;
__builtin_memcpy(&u, &x, 4);
u += 0x7FFF + ((u >> 16) & 1u);
return (uint16_t)(u >> 16);
}
static __device__ __forceinline__ uint8_t float_to_fp4_e2m1(float x) {
uint32_t u;
__builtin_memcpy(&u, &x, 4);
uint32_t sign = (u >> 28) & 0x8u;
uint32_t e = (u >> 23) & 0xFFu;
uint32_t m = u & 0x7FFFFFu;
if (e == 0) return (uint8_t)sign;
if (e < 127u) {
uint32_t adj = 126u - e;
m = (adj < 23u) ? ((0x400000u | (m >> 1u)) >> adj) : 0u;
e = 0;
} else {
e = e - 126u;
}
uint32_t combined = (e << 2) | (m >> 21);
uint32_t e2m1 = (combined + 1u) >> 1;
if (e2m1 > 7u) e2m1 = 7u;
return (uint8_t)(sign | e2m1);
}
typedef int __attribute__((ext_vector_type(8))) i32x8_t;
typedef float __attribute__((ext_vector_type(4))) f32x4_t;
// ===================== Kernel 1: Quantize A =====================
__global__ void __launch_bounds__(256)
quant_a_kernel(
const uint16_t* __restrict__ A,
uint8_t* __restrict__ A_q,
uint8_t* __restrict__ A_scale,
int M, int K
) {
int gidx = blockIdx.x * 256 + threadIdx.x;
int k_groups = K / 32;
int total = M * k_groups;
if (gidx >= total) return;
int row = gidx / k_groups;
int kg = gidx % k_groups;
const uint16_t* ap = A + (long)row * K + kg * 32;
int4 v0, v1, v2, v3;
__builtin_memcpy(&v0, ap, 16);
__builtin_memcpy(&v1, ap + 8, 16);
__builtin_memcpy(&v2, ap + 16, 16);
__builtin_memcpy(&v3, ap + 24, 16);
const uint16_t* s0 = (const uint16_t*)&v0;
const uint16_t* s1 = (const uint16_t*)&v1;
const uint16_t* s2 = (const uint16_t*)&v2;
const uint16_t* s3 = (const uint16_t*)&v3;
float vals[32];
float mx = 0.0f;
#pragma unroll
for (int i = 0; i < 8; i++) {
vals[i] = bf16_to_f32(s0[i]);
vals[i+8] = bf16_to_f32(s1[i]);
vals[i+16] = bf16_to_f32(s2[i]);
vals[i+24] = bf16_to_f32(s3[i]);
}
#pragma unroll
for (int i = 0; i < 32; i++) mx = fmaxf(mx, fabsf(vals[i]));
uint8_t sc;
float qs;
if (mx == 0.0f) { sc = 0; qs = 0.0f; }
else {
uint32_t mx_u;
__builtin_memcpy(&mx_u, &mx, 4);
mx_u = (mx_u + 0x200000u) & 0xFF800000u;
float amr;
__builtin_memcpy(&amr, &mx_u, 4);
float su = __builtin_floorf(__builtin_log2f(amr)) - 2.0f;
su = fmaxf(-127.0f, fminf(127.0f, su));
sc = (uint8_t)((int)su + 127);
qs = __builtin_exp2f(-su);
}
uint8_t fp4[16];
#pragma unroll
for (int i = 0; i < 16; i++) {
uint8_t lo = float_to_fp4_e2m1(vals[2*i] * qs);
uint8_t hi = float_to_fp4_e2m1(vals[2*i+1] * qs);
fp4[i] = (lo & 0xF) | ((hi & 0xF) << 4);
}
__builtin_memcpy(A_q + (long)row * (K / 2) + kg * 16, fp4, 16);
A_scale[row * k_groups + kg] = sc;
}
// ===================== Kernel 2: Register-only GEMM =====================
template<int WM, int WN, bool DIRECT_BF16>
__global__ void __launch_bounds__(WM * WN * WARP_SIZE)
mxfp4_gemm_reg(
const uint8_t* __restrict__ A_q,
const uint8_t* __restrict__ A_scale,
const uint8_t* __restrict__ B_shuf,
const uint8_t* __restrict__ B_sc_shuf,
void* __restrict__ C_out,
int M, int N, int K,
int sB, int sSC,
int k_per_split
) {
const int tid = threadIdx.x;
const int wid = tid / WARP_SIZE;
const int lid = tid % WARP_SIZE;
const int wm = wid / WN;
const int wn = wid % WN;
const int tile_m = blockIdx.y * (WM * 16) + wm * 16;
const int tile_n = blockIdx.x * (WN * 16) + wn * 16;
if (tile_m >= M || tile_n >= N) return;
const int a_row = tile_m + (lid & 15);
const int b_col = tile_n + (lid & 15);
const int kg = lid >> 4;
const int b_rg = b_col >> 4;
const int b_i2 = b_col & 15;
const int b_i0 = b_col >> 5;
const int b_i1 = (b_col >> 4) & 1;
const long b_base = (long)b_rg * (sB * 16);
const bool a_ok = (a_row < M);
const bool b_ok = (b_col < N);
const int split_id = blockIdx.z;
const int k_start = split_id * k_per_split;
const int k_end = min(k_start + k_per_split, K);
const int half_K = K / 2;
const int sc_K = K / 32;
f32x4_t acc = {0.f, 0.f, 0.f, 0.f};
for (int kb = k_start; kb < k_end; kb += 128) {
i32x8_t a_frag = {0,0,0,0,0,0,0,0};
int a_sv = 0;
if (a_ok) {
const uint8_t* ap = A_q + (long)a_row * half_K + (kb >> 1) + (kg << 4);
int4 tmp; __builtin_memcpy(&tmp, ap, 16);
a_frag[0] = tmp.x; a_frag[1] = tmp.y;
a_frag[2] = tmp.z; a_frag[3] = tmp.w;
a_sv = (int)A_scale[a_row * sc_K + (kb >> 5) + kg];
}
i32x8_t b_frag = {0,0,0,0,0,0,0,0};
int b_sv = 0;
if (b_ok) {
int bk = (kb >> 1) + (kg << 4);
int i3 = bk >> 5; int i4 = (bk >> 4) & 1;
const uint8_t* bp = B_shuf + b_base + i3 * 512 + i4 * 256 + b_i2 * 16;
int4 tmp; __builtin_memcpy(&tmp, bp, 16);
b_frag[0] = tmp.x; b_frag[1] = tmp.y;
b_frag[2] = tmp.z; b_frag[3] = tmp.w;
int sg = (kb >> 5) + kg;
b_sv = (int)B_sc_shuf[b_i0 * (sSC * 32) + (sg >> 3) * 256 + (sg & 3) * 64 + b_i2 * 4 + ((sg >> 2) & 1) * 2 + b_i1];
}
#if defined(__gfx950__)
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a_frag, b_frag, acc, 4, 4, 0, a_sv, 0, b_sv);
#endif
}
if (b_ok) {
if constexpr (DIRECT_BF16) {
uint16_t* out = reinterpret_cast<uint16_t*>(C_out);
#pragma unroll
for (int r = 0; r < 4; r++) {
int mr = tile_m + (lid >> 4) * 4 + r;
if (mr < M) out[(long)mr * N + b_col] = f32_to_bf16(acc[r]);
}
} else {
float* out = reinterpret_cast<float*>(C_out);
long off = (long)split_id * M * N;
#pragma unroll
for (int r = 0; r < 4; r++) {
int mr = tile_m + (lid >> 4) * 4 + r;
if (mr < M) out[off + (long)mr * N + b_col] = acc[r];
}
}
}
}
// ===================== Kernel 3: LDS GEMM with fixes =====================
// vs v8: (1) bank conflict padding, (2) coalesced A load, (3) scale loading bugfix
template<int BM, bool DIRECT_BF16>
__global__ void __launch_bounds__(256)
mxfp4_gemm_lds(
const uint8_t* __restrict__ A_q,
const uint8_t* __restrict__ A_scale,
const uint8_t* __restrict__ B_shuf,
const uint8_t* __restrict__ B_sc_shuf,
void* __restrict__ C_out,
int M, int N, int K,
int sB, int sSC,
int k_per_split
) {
constexpr int BN = 128;
constexpr int BK_EXT = 512;
constexpr int BK_HALF = BK_EXT / 2; // 256 bytes per row (data)
constexpr int BK_HALF_PAD = BK_HALF + 4; // 260 bytes stride (bank conflict fix: 260/4=65, 65%32=1)
constexpr int BK_SC = BK_EXT / 32; // 16 scale groups per row
constexpr int MXdl = BM / 16;
constexpr int LOADS_PER_ROW = BK_HALF / 16; // 16
// Double-buffered LDS with padded stride
__shared__ uint8_t A_lds[2][BM * BK_HALF_PAD];
__shared__ uint8_t A_sc_lds[2][BM * BK_SC];
const int tid = threadIdx.x;
const int wid = tid / WARP_SIZE;
const int lid = tid % WARP_SIZE;
const int tile_m = blockIdx.y * BM;
const int tile_n = blockIdx.x * BN;
if (tile_m >= M) return;
const int split_id = blockIdx.z;
const int k_start = split_id * k_per_split;
const int k_end = min(k_start + k_per_split, K);
const int half_K = K / 2;
const int sc_K = K / 32;
// Each wave handles 32 N-cols, 2 MFMA tiles of 16 cols
const int wave_n = tile_n + wid * 32;
const int b_kg = lid >> 4;
int b_col[2];
b_col[0] = wave_n + (lid & 15);
b_col[1] = wave_n + 16 + (lid & 15);
int b_i2[2], b_i0[2], b_i1[2];
long b_base_addr[2];
bool b_ok[2];
#pragma unroll
for (int nt = 0; nt < 2; nt++) {
b_ok[nt] = (b_col[nt] < N);
b_i2[nt] = b_col[nt] & 15;
b_i0[nt] = b_col[nt] >> 5;
b_i1[nt] = (b_col[nt] >> 4) & 1;
b_base_addr[nt] = (long)(b_col[nt] >> 4) * (sB * 16);
}
// Accumulators
f32x4_t acc[MXdl][2];
#pragma unroll
for (int mt = 0; mt < MXdl; mt++)
for (int nt = 0; nt < 2; nt++)
acc[mt][nt] = {0.f, 0.f, 0.f, 0.f};
// === Cooperative A load with coalesced access + padded LDS ===
constexpr int TOTAL_LOADS = BM * LOADS_PER_ROW;
constexpr int TOTAL_SC = BM * BK_SC;
auto load_a = [&](int buf, int sk) {
int sk_half = sk >> 1;
int sk_k32 = sk >> 5;
// FIX 2: Coalesced loading — consecutive threads load consecutive 16B chunks
for (int li = tid; li < TOTAL_LOADS; li += 256) {
int row = li / LOADS_PER_ROW;
int col16 = li % LOADS_PER_ROW;
int global_row = tile_m + row;
int global_k_byte = sk_half + col16 * 16;
int4 val = {0, 0, 0, 0};
if (global_row < M && global_k_byte + 16 <= half_K) {
__builtin_memcpy(&val, A_q + (long)global_row * half_K + global_k_byte, 16);
}
// FIX 1: Write to padded LDS stride (BK_HALF_PAD = 260, avoids bank conflicts)
__builtin_memcpy(&A_lds[buf][row * BK_HALF_PAD + col16 * 16], &val, 16);
}
// FIX 3: Scale loading — loop covers ALL BM*BK_SC entries (was only 256 of 1024)
for (int s = tid; s < TOTAL_SC; s += 256) {
int row = s / BK_SC;
int sg = s % BK_SC;
int global_row = tile_m + row;
int global_sg = sk_k32 + sg;
uint8_t v = 0;
if (global_row < M && global_sg < sc_K)
v = A_scale[global_row * sc_K + global_sg];
A_sc_lds[buf][s] = v;
}
};
// === Prologue: load first super-K-step ===
load_a(0, k_start);
__syncthreads();
int cur_buf = 0;
for (int sk = k_start; sk < k_end; sk += BK_EXT) {
int next_buf = cur_buf ^ 1;
int sk_len = min(BK_EXT, k_end - sk);
// Prefetch next super-K-step
int next_sk = sk + BK_EXT;
if (next_sk < k_end) {
load_a(next_buf, next_sk);
}
// Inner K-loop: up to 4 iterations of 128 FP4 each
for (int kk = 0; kk < sk_len; kk += 128) {
int kb = sk + kk;
// Load B from global (per-wave, 2 N-tiles)
i32x8_t b_frag[2];
int b_sv[2];
#pragma unroll
for (int nt = 0; nt < 2; nt++) {
b_frag[nt] = {0,0,0,0,0,0,0,0};
b_sv[nt] = 0;
if (b_ok[nt]) {
int bk = (kb >> 1) + (b_kg << 4);
int i3 = bk >> 5;
int i4 = (bk >> 4) & 1;
const uint8_t* bp = B_shuf + b_base_addr[nt] + i3 * 512 + i4 * 256 + b_i2[nt] * 16;
int4 tmp;
__builtin_memcpy(&tmp, bp, 16);
b_frag[nt][0] = tmp.x; b_frag[nt][1] = tmp.y;
b_frag[nt][2] = tmp.z; b_frag[nt][3] = tmp.w;
int sg = (kb >> 5) + b_kg;
int si3 = sg >> 3;
int si4 = (sg >> 2) & 1;
int si5 = sg & 3;
b_sv[nt] = (int)B_sc_shuf[b_i0[nt] * (sSC * 32) + si3 * 256 + si5 * 64 + b_i2[nt] * 4 + si4 * 2 + b_i1[nt]];
}
}
// Compute MXdl × 2 MFMA tiles
#pragma unroll
for (int mt = 0; mt < MXdl; mt++) {
// Read A from LDS (using padded stride BK_HALF_PAD)
i32x8_t a_frag = {0,0,0,0,0,0,0,0};
int a_sv = 0;
{
int a_row_lds = mt * 16 + (lid & 15);
int a_k_off = (kk >> 1) + (lid >> 4) * 16;
if (tile_m + a_row_lds < M) {
const uint8_t* ap = &A_lds[cur_buf][a_row_lds * BK_HALF_PAD + a_k_off];
int4 tmp;
__builtin_memcpy(&tmp, ap, 16);
a_frag[0] = tmp.x; a_frag[1] = tmp.y;
a_frag[2] = tmp.z; a_frag[3] = tmp.w;
a_sv = (int)A_sc_lds[cur_buf][a_row_lds * BK_SC + (kk >> 5) + (lid >> 4)];
}
}
#pragma unroll
for (int nt = 0; nt < 2; nt++) {
#if defined(__gfx950__)
acc[mt][nt] = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
a_frag, b_frag[nt], acc[mt][nt], 4, 4, 0, a_sv, 0, b_sv[nt]);
#endif
}
}
}
__syncthreads();
cur_buf = next_buf;
}
// === Write output ===
#pragma unroll
for (int mt = 0; mt < MXdl; mt++) {
#pragma unroll
for (int nt = 0; nt < 2; nt++) {
if (!b_ok[nt]) continue;
if constexpr (DIRECT_BF16) {
uint16_t* out = reinterpret_cast<uint16_t*>(C_out);
#pragma unroll
for (int r = 0; r < 4; r++) {
int mr = tile_m + mt * 16 + (lid >> 4) * 4 + r;
if (mr < M) out[(long)mr * N + b_col[nt]] = f32_to_bf16(acc[mt][nt][r]);
}
} else {
float* out = reinterpret_cast<float*>(C_out);
long off = (long)split_id * M * N;
#pragma unroll
for (int r = 0; r < 4; r++) {
int mr = tile_m + mt * 16 + (lid >> 4) * 4 + r;
if (mr < M) out[off + (long)mr * N + b_col[nt]] = acc[mt][nt][r];
}
}
}
}
}
// ===================== Kernel 4: Reduce + convert to bf16 =====================
__global__ void __launch_bounds__(256)
reduce_bf16(
const float* __restrict__ C_f32,
uint16_t* __restrict__ C,
int M, int N, int P
) {
int idx = blockIdx.x * 256 + threadIdx.x;
if (idx >= M * N) return;
float sum = C_f32[idx];
long stride = (long)M * N;
for (int p = 1; p < P; p++) sum += C_f32[p * stride + idx];
uint32_t u;
__builtin_memcpy(&u, &sum, 4);
u += 0x7FFF + ((u >> 16) & 1u);
C[idx] = (uint16_t)(u >> 16);
}
// ===================== Host dispatch =====================
void mxfp4_fused_gemm(
torch::Tensor A, torch::Tensor B_shuf, torch::Tensor B_sc_shuf,
torch::Tensor C,
torch::Tensor A_q, torch::Tensor A_scale, torch::Tensor ws_buf,
int M, int N, int K)
{
auto* bs = reinterpret_cast<const uint8_t*>(B_shuf.data_ptr());
auto* bsc = reinterpret_cast<const uint8_t*>(B_sc_shuf.data_ptr());
auto* c = reinterpret_cast<uint16_t*>(C.data_ptr());
int sB = K / 2, sSC = ((K / 32 + 7) / 8) * 8;
// 1. Quant A (into pre-allocated buffers)
auto* aq = A_q.data_ptr<uint8_t>();
auto* asc = A_scale.data_ptr<uint8_t>();
{
int total = M * (K / 32);
int blocks = (total + 255) / 256;
quant_a_kernel<<<blocks, 256>>>(
reinterpret_cast<const uint16_t*>(A.data_ptr()),
aq, asc, M, K);
}
// 2. GEMM — register-only for all shapes (higher occupancy, no LDS sync overhead)
bool use_lds = false;
if (!use_lds) {
// === Register-only kernel (best for small M, small K) ===
int WM_val, BM, BN;
if (M <= 16) { WM_val = 1; BM = 16; BN = 64; }
else { WM_val = 2; BM = 32; BN = 64; }
const int WN_val = 4;
int grid_x = (N + BN - 1) / BN;
int grid_y = (M + BM - 1) / BM;
int grid_mn = grid_x * grid_y;
int P = 1, max_splits = K / 128;
while (grid_mn * P < 256 && P * 2 <= max_splits) P *= 2;
int k_per_split = ((K / P + 127) / 128) * 128;
if (P == 1) {
dim3 grid(grid_x, grid_y, 1);
if (M <= 16) mxfp4_gemm_reg<1,4,true><<<grid, 4*WARP_SIZE>>>(aq,asc,bs,bsc,(void*)c,M,N,K,sB,sSC,K);
else mxfp4_gemm_reg<2,4,true><<<grid, 8*WARP_SIZE>>>(aq,asc,bs,bsc,(void*)c,M,N,K,sB,sSC,K);
} else {
float* ws = ws_buf.data_ptr<float>();
dim3 grid(grid_x, grid_y, P);
if (M <= 16) mxfp4_gemm_reg<1,4,false><<<grid, 4*WARP_SIZE>>>(aq,asc,bs,bsc,(void*)ws,M,N,K,sB,sSC,k_per_split);
else mxfp4_gemm_reg<2,4,false><<<grid, 8*WARP_SIZE>>>(aq,asc,bs,bsc,(void*)ws,M,N,K,sB,sSC,k_per_split);
int rblocks = (M * N + 255) / 256;
reduce_bf16<<<rblocks, 256>>>(ws, c, M, N, P);
}
} else {
// === LDS kernel (best for larger M or long K) ===
int BM;
if (M <= 16) BM = 16;
else if (M <= 32) BM = 32;
else BM = 64;
int grid_x = (N + 127) / 128;
int grid_y = (M + BM - 1) / BM;
int grid_mn = grid_x * grid_y;
// SplitK: target good wave occupancy
int P = 1, max_splits = K / 128;
int total_waves = grid_mn * 4;
if (total_waves < 384) {
while (grid_mn * P < 192 && P * 2 <= max_splits) P *= 2;
}
int k_per_split = ((K / P + 127) / 128) * 128;
#define LAUNCH_LDS_K(BM_V, DIRECT, DST) do { \
dim3 grid(grid_x, grid_y, (DIRECT) ? 1 : P); \
mxfp4_gemm_lds<BM_V, DIRECT><<<grid, 256>>>( \
aq, asc, bs, bsc, (void*)(DST), M, N, K, sB, sSC, \
(DIRECT) ? K : k_per_split); \
} while(0)
if (P == 1) {
if (BM == 16) { LAUNCH_LDS_K(16, true, c); }
else if (BM == 32) { LAUNCH_LDS_K(32, true, c); }
else { LAUNCH_LDS_K(64, true, c); }
} else {
float* ws = ws_buf.data_ptr<float>();
if (BM == 16) { LAUNCH_LDS_K(16, false, ws); }
else if (BM == 32) { LAUNCH_LDS_K(32, false, ws); }
else { LAUNCH_LDS_K(64, false, ws); }
int rblocks = (M * N + 255) / 256;
reduce_bf16<<<rblocks, 256>>>(ws, c, M, N, P);
}
#undef LAUNCH_LDS_K
}
}
"""
_module = load_inline(
name="mxfp4_hip_v11",
cpp_sources=[CPP_WRAPPER],
cuda_sources=[CUDA_SRC],
functions=["mxfp4_fused_gemm"],
verbose=True,
extra_cuda_cflags=["--offload-arch=gfx950", "-O3", "-std=c++17"],
)
# Pre-allocated buffer cache: (m, n, k) → (A_q, A_scale, ws, C)
_buf_cache = {}
def _get_buffers(m, n, k, device):
key = (m, n, k)
if key in _buf_cache:
return _buf_cache[key]
A_q = torch.empty((m, k // 2), dtype=torch.uint8, device=device)
A_scale = torch.empty((m, k // 32), dtype=torch.uint8, device=device)
C = torch.empty((m, n), dtype=torch.bfloat16, device=device)
# Compute P to know ws size
use_lds = (m >= 32 and k >= 1024) or m >= 64
if not use_lds:
BN = 64
BM = 16 if m <= 16 else 32
grid_mn = ((n + BN - 1) // BN) * ((m + BM - 1) // BM)
P = 1
max_splits = k // 128
while grid_mn * P < 256 and P * 2 <= max_splits:
P *= 2
else:
BM = 16 if m <= 16 else (32 if m <= 32 else 64)
grid_mn = ((n + 127) // 128) * ((m + BM - 1) // BM)
P = 1
max_splits = k // 128
total_waves = grid_mn * 4
if total_waves < 384:
while grid_mn * P < 192 and P * 2 <= max_splits:
P *= 2
if P > 1:
ws = torch.empty((P, m, n), dtype=torch.float32, device=device)
else:
ws = torch.empty(1, dtype=torch.float32, device=device) # dummy
_buf_cache[key] = (A_q, A_scale, ws, C)
return A_q, A_scale, ws, C
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
m, k = A.shape
n = B.shape[0]
A_q, A_scale, ws, C = _get_buffers(m, n, k, A.device)
_module.mxfp4_fused_gemm(A, B_shuffle, B_scale_sh, C, A_q, A_scale, ws, m, n, k)
return C
scrolls · 602 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