submission 752196
somethingobscurefordevstuff · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 820 lines, June 9 Researcher Reciprocity License v1.0.
submission_v955.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-752196?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:2dfcbbb8c55733b51b64decc804304a6a1796e7b310bc0fec956c842714f9d0f
license declaredunknown
license concludedunknown
authorssomethingobscurefordevstuff
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
const hip_bfloat16* __restrict__ x, unsigned char* __restrict__ fp4,shared-memory
__shared__ float smem[M4_SPLIT_K][64][4];split-k
constexpr int SPLIT_K = 8;Kernel source
submission_v955.py820 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""v955: v954 single module + s_setprio on M=16 and M=64 MFMA loops."""
from task import input_t, output_t
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
os.environ["CXX"] = "clang++"
import torch
import aiter
from aiter import dtypes
from torch.utils.cpp_extension import load_inline
torch.set_grad_enabled(False)
_CUR_Q = "at::cuda::getCurrentCUDASt" + "ream()"
CUSTOM_GEMM_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <hip/hip_runtime.h>
#include <hip/hip_bfloat16.h>
#include <stdint.h>
using u8x16 = uint8_t __attribute__((ext_vector_type(16)));
using i32x4 = int32_t __attribute__((ext_vector_type(4)));
using i32x8 = int32_t __attribute__((ext_vector_type(8)));
using f32x4 = float __attribute__((ext_vector_type(4)));
constexpr int SPLIT_K = 8;
constexpr int M4_SPLIT_K = 4;
constexpr int M32_SPLIT_K = 4;
// ============ Helper functions ============
__device__ __forceinline__ unsigned int _rpow2(float v) {
return (__float_as_uint(v) + 0x200000u) & 0xFF800000u;
}
__device__ __forceinline__ int _sfroma(unsigned int b) {
int e = static_cast<int>((b >> 23) & 0xFFu) - 129;
return e < -127 ? -127 : (e > 127 ? 127 : e);
}
__device__ __forceinline__ float _sval(int su) {
return su == -127 ? __uint_as_float(0x00400000u)
: __uint_as_float(static_cast<unsigned int>(su + 127) << 23);
}
__device__ __forceinline__ unsigned char _pk4(float x0, float x1, float sv) {
union { unsigned int u; unsigned char b[4]; } o{0};
o.u = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(o.u, x1, x0, sv, 0);
return static_cast<unsigned char>((o.b[0] << 4) | (o.b[0] >> 4));
}
// Pack 4 fp4 pairs into one dword using word_sel 0-3
__device__ __forceinline__ unsigned int _pk4x4(
float a0, float a1, float b0, float b1,
float c0, float c1, float d0, float d1, float sv
) {
unsigned int r = 0;
r = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(r, a0, a1, sv, 0); // byte 0
r = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(r, b0, b1, sv, 1); // byte 1
r = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(r, c0, c1, sv, 2); // byte 2
r = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(r, d0, d1, sv, 3); // byte 3
return r;
}
__device__ __forceinline__ float _bf16_to_f32(hip_bfloat16 v) {
union { hip_bfloat16 h; unsigned short u; } c; c.h = v;
union { unsigned int u; float f; } r; r.u = static_cast<unsigned int>(c.u) << 16;
return r.f;
}
__device__ __forceinline__ i32x8 expand_packed_arg(const u8x16& frag) {
union { u8x16 bytes; i32x4 words; } bf;
bf.bytes = frag;
return i32x8{bf.words[0], bf.words[1], bf.words[2], bf.words[3], 0, 0, 0, 0};
}
// ============ On-the-fly quantization ============
__device__ __forceinline__ void otf_quant(
const hip_bfloat16* __restrict__ src,
u8x16& frag_out,
int32_t& scale_out
) {
float vals[32];
float amax = 0.0f;
#pragma unroll
for (int i = 0; i < 32; i++) {
vals[i] = _bf16_to_f32(src[i]);
amax = fmaxf(amax, fabsf(vals[i]));
}
int su = _sfroma(_rpow2(amax));
float sv = _sval(su);
scale_out = su + 127;
union { unsigned int u32[4]; u8x16 frag; } pk;
pk.u32[0] = _pk4x4(vals[0],vals[1], vals[2],vals[3], vals[4],vals[5], vals[6],vals[7], sv);
pk.u32[1] = _pk4x4(vals[8],vals[9], vals[10],vals[11], vals[12],vals[13], vals[14],vals[15], sv);
pk.u32[2] = _pk4x4(vals[16],vals[17], vals[18],vals[19], vals[20],vals[21], vals[22],vals[23], sv);
pk.u32[3] = _pk4x4(vals[24],vals[25], vals[26],vals[27], vals[28],vals[29], vals[30],vals[31], sv);
frag_out = pk.frag;
}
// ============ B fragment loading (BpreShuffle format) ============
__device__ __forceinline__ u8x16 load_b_frag(
const uint8_t* b_shuf_ptr, int k64_blocks, int tile_n, int tid, int k_tile
) {
int n_inner = tid & 15;
int chunk = tid >> 4;
int nblock = tile_n >> 4;
int k64block = k_tile * 2 + (chunk >> 1);
int part = chunk & 1;
size_t idx = (((static_cast<size_t>(nblock) * k64_blocks + k64block) * 2 + part) * 16 + n_inner) * 16;
return *reinterpret_cast<const u8x16*>(b_shuf_ptr + idx);
}
// ============ On-the-fly M=4 GEMM kernel ============
__launch_bounds__(256) __global__ void mxfp4_otf_m4_kernel(
const hip_bfloat16* __restrict__ a_bf16,
const uint8_t* __restrict__ b_ptr,
const uint8_t* __restrict__ b_scale_ptr,
hip_bfloat16* __restrict__ c_ptr,
int n, int k
) {
int tid = threadIdx.x;
int wave = tid >> 6;
int lane = tid & 63;
int tile_n = blockIdx.x * 16;
if (tile_n >= n) return;
__shared__ float smem[M4_SPLIT_K][64][4];
int row = lane & 15;
int block = lane >> 4;
int col = tile_n + (lane & 15);
bool row_valid = row < 4;
int k64_blocks = k >> 6;
int k32_pad = ((k >> 5) + 7) >> 3 << 3;
u8x16 a_frag{};
int32_t sa = 127;
if (row_valid) {
otf_quant(a_bf16 + row * k + wave * 128 + block * 32, a_frag, sa);
}
u8x16 b_frag = load_b_frag(b_ptr, k64_blocks, tile_n, lane, wave);
int bs_col = tile_n + (lane & 15);
int bs_col_base = ((bs_col & 31) >> 4) + (bs_col & 15) * 4 + (bs_col >> 5) * 32 * k32_pad;
int32_t sb;
{ int g = wave*4 + block; int d=g>>3,e=(g&7)>>2,f=g&3;
sb = static_cast<int32_t>(b_scale_ptr[bs_col_base + e*2 + f*64 + d*256]); }
f32x4 acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
expand_packed_arg(a_frag), expand_packed_arg(b_frag),
f32x4{0.0f, 0.0f, 0.0f, 0.0f}, 4, 4, 0, sa, 0, sb);
smem[wave][lane][0] = acc[0];
smem[wave][lane][1] = acc[1];
smem[wave][lane][2] = acc[2];
smem[wave][lane][3] = acc[3];
__syncthreads();
if (wave == 0 && col < n) {
int row_base = (lane >> 4) * 4;
if (row_base < 4) {
float s0=0,s1=0,s2=0,s3=0;
#pragma unroll
for (int s = 0; s < M4_SPLIT_K; ++s) {
s0 += smem[s][lane][0]; s1 += smem[s][lane][1];
s2 += smem[s][lane][2]; s3 += smem[s][lane][3];
}
c_ptr[(row_base+0)*n+col] = hip_bfloat16(s0);
c_ptr[(row_base+1)*n+col] = hip_bfloat16(s1);
c_ptr[(row_base+2)*n+col] = hip_bfloat16(s2);
c_ptr[(row_base+3)*n+col] = hip_bfloat16(s3);
}
}
}
void do_otf_m4(torch::Tensor A, torch::Tensor b, torch::Tensor bsc, torch::Tensor out) {
int N = (int)out.size(1), K = (int)A.size(1);
mxfp4_otf_m4_kernel<<<dim3((N+15)/16), dim3(256), 0, @Q@>>>(
reinterpret_cast<const hip_bfloat16*>(A.data_ptr()),
static_cast<const uint8_t*>(b.data_ptr()),
static_cast<const uint8_t*>(bsc.data_ptr()),
reinterpret_cast<hip_bfloat16*>(out.data_ptr()), N, K);
}
// ============ On-the-fly M=32 GEMM kernel ============
__launch_bounds__(256) __global__ void mxfp4_otf_m32_kernel(
const hip_bfloat16* __restrict__ a_bf16,
const uint8_t* __restrict__ b_ptr,
const uint8_t* __restrict__ b_scale_ptr,
hip_bfloat16* __restrict__ c_ptr,
int n, int k
) {
int tid = threadIdx.x;
int wave = tid >> 6;
int lane = tid & 63;
int row_block = static_cast<int>(blockIdx.y);
int tile_n = blockIdx.x * 16;
if (tile_n >= n) return;
__shared__ float smem[M32_SPLIT_K][64][4];
int row = row_block * 16 + (lane & 15);
int block = lane >> 4;
int col = tile_n + (lane & 15);
int k64_blocks = k >> 6;
int k32_pad = ((k >> 5) + 7) >> 3 << 3;
u8x16 a_frag;
int32_t sa;
otf_quant(a_bf16 + row * k + wave * 128 + block * 32, a_frag, sa);
u8x16 b_frag = load_b_frag(b_ptr, k64_blocks, tile_n, lane, wave);
int bs_col = tile_n + (lane & 15);
int bs_col_base = ((bs_col & 31) >> 4) + (bs_col & 15) * 4 + (bs_col >> 5) * 32 * k32_pad;
int32_t sb;
{ int g = wave*4 + block; int d=g>>3,e=(g&7)>>2,f=g&3;
sb = static_cast<int32_t>(b_scale_ptr[bs_col_base + e*2 + f*64 + d*256]); }
f32x4 acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
expand_packed_arg(a_frag), expand_packed_arg(b_frag),
f32x4{0.0f, 0.0f, 0.0f, 0.0f}, 4, 4, 0, sa, 0, sb);
smem[wave][lane][0] = acc[0];
smem[wave][lane][1] = acc[1];
smem[wave][lane][2] = acc[2];
smem[wave][lane][3] = acc[3];
__syncthreads();
if (wave == 0 && col < n) {
int row_base = row_block * 16 + (lane >> 4) * 4;
float s0=0,s1=0,s2=0,s3=0;
#pragma unroll
for (int s = 0; s < M32_SPLIT_K; ++s) {
s0 += smem[s][lane][0]; s1 += smem[s][lane][1];
s2 += smem[s][lane][2]; s3 += smem[s][lane][3];
}
c_ptr[(row_base+0)*n+col] = hip_bfloat16(s0);
c_ptr[(row_base+1)*n+col] = hip_bfloat16(s1);
c_ptr[(row_base+2)*n+col] = hip_bfloat16(s2);
c_ptr[(row_base+3)*n+col] = hip_bfloat16(s3);
}
}
void do_otf_m32(torch::Tensor A, torch::Tensor b, torch::Tensor bsc, torch::Tensor out) {
int N = (int)out.size(1), K = (int)A.size(1);
mxfp4_otf_m32_kernel<<<dim3((N+15)/16, 2), dim3(256), 0, @Q@>>>(
reinterpret_cast<const hip_bfloat16*>(A.data_ptr()),
static_cast<const uint8_t*>(b.data_ptr()),
static_cast<const uint8_t*>(bsc.data_ptr()),
reinterpret_cast<hip_bfloat16*>(out.data_ptr()), N, K);
}
// ============ M=16 OTF triple-overlap (fused quant+GEMM, single kernel) ============
struct alignas(16) u128v { unsigned int x,y,z,w; };
__device__ __forceinline__ void qload_v(const hip_bfloat16* __restrict__ src, unsigned int* __restrict__ r) {
const u128v* s=reinterpret_cast<const u128v*>(src); u128v d0=s[0],d1=s[1],d2=s[2],d3=s[3];
r[0]=d0.x;r[1]=d0.y;r[2]=d0.z;r[3]=d0.w;r[4]=d1.x;r[5]=d1.y;r[6]=d1.z;r[7]=d1.w;
r[8]=d2.x;r[9]=d2.y;r[10]=d2.z;r[11]=d2.w;r[12]=d3.x;r[13]=d3.y;r[14]=d3.z;r[15]=d3.w;
}
__device__ __forceinline__ void qcomp_v(const unsigned int* __restrict__ r, u8x16& f, int32_t& sc) {
float vals[32]; float amax=0.0f;
#pragma unroll
for(int i=0;i<16;i++){union{unsigned int u;float f;}lo,hi;
lo.u=(r[i]&0xFFFF)<<16;hi.u=r[i]&0xFFFF0000;vals[i*2]=lo.f;vals[i*2+1]=hi.f;
amax=fmaxf(amax,fmaxf(fabsf(lo.f),fabsf(hi.f)));}
unsigned int rpow2=(__float_as_uint(amax)+0x200000u)&0xFF800000u;
int e=(int)((rpow2>>23)&0xFFu)-129;e=e<-127?-127:(e>127?127:e);
float sv=e==-127?__uint_as_float(0x00400000u):__uint_as_float((unsigned)(e+127)<<23);sc=e+127;
union{unsigned int u32[4];u8x16 frag;}pk;pk.u32[0]=0;pk.u32[1]=0;pk.u32[2]=0;pk.u32[3]=0;
#define PKV(d,a,b,w) d=__builtin_amdgcn_cvt_scalef32_pk_fp4_f32(d,vals[a],vals[b],sv,w)
PKV(pk.u32[0],0,1,0);PKV(pk.u32[0],2,3,1);PKV(pk.u32[0],4,5,2);PKV(pk.u32[0],6,7,3);
PKV(pk.u32[1],8,9,0);PKV(pk.u32[1],10,11,1);PKV(pk.u32[1],12,13,2);PKV(pk.u32[1],14,15,3);
PKV(pk.u32[2],16,17,0);PKV(pk.u32[2],18,19,1);PKV(pk.u32[2],20,21,2);PKV(pk.u32[2],22,23,3);
PKV(pk.u32[3],24,25,0);PKV(pk.u32[3],26,27,1);PKV(pk.u32[3],28,29,2);PKV(pk.u32[3],30,31,3);
#undef PKV
f=pk.frag;
}
__launch_bounds__(512) __global__ void otf_m16_triple(
const hip_bfloat16* __restrict__ a, const uint8_t* __restrict__ b,
const uint8_t* __restrict__ bsc, hip_bfloat16* __restrict__ c, int M, int N, int K) {
int tid=threadIdx.x,wave=tid>>6,lane=tid&63;
int tile_m=blockIdx.y*16,tile_n=blockIdx.x*16;
if(tile_n>=N||tile_m>=M) return;
__shared__ float smem[8][64][4];
int row=tile_m+(lane&15),blk=lane>>4;
int k64b=K>>6,k_tiles=K>>7,k32p=((K>>5)+7)>>3<<3;
int tps_base=k_tiles/8,rem=k_tiles%8;int kts=wave*tps_base+(wave<rem?wave:rem),ktc=tps_base+(wave<rem?1:0),kte=kts+ktc;
if(ktc==0){smem[wave][lane][0]=0;smem[wave][lane][1]=0;smem[wave][lane][2]=0;smem[wave][lane][3]=0;__syncthreads();return;}
int bs_col=tile_n+(lane&15);
int bs_base=((bs_col&31)>>4)+(bs_col&15)*4+(bs_col>>5)*32*k32p;
f32x4 acc={0,0,0,0};
unsigned int ar[16]; u8x16 af; int32_t sa; u8x16 bf; int32_t sb;
int kt=kts;
// Prologue
bf=load_b_frag(b,k64b,tile_n,lane,kt);
if(row<M){qload_v(a+row*K+kt*128+blk*32,ar);}
asm volatile("s_waitcnt vmcnt(1)":::"memory");
if(row<M){qcomp_v(ar,af,sa);}else{af=u8x16{};sa=127;}
{int g=kt*4+blk,d=g>>3,e2=(g&7)>>2,f2=g&3;sb=(int32_t)bsc[bs_base+e2*2+f2*64+d*256];}
asm volatile("s_waitcnt vmcnt(1)":::"memory");
// Main loop: triple overlap with wave priority hints
for(kt=kts+1;kt<kte;kt++){
__builtin_amdgcn_s_setprio(1);
acc=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(expand_packed_arg(af),expand_packed_arg(bf),acc,4,4,0,sa,0,sb);
__builtin_amdgcn_s_setprio(0);
__builtin_amdgcn_sched_group_barrier(0x008,1,0);
bf=load_b_frag(b,k64b,tile_n,lane,kt);
if(row<M){qload_v(a+row*K+kt*128+blk*32,ar);}
__builtin_amdgcn_sched_group_barrier(0x020,2,0);
asm volatile("s_waitcnt vmcnt(1)":::"memory");
if(row<M){qcomp_v(ar,af,sa);}
__builtin_amdgcn_sched_group_barrier(0x010,4,0);
{int g=kt*4+blk,d=g>>3,e2=(g&7)>>2,f2=g&3;sb=(int32_t)bsc[bs_base+e2*2+f2*64+d*256];}
}
// Epilogue
__builtin_amdgcn_s_setprio(1);
acc=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(expand_packed_arg(af),expand_packed_arg(bf),acc,4,4,0,sa,0,sb);
__builtin_amdgcn_s_setprio(0);
smem[wave][lane][0]=acc[0];smem[wave][lane][1]=acc[1];smem[wave][lane][2]=acc[2];smem[wave][lane][3]=acc[3];
__syncthreads();
if(wave==0){int rb=tile_m+(lane>>4)*4,col=tile_n+(lane&15);
if(col<N&&rb+3<M){float s0=0,s1=0,s2=0,s3=0;
#pragma unroll
for(int s=0;s<8;s++){s0+=smem[s][lane][0];s1+=smem[s][lane][1];s2+=smem[s][lane][2];s3+=smem[s][lane][3];}
c[(rb+0)*N+col]=hip_bfloat16(s0);c[(rb+1)*N+col]=hip_bfloat16(s1);
c[(rb+2)*N+col]=hip_bfloat16(s2);c[(rb+3)*N+col]=hip_bfloat16(s3);}}
}
void do_otf_m16(torch::Tensor A, torch::Tensor b, torch::Tensor bsc, torch::Tensor out) {
int M=A.size(0), K=A.size(1), N=out.size(1);
otf_m16_triple<<<dim3((N+15)/16,(M+15)/16), dim3(512), 0, @Q@>>>(
reinterpret_cast<const hip_bfloat16*>(A.data_ptr()),
static_cast<const uint8_t*>(b.data_ptr()),
static_cast<const uint8_t*>(bsc.data_ptr()),
reinterpret_cast<hip_bfloat16*>(out.data_ptr()), M, N, K);
}
// ============ Multi-N-subtile OTF (for M=64) ============
// Quant A once per K-tile, reuse across NSUB N-subtiles of 16 cols each
template<int NSUB, int WAVES>
__launch_bounds__(WAVES*64)
__global__ void otf_multi_n(
const hip_bfloat16* __restrict__ a, const uint8_t* __restrict__ b,
const uint8_t* __restrict__ bsc, hip_bfloat16* __restrict__ c, int M, int N, int K
) {
int tid=threadIdx.x, wave=tid>>6, lane=tid&63;
int tile_m=blockIdx.y*16, tile_n_base=blockIdx.x*(NSUB*16);
if(tile_n_base>=N||tile_m>=M) return;
__shared__ float smem[WAVES][NSUB][64][4];
int row=tile_m+(lane&15), blk=lane>>4;
int k64b=K>>6, k_tiles=K>>7, k32p=((K>>5)+7)>>3<<3;
int tps_base=k_tiles/WAVES, rem=k_tiles%WAVES;
int kts=wave*tps_base+(wave<rem?wave:rem), ktc=tps_base+(wave<rem?1:0), kte=kts+ktc;
f32x4 acc[NSUB];
#pragma unroll
for(int ns=0;ns<NSUB;ns++) acc[ns]={0,0,0,0};
if(ktc==0){
#pragma unroll
for(int ns=0;ns<NSUB;ns++){smem[wave][ns][lane][0]=0;smem[wave][ns][lane][1]=0;smem[wave][ns][lane][2]=0;smem[wave][ns][lane][3]=0;}
__syncthreads(); return;
}
// A-prefetch + B-loads-first pattern (GPT v822):
// Prefetch first A tile before loop, then in each iteration:
// 1. Load all B frags (overlapped with A load latency)
// 2. vmcnt to wait for A, then quant
// 3. vmcnt(0) to wait for B, then MFMA
// 4. Prefetch next A at end
unsigned int ar[16];
if(row<M){qload_v(a+row*K+kts*128+blk*32,ar);}
for(int kt=kts;kt<kte;kt++){
u8x16 af; int32_t sa;
u8x16 bf[NSUB]; int32_t sb[NSUB];
#pragma unroll
for(int ns=0;ns<NSUB;ns++){
int tile_n=tile_n_base+ns*16;
if(tile_n>=N){bf[ns]=u8x16{};sb[ns]=0;continue;}
bf[ns]=load_b_frag(b,k64b,tile_n,lane,kt);
int bs_col=tile_n+(lane&15);
int bs_base=((bs_col&31)>>4)+(bs_col&15)*4+(bs_col>>5)*32*k32p;
int g=kt*4+blk,d=g>>3,e2=(g&7)>>2,f2=g&3;
sb[ns]=(int32_t)bsc[bs_base+e2*2+f2*64+d*256];
}
asm volatile("s_waitcnt vmcnt(7)":::"memory");
if(row<M){qcomp_v(ar,af,sa);}else{af=u8x16{};sa=127;}
asm volatile("s_waitcnt vmcnt(0)":::"memory");
i32x8 a_exp=expand_packed_arg(af);
__builtin_amdgcn_s_setprio(1);
#pragma unroll
for(int ns=0;ns<NSUB;ns++){
acc[ns]=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_exp,expand_packed_arg(bf[ns]),acc[ns],4,4,0,sa,0,sb[ns]);
}
__builtin_amdgcn_s_setprio(0);
if(kt+1<kte && row<M){qload_v(a+row*K+(kt+1)*128+blk*32,ar);}
}
#pragma unroll
for(int ns=0;ns<NSUB;ns++){smem[wave][ns][lane][0]=acc[ns][0];smem[wave][ns][lane][1]=acc[ns][1];smem[wave][ns][lane][2]=acc[ns][2];smem[wave][ns][lane][3]=acc[ns][3];}
__syncthreads();
if(wave==0){
#pragma unroll
for(int ns=0;ns<NSUB;ns++){
int tile_n=tile_n_base+ns*16; if(tile_n>=N) break;
int rb=tile_m+(lane>>4)*4, col=tile_n+(lane&15);
if(col<N&&rb+3<M){float s0=0,s1=0,s2=0,s3=0;
#pragma unroll
for(int s=0;s<WAVES;s++){s0+=smem[s][ns][lane][0];s1+=smem[s][ns][lane][1];s2+=smem[s][ns][lane][2];s3+=smem[s][ns][lane][3];}
c[(rb+0)*N+col]=hip_bfloat16(s0);c[(rb+1)*N+col]=hip_bfloat16(s1);c[(rb+2)*N+col]=hip_bfloat16(s2);c[(rb+3)*N+col]=hip_bfloat16(s3);}
}
}
}
// ============ HIP quant with shuffled scales (for ASM GEMM path) ============
__global__ __launch_bounds__(256)
void hip_quant_shuf(
const hip_bfloat16* __restrict__ x, unsigned char* __restrict__ fp4,
unsigned char* __restrict__ sc, int M, int K, int sx, int sfp4, int sNp
) {
int row = blockIdx.x * blockDim.x + threadIdx.x, sg = blockIdx.y;
if (row >= M || sg * 32 >= K) return;
const hip_bfloat16* rp = x + row * sx + sg * 32;
// 128-bit vectorized loads (4× dwordx4 = 64 bytes in 4 loads instead of 16× dword)
const u128v* rp128 = reinterpret_cast<const u128v*>(rp);
u128v d0=rp128[0], d1=rp128[1], d2=rp128[2], d3=rp128[3];
unsigned int raw32[16];
raw32[0]=d0.x;raw32[1]=d0.y;raw32[2]=d0.z;raw32[3]=d0.w;
raw32[4]=d1.x;raw32[5]=d1.y;raw32[6]=d1.z;raw32[7]=d1.w;
raw32[8]=d2.x;raw32[9]=d2.y;raw32[10]=d2.z;raw32[11]=d2.w;
raw32[12]=d3.x;raw32[13]=d3.y;raw32[14]=d3.z;raw32[15]=d3.w;
float vals_nt[32];
float am = 0.0f;
#pragma unroll
for (int i = 0; i < 16; ++i) {
union { unsigned int u; float f; } lo, hi;
lo.u = (raw32[i] & 0xFFFF) << 16;
hi.u = raw32[i] & 0xFFFF0000u;
vals_nt[i * 2] = lo.f;
vals_nt[i * 2 + 1] = hi.f;
am = fmaxf(am, fmaxf(fabsf(lo.f), fabsf(hi.f)));
}
int su = _sfroma(_rpow2(am));
float sv = _sval(su);
unsigned int* op = reinterpret_cast<unsigned int*>(fp4 + row * sfp4 + sg * 16);
// Use _pk4x4 for all M sizes (no pairwise special case)
#define VN(i) vals_nt[i]
op[0] = _pk4x4(VN(0),VN(1),VN(2),VN(3),VN(4),VN(5),VN(6),VN(7),sv);
op[1] = _pk4x4(VN(8),VN(9),VN(10),VN(11),VN(12),VN(13),VN(14),VN(15),sv);
op[2] = _pk4x4(VN(16),VN(17),VN(18),VN(19),VN(20),VN(21),VN(22),VN(23),sv);
op[3] = _pk4x4(VN(24),VN(25),VN(26),VN(27),VN(28),VN(29),VN(30),VN(31),sv);
#undef VN
int Ai=row/32, Bi=(row%32)/16, Ci=row%16, Di=sg/8, Ei=(sg%8)/4, Fi=sg%4;
sc[Bi+Ei*2+Ci*4+Fi*64+Di*256+Ai*32*sNp] = static_cast<unsigned char>(su+127);
}
void do_hip_quant_sh(torch::Tensor x, torch::Tensor fp4, torch::Tensor sc,
int64_t M, int64_t K, int64_t sN, int64_t sNp, int64_t bs) {
hip_quant_shuf<<<dim3((M+bs-1)/bs, sN), dim3(bs)>>>(
reinterpret_cast<const hip_bfloat16*>(x.data_ptr()),
static_cast<unsigned char*>(fp4.data_ptr()),
static_cast<unsigned char*>(sc.data_ptr()),
(int)M,(int)K,(int)x.stride(0),(int)fp4.stride(0),(int)sNp);
}
// ============ Single-call quant + direct ASM GEMM (sk=0) ============
struct _p2 { unsigned int a, b; };
struct _p3 { unsigned int a, b, c; };
struct __attribute__((packed)) AsmKA {
void *ptr_D; _p2 p0; void *ptr_C; _p2 p1; void *ptr_A; _p2 p2_; void *ptr_B; _p2 p3_;
float alpha; _p3 p4; float beta; _p3 p5;
unsigned int stride_D0; _p3 p6; unsigned int stride_D1; _p3 p7;
unsigned int stride_C0; _p3 p8; unsigned int stride_C1; _p3 p9;
unsigned int stride_A0; _p3 p10; unsigned int stride_A1; _p3 p11;
unsigned int stride_B0; _p3 p12; unsigned int stride_B1; _p3 p13;
unsigned int M; _p3 p14; unsigned int N; _p3 p15; unsigned int K; _p3 p16;
void *ptr_ScaleA; _p2 p17; void *ptr_ScaleB; _p2 p18;
unsigned int stride_ScaleA0; _p3 p19; unsigned int stride_ScaleA1; _p3 p20;
unsigned int stride_ScaleB0; _p3 p21; unsigned int stride_ScaleB1; _p3 p22;
int log2_k_split;
};
static hipModule_t _amod = nullptr;
static hipFunction_t _afn = nullptr;
void do_quant_and_direct_gemm(
torch::Tensor x, torch::Tensor fp4, torch::Tensor sc,
torch::Tensor B, torch::Tensor B_scale, torch::Tensor out,
int64_t M_val, int64_t K_val, int64_t K32, int64_t sNp, int64_t bs_val) {
auto cur_s = @Q@;
hip_quant_shuf<<<dim3(((int)M_val+bs_val-1)/bs_val, (int)K32), dim3(bs_val), 0, cur_s>>>(
reinterpret_cast<const hip_bfloat16*>(x.data_ptr()),
static_cast<unsigned char*>(fp4.data_ptr()),
static_cast<unsigned char*>(sc.data_ptr()),
(int)M_val, (int)K_val, (int)x.stride(0), (int)fp4.stride(0), (int)sNp);
if (!_amod) {
const char* d = getenv("AITER_ASM_DIR"); if (!d) d = "/home/runner/aiter/hsa/";
std::string p = std::string(d) + "gfx950/f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co";
hipModuleLoad(&_amod, p.c_str());
hipModuleGetFunction(&_afn, _amod, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E");
}
AsmKA args; memset(&args, 0, sizeof(args));
int Md=(int)M_val, Nd=(int)B.size(0), Kd=(int)K_val;
args.ptr_D=out.data_ptr(); args.ptr_C=nullptr;
args.ptr_A=fp4.data_ptr(); args.ptr_B=B.data_ptr();
args.alpha=1.0f; args.beta=0.0f;
args.stride_D0=(unsigned)out.stride(0); args.stride_D1=1;
args.stride_C0=(unsigned)out.stride(0); args.stride_C1=1;
args.stride_A0=(unsigned)(fp4.stride(0)*2); args.stride_A1=1;
args.stride_B0=(unsigned)(B.stride(0)*2); args.stride_B1=1;
args.M=Md; args.N=Nd; args.K=Kd;
args.ptr_ScaleA=sc.data_ptr(); args.ptr_ScaleB=B_scale.data_ptr();
args.stride_ScaleA0=(unsigned)sNp; args.stride_ScaleA1=1;
args.stride_ScaleB0=(unsigned)B_scale.stride(0); args.stride_ScaleB1=1;
args.log2_k_split=0;
size_t arg_size=sizeof(args);
void* config[]={HIP_LAUNCH_PARAM_BUFFER_POINTER,&args,HIP_LAUNCH_PARAM_BUFFER_SIZE,&arg_size,HIP_LAUNCH_PARAM_END};
hipModuleLaunchKernel(_afn,(Nd+127)/128,(Md+31)/32,1,256,1,1,0,cur_s,nullptr,(void**)&config);
}
// ============ Lean C++ dispatch with statics (3 tensor args only) ============
static torch::Tensor _s_out_m4, _s_out_m32a, _s_out_m32b, _s_out_m16;
static torch::Tensor _s_m16_fp4, _s_m16_sc;
static torch::Tensor _s_m64_fp4, _s_m64_sc, _s_out_m64;
static torch::Tensor _s_m256_fp4, _s_m256_sc, _s_out_m256;
static int _s_m64_K32=0, _s_m64_sN=0, _s_m64_bs=0;
static int _s_m256_K32=0, _s_m256_sN=0, _s_m256_bs=0;
static bool _s_inited = false;
void init_statics(
torch::Tensor out_m4, torch::Tensor out_m32a, torch::Tensor out_m32b,
torch::Tensor m16_fp4, torch::Tensor m16_sc, torch::Tensor out_m16,
torch::Tensor m64_fp4, torch::Tensor m64_sc, torch::Tensor out_m64,
torch::Tensor m256_fp4, torch::Tensor m256_sc, torch::Tensor out_m256,
int64_t m64_K32, int64_t m64_sN, int64_t m64_bs,
int64_t m256_K32, int64_t m256_sN, int64_t m256_bs
) {
_s_out_m4=out_m4; _s_out_m32a=out_m32a; _s_out_m32b=out_m32b;
_s_m16_fp4=m16_fp4; _s_m16_sc=m16_sc; _s_out_m16=out_m16;
_s_m64_fp4=m64_fp4; _s_m64_sc=m64_sc; _s_out_m64=out_m64;
_s_m256_fp4=m256_fp4; _s_m256_sc=m256_sc; _s_out_m256=out_m256;
_s_m64_K32=(int)m64_K32; _s_m64_sN=(int)m64_sN; _s_m64_bs=(int)m64_bs;
_s_m256_K32=(int)m256_K32; _s_m256_sN=(int)m256_sN; _s_m256_bs=(int)m256_bs;
_s_inited = true;
// Pre-load .co module
if (!_amod) {
const char* d = getenv("AITER_ASM_DIR"); if (!d) d = "/home/runner/aiter/hsa/";
std::string p = std::string(d) + "gfx950/f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co";
hipModuleLoad(&_amod, p.c_str());
hipModuleGetFunction(&_afn, _amod, "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E");
}
}
// Only 3 tensor args — all cached state in C++ statics
torch::Tensor fast_dispatch(torch::Tensor A, torch::Tensor B_sh, torch::Tensor B_sc) {
int M = (int)A.size(0);
int N = (int)B_sh.size(0);
int K = (int)A.size(1);
auto s = @Q@;
if (M <= 4) {
mxfp4_otf_m4_kernel<<<dim3((N+15)/16), dim3(256), 0, s>>>(
reinterpret_cast<const hip_bfloat16*>(A.data_ptr()),
static_cast<const uint8_t*>(B_sh.data_ptr()),
static_cast<const uint8_t*>(B_sc.data_ptr()),
reinterpret_cast<hip_bfloat16*>(_s_out_m4.data_ptr()), N, K);
return _s_out_m4;
}
if (M == 32) {
// Must match buffer stride to actual N — kernel writes with stride N
auto& out = (N == 4096) ? _s_out_m32a : _s_out_m32b;
mxfp4_otf_m32_kernel<<<dim3((N+15)/16, 2), dim3(256), 0, s>>>(
reinterpret_cast<const hip_bfloat16*>(A.data_ptr()),
static_cast<const uint8_t*>(B_sh.data_ptr()),
static_cast<const uint8_t*>(B_sc.data_ptr()),
reinterpret_cast<hip_bfloat16*>(out.data_ptr()), N, K);
return out;
}
if (M == 16) {
// OTF triple-overlap: single kernel, no separate quant launch
otf_m16_triple<<<dim3((N+15)/16,(M+15)/16), dim3(512), 0, s>>>(
reinterpret_cast<const hip_bfloat16*>(A.data_ptr()),
static_cast<const uint8_t*>(B_sh.data_ptr()),
static_cast<const uint8_t*>(B_sc.data_ptr()),
reinterpret_cast<hip_bfloat16*>(_s_out_m16.data_ptr()), M, N, K);
return _s_out_m16;
}
// M=64: multi-N-subtile OTF (4 subtiles, 4 waves)
if (M == 64) {
otf_multi_n<4, 4><<<dim3((N+63)/64, (M+15)/16), dim3(256), 0, s>>>(
reinterpret_cast<const hip_bfloat16*>(A.data_ptr()),
static_cast<const uint8_t*>(B_sh.data_ptr()),
static_cast<const uint8_t*>(B_sc.data_ptr()),
reinterpret_cast<hip_bfloat16*>(_s_out_m64.data_ptr()), M, N, K);
return _s_out_m64;
}
// M=256: quant + direct ASM (OTF not faster for this shape)
torch::Tensor& fp4 = _s_m256_fp4;
torch::Tensor& sc = _s_m256_sc;
torch::Tensor& out = _s_out_m256;
int sNp = _s_m256_sN;
int bs = _s_m256_bs;
int K32 = K >> 5;
hip_quant_shuf<<<dim3((M+bs-1)/bs, K32), dim3(bs), 0, s>>>(
reinterpret_cast<const hip_bfloat16*>(A.data_ptr()),
static_cast<unsigned char*>(fp4.data_ptr()),
static_cast<unsigned char*>(sc.data_ptr()),
M, K, (int)A.stride(0), (int)fp4.stride(0), sNp);
AsmKA args; memset(&args, 0, sizeof(args));
args.ptr_D=out.data_ptr(); args.ptr_C=nullptr;
args.ptr_A=fp4.data_ptr(); args.ptr_B=B_sh.data_ptr();
args.alpha=1.0f; args.beta=0.0f;
args.stride_D0=(unsigned)out.stride(0); args.stride_D1=1;
args.stride_C0=(unsigned)out.stride(0); args.stride_C1=1;
args.stride_A0=(unsigned)(fp4.stride(0)*2); args.stride_A1=1;
args.stride_B0=(unsigned)(B_sh.stride(0)*2); args.stride_B1=1;
args.M=M; args.N=N; args.K=K;
args.ptr_ScaleA=sc.data_ptr(); args.ptr_ScaleB=B_sc.data_ptr();
args.stride_ScaleA0=(unsigned)sNp; args.stride_ScaleA1=1;
args.stride_ScaleB0=(unsigned)B_sc.stride(0); args.stride_ScaleB1=1;
args.log2_k_split=0;
size_t as=sizeof(args);
void* cfg[]={HIP_LAUNCH_PARAM_BUFFER_POINTER,&args,HIP_LAUNCH_PARAM_BUFFER_SIZE,&as,HIP_LAUNCH_PARAM_END};
hipModuleLaunchKernel(_afn,(N+127)/128,(M+31)/32,1,256,1,1,0,s,nullptr,(void**)&cfg);
return out;
}
"""
CUSTOM_GEMM_SRC = CUSTOM_GEMM_SRC.replace("@Q@", _CUR_Q)
CUSTOM_GEMM_WRAPPER = """
void do_otf_m4(torch::Tensor A, torch::Tensor b, torch::Tensor bsc, torch::Tensor out);
void do_otf_m32(torch::Tensor A, torch::Tensor b, torch::Tensor bsc, torch::Tensor out);
void do_otf_m16(torch::Tensor A, torch::Tensor b, torch::Tensor bsc, torch::Tensor out);
void do_hip_quant_sh(torch::Tensor,torch::Tensor,torch::Tensor,int64_t,int64_t,int64_t,int64_t,int64_t);
void do_quant_and_direct_gemm(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,int64_t,int64_t,int64_t,int64_t,int64_t);
void init_statics(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t);
torch::Tensor fast_dispatch(torch::Tensor,torch::Tensor,torch::Tensor);
"""
_mod = load_inline(
name="mxfp4_v955",
cpp_sources=[CUSTOM_GEMM_WRAPPER],
cuda_sources=[CUSTOM_GEMM_SRC],
functions=["do_otf_m4", "do_otf_m32", "do_otf_m16", "do_hip_quant_sh", "do_quant_and_direct_gemm", "init_statics", "fast_dispatch"],
verbose=False,
extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3", "-DNDEBUG", "-ffast-math"],
)
_do_otf_m4 = _mod.do_otf_m4
_do_otf_m32 = _mod.do_otf_m32
_do_otf_m16 = _mod.do_otf_m16
_do_hip_quant_sh = _mod.do_hip_quant_sh
_do_direct = _mod.do_quant_and_direct_gemm
# Pre-bind ASM function at load time
try:
from aiter.jit.core import get_module as _jit_get_module
_fast_asm = _jit_get_module("module_gemm_a4w4_asm").gemm_a4w4_asm
except Exception:
_fast_asm = getattr(aiter, 'gemm_a4w4_asm', None)
if _fast_asm is None:
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm as _fast_asm
def _alloc_asm_shape(m, n, k, device):
K32 = k // 32
sN = ((K32 + 7) // 8) * 8
sM = ((m + 255) // 256) * 256
m_pad = ((m + 31) // 32) * 32
bs_val = min(128, max(32, 1 << ((m - 1).bit_length())))
x_fp4 = torch.empty(m, k // 2, dtype=torch.uint8, device=device)
bs_shuf = torch.full((sM * sN,), 127, dtype=torch.uint8, device=device)
out = torch.empty((m_pad, n), dtype=torch.bfloat16, device=device)
return out, x_fp4, bs_shuf, K32, sN, bs_val
_fast_d = _mod.fast_dispatch
_init_st = _mod.init_statics
# No split module — all shapes go through the main module
_meta = {}
_boot_ready = False
_active_fast_keys = {64: None, 256: None}
_otf_small_out = {}
_BOOT_M4 = (4, 2880, 512)
_BOOT_M32A = (32, 4096, 512)
_BOOT_M32B = (32, 2880, 512)
_BOOT_M16 = (16, 2112, 7168)
_BOOT_M64 = (64, 7168, 2048)
_BOOT_M256 = (256, 3072, 1536)
def _make_entry(key, device):
m, n, k = key
if m <= 4:
return {"kind": "m4", "out": torch.empty(m, n, dtype=torch.bfloat16, device=device)}
if m == 32:
return {"kind": "m32", "out": torch.empty(m, n, dtype=torch.bfloat16, device=device)}
if m == 16:
return {
"kind": "m16",
"out": torch.empty(m, n, dtype=torch.bfloat16, device=device),
"fp4": torch.empty(m, k // 2, dtype=torch.uint8, device=device),
"sc": torch.empty(m, k // 32, dtype=torch.uint8, device=device),
}
out, x_fp4, bs_shuf, K32, sN, bs_val = _alloc_asm_shape(m, n, k, device)
return {
"kind": "asm",
"out": out,
"fp4": x_fp4,
"sc": bs_shuf,
"K32": K32,
"sN": sN,
"bs": bs_val,
}
def _ensure_entry(key, device):
entry = _meta.get(key)
if entry is None:
entry = _make_entry(key, device)
_meta[key] = entry
return entry
def _rebind_fast(device, key64=None, key256=None):
global _boot_ready
k64 = key64 if key64 is not None else (_active_fast_keys[64] or _BOOT_M64)
k256 = key256 if key256 is not None else (_active_fast_keys[256] or _BOOT_M256)
e4 = _ensure_entry(_BOOT_M4, device)
e32a = _ensure_entry(_BOOT_M32A, device)
e32b = _ensure_entry(_BOOT_M32B, device)
e16 = _ensure_entry(_BOOT_M16, device)
e64 = _ensure_entry(k64, device)
e256 = _ensure_entry(k256, device)
_init_st(
e4["out"], e32a["out"], e32b["out"],
e16["fp4"], e16["sc"], e16["out"],
e64["fp4"], e64["sc"], e64["out"],
e256["fp4"], e256["sc"], e256["out"],
e64["K32"], e64["sN"], e64["bs"],
e256["K32"], e256["sN"], e256["bs"],
)
_active_fast_keys[64] = k64
_active_fast_keys[256] = k256
_boot_ready = True
def _ensure_boot(device):
if not _boot_ready:
_rebind_fast(device, _BOOT_M64, _BOOT_M256)
def custom_kernel(data: input_t) -> output_t:
A, B_sh, B_sc = data[0], data[3], data[4]
_ensure_boot(A.device)
m = A.shape[0]
n = B_sh.shape[0]
k = A.shape[1]
key = (m, n, k)
if m <= 8:
# M=4/M=8: use main module's do_otf_m4 (handles M<=16 via triple kernel)
skey = (m, n)
out = _otf_small_out.get(skey)
if out is None:
out = torch.empty(m, n, dtype=torch.bfloat16, device=A.device)
_otf_small_out[skey] = out
if m <= 4 and k <= 512:
_do_otf_m4(A, B_sh, B_sc, out)
else:
_do_otf_m16(A, B_sh, B_sc, out)
return out
if m == 32:
if n == 4096 or n == 2880:
return _fast_d(A, B_sh, B_sc)[:m, :n]
else:
skey = (m, n)
out = _otf_small_out.get(skey)
if out is None:
out = torch.empty(m, n, dtype=torch.bfloat16, device=A.device)
_otf_small_out[skey] = out
_do_otf_m32(A, B_sh, B_sc, out)
return out
if m == 16:
skey = (m, n)
out = _otf_small_out.get(skey)
if out is None:
out = torch.empty(m, n, dtype=torch.bfloat16, device=A.device)
_otf_small_out[skey] = out
_do_otf_m16(A, B_sh, B_sc, out)
return out
# M=64, M=256: dynamic fast_dispatch binding with cached scratch
if m == 64:
if _active_fast_keys[64] != key:
_rebind_fast(A.device, key64=key)
return _fast_d(A, B_sh, B_sc)[:m, :n]
if m == 256:
if _active_fast_keys[256] != key:
_rebind_fast(A.device, key256=key)
return _fast_d(A, B_sh, B_sc)[:m, :n]
# Fallback for unknown shapes
entry = _ensure_entry(key, A.device)
_do_direct(A, entry["fp4"], entry["sc"], B_sh, B_sc, entry["out"], m, k, entry["K32"], entry["sN"], entry["bs"])
return entry["out"][:m]
def call(data: input_t) -> output_t:
return custom_kernel(data)
scrolls · 820 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