submission 725123
anAirdrop · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 135 lines, June 9 Researcher Reciprocity License v1.0.
submission_s45.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-725123?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:0cbe2a7f0eaf8221635654b886ddb9a754486a71e40684c18cd7fcd20dbab548
license declaredunknown
license concludedunknown
authorsanAirdrop
imported2026-08-26
Kernel source
submission_s45.py135 lines
"""Hybrid: fused HIP MFMA for K≤512, AITER for K>512. Works in leaderboard mode!
Fused kernel: single launch, no quant caching needed (~14µs for K=512)
AITER path: quant + shuffle + GEMM (3 launches, ~20-34µs for large K)
"""
from task import input_t, output_t
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
import torch
import sys
_hip_module = None
def _get_module():
global _hip_module
if _hip_module is not None:
return _hip_module
from torch.utils.cpp_extension import load_inline
hip_src = r"""
#include <hip/hip_runtime.h>
typedef int __attribute__((ext_vector_type(4))) int4_vec;
typedef int __attribute__((ext_vector_type(8))) int8_vec;
typedef float __attribute__((ext_vector_type(4))) float4_vec;
__device__ __forceinline__ unsigned char load_b_scale(
const unsigned char* bs, int row, int col, int sn) {
int d0=row/32, d1=(row%32)/16, d2=row%16, d3=col/8, d4=(col%8)/4, d5=col%4;
return bs[d0*(sn/8*256)+d3*256+d5*64+d2*4+d4*2+d1];
}
__device__ __forceinline__ unsigned char fp4_code(float ax) {
return (unsigned char)(
(ax > 0.25f) + (ax >= 0.75f) + (ax > 1.25f) + (ax >= 1.75f) +
(ax > 2.5f) + (ax >= 3.5f) + (ax > 5.0f));
}
extern "C" __global__ void fused_mfma_gemm(
const unsigned short* __restrict__ A_bf16,
const unsigned char* __restrict__ B_q,
const unsigned char* __restrict__ B_scale_sh,
unsigned short* __restrict__ C,
int M, int N, int K,
int stride_a, int stride_bq, int pad_Ks, int stride_c) {
int m_tile=blockIdx.x*16, n_tile=blockIdx.y*16;
int lane=threadIdx.x, m_lane=lane&15, k_lane=lane>>4;
int a_row=m_tile+m_lane, b_row=n_tile+m_lane;
float4_vec acc={0,0,0,0};
int Kg=K/32;
for(int k128=0; k128<K; k128+=128){
int a_k_start=k128+k_lane*32;
float max_abs=0.0f; float a_f[32];
if(a_row<M){
const unsigned short* ap=A_bf16+a_row*stride_a+a_k_start;
for(int i=0;i<32;i++){
float v=__uint_as_float((unsigned int)ap[i]<<16);
a_f[i]=v; max_abs=fmaxf(max_abs,fabsf(v));
}
} else { for(int i=0;i<32;i++) a_f[i]=0.0f; }
max_abs=fmaxf(max_abs,1.175494e-38f);
unsigned int mi=__float_as_uint(max_abs);
mi=(mi+0x200000u)&0x7F800000u;
int su=(int)(mi>>23)-127-2;
su=min(max(su,-127),127);
float qs=exp2f((float)(-su));
unsigned char asc=(unsigned char)(su+127);
unsigned char pk[16];
for(int i=0;i<16;i++){
float v0=a_f[2*i]*qs, v1=a_f[2*i+1]*qs;
unsigned char s0=(v0<0.0f)?8:0, s1=(v1<0.0f)?8:0;
pk[i]=(s0|fp4_code(fabsf(v0)))|((s1|fp4_code(fabsf(v1)))<<4);
}
int4_vec a_reg; __builtin_memcpy(&a_reg,pk,16);
int8_vec a_pad={a_reg[0],a_reg[1],a_reg[2],a_reg[3],0,0,0,0};
int4_vec b_reg;
if(b_row<N) b_reg=*reinterpret_cast<const int4_vec*>(B_q+b_row*stride_bq+(k128/2)+k_lane*16);
else b_reg={0,0,0,0};
int8_vec b_pad={b_reg[0],b_reg[1],b_reg[2],b_reg[3],0,0,0,0};
int akg=k128/32+k_lane;
unsigned char bsc=0x7F;
if(b_row<N&&akg<Kg) bsc=load_b_scale(B_scale_sh,b_row,akg,pad_Ks);
acc=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_pad,b_pad,acc,4,4,0,(int)asc,0,(int)bsc);
}
int cm=m_tile+k_lane*4, cn=n_tile+m_lane;
for(int i=0;i<4;i++){
if(cm+i<M&&cn<N) C[(cm+i)*stride_c+cn]=(unsigned short)(__float_as_uint(acc[i])>>16);
}
}
torch::Tensor launch_fused(torch::Tensor A, torch::Tensor B_q, torch::Tensor B_scale_sh,
int M, int N, int K, int pad_Ks) {
auto C=torch::empty({M,N},torch::dtype(torch::kBFloat16).device(A.device()));
fused_mfma_gemm<<<dim3((M+15)/16,(N+15)/16),dim3(64)>>>(
reinterpret_cast<const unsigned short*>(A.data_ptr<at::BFloat16>()),
B_q.data_ptr<unsigned char>(),B_scale_sh.data_ptr<unsigned char>(),
reinterpret_cast<unsigned short*>(C.data_ptr<at::BFloat16>()),
M,N,K,(int)A.stride(0),(int)B_q.stride(0),pad_Ks,(int)C.stride(0));
return C;
}
"""
cpp_src = "torch::Tensor launch_fused(torch::Tensor,torch::Tensor,torch::Tensor,int,int,int,int);"
try:
_hip_module = load_inline(name='fused_mfma_rtne', cpp_sources=[cpp_src],
cuda_sources=[hip_src], functions=['launch_fused'],
extra_cuda_cflags=['-O3','--offload-arch=gfx950'], verbose=False)
except Exception as e:
print(f"[s45] compile FAILED: {e}", file=sys.stderr)
return _hip_module
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
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]
mod = _get_module()
if k <= 512 and mod is not None:
# Fused: single launch, ~14µs for K=512 (vs ~20µs AITER 3-launch)
pad_Ks = B_scale_sh.view(torch.uint8).shape[1]
return mod.launch_fused(A, B_q.view(torch.uint8), B_scale_sh.view(torch.uint8), m, n, k, pad_Ks)
else:
# AITER: 3 launches but fast ASM GEMM for large K
A_q_fp4, A_scale = dynamic_mxfp4_quant(A)
A_q = A_q_fp4.view(dtypes.fp4x2)
scale = A_scale; sc_m, sc_n = scale.shape
pad_m = ((sc_m+255)//256)*256; pad_n = ((sc_n+7)//8)*8
if pad_m != sc_m or pad_n != sc_n:
sp = torch.empty(pad_m, pad_n, dtype=scale.dtype, device=scale.device)
sp[:sc_m,:sc_n] = scale; scale = sp
sm, sn = scale.shape
A_scale_sh = scale.view(sm//32,2,16,sn//8,2,4).permute(0,3,5,2,4,1).contiguous().view(sm,sn).view(dtypes.fp8_e8m0)
return aiter.gemm_a4w4(A_q, B_shuffle, A_scale_sh, B_scale_sh, dtype=dtypes.bf16, bpreshuffle=True)
scrolls · 135 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