submission 661199
JIAQI PAN · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 646 lines, June 9 Researcher Reciprocity License v1.0.
submission_hip_vG.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-661199?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:e3ae6b7de46ccb4e2281c8988a9318853093bacd05c4f22afc61f6054c127e32
license declaredunknown
license concludedunknown
authorsJIAQI PAN
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ __half s_a[2][GEN_BLOCK_M * GEN_A_LDS_STRIDE];tile-k = 32
constexpr int BLOCK_K = 32;tile-m = 8
constexpr int GEN_BLOCK_M = 8;tile-n = 32
constexpr int GEN_BLOCK_N = 32;vector-width = float2
float2 p0 = __half22float2(__hmul2(a_row2[kk2 + 0], b_row2[kk2 + 0]));Kernel source
submission_hip_vG.py646 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
Shape-dispatched fused-A native HIP path.
Design:
- keeps fused A quantization
- keeps ping-pong LDS staging
- adds specialized wide-N kernels for the benchmark's small-M regimes
- keeps a general fallback kernel for non-matched shapes
Specialized dispatch:
- M == 4 and N % 64 == 0 -> 4x64 kernel
- M == 16 or M == 32 and N % 64 == 0 -> 8x64 kernel
- otherwise -> general 8x32 kernel
"""
from __future__ import annotations
import os
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
os.environ.setdefault("CXX", "clang++")
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CPP_SRC = """
void quant_gemm_mxfp4_vG(
torch::Tensor a,
torch::Tensor b_q,
torch::Tensor b_scale_sh,
torch::Tensor out
);
"""
CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/BFloat16.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <hip/hip_fp16.h>
#include <hip/hip_runtime.h>
#include <cstdint>
#include <stdexcept>
namespace {
constexpr int SCALE_GROUP_SIZE = 32;
constexpr int BLOCK_K = 32;
constexpr int GEN_BLOCK_M = 8;
constexpr int GEN_BLOCK_N = 32;
constexpr int GEN_A_LDS_STRIDE = BLOCK_K + 2;
constexpr int GEN_B_LDS_STRIDE = BLOCK_K + 2;
constexpr int SPEC_BLOCK_N = 64;
constexpr int SPEC_A_LDS_STRIDE = BLOCK_K + 2;
constexpr int SPEC_B_LDS_STRIDE = BLOCK_K + 2;
__device__ __constant__ uint32_t FP4_TO_F32_BITS[16] = {
0x00000000u, 0x3F000000u, 0x3F800000u, 0x3FC00000u,
0x40000000u, 0x40400000u, 0x40800000u, 0x40C00000u,
0x80000000u, 0xBF000000u, 0xBF800000u, 0xBFC00000u,
0xC0000000u, 0xC0400000u, 0xC0800000u, 0xC0C00000u,
};
__device__ __forceinline__ uint16_t f32_to_bf16_rne(float x) {
uint32_t bits = __float_as_uint(x);
uint32_t lsb = (bits >> 16) & 1u;
uint32_t bias = 0x7FFFu + lsb;
return static_cast<uint16_t>((bits + bias) >> 16);
}
__device__ __forceinline__ float e8m0_bits_to_f32(uint8_t e8m0) {
if (e8m0 == static_cast<uint8_t>(255)) {
return NAN;
}
uint32_t bits = e8m0 == 0 ? 0x00400000u : (static_cast<uint32_t>(e8m0) << 23);
return __uint_as_float(bits);
}
__device__ __forceinline__ float fp4_bits_to_f32(uint8_t q) {
return __uint_as_float(FP4_TO_F32_BITS[q & 0xFu]);
}
__device__ __forceinline__ void unpack_fp4x8_half2(uint32_t packed, __half scale_h, __half2* dst) {
__half2 scale2 = __halves2half2(scale_h, scale_h);
#pragma unroll
for (int i = 0; i < 4; ++i) {
uint8_t lo = static_cast<uint8_t>((packed >> (8 * i)) & 0xFu);
uint8_t hi = static_cast<uint8_t>((packed >> (8 * i + 4)) & 0xFu);
__half2 vals = __floats2half2_rn(fp4_bits_to_f32(lo), fp4_bits_to_f32(hi));
dst[i] = __hmul2(vals, scale2);
}
}
__device__ __forceinline__ uint32_t wave_reduce_max_u32(uint32_t v) {
uint32_t t = static_cast<uint32_t>(__shfl_down(static_cast<int>(v), 16, 32));
v = v > t ? v : t;
t = static_cast<uint32_t>(__shfl_down(static_cast<int>(v), 8, 32));
v = v > t ? v : t;
t = static_cast<uint32_t>(__shfl_down(static_cast<int>(v), 4, 32));
v = v > t ? v : t;
t = static_cast<uint32_t>(__shfl_down(static_cast<int>(v), 2, 32));
v = v > t ? v : t;
t = static_cast<uint32_t>(__shfl_down(static_cast<int>(v), 1, 32));
v = v > t ? v : t;
return v;
}
__device__ __forceinline__ uint8_t scale_bits_from_bf16_abs(uint16_t abs_bits) {
uint32_t is_zero = abs_bits == 0u;
uint32_t exp = (abs_bits >> 7) & 0xFFu;
uint32_t mant = abs_bits & 0x7Fu;
uint32_t is_subnormal = exp == 0u;
uint32_t g = (mant >> 6) & 1u;
uint32_t r = (mant >> 5) & 1u;
uint32_t s = (mant & 0x1Fu) != 0u;
uint32_t is_special = (exp == 0xFFu);
uint32_t round_up = (1u - is_special) & (g & (r | s | 1u));
uint32_t e8 = exp + round_up;
int scale_bits = static_cast<int>(e8) - 2;
if (scale_bits < 0) scale_bits = 0;
if (scale_bits > 254) scale_bits = 254;
uint32_t normal_val = static_cast<uint32_t>(scale_bits);
uint32_t out =
is_zero * 127u +
(1u - is_zero) * ((1u - is_subnormal) * normal_val);
return static_cast<uint8_t>(out);
}
__device__ __forceinline__ uint8_t bf16_to_fp4_bits_scaled(uint16_t x_bits, uint8_t scale_bits) {
uint32_t sign = (x_bits >> 12) & 0x8u;
uint32_t abs_bits = x_bits & 0x7FFFu;
uint32_t exp = (abs_bits >> 7) & 0xFFu;
uint32_t mant = abs_bits & 0x7Fu;
uint32_t is_zero = abs_bits == 0u;
uint32_t is_subnormal = exp == 0u;
uint32_t is_special = exp == 0xFFu;
uint32_t is_normal = (1u - is_zero) & (1u - is_subnormal) & (1u - is_special);
int de = static_cast<int>(exp) - static_cast<int>(scale_bits);
int d = de;
if (d < -2) d = -2;
if (d > 2) d = 2;
uint32_t gt32 = mant > 32u;
uint32_t ge64 = mant >= 64u;
uint32_t ge96 = mant >= 96u;
uint32_t nz = mant != 0u;
uint32_t is_m2 = d == -2;
uint32_t is_m1 = d == -1;
uint32_t is_0 = d == 0;
uint32_t is_p1 = d == 1;
uint32_t is_p2 = d == 2;
uint32_t mid_mag =
is_m2 * nz +
is_m1 * (1u + ge64) +
is_0 * (2u + gt32 + ge96) +
is_p1 * (4u + gt32 + ge96) +
is_p2 * (6u + gt32);
uint32_t ge3 = de >= 3;
uint32_t mid_sel = (de > -3) & (de < 3);
uint32_t mag = ge3 * 7u + mid_sel * mid_mag;
uint32_t out_mag = is_special * 7u + is_normal * mag;
return static_cast<uint8_t>(sign | out_mag);
}
template <int BLOCK_M_T, int A_STRIDE_T, bool CHECK_M>
__device__ __forceinline__ void load_a_tile(
const uint16_t* __restrict__ a_bf16,
__half* __restrict__ s_a_buf,
int global_m,
int local_m,
int local_n,
int M,
int K,
int kb
) {
const __half zero_h = __float2half_rn(0.0f);
const __half2 zero2 = __halves2half2(zero_h, zero_h);
uint16_t a_bits0 = 0;
uint16_t a_bits1 = 0;
uint32_t local_abs = 0;
bool row_valid = !CHECK_M || (global_m < M);
if ((local_n & 1) == 0 && local_n < 32 && row_valid) {
int gk = kb * BLOCK_K + local_n;
const uint32_t* a_pair_ptr = reinterpret_cast<const uint32_t*>(
a_bf16 + global_m * K + gk
);
uint32_t a_pair = *a_pair_ptr;
a_bits0 = static_cast<uint16_t>(a_pair & 0xFFFFu);
a_bits1 = static_cast<uint16_t>(a_pair >> 16);
uint32_t abs0 = static_cast<uint32_t>(a_bits0 & 0x7FFFu);
uint32_t abs1 = static_cast<uint32_t>(a_bits1 & 0x7FFFu);
local_abs = abs0 > abs1 ? abs0 : abs1;
}
uint32_t amax_bits = wave_reduce_max_u32(local_abs);
uint32_t scale_bits_local =
local_n == 0 ? static_cast<uint32_t>(row_valid
? scale_bits_from_bf16_abs(static_cast<uint16_t>(amax_bits))
: static_cast<uint8_t>(127))
: 0u;
uint32_t scale_bits_u32 = static_cast<uint32_t>(__shfl(static_cast<int>(scale_bits_local), 0, 32));
if ((local_n & 1) == 0 && local_n < 32) {
__half2 a_pair_h2 = zero2;
if (row_valid) {
__half scale_h = __float2half_rn(e8m0_bits_to_f32(static_cast<uint8_t>(scale_bits_u32)));
__half2 scale2 = __halves2half2(scale_h, scale_h);
uint8_t q0 = bf16_to_fp4_bits_scaled(a_bits0, static_cast<uint8_t>(scale_bits_u32));
uint8_t q1 = bf16_to_fp4_bits_scaled(a_bits1, static_cast<uint8_t>(scale_bits_u32));
__half2 vals = __floats2half2_rn(fp4_bits_to_f32(q0), fp4_bits_to_f32(q1));
a_pair_h2 = __hmul2(vals, scale2);
}
reinterpret_cast<__half2*>(s_a_buf + local_m * A_STRIDE_T)[local_n >> 1] = a_pair_h2;
}
}
template <int BLOCK_N_T, int B_STRIDE_T>
__device__ __forceinline__ void load_b_tile_loop(
const uint8_t* __restrict__ b_q,
const uint8_t* __restrict__ b_scale_sh,
__half* __restrict__ s_b_buf,
int linear_tid,
int block_threads,
int tile_n_base,
int kb,
int K,
int k_scale,
int b_q_stride
) {
constexpr int B_TILE_LOADS = BLOCK_N_T * (BLOCK_K / 8);
for (int load_idx = linear_tid; load_idx < B_TILE_LOADS; load_idx += block_threads) {
int ln = load_idx / (BLOCK_K / 8);
int chunk = load_idx % (BLOCK_K / 8);
int gn = tile_n_base + ln;
int gn_group = gn >> 5;
int lane_row = gn & 31;
int kb_group = kb >> 3;
int kb_sub = kb & 7;
int kb_c = kb_sub & 3;
int kb_e = kb_sub >> 2;
int blocks_k = k_scale >> 3;
int b_scale_base =
((((gn_group * blocks_k + kb_group) * 4 + kb_c) << 4) << 2) + (kb_e << 1);
int scale_src = b_scale_base + ((lane_row & 15) << 2) + (lane_row >> 4);
__half scale_h = __float2half_rn(e8m0_bits_to_f32(b_scale_sh[scale_src]));
int gk0 = kb * BLOCK_K + chunk * 8;
__half2* dst = reinterpret_cast<__half2*>(s_b_buf + ln * B_STRIDE_T + chunk * 8);
const uint32_t* packed_ptr = reinterpret_cast<const uint32_t*>(
b_q + gn * b_q_stride + (gk0 >> 1)
);
unpack_fp4x8_half2(*packed_ptr, scale_h, dst);
}
}
template <int A_STRIDE_T, int B_STRIDE_T>
__device__ __forceinline__ void compute_tile_1col(
const __half* __restrict__ s_a_buf,
const __half* __restrict__ s_b_buf,
int local_m,
int local_n,
float& acc0,
float& acc1,
float& acc2,
float& acc3
) {
const __half2* a_row2 = reinterpret_cast<const __half2*>(s_a_buf + local_m * A_STRIDE_T);
const __half2* b_row2 = reinterpret_cast<const __half2*>(s_b_buf + local_n * B_STRIDE_T);
#pragma unroll
for (int kk2 = 0; kk2 < BLOCK_K / 2; kk2 += 4) {
float2 p0 = __half22float2(__hmul2(a_row2[kk2 + 0], b_row2[kk2 + 0]));
float2 p1 = __half22float2(__hmul2(a_row2[kk2 + 1], b_row2[kk2 + 1]));
float2 p2 = __half22float2(__hmul2(a_row2[kk2 + 2], b_row2[kk2 + 2]));
float2 p3 = __half22float2(__hmul2(a_row2[kk2 + 3], b_row2[kk2 + 3]));
acc0 += p0.x + p0.y;
acc1 += p1.x + p1.y;
acc2 += p2.x + p2.y;
acc3 += p3.x + p3.y;
}
}
template <int A_STRIDE_T, int B_STRIDE_T>
__device__ __forceinline__ void compute_tile_2col(
const __half* __restrict__ s_a_buf,
const __half* __restrict__ s_b_buf,
int local_m,
int local_n,
float& acc0a,
float& acc1a,
float& acc2a,
float& acc3a,
float& acc0b,
float& acc1b,
float& acc2b,
float& acc3b
) {
const __half2* a_row2 = reinterpret_cast<const __half2*>(s_a_buf + local_m * A_STRIDE_T);
const __half2* b_row2a = reinterpret_cast<const __half2*>(s_b_buf + local_n * B_STRIDE_T);
const __half2* b_row2b = reinterpret_cast<const __half2*>(s_b_buf + (local_n + 32) * B_STRIDE_T);
#pragma unroll
for (int kk2 = 0; kk2 < BLOCK_K / 2; kk2 += 4) {
float2 pa0 = __half22float2(__hmul2(a_row2[kk2 + 0], b_row2a[kk2 + 0]));
float2 pa1 = __half22float2(__hmul2(a_row2[kk2 + 1], b_row2a[kk2 + 1]));
float2 pa2 = __half22float2(__hmul2(a_row2[kk2 + 2], b_row2a[kk2 + 2]));
float2 pa3 = __half22float2(__hmul2(a_row2[kk2 + 3], b_row2a[kk2 + 3]));
acc0a += pa0.x + pa0.y;
acc1a += pa1.x + pa1.y;
acc2a += pa2.x + pa2.y;
acc3a += pa3.x + pa3.y;
float2 pb0 = __half22float2(__hmul2(a_row2[kk2 + 0], b_row2b[kk2 + 0]));
float2 pb1 = __half22float2(__hmul2(a_row2[kk2 + 1], b_row2b[kk2 + 1]));
float2 pb2 = __half22float2(__hmul2(a_row2[kk2 + 2], b_row2b[kk2 + 2]));
float2 pb3 = __half22float2(__hmul2(a_row2[kk2 + 3], b_row2b[kk2 + 3]));
acc0b += pb0.x + pb0.y;
acc1b += pb1.x + pb1.y;
acc2b += pb2.x + pb2.y;
acc3b += pb3.x + pb3.y;
}
}
template <bool CHECK_M>
__global__ __launch_bounds__(GEN_BLOCK_M * GEN_BLOCK_N, 2) void gemm_general_kernel(
const uint16_t* __restrict__ a_bf16,
const uint8_t* __restrict__ b_q,
const uint8_t* __restrict__ b_scale_sh,
uint16_t* __restrict__ out,
int M,
int N,
int K,
int k_scale,
int row_offset
) {
int local_n = threadIdx.x;
int local_m = threadIdx.y;
int linear_tid = local_m * GEN_BLOCK_N + local_n;
int global_m = row_offset + blockIdx.y * GEN_BLOCK_M + local_m;
int global_n = blockIdx.x * GEN_BLOCK_N + local_n;
int tile_n_base = blockIdx.x * GEN_BLOCK_N;
__shared__ __half s_a[2][GEN_BLOCK_M * GEN_A_LDS_STRIDE];
__shared__ __half s_b[2][GEN_BLOCK_N * GEN_B_LDS_STRIDE];
float acc0 = 0.0f;
float acc1 = 0.0f;
float acc2 = 0.0f;
float acc3 = 0.0f;
int k_blocks = K / BLOCK_K;
int b_q_stride = K / 2;
int block_threads = GEN_BLOCK_M * GEN_BLOCK_N;
load_a_tile<GEN_BLOCK_M, GEN_A_LDS_STRIDE, CHECK_M>(
a_bf16, s_a[0], global_m, local_m, local_n, M, K, 0
);
if (linear_tid < GEN_BLOCK_N * (BLOCK_K / 8)) {
load_b_tile_loop<GEN_BLOCK_N, GEN_B_LDS_STRIDE>(
b_q, b_scale_sh, s_b[0], linear_tid, block_threads, tile_n_base, 0, K, k_scale, b_q_stride
);
}
__syncthreads();
for (int kb = 0; kb < k_blocks - 1; ++kb) {
int cur = kb & 1;
int nxt = cur ^ 1;
compute_tile_1col<GEN_A_LDS_STRIDE, GEN_B_LDS_STRIDE>(
s_a[cur], s_b[cur], local_m, local_n, acc0, acc1, acc2, acc3
);
load_a_tile<GEN_BLOCK_M, GEN_A_LDS_STRIDE, CHECK_M>(
a_bf16, s_a[nxt], global_m, local_m, local_n, M, K, kb + 1
);
if (linear_tid < GEN_BLOCK_N * (BLOCK_K / 8)) {
load_b_tile_loop<GEN_BLOCK_N, GEN_B_LDS_STRIDE>(
b_q, b_scale_sh, s_b[nxt], linear_tid, block_threads, tile_n_base, kb + 1, K, k_scale, b_q_stride
);
}
__syncthreads();
}
compute_tile_1col<GEN_A_LDS_STRIDE, GEN_B_LDS_STRIDE>(
s_a[(k_blocks - 1) & 1], s_b[(k_blocks - 1) & 1], local_m, local_n, acc0, acc1, acc2, acc3
);
if (global_n < N) {
float acc = (acc0 + acc1) + (acc2 + acc3);
if constexpr (CHECK_M) {
if (global_m < M) {
out[global_m * N + global_n] = f32_to_bf16_rne(acc);
}
} else {
out[global_m * N + global_n] = f32_to_bf16_rne(acc);
}
}
}
template <int BLOCK_M_T, bool CHECK_M>
__global__ __launch_bounds__(BLOCK_M_T * 32, 2) void gemm_smallm_widen_kernel(
const uint16_t* __restrict__ a_bf16,
const uint8_t* __restrict__ b_q,
const uint8_t* __restrict__ b_scale_sh,
uint16_t* __restrict__ out,
int M,
int N,
int K,
int k_scale,
int row_offset
) {
int local_n = threadIdx.x;
int local_m = threadIdx.y;
int linear_tid = local_m * 32 + local_n;
int global_m = row_offset + blockIdx.y * BLOCK_M_T + local_m;
int tile_n_base = blockIdx.x * SPEC_BLOCK_N;
int global_n0 = tile_n_base + local_n;
int global_n1 = global_n0 + 32;
__shared__ __half s_a[2][BLOCK_M_T * SPEC_A_LDS_STRIDE];
__shared__ __half s_b[2][SPEC_BLOCK_N * SPEC_B_LDS_STRIDE];
float acc0a = 0.0f;
float acc1a = 0.0f;
float acc2a = 0.0f;
float acc3a = 0.0f;
float acc0b = 0.0f;
float acc1b = 0.0f;
float acc2b = 0.0f;
float acc3b = 0.0f;
int k_blocks = K / BLOCK_K;
int b_q_stride = K / 2;
int block_threads = BLOCK_M_T * 32;
load_a_tile<BLOCK_M_T, SPEC_A_LDS_STRIDE, CHECK_M>(
a_bf16, s_a[0], global_m, local_m, local_n, M, K, 0
);
load_b_tile_loop<SPEC_BLOCK_N, SPEC_B_LDS_STRIDE>(
b_q, b_scale_sh, s_b[0], linear_tid, block_threads, tile_n_base, 0, K, k_scale, b_q_stride
);
__syncthreads();
for (int kb = 0; kb < k_blocks - 1; ++kb) {
int cur = kb & 1;
int nxt = cur ^ 1;
compute_tile_2col<SPEC_A_LDS_STRIDE, SPEC_B_LDS_STRIDE>(
s_a[cur], s_b[cur], local_m, local_n,
acc0a, acc1a, acc2a, acc3a,
acc0b, acc1b, acc2b, acc3b
);
load_a_tile<BLOCK_M_T, SPEC_A_LDS_STRIDE, CHECK_M>(
a_bf16, s_a[nxt], global_m, local_m, local_n, M, K, kb + 1
);
load_b_tile_loop<SPEC_BLOCK_N, SPEC_B_LDS_STRIDE>(
b_q, b_scale_sh, s_b[nxt], linear_tid, block_threads, tile_n_base, kb + 1, K, k_scale, b_q_stride
);
__syncthreads();
}
compute_tile_2col<SPEC_A_LDS_STRIDE, SPEC_B_LDS_STRIDE>(
s_a[(k_blocks - 1) & 1], s_b[(k_blocks - 1) & 1], local_m, local_n,
acc0a, acc1a, acc2a, acc3a,
acc0b, acc1b, acc2b, acc3b
);
float acca = (acc0a + acc1a) + (acc2a + acc3a);
float accb = (acc0b + acc1b) + (acc2b + acc3b);
if constexpr (CHECK_M) {
if (global_m < M) {
out[global_m * N + global_n0] = f32_to_bf16_rne(acca);
out[global_m * N + global_n1] = f32_to_bf16_rne(accb);
}
} else {
out[global_m * N + global_n0] = f32_to_bf16_rne(acca);
out[global_m * N + global_n1] = f32_to_bf16_rne(accb);
}
}
} // namespace
void quant_gemm_mxfp4_vG(
torch::Tensor a,
torch::Tensor b_q,
torch::Tensor b_scale_sh,
torch::Tensor out
) {
TORCH_CHECK(a.is_cuda(), "a must be CUDA/HIP");
TORCH_CHECK(b_q.is_cuda(), "b_q must be CUDA/HIP");
TORCH_CHECK(b_scale_sh.is_cuda(), "b_scale_sh must be CUDA/HIP");
TORCH_CHECK(out.is_cuda(), "out must be CUDA/HIP");
TORCH_CHECK(a.dim() == 2, "a must be 2D");
TORCH_CHECK(b_q.dim() == 2, "b_q must be 2D");
TORCH_CHECK(b_scale_sh.dim() == 2, "b_scale_sh must be 2D");
TORCH_CHECK(out.dim() == 2, "out must be 2D");
TORCH_CHECK(a.scalar_type() == torch::kBFloat16, "a must be bf16");
TORCH_CHECK(b_q.element_size() == 1, "b_q elements must be 1 byte");
TORCH_CHECK(b_scale_sh.element_size() == 1, "b_scale_sh elements must be 1 byte");
TORCH_CHECK(out.scalar_type() == torch::kBFloat16, "out must be bf16");
TORCH_CHECK(a.is_contiguous(), "a must be contiguous");
TORCH_CHECK(b_q.is_contiguous(), "b_q must be contiguous");
TORCH_CHECK(b_scale_sh.is_contiguous(), "b_scale_sh must be contiguous");
TORCH_CHECK(out.is_contiguous(), "out must be contiguous");
int M = static_cast<int>(a.size(0));
int K = static_cast<int>(a.size(1));
int N = static_cast<int>(b_q.size(0));
int k_scale = K / SCALE_GROUP_SIZE;
TORCH_CHECK(K % 64 == 0, "K must be divisible by 64");
TORCH_CHECK(b_q.size(1) == K / 2, "b_q shape mismatch");
TORCH_CHECK(b_scale_sh.size(1) == k_scale, "b_scale_sh col mismatch");
TORCH_CHECK(b_scale_sh.size(0) >= N, "b_scale_sh row mismatch");
TORCH_CHECK(b_scale_sh.size(0) % 32 == 0, "b_scale_sh padded rows must be divisible by 32");
TORCH_CHECK(out.size(0) == M && out.size(1) == N, "out shape mismatch");
auto a_ptr = reinterpret_cast<const uint16_t*>(a.data_ptr<at::BFloat16>());
auto out_ptr = reinterpret_cast<uint16_t*>(out.data_ptr<at::BFloat16>());
auto b_q_ptr = reinterpret_cast<const uint8_t*>(b_q.data_ptr());
auto b_scale_sh_ptr = reinterpret_cast<const uint8_t*>(b_scale_sh.data_ptr());
if (N % 64 == 0 && M == 4) {
dim3 grid(N / 64, 1);
dim3 block(32, 4);
gemm_smallm_widen_kernel<4, false><<<grid, block>>>(
a_ptr, b_q_ptr, b_scale_sh_ptr, out_ptr, M, N, K, k_scale, 0
);
} else if (N % 64 == 0 && (M == 16 || M == 32)) {
dim3 grid(N / 64, M / 8);
dim3 block(32, 8);
gemm_smallm_widen_kernel<8, false><<<grid, block>>>(
a_ptr, b_q_ptr, b_scale_sh_ptr, out_ptr, M, N, K, k_scale, 0
);
} else {
int full_m_tiles = M / GEN_BLOCK_M;
int tail_m = M - full_m_tiles * GEN_BLOCK_M;
dim3 block(GEN_BLOCK_N, GEN_BLOCK_M);
if (full_m_tiles > 0) {
dim3 full_grid((N + GEN_BLOCK_N - 1) / GEN_BLOCK_N, full_m_tiles);
gemm_general_kernel<false><<<full_grid, block>>>(
a_ptr, b_q_ptr, b_scale_sh_ptr, out_ptr, M, N, K, k_scale, 0
);
}
if (tail_m > 0) {
dim3 tail_grid((N + GEN_BLOCK_N - 1) / GEN_BLOCK_N, 1);
gemm_general_kernel<true><<<tail_grid, block>>>(
a_ptr, b_q_ptr, b_scale_sh_ptr, out_ptr, M, N, K, k_scale, full_m_tiles * GEN_BLOCK_M
);
}
}
hipError_t err = hipGetLastError();
if (err != hipSuccess) {
throw std::runtime_error(hipGetErrorString(err));
}
}
"""
try:
_MODULE = load_inline(
name="mxfp4_quant_gemm_vG",
cpp_sources=[CPP_SRC],
cuda_sources=[CUDA_SRC],
functions=["quant_gemm_mxfp4_vG"],
extra_cflags=["-O3"],
extra_cuda_cflags=["-O3", "-std=c++20", "--offload-arch=gfx950"],
verbose=False,
)
except Exception:
_MODULE = None
_OUT_CACHE: dict[tuple[int, int, int, str, int | None], torch.Tensor] = {}
def _get_output_buffer(
m: int,
n: int,
k: int,
device: torch.device,
) -> torch.Tensor:
key = (m, n, k, device.type, device.index)
cached = _OUT_CACHE.get(key)
if cached is not None:
return cached
out = torch.empty((m, n), device=device, dtype=torch.bfloat16)
_OUT_CACHE[key] = out
return out
def _fallback_kernel(data: input_t) -> output_t:
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
a, _b, _b_q, b_shuffle, b_scale_sh = data
a_q, a_scale_sh = dynamic_mxfp4_quant(a.contiguous())
a_q = a_q.view(dtypes.fp4x2)
a_scale_sh = e8m0_shuffle(a_scale_sh).view(dtypes.fp8_e8m0)
return aiter.gemm_a4w4(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
def custom_kernel(data: input_t) -> output_t:
if _MODULE is None:
return _fallback_kernel(data)
a, _b, b_q, _b_shuffle, b_scale_sh = data
if not a.is_contiguous():
a = a.contiguous()
if not b_q.is_contiguous():
b_q = b_q.contiguous()
if not b_scale_sh.is_contiguous():
b_scale_sh = b_scale_sh.contiguous()
out = _get_output_buffer(a.shape[0], b_q.shape[0], a.shape[1], a.device)
_MODULE.quant_gemm_mxfp4_vG(a, b_q, b_scale_sh, out)
return out
scrolls · 646 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