submission 683418
AlexNebula · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 160 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-683418?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:13a0a7a189bf4b25ade2e7a4c9a908550be57408b37ca904f8d9d0a0bfc31c8f
license declaredunknown
license concludedunknown
authorsAlexNebula
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4-MM v56: CVT(val, scale) 直接用 — 去掉手动乘+clamp。vector-width = uint4
const uint4* p128=(const uint4*)(A+row*k+col*32+quarter*8);Kernel source
submission.py160 lines
"""
MXFP4-MM v56: CVT(val, scale) 直接用 — 去掉手动乘+clamp。
CVT scale = 除法,所以传 2^(scale_unbiased) 让 CVT 做 val / 2^su = val * 2^(-su)。
之前 v14 失败因为 scale 传反了。现在传正确的 divisor。
"""
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_SRC = """
#include <torch/extension.h>
torch::Tensor fused_single(torch::Tensor A, torch::Tensor B, torch::Tensor Bq,
torch::Tensor C, int m, int n, int k);
void hip_quant_a(torch::Tensor A, torch::Tensor Aq, torch::Tensor As,
int m, int k, int sm, int sn);
"""
HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
typedef uint8_t fp4x2_t;
typedef fp4x2_t fp4_vec32 __attribute__((ext_vector_type(32)));
typedef float f32_vec4 __attribute__((ext_vector_type(4)));
#define FP4_TYPE 4
// Single kernel for small K: CVT(val, divisor) — no manual multiply or clamp
__global__ __launch_bounds__(64)
void single_kernel(const __hip_bfloat16* A, const __hip_bfloat16* B,
const uint8_t* Bq, __hip_bfloat16* C, int m, int n, int k, int bqs) {
int bm=blockIdx.x, bn=blockIdx.y, lane=threadIdx.x, rg=lane/16, rit=lane%16;
int a_row=bm*16+rit, b_col=bn*16+rit;
f32_vec4 c={};
for(int kt=0; kt<k/128; kt++){
fp4_vec32 ar={}; uint8_t ae=0;
if(a_row<m){
const uint32_t* ap32=(const uint32_t*)(A+a_row*k+kt*128+rg*32);
uint32_t mx=0; float v[32];
#pragma unroll
for(int i=0;i<16;i++){
uint32_t pair=ap32[i]; mx=max(mx,max(pair&0x7FFFu,(pair>>16)&0x7FFFu));
uint16_t lo=(uint16_t)(pair&0xFFFF), hi=(uint16_t)(pair>>16);
v[i*2]=__bfloat162float(*(__hip_bfloat16*)&lo);
v[i*2+1]=__bfloat162float(*(__hip_bfloat16*)&hi);
}
if(mx>0){
uint16_t mb=(uint16_t)mx; float ax=__bfloat162float(*(__hip_bfloat16*)&mb);
uint32_t ab=__float_as_uint(ax); uint32_t rb=(ab+0x200000u)&0xFF800000u;
int su=max(-127,min(127,(int)((rb>>23)&0xFF)-129)); ae=(uint8_t)(su+127);
// CVT divisor = 2^(su) — CVT does val / divisor = val * 2^(-su) = val * quant_scale
float divisor=__uint_as_float((uint32_t)(su+127)<<23);
uint32_t pk[4];
#pragma unroll
for(int g=0;g<4;g++){uint32_t p=0;
// CVT handles clamp internally when using proper scale
p=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p,v[g*8+0],v[g*8+1],divisor,0);
p=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p,v[g*8+2],v[g*8+3],divisor,1);
p=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p,v[g*8+4],v[g*8+5],divisor,2);
p=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p,v[g*8+6],v[g*8+7],divisor,3);pk[g]=p;}
*(uint4*)(&ar)=*(uint4*)(pk);
}
}
fp4_vec32 br={}; uint8_t be=127;
if(b_col<n){
*(uint4*)(&br)=*(const uint4*)(Bq+b_col*bqs+kt*64+rg*16);
const uint32_t* bp=(const uint32_t*)(B+b_col*k+kt*128+rg*32);
uint32_t bmx=0;
#pragma unroll
for(int i=0;i<16;i++){uint32_t bv=bp[i];bmx=max(bmx,max(bv&0x7FFFu,(bv>>16)&0x7FFFu));}
if(bmx>0){uint16_t mb=(uint16_t)bmx;float bax=__bfloat162float(*(__hip_bfloat16*)&mb);
uint32_t ab=__float_as_uint(bax);uint32_t rb=(ab+0x200000u)&0xFF800000u;
be=(uint8_t)(max(-127,min(127,(int)((rb>>23)&0xFF)-129))+127);}
}
c=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(ar,br,c,FP4_TYPE,FP4_TYPE,0,(int)ae,0,(int)be);
}
for(int i=0;i<4;i++){int gr=bm*16+rg*4+i,gc=bn*16+rit;
if(gr<m&&gc<n) C[gr*n+gc]=(__hip_bfloat16)c[i];}
}
torch::Tensor fused_single(torch::Tensor A,torch::Tensor B,torch::Tensor Bq,
torch::Tensor C,int m,int n,int k){
single_kernel<<<dim3((m+15)/16,(n+15)/16),64>>>(
(__hip_bfloat16*)A.data_ptr(),(__hip_bfloat16*)B.data_ptr(),
(uint8_t*)Bq.data_ptr(),(__hip_bfloat16*)C.data_ptr(),m,n,k,(int)Bq.stride(0));
return C;}
// Large K: HIP A quant with CVT(val, divisor) — no manual multiply or clamp
__global__ void quant_a_kernel(const __hip_bfloat16* A, uint8_t* Aq, uint8_t* As,
int m, int k, int aq_stride, int sm, int sn) {
int idx=blockIdx.x*blockDim.x+threadIdx.x;
int total=m*(k/32)*4; if(idx>=total) return;
int block_id=idx/4, quarter=idx%4;
int row=block_id/(k/32), col=block_id%(k/32);
const uint4* p128=(const uint4*)(A+row*k+col*32+quarter*8);
uint4 data=*p128; uint32_t vals[4]={data.x,data.y,data.z,data.w};
uint32_t partial_max=0; float v[8];
#pragma unroll
for(int i=0;i<4;i++){uint32_t pair=vals[i];
partial_max=max(partial_max,max(pair&0x7FFFu,(pair>>16)&0x7FFFu));
uint16_t lo=(uint16_t)(pair&0xFFFF),hi=(uint16_t)(pair>>16);
v[i*2]=__bfloat162float(*(__hip_bfloat16*)&lo);v[i*2+1]=__bfloat162float(*(__hip_bfloat16*)&hi);}
uint32_t mx=partial_max; mx=max(mx,__shfl_xor(mx,1)); mx=max(mx,__shfl_xor(mx,2));
uint8_t e=0; float divisor=1.0f;
if(mx>0){uint16_t mb=(uint16_t)mx;float ax=__bfloat162float(*(__hip_bfloat16*)&mb);
uint32_t ab=__float_as_uint(ax);uint32_t rb=(ab+0x200000u)&0xFF800000u;
int su=max(-127,min(127,(int)((rb>>23)&0xFF)-129));e=(uint8_t)(su+127);
divisor=__uint_as_float((uint32_t)(su+127)<<23);}
if(quarter==0){int d0=row/32,d1=(row/16)%2,d2=row%16,d3=col/8,d4=(col/4)%2,d5=col%4;
int si=d0*(sn/8*256)+d3*256+d5*64+d2*4+d4*2+d1;
if(si<sm*sn) As[si]=e;}
// CVT(val, divisor) — no manual multiply or clamp needed
uint32_t pk=0;
pk=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk,v[0],v[1],divisor,0);
pk=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk,v[2],v[3],divisor,1);
pk=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk,v[4],v[5],divisor,2);
pk=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk,v[6],v[7],divisor,3);
*(uint32_t*)(Aq+row*aq_stride+col*16+quarter*4)=pk;
}
void hip_quant_a(torch::Tensor A,torch::Tensor Aq,torch::Tensor As,int m,int k,int sm,int sn){
int total=m*(k/32)*4;
quant_a_kernel<<<(total+255)/256,256>>>((__hip_bfloat16*)A.data_ptr(),(uint8_t*)Aq.data_ptr(),
(uint8_t*)As.data_ptr(),m,k,(int)Aq.stride(0),sm,sn);}
"""
module = load_inline(name='mxfp4_mm', cpp_sources=[CPP_SRC], cuda_sources=[HIP_SRC],
functions=['fused_single','hip_quant_a'], verbose=True,
extra_cuda_cflags=["--offload-arch=gfx950","-std=c++20","-O3","-mllvm","-amdgpu-early-inline-all=true"])
_bufs = {}
def custom_kernel(data: input_t) -> output_t:
import aiter
from aiter import dtypes
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape; n = B.shape[0]
bk = (m, n, k)
if bk not in _bufs:
_bufs.clear()
sm = ((m+255)//256)*256; sn = ((k//32+7)//8)*8
_bufs[bk] = {
'C': torch.empty((m,n), dtype=torch.bfloat16, device=A.device),
'Aq': torch.empty((m,k//2), dtype=torch.uint8, device=A.device),
'As': torch.zeros((sm,sn), dtype=torch.uint8, device=A.device),
'sm': sm, 'sn': sn,
}
C = _bufs[bk]['C']
if k <= 1024:
return module.fused_single(A, B, B_q.view(torch.uint8), C, m, n, k)
else:
sm, sn = _bufs[bk]['sm'], _bufs[bk]['sn']
if sm > m or sn > k//32:
_bufs[bk]['As'].zero_()
module.hip_quant_a(A, _bufs[bk]['Aq'], _bufs[bk]['As'], m, k, sm, sn)
return aiter.gemm_a4w4(
_bufs[bk]['Aq'].view(dtypes.fp4x2), B_shuffle,
_bufs[bk]['As'].view(dtypes.fp8_e8m0), B_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True)
scrolls · 160 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