submission 610306
Zaber · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 251 lines, June 9 Researcher Reciprocity License v1.0.
submission_v66.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-610306?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:b777950aac5507b3bb7f629a5b62cc7f0285e2795c00d9b69c5759e9c6c1b265
license declaredunknown
license concludedunknown
authorsZaber
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
int m, k, n, split_k;Kernel source
submission_v66.py251 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ['CXX'] = 'clang++'
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CPP_SRC = """
#include <torch/extension.h>
#include <c10/core/DispatchKeySet.h>
#include <ATen/core/dispatch/Dispatcher.h>
#include <tuple>
#include <string>
void init_kernels(std::string hsa_dir);
torch::Tensor monolithic_gemm(
torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh,
int m, int k, int n);
"""
CUDA_SRC = """
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <unordered_map>
#include <string>
#include <c10/core/DispatchKeySet.h>
#include <ATen/core/dispatch/Dispatcher.h>
struct CacheKey {
int m, k, n, split_k;
bool operator==(const CacheKey& o) const {
return m == o.m && k == o.k && n == o.n && split_k == o.split_k;
}
};
struct CacheKeyHash {
size_t operator()(const CacheKey& k) const {
return ((size_t)k.m << 48) ^ ((size_t)k.split_k << 32) ^ ((size_t)k.k << 16) ^ k.n;
}
};
static std::unordered_map<CacheKey, std::tuple<torch::Tensor, torch::Tensor, torch::Tensor>, CacheKeyHash> g_tensor_cache;
static std::string cached_hsa_dir = "";
std::pair<const char*, int> get_dispatch(int m, int k, int n) {
if (m <= 4 && n == 2880 && k == 512) return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 0};
if (m <= 32 && n == 4096 && k == 512) return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 0};
if (m <= 32 && n == 2880 && k == 512) return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 0};
if (m <= 16 && n == 2112 && k == 7168) return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 2};
if (m <= 64 && n == 7168 && k == 2048) return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 0};
if (m <= 256 && n == 3072 && k == 1536) return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 1};
if (m <= 8 && n == 2112 && k == 7168) return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 2};
if (m <= 16 && n == 3072 && k == 1536) return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 1};
if (m <= 64 && n == 3072 && k == 1536) return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 1};
return {"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 0};
}
void init_kernels(std::string hsa_dir) {
cached_hsa_dir = hsa_dir;
}
__global__ void prepare_a_kernel(const __nv_bfloat16* __restrict__ A,
uint8_t* __restrict__ A_q,
uint8_t* __restrict__ bs_e8m0,
int M, int K, int scaleN, int scaleM_pad,
int stride_am, int stride_ak) {
int global_tid = blockIdx.x * blockDim.x + threadIdx.x;
int subgroup_id = global_tid / 4;
int lane_id = threadIdx.x % 4;
int total_subgroups = M * (K / 32);
if (subgroup_id >= total_subgroups) return;
int m = subgroup_id / (K / 32);
int n = subgroup_id % (K / 32);
int base_idx = m * stride_am + n * 32 + lane_id * 8;
union {
ulonglong2 vec;
uint16_t u16[8];
} a_vec;
a_vec.vec = *reinterpret_cast<const ulonglong2*>(&A[base_idx]);
uint16_t thread_max = 0;
for (int i = 0; i < 8; ++i) {
uint16_t abs_val = a_vec.u16[i] & 0x7FFF;
if (abs_val > thread_max) thread_max = abs_val;
}
for (int offset = 2; offset > 0; offset /= 2) {
int other_i = __shfl_down((int)thread_max, offset, 64);
uint16_t other = (uint16_t)other_i;
if ((threadIdx.x % 4) + offset < 4) {
if (other > thread_max) thread_max = other;
}
}
int max_i = __shfl((int)thread_max, (threadIdx.x / 4) * 4, 64);
uint16_t max_abs = (uint16_t)max_i;
uint16_t amax_rounded = (max_abs + 0x0020) & 0xFF80;
int exp = amax_rounded >> 7;
int scale_e8m0_unbiased = (exp == 0) ? -127 : exp - 129;
scale_e8m0_unbiased = max(-127, min(127, scale_e8m0_unbiased));
uint32_t quant_scale_u = (127 - scale_e8m0_unbiased) << 23;
float quant_scale = __uint_as_float(quant_scale_u);
uint32_t out_packed = 0;
#pragma unroll
for(int i = 0; i < 4; ++i) {
__nv_bfloat16 a_val0, a_val1;
*(uint16_t*)&a_val0 = a_vec.u16[i * 2];
*(uint16_t*)&a_val1 = a_vec.u16[i * 2 + 1];
float val0 = __bfloat162float(a_val0) * quant_scale;
float val1 = __bfloat162float(a_val1) * quant_scale;
uint32_t cvt = 0;
cvt = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(cvt, val0, val1, 1.0f, 0);
out_packed |= ((cvt & 0xFF) << (i * 8));
}
uint32_t* A_q_u32 = (uint32_t*)A_q;
int aq_idx = m * (K / 8) + n * 4 + lane_id;
A_q_u32[aq_idx] = out_packed;
if (lane_id == 0) {
uint8_t bs_val = (uint8_t)(scale_e8m0_unbiased + 127);
int m_mod_32 = m % 32;
int bs_0 = m / 32;
int bs_1 = m_mod_32 / 16;
int bs_2 = m_mod_32 % 16;
int n_mod_8 = n % 8;
int bs_3 = n / 8;
int bs_4 = n_mod_8 / 4;
int bs_5 = n_mod_8 % 4;
int bs_offs = bs_1 + (bs_4 * 2) + (bs_2 * 4) + (bs_5 * 64) + (bs_3 * 256) + (bs_0 * 32 * scaleN);
bs_e8m0[bs_offs] = bs_val;
}
}
torch::Tensor monolithic_gemm(
torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh,
int m, int k, int n) {
auto dispatch = get_dispatch(m, k, n);
const char* kname_ptr = dispatch.first;
int split_k = dispatch.second;
CacheKey key = {m, k, n, split_k};
torch::Tensor x_fp4, bs_e8m0, out;
int scaleN_valid = k / 32;
int scaleN = ((scaleN_valid + 7) / 8) * 8;
int scaleM_pad = ((m + 31) / 32) * 32;
int scaleM_256 = ((m + 255) / 256) * 256;
int padded_m = ((m + 31) / 32) * 32;
if (g_tensor_cache.find(key) == g_tensor_cache.end()) {
auto options_u8 = torch::TensorOptions().dtype(torch::kUInt8).device(A.device());
auto options_bf16 = torch::TensorOptions().dtype(torch::kBFloat16).device(A.device());
x_fp4 = torch::empty({m, k / 2}, options_u8);
bs_e8m0 = torch::empty({scaleM_256, scaleN}, options_u8);
bs_e8m0.fill_(127);
if (split_k > 0) out = torch::zeros({padded_m, n}, options_bf16);
else out = torch::empty({padded_m, n}, options_bf16);
g_tensor_cache[key] = std::make_tuple(x_fp4, bs_e8m0, out);
} else {
auto& tuple_val = g_tensor_cache[key];
x_fp4 = std::get<0>(tuple_val);
bs_e8m0 = std::get<1>(tuple_val);
out = std::get<2>(tuple_val);
if (split_k > 0) out.zero_();
}
int total_subgroups = m * (k / 32);
int threads = 256;
int blocks = (total_subgroups * 4 + threads - 1) / threads;
if (blocks > 0) {
prepare_a_kernel<<<blocks, threads>>>(
reinterpret_cast<const __nv_bfloat16*>(A.data_ptr()),
x_fp4.data_ptr<uint8_t>(), bs_e8m0.data_ptr<uint8_t>(),
m, k, scaleN, scaleM_pad, A.stride(0), A.stride(1)
);
}
// Bypass ATen dtype enforcement dynamically using the B_shuffle and B_scale types
at::Tensor x_view = at::from_blob(x_fp4.data_ptr(), {m, k / 2}, B_shuffle.options());
at::Tensor bs_view = at::from_blob(bs_e8m0.data_ptr(), {scaleM_256, scaleN}, B_scale_sh.options());
static auto op = c10::Dispatcher::singleton().findSchemaOrThrow("aiter::gemm_a4w4_asm", "");
torch::jit::Stack stack;
stack.push_back(x_view);
stack.push_back(B_shuffle);
stack.push_back(bs_view);
stack.push_back(B_scale_sh);
stack.push_back(out);
stack.push_back(std::string(kname_ptr));
stack.push_back(c10::IValue()); // bias=None
stack.push_back(1.0); // alpha=1.0
stack.push_back(0.0); // beta=0.0
stack.push_back(true); // bpreshuffle=True
stack.push_back(c10::optional<int64_t>(split_k)); // log2_k_split
op.callBoxed(&stack);
return out;
}
"""
_module_cache = None
def get_module():
global _module_cache
if _module_cache is None:
_module_cache = load_inline(
name='prepare_a_module_v66',
cpp_sources=[CPP_SRC],
cuda_sources=[CUDA_SRC],
functions=['monolithic_gemm', 'init_kernels'],
verbose=False,
extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3"],
)
import aiter
hsa_dir = os.path.normpath(os.path.join(os.path.dirname(aiter.__file__), "..", "hsa", "gfx950", "f4gemm"))
_module_cache.init_kernels(hsa_dir)
return _module_cache
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape
n = B.shape[0]
module = get_module()
out = module.monolithic_gemm(A, B_shuffle, B_scale_sh, m, k, n)
return out[:m]
scrolls · 251 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