submission 696489
bill_97933 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 539 lines, June 9 Researcher Reciprocity License v1.0.
v594_triton_bn32_s5.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-696489?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:e14f231602ee88c75c4ad5914294b030df001c67b582310b113cc3a64fa81846
license declaredunknown
license concludedunknown
authorsbill_97933
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
k_scale = K_packed // 16 # K=1024 fp4 → 32 scale groupsnum-warps = 4
num_warps=4, num_stages=2, waves_per_eu=4, matrix_instr_nonkdim=16,shared-memory
__shared__ float lds[4][8][256];split-k
- S2 (M=16, K=7168): ext_splitk ks=16 → 12.60µsstages = 2
num_warps=4, num_stages=2, waves_per_eu=4, matrix_instr_nonkdim=16,tile-k = 1024
v594_triton_bn32_s5: v590 but S5 Triton uses AITer's tuned config (BN=32, BK=1024)tile-m = 32
SPLITK_BLOCK_SIZE=2*K_packed, BLOCK_SIZE_M=32, BLOCK_SIZE_N=32,tile-n = 32
v594_triton_bn32_s5: v590 but S5 Triton uses AITer's tuned config (BN=32, BK=1024)vector-width = uint4
uint4 a0,a1,a2,a3;Kernel source
v594_triton_bn32_s5.py539 lines
"""
v594_triton_bn32_s5: v590 but S5 Triton uses AITer's tuned config (BN=32, BK=1024)
- S5 Triton: BN=16→32, BK=512→1024, warps=2→4, GROUP_SIZE_M=4→2
448 CTAs (1.75 waves) vs v590's 896 CTAs (3.5 waves)
Matches gfx950-GEMM-AFP4WFP4-N=7168-K=2048.json M_LEQ_64 config
- Hypothesis: fewer k-iters (2 vs 4) + larger BN reduces overhead per CTA
per-shape dispatch (expected):
- S1 (M=4): gemm_fused_16x128 → 6.12µs
- S2 (M=16, K=7168): ext_splitk ks=16 → 12.60µs
- S3 (M=32, N=4096, K=512): gemm_fused_16x16 → 6.73µs
- S4 (M=32, N=2880, K=512): gemm_fused_16x16 → 6.65µs
- S5 (M=64): quant_raw_v2 + Triton BN=32 → ???µs (was 12.9µs with BN=16)
- S6 (M=256): quant_a + ASM GEMM 32x128 → 12.5µs
"""
import os
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
os.environ.setdefault("CXX", "clang++")
import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
import math
SCALE_GROUP_SIZE = 32
_HIP_SRC = r"""
#include <hip/hip_runtime.h>
#include <hip/hip_ext_ocp.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <torch/extension.h>
#include <cstdint>
typedef int __attribute__((ext_vector_type(8))) i32x8;
typedef float __attribute__((ext_vector_type(4))) f32x4;
typedef __bf16 bf16v2_t __attribute__((ext_vector_type(2)));
__device__ __forceinline__ int pack8_hw(const uint16_t* src, float hs) {
unsigned int d = 0;
d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(d, *reinterpret_cast<const bf16v2_t*>(&src[0]), hs, 0);
d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(d, *reinterpret_cast<const bf16v2_t*>(&src[2]), hs, 1);
d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(d, *reinterpret_cast<const bf16v2_t*>(&src[4]), hs, 2);
d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(d, *reinterpret_cast<const bf16v2_t*>(&src[6]), hs, 3);
return (int)d;
}
__device__ __forceinline__ void quant_32(const uint16_t* ap, int out[4], int32_t& spk) {
uint16_t mx = 0;
#pragma unroll
for (int j=0;j<32;j++) mx=max(mx,(uint16_t)(ap[j]&0x7FFF));
uint32_t au=(((uint32_t)mx<<16)+0x200000u)&0xFF800000u;
int ef=(au>>23u)&0xFFu;
int su=(au==0u)?-127:max(-127,min(127,ef-127-2));
float hs=(su>=-126)?__uint_as_float((uint32_t)(su+127)<<23):0.0f;
#pragma unroll
for (int j=0;j<4;j++) out[j]=pack8_hw(&ap[j*8],hs);
spk=(int32_t)(uint8_t)(su+127);
}
__device__ __forceinline__ int32_t load_b_scale_shuf(const uint8_t* Bsc,int ng,int bb,int SNG,int N) {
if (ng>=N) return 127;
int ifl=(ng/32)*(SNG*256)+(bb/8)*256+(bb%4)*64+(ng%16)*4+((bb%8)/4)*2+(ng%32)/16;
return (int32_t)Bsc[ifl];
}
// ══════ 16×128 fused kernel ══════
__global__ void __launch_bounds__(256,1)
gemm_fused_16x128_kernel(
const uint16_t* __restrict__ A, const uint8_t* __restrict__ Bq,
const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C,
int M, int N, int K, int strA, int strBq, int BscSN)
{
const int mt=blockIdx.x, nt=blockIdx.y;
const int wid=threadIdx.x>>6, lid=threadIdx.x&63;
const int t_row=lid&15, kpart=lid>>4;
const int SNG=BscSN>>3;
const int m_row=mt*16+t_row, a_row_off=m_row*strA;
const int n_base=nt*128;
int total_steps=K>>7, steps_per_warp=(total_steps+3)>>2;
int k_start=wid*steps_per_warp*128, k_end=min(k_start+steps_per_warp*128,K);
f32x4 c_acc[8];
#pragma unroll
for (int ni=0;ni<8;ni++) c_acc[ni]={0,0,0,0};
for (int kb=k_start;kb<k_end;kb+=128) {
int k_off=kb+kpart*32, bk=(kb>>1)+kpart*16, b_scale_k=(kb>>5)+kpart;
uint4 a0,a1,a2,a3;
if (m_row<M&&k_off+31<K) {
const uint4*src=reinterpret_cast<const uint4*>(&A[a_row_off+k_off]);
a0=src[0];a1=src[1];a2=src[2];a3=src[3];
} else {a0={0,0,0,0};a1={0,0,0,0};a2={0,0,0,0};a3={0,0,0,0};}
uint16_t a_local[32]; uint4*adst=reinterpret_cast<uint4*>(a_local);
adst[0]=a0;adst[1]=a1;adst[2]=a2;adst[3]=a3;
int a_i32[4]; int32_t a_spk;
quant_32(a_local,a_i32,a_spk);
i32x8 am={a_i32[0],a_i32[1],a_i32[2],a_i32[3],0,0,0,0};
#pragma unroll
for (int ni=0;ni<8;ni++) {
int b_ng=n_base+ni*16+t_row;
int b_i32[4];
if (b_ng<N&&bk+15<(K>>1)) {
uint4 bd=*reinterpret_cast<const uint4*>(&Bq[b_ng*strBq+bk]);
b_i32[0]=((int*)&bd)[0];b_i32[1]=((int*)&bd)[1];
b_i32[2]=((int*)&bd)[2];b_i32[3]=((int*)&bd)[3];
} else {b_i32[0]=0;b_i32[1]=0;b_i32[2]=0;b_i32[3]=0;}
int32_t b_spk=load_b_scale_shuf(Bsc,b_ng,b_scale_k,SNG,N);
i32x8 bm={b_i32[0],b_i32[1],b_i32[2],b_i32[3],0,0,0,0};
c_acc[ni]=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(am,bm,c_acc[ni],4,4,0,a_spk,0,b_spk);
}
}
__shared__ float lds[4][8][256];
#pragma unroll
for (int ni=0;ni<8;ni++)
#pragma unroll
for (int j=0;j<4;j++) lds[wid][ni][lid*4+j]=c_acc[ni][j];
__syncthreads();
#pragma unroll
for (int ni=0;ni<8;ni++) {
#pragma unroll
for (int j=0;j<4;j++) {
float sum=lds[0][ni][lid*4+j]+lds[1][ni][lid*4+j]+lds[2][ni][lid*4+j]+lds[3][ni][lid*4+j];
int mo=mt*16+kpart*4+j, no=n_base+ni*16+t_row;
if (mo<M&&no<N) {
uint32_t fp=__float_as_uint(sum); fp+=0x7FFFu+((fp>>16)&1u);
C[mo*N+no]=(uint16_t)(fp>>16u);
}
}
}
}
// ══════ 16×16 fused kernel ══════
__global__ void __launch_bounds__(256,3)
gemm_fused_16x16_kernel(
const uint16_t* __restrict__ A, const uint8_t* __restrict__ Bq,
const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C,
int M, int N, int K, int strA, int strBq, int BscSN)
{
const int mt=blockIdx.x,nt=blockIdx.y;
const int wid=threadIdx.x>>6,lid=threadIdx.x&63;
const int t_row=lid&15,kpart=lid>>4,SNG=BscSN>>3;
const int m_row=mt*16+t_row,a_row_off=m_row*strA,b_ng=nt*16+t_row;
int total_steps=K>>7,steps_per_warp=(total_steps+3)>>2;
int k_start=wid*steps_per_warp*128,k_end=min(k_start+steps_per_warp*128,K);
f32x4 c_acc={0,0,0,0};
for (int kb=k_start;kb<k_end;kb+=128) {
int k_off=kb+kpart*32,bk=(kb>>1)+kpart*16;
uint4 a0,a1,a2,a3;
if (m_row<M&&k_off+31<K){const uint4*s=reinterpret_cast<const uint4*>(&A[a_row_off+k_off]);a0=s[0];a1=s[1];a2=s[2];a3=s[3];}
else{a0={0,0,0,0};a1={0,0,0,0};a2={0,0,0,0};a3={0,0,0,0};}
int b_i32[4];
if (b_ng<N&&bk+15<(K>>1)){uint4 bd=*reinterpret_cast<const uint4*>(&Bq[b_ng*strBq+bk]);b_i32[0]=((int*)&bd)[0];b_i32[1]=((int*)&bd)[1];b_i32[2]=((int*)&bd)[2];b_i32[3]=((int*)&bd)[3];}
else{b_i32[0]=0;b_i32[1]=0;b_i32[2]=0;b_i32[3]=0;}
int32_t b_spk=load_b_scale_shuf(Bsc,b_ng,(kb>>5)+kpart,SNG,N);
uint16_t a_local[32];uint4*d=reinterpret_cast<uint4*>(a_local);d[0]=a0;d[1]=a1;d[2]=a2;d[3]=a3;
int a_i32[4];int32_t a_spk;quant_32(a_local,a_i32,a_spk);
i32x8 am={a_i32[0],a_i32[1],a_i32[2],a_i32[3],0,0,0,0};
i32x8 bm={b_i32[0],b_i32[1],b_i32[2],b_i32[3],0,0,0,0};
c_acc=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(am,bm,c_acc,4,4,0,a_spk,0,b_spk);
}
__shared__ float r16[4][64*4];
#pragma unroll
for (int i=0;i<4;i++) r16[wid][lid*4+i]=c_acc[i];
__syncthreads();
if (wid==0) {
for (int j=0;j<4;j++) {
float sum=r16[0][lid*4+j]+r16[1][lid*4+j]+r16[2][lid*4+j]+r16[3][lid*4+j];
int mo=mt*16+kpart*4+j,no=nt*16+t_row;
if (mo<M&&no<N){uint32_t fp=__float_as_uint(sum);fp+=0x7FFFu+((fp>>16)&1u);C[mo*N+no]=(uint16_t)(fp>>16u);}
}
}
}
// ══════ ext splitk for S2 ══════
__global__ void __launch_bounds__(64,8)
gemm_ext_splitk_16x16_bsh_kernel(
const uint16_t* __restrict__ A,const uint8_t* __restrict__ Bsh,
const uint8_t* __restrict__ Bsc,float* __restrict__ Cfp32,
int M,int N,int K,int strA,int BscSN,int kper)
{
const int mt=blockIdx.x,nt=blockIdx.y,ks=blockIdx.z;
const int lid=threadIdx.x,t_row=lid&15,kpart=lid>>4,SNG=BscSN>>3;
const int m_row=mt*16+t_row,a_row_off=m_row*strA,b_ng=nt*16+t_row;
int kst=(ks*kper/128)*128,ken=min(((ks+1)*kper+127)/128*128,K);
if (kst>=ken) return;
const int nkt=K/32;
const int bsh_base=(b_ng<N)?((b_ng/16)*nkt*256+(b_ng%16)*16):0;
f32x4 c_acc={0,0,0,0};
for (int kb=kst;kb<ken;kb+=128) {
int k_off=kb+kpart*32;
uint16_t a_local[32];
if (m_row<M&&k_off+31<K){const uint4*s=reinterpret_cast<const uint4*>(&A[a_row_off+k_off]);uint4*d=reinterpret_cast<uint4*>(a_local);d[0]=s[0];d[1]=s[1];d[2]=s[2];d[3]=s[3];}
else{
#pragma unroll
for(int j=0;j<32;j++) a_local[j]=(m_row<M&&k_off+j<K)?A[a_row_off+k_off+j]:0;
}
int a_i32[4];int32_t a_spk;quant_32(a_local,a_i32,a_spk);
int bk=(kb>>1)+kpart*16;
int b_i32[4];
if (b_ng<N&&bk+15<K/2){uint4 bd=*reinterpret_cast<const uint4*>(&Bsh[bsh_base+(bk>>4)*256]);b_i32[0]=((int*)&bd)[0];b_i32[1]=((int*)&bd)[1];b_i32[2]=((int*)&bd)[2];b_i32[3]=((int*)&bd)[3];}
else{b_i32[0]=0;b_i32[1]=0;b_i32[2]=0;b_i32[3]=0;}
int32_t b_spk=load_b_scale_shuf(Bsc,b_ng,(kb>>5)+kpart,SNG,N);
i32x8 am={a_i32[0],a_i32[1],a_i32[2],a_i32[3],0,0,0,0};
i32x8 bm={b_i32[0],b_i32[1],b_i32[2],b_i32[3],0,0,0,0};
c_acc=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(am,bm,c_acc,4,4,0,a_spk,0,b_spk);
}
#pragma unroll
for (int j=0;j<4;j++){int mo=mt*16+kpart*4+j,no=nt*16+t_row;if(mo<M&&no<N)atomicAdd(&Cfp32[mo*N+no],c_acc[j]);}
}
__global__ void __launch_bounds__(1024)
fp32_to_bf16(float* __restrict__ src,uint16_t* __restrict__ dst,int n){
int i=blockIdx.x*1024+threadIdx.x;
if(i<n){float val=src[i];src[i]=0.0f;uint32_t fp=__float_as_uint(val);fp+=0x7FFFu+((fp>>16)&1u);dst[i]=(uint16_t)(fp>>16u);}
}
// ══════ quant kernels ══════
__global__ void __launch_bounds__(64,16)
quant_a_kernel(const uint16_t* __restrict__ A,uint8_t* __restrict__ Aq,
uint8_t* __restrict__ Asc,int M,int K,int strA,int sm){
int idx=blockIdx.x*blockDim.x+threadIdx.x;
int n_scales=K/32,total_blocks=sm*n_scales;
if(idx>=total_blocks)return;
int row=idx/n_scales,blk=idx%n_scales,k_off=blk*32;
uint16_t vals[32];
if(row<M){const uint4*s=reinterpret_cast<const uint4*>(&A[row*strA+k_off]);uint4*d=reinterpret_cast<uint4*>(vals);d[0]=s[0];d[1]=s[1];d[2]=s[2];d[3]=s[3];}
else{
#pragma unroll
for(int j=0;j<32;j++)vals[j]=0;
}
uint16_t mx=0;
#pragma unroll
for(int j=0;j<32;j++)mx=max(mx,(uint16_t)(vals[j]&0x7FFF));
uint32_t au=(((uint32_t)mx<<16)+0x200000u)&0xFF800000u;
int ef=(au>>23u)&0xFFu,su=(au==0u)?-127:max(-127,min(127,ef-127-2));
float hs=(su>=-126)?__uint_as_float((uint32_t)(su+127)<<23):0.0f;
uint8_t scale_val=(uint8_t)(su+127);
unsigned int packed[4];
packed[0]=0;packed[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[0],*reinterpret_cast<bf16v2_t*>(&vals[0]),hs,0);
packed[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[0],*reinterpret_cast<bf16v2_t*>(&vals[2]),hs,1);
packed[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[0],*reinterpret_cast<bf16v2_t*>(&vals[4]),hs,2);
packed[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[0],*reinterpret_cast<bf16v2_t*>(&vals[6]),hs,3);
packed[1]=0;packed[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[1],*reinterpret_cast<bf16v2_t*>(&vals[8]),hs,0);
packed[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[1],*reinterpret_cast<bf16v2_t*>(&vals[10]),hs,1);
packed[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[1],*reinterpret_cast<bf16v2_t*>(&vals[12]),hs,2);
packed[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[1],*reinterpret_cast<bf16v2_t*>(&vals[14]),hs,3);
packed[2]=0;packed[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[2],*reinterpret_cast<bf16v2_t*>(&vals[16]),hs,0);
packed[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[2],*reinterpret_cast<bf16v2_t*>(&vals[18]),hs,1);
packed[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[2],*reinterpret_cast<bf16v2_t*>(&vals[20]),hs,2);
packed[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[2],*reinterpret_cast<bf16v2_t*>(&vals[22]),hs,3);
packed[3]=0;packed[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[3],*reinterpret_cast<bf16v2_t*>(&vals[24]),hs,0);
packed[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[3],*reinterpret_cast<bf16v2_t*>(&vals[26]),hs,1);
packed[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[3],*reinterpret_cast<bf16v2_t*>(&vals[28]),hs,2);
packed[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[3],*reinterpret_cast<bf16v2_t*>(&vals[30]),hs,3);
if(row<M)*reinterpret_cast<uint4*>(&Aq[row*(K/2)+blk*16])=*reinterpret_cast<uint4*>(packed);
int sn=K/32,i0=row/32,i1=(row%32)/16,i2=row%16,i3=blk/8,i4=(blk%8)/4,i5=blk%4;
Asc[i0*(32*sn)+i3*256+i5*64+i2*4+i4*2+i1]=scale_val;
}
__global__ void __launch_bounds__(64,16)
quant_a_raw_v2_kernel(const uint16_t* __restrict__ A,uint8_t* __restrict__ Aq,
uint8_t* __restrict__ Asc,int M,int K,int strA){
int idx=blockIdx.x*blockDim.x+threadIdx.x;
int n_scales=K/32,total=M*n_scales;
if(idx>=total)return;
int row=idx/n_scales,blk=idx%n_scales,k_off=blk*32;
uint16_t vals[32];
if(row<M){const uint4*s=reinterpret_cast<const uint4*>(&A[row*strA+k_off]);uint4*d=reinterpret_cast<uint4*>(vals);d[0]=s[0];d[1]=s[1];d[2]=s[2];d[3]=s[3];}
else{
#pragma unroll
for(int j=0;j<32;j++)vals[j]=0;
}
float fmax_val=0.0f;
const __hip_bfloat16*bf16v=reinterpret_cast<const __hip_bfloat16*>(vals);
#pragma unroll
for(int j=0;j<32;j++){float fv=__bfloat162float(bf16v[j]);fmax_val=fmaxf(fmax_val,fabsf(fv));}
int scale_exp=-127;
if(fmax_val>0.0f){uint32_t ab;memcpy(&ab,&fmax_val,4);ab=(ab+0x00200000u)&0xFF800000u;float r;memcpy(&r,&ab,4);scale_exp=max(-127,min(127,(int)floorf(log2f(r))-2));}
uint8_t scale_byte=(uint8_t)(scale_exp+127);
float hs=(scale_exp>=-126)?__uint_as_float((uint32_t)(scale_exp+127)<<23):0.0f;
unsigned int packed[4];
packed[0]=0;packed[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[0],*reinterpret_cast<bf16v2_t*>(&vals[0]),hs,0);
packed[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[0],*reinterpret_cast<bf16v2_t*>(&vals[2]),hs,1);
packed[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[0],*reinterpret_cast<bf16v2_t*>(&vals[4]),hs,2);
packed[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[0],*reinterpret_cast<bf16v2_t*>(&vals[6]),hs,3);
packed[1]=0;packed[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[1],*reinterpret_cast<bf16v2_t*>(&vals[8]),hs,0);
packed[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[1],*reinterpret_cast<bf16v2_t*>(&vals[10]),hs,1);
packed[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[1],*reinterpret_cast<bf16v2_t*>(&vals[12]),hs,2);
packed[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[1],*reinterpret_cast<bf16v2_t*>(&vals[14]),hs,3);
packed[2]=0;packed[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[2],*reinterpret_cast<bf16v2_t*>(&vals[16]),hs,0);
packed[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[2],*reinterpret_cast<bf16v2_t*>(&vals[18]),hs,1);
packed[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[2],*reinterpret_cast<bf16v2_t*>(&vals[20]),hs,2);
packed[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[2],*reinterpret_cast<bf16v2_t*>(&vals[22]),hs,3);
packed[3]=0;packed[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[3],*reinterpret_cast<bf16v2_t*>(&vals[24]),hs,0);
packed[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[3],*reinterpret_cast<bf16v2_t*>(&vals[26]),hs,1);
packed[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[3],*reinterpret_cast<bf16v2_t*>(&vals[28]),hs,2);
packed[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[3],*reinterpret_cast<bf16v2_t*>(&vals[30]),hs,3);
if(row<M)*reinterpret_cast<uint4*>(&Aq[row*(K/2)+blk*16])=*reinterpret_cast<uint4*>(packed);
Asc[row*n_scales+blk]=scale_byte;
}
static hipModule_t _asm_mod=nullptr;
static hipFunction_t _asm_fn=nullptr;
struct __attribute__((packed)) AsmArgs{
void*D;char _0[8];void*C;char _1[8];void*A;char _2[8];void*B;char _3[8];
float alpha;char _4[12];float beta;char _5[12];
unsigned int sD0;char _6[12];unsigned int sD1;char _7[12];
unsigned int sC0;char _8[12];unsigned int sC1;char _9[12];
unsigned int sA0;char _10[12];unsigned int sA1;char _11[12];
unsigned int sB0;char _12[12];unsigned int sB1;char _13[12];
unsigned int M;char _14[12];unsigned int N;char _15[12];unsigned int K;char _16[12];
void*SA;char _17[8];void*SB;char _18[8];
unsigned int sSA0;char _19[12];unsigned int sSA1;char _20[12];
unsigned int sSB0;char _21[12];unsigned int sSB1;char _22[12];
int log2ks;
};
torch::Tensor asm_gemm_a4w4(torch::Tensor Aq,torch::Tensor Bsh,torch::Tensor Asc,torch::Tensor Bsc,int64_t m,int64_t n,int64_t k,torch::Tensor out){
if(!_asm_mod){hipModuleLoad(&_asm_mod,"/home/runner/aiter/hsa//gfx950/f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co");hipModuleGetFunction(&_asm_fn,_asm_mod,"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E");}
int mp=((int)m+31)/32*32;
AsmArgs a={};a.D=out.data_ptr();a.C=out.data_ptr();a.A=Aq.data_ptr();a.B=Bsh.data_ptr();
a.alpha=1.0f;a.beta=0.0f;a.sC0=(unsigned)out.stride(0);a.sC1=1;
a.sA0=(unsigned)(Aq.stride(0)*2);a.sA1=1;a.sB0=(unsigned)(Bsh.stride(0)*2);a.sB1=1;
a.M=(unsigned)m;a.N=(unsigned)n;a.K=(unsigned)k;a.SA=Asc.data_ptr();a.SB=Bsc.data_ptr();
a.sSA0=(unsigned)Asc.stride(0);a.sSA1=1;a.sSB0=(unsigned)Bsc.stride(0);a.sSB1=1;a.log2ks=0;
size_t asz=sizeof(a);
void*cfg[]={HIP_LAUNCH_PARAM_BUFFER_POINTER,&a,HIP_LAUNCH_PARAM_BUFFER_SIZE,&asz,HIP_LAUNCH_PARAM_END};
hipModuleLaunchKernel(_asm_fn,((unsigned)n+127)/128,(mp+31)/32,1,256,1,1,0,0,nullptr,(void**)cfg);
if(mp>(int)m)return out.slice(0,0,(int)m);return out;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME,m){
m.def("gemm_fused_16x128",[](torch::Tensor A,torch::Tensor Bq,torch::Tensor Bsc,int64_t N,torch::Tensor C)->torch::Tensor{
int M=(int)A.size(0),K=(int)A.size(1);
dim3 grid((M+15)/16,((int)N+127)/128);
gemm_fused_16x128_kernel<<<grid,256>>>(reinterpret_cast<const uint16_t*>(A.data_ptr()),reinterpret_cast<const uint8_t*>(Bq.data_ptr()),reinterpret_cast<const uint8_t*>(Bsc.data_ptr()),reinterpret_cast<uint16_t*>(C.data_ptr()),M,(int)N,K,(int)A.stride(0),(int)(Bq.stride(0)*Bq.element_size()),(int)Bsc.size(1));
return C;
});
m.def("gemm_fused_16x16",[](torch::Tensor A,torch::Tensor Bq,torch::Tensor Bsc,int64_t N,torch::Tensor C)->torch::Tensor{
int M=(int)A.size(0),K=(int)A.size(1);
dim3 grid((M+15)/16,(N+15)/16);
gemm_fused_16x16_kernel<<<grid,256>>>(reinterpret_cast<const uint16_t*>(A.data_ptr()),reinterpret_cast<const uint8_t*>(Bq.data_ptr()),reinterpret_cast<const uint8_t*>(Bsc.data_ptr()),reinterpret_cast<uint16_t*>(C.data_ptr()),M,(int)N,K,(int)A.stride(0),(int)(Bq.stride(0)*Bq.element_size()),(int)Bsc.size(1));
return C;
});
m.def("gemm_ext_splitk_16x16_bsh_cached",[](torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,int64_t N,int64_t ks,torch::Tensor Cfp,torch::Tensor Cbf)->torch::Tensor{
int M=(int)A.size(0),K=(int)A.size(1);
int kper=((K/128+(int)ks-1)/(int)ks)*128,aks=(K+kper-1)/kper;
dim3 grid((M+15)/16,((int)N+15)/16,aks);
gemm_ext_splitk_16x16_bsh_kernel<<<grid,64>>>(reinterpret_cast<const uint16_t*>(A.data_ptr()),reinterpret_cast<const uint8_t*>(Bsh.data_ptr()),reinterpret_cast<const uint8_t*>(Bsc.data_ptr()),Cfp.data_ptr<float>(),M,(int)N,K,(int)A.stride(0),(int)Bsc.size(1),kper);
int total=M*(int)N;
fp32_to_bf16<<<(total+1023)/1024,1024>>>(Cfp.data_ptr<float>(),reinterpret_cast<uint16_t*>(Cbf.data_ptr()),total);
return Cbf;
});
m.def("asm_gemm_a4w4",&asm_gemm_a4w4);
m.def("warmup_asm",[](){if(!_asm_mod){hipModuleLoad(&_asm_mod,"/home/runner/aiter/hsa//gfx950/f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co");hipModuleGetFunction(&_asm_fn,_asm_mod,"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E");}});
m.def("quant_and_asm_gemm",[](torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,int64_t N,torch::Tensor Aq,torch::Tensor Asc,torch::Tensor out)->torch::Tensor{
int M=(int)A.size(0),K=(int)A.size(1),sm=(int)Asc.size(0);
quant_a_kernel<<<(sm*(K/32)+63)/64,64>>>(reinterpret_cast<const uint16_t*>(A.data_ptr()),reinterpret_cast<uint8_t*>(Aq.data_ptr()),reinterpret_cast<uint8_t*>(Asc.data_ptr()),M,K,(int)A.stride(0),sm);
return asm_gemm_a4w4(Aq,Bsh,Asc,Bsc,M,(int)N,K,out);
});
m.def("quant_raw_v2",[](torch::Tensor A,torch::Tensor Aq,torch::Tensor Asc){
int M=(int)A.size(0),K=(int)A.size(1);
quant_a_raw_v2_kernel<<<(M*(K/32)+63)/64,64>>>(reinterpret_cast<const uint16_t*>(A.data_ptr()),Aq.data_ptr<uint8_t>(),Asc.data_ptr<uint8_t>(),M,K,(int)A.stride(0));
});
}
"""
_EXT = None
def _get_ext():
global _EXT
if _EXT is None:
_EXT = load_inline(
name="fused_mxfp4_v594",
cpp_sources=[""],
cuda_sources=[_HIP_SRC],
extra_cuda_cflags=["-O3", "-std=c++17", "--offload-arch=gfx950"],
verbose=False,
)
return _EXT
@triton.jit
def remap_xcd(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):
pids_per_xcd=(GRID_MN+NUM_XCDS-1)//NUM_XCDS
tall_xcds=GRID_MN%NUM_XCDS; tall_xcds=NUM_XCDS if tall_xcds==0 else tall_xcds
xcd=pid%NUM_XCDS; local_pid=pid//NUM_XCDS
if xcd<tall_xcds: pid=xcd*pids_per_xcd+local_pid
else: pid=(tall_xcds*pids_per_xcd+(xcd-tall_xcds)*(pids_per_xcd-1)+local_pid)
return pid
@triton.jit
def pid_grid(pid,num_pid_m,num_pid_n,GROUP_SIZE_M:tl.constexpr=1):
if GROUP_SIZE_M==1: pid_m=pid//num_pid_n; pid_n=pid%num_pid_n
else:
num_pid_in_group=GROUP_SIZE_M*num_pid_n; group_id=pid//num_pid_in_group
first_pid_m=group_id*GROUP_SIZE_M; group_size_m=min(num_pid_m-first_pid_m,GROUP_SIZE_M)
tl.assume(group_size_m>=0)
pid_m=first_pid_m+(pid%group_size_m); pid_n=(pid%num_pid_in_group)//group_size_m
return pid_m,pid_n
@triton.heuristics({"EVEN_K":lambda args:(args["K"]%(args["BLOCK_SIZE_K"]//2)==0)and(args["SPLITK_BLOCK_SIZE"]%args["BLOCK_SIZE_K"]==0)and(args["K"]%(args["SPLITK_BLOCK_SIZE"]//2)==0)})
@triton.jit
def _gemm_fp4_kernel(
a_ptr,b_ptr,c_ptr,a_scales_ptr,b_scales_ptr,M,N,K,
stride_am,stride_ak,stride_bk,stride_bn,stride_ck,stride_cm,stride_cn,
stride_asm,stride_ask,stride_bsn,stride_bsk,B_SCALE_SN_PAD,
BLOCK_SIZE_M:tl.constexpr,BLOCK_SIZE_N:tl.constexpr,BLOCK_SIZE_K:tl.constexpr,
GROUP_SIZE_M:tl.constexpr,NUM_KSPLIT:tl.constexpr,SPLITK_BLOCK_SIZE:tl.constexpr,
EVEN_K:tl.constexpr,num_warps:tl.constexpr,num_stages:tl.constexpr,
waves_per_eu:tl.constexpr,matrix_instr_nonkdim:tl.constexpr,
):
tl.assume(stride_am>0);tl.assume(stride_ak>0);tl.assume(stride_bk>0);tl.assume(stride_bn>0)
tl.assume(stride_cm>0);tl.assume(stride_cn>0);tl.assume(stride_asm>0);tl.assume(stride_ask>0)
tl.assume(stride_bsk>0);tl.assume(stride_bsn>0)
GRID_MN=tl.cdiv(M,BLOCK_SIZE_M)*tl.cdiv(N,BLOCK_SIZE_N)
pid_unified=tl.program_id(axis=0); pid_unified=remap_xcd(pid_unified,GRID_MN*NUM_KSPLIT,NUM_XCDS=8)
pid_k=pid_unified%NUM_KSPLIT; pid=pid_unified//NUM_KSPLIT
num_pid_m=tl.cdiv(M,BLOCK_SIZE_M); num_pid_n=tl.cdiv(N,BLOCK_SIZE_N)
if NUM_KSPLIT==1: pid_m,pid_n=pid_grid(pid,num_pid_m,num_pid_n,GROUP_SIZE_M=GROUP_SIZE_M)
else: pid_m=pid//num_pid_n; pid_n=pid%num_pid_n
tl.assume(pid_m>=0); tl.assume(pid_n>=0)
SCALE_GROUP_SIZE:tl.constexpr=32
if (pid_k*SPLITK_BLOCK_SIZE//2)<K:
num_k_iter=tl.cdiv(SPLITK_BLOCK_SIZE//2,BLOCK_SIZE_K//2)
offs_k=tl.arange(0,BLOCK_SIZE_K//2); offs_k_split=pid_k*(SPLITK_BLOCK_SIZE//2)+offs_k
offs_am=(pid_m*BLOCK_SIZE_M+tl.arange(0,BLOCK_SIZE_M))%M
offs_bn=(pid_n*BLOCK_SIZE_N+tl.arange(0,BLOCK_SIZE_N))%N
a_ptrs=a_ptr+(offs_am[:,None]*stride_am+offs_k_split[None,:]*stride_ak)
b_ptrs=b_ptr+(offs_k_split[:,None]*stride_bk+offs_bn[None,:]*stride_bn)
offs_ks=(pid_k*(SPLITK_BLOCK_SIZE//SCALE_GROUP_SIZE))+tl.arange(0,BLOCK_SIZE_K//SCALE_GROUP_SIZE)
a_scale_ptrs=a_scales_ptr+offs_am[:,None]*stride_asm+offs_ks[None,:]*stride_ask
bn=offs_bn[:,None];bk=offs_ks[None,:]
b_i0=bn//32;b_i1=(bn%32)//16;b_i2=bn%16;b_i3=bk//8;b_i4=(bk%8)//4;b_i5=bk%4
b_sh_idx=b_i0*(32*B_SCALE_SN_PAD)+b_i3*256+b_i5*64+b_i2*4+b_i4*2+b_i1
accumulator=tl.zeros((BLOCK_SIZE_M,BLOCK_SIZE_N),dtype=tl.float32)
for k in range(pid_k*num_k_iter,(pid_k+1)*num_k_iter):
a_scales=tl.load(a_scale_ptrs); b_scales=tl.load(b_scales_ptr+b_sh_idx,cache_modifier=".cg")
if EVEN_K: a=tl.load(a_ptrs);b=tl.load(b_ptrs,cache_modifier=".cg")
else:
a=tl.load(a_ptrs,mask=offs_k[None,:]<K-k*(BLOCK_SIZE_K//2),other=0)
b=tl.load(b_ptrs,mask=offs_k[:,None]<K-k*(BLOCK_SIZE_K//2),other=0,cache_modifier=".cg")
accumulator=tl.dot_scaled(a,a_scales,"e2m1",b,b_scales,"e2m1",accumulator)
a_ptrs+=(BLOCK_SIZE_K//2)*stride_ak;b_ptrs+=(BLOCK_SIZE_K//2)*stride_bk
a_scale_ptrs+=(BLOCK_SIZE_K//SCALE_GROUP_SIZE)*stride_ask
bk=bk+(BLOCK_SIZE_K//SCALE_GROUP_SIZE);b_i3=bk//8;b_i4=(bk%8)//4;b_i5=bk%4
b_sh_idx=b_i0*(32*B_SCALE_SN_PAD)+b_i3*256+b_i5*64+b_i2*4+b_i4*2+b_i1
c=accumulator.to(c_ptr.type.element_ty)
offs_cm=(pid_m*BLOCK_SIZE_M+tl.arange(0,BLOCK_SIZE_M)).to(tl.int64)
offs_cn=(pid_n*BLOCK_SIZE_N+tl.arange(0,BLOCK_SIZE_N)).to(tl.int64)
c_ptrs=c_ptr+stride_cm*offs_cm[:,None]+stride_cn*offs_cn[None,:]+pid_k*stride_ck
tl.store(c_ptrs,c,mask=(offs_cm[:,None]<M)&(offs_cn[None,:]<N))
_ws_cache: dict = {}
_warmed = False
def _prewarm_triton():
"""Pre-compile Triton kernel with BN=32, BK=1024, warps=4 config."""
try:
import torch
dev = torch.device('cuda', 0)
# Use minimal sizes matching BN=32, BK=1024
# K_packed=512 (K=1024 fp4 elems), BLOCK_SIZE_K=1024 → 1 k-iter
M, N, K_packed = 32, 32, 512
k_scale = K_packed // 16 # K=1024 fp4 → 32 scale groups
aq = torch.zeros(M, K_packed, dtype=torch.uint8, device=dev)
bt = torch.zeros(K_packed, N, dtype=torch.uint8, device=dev)
y = torch.zeros(M, N, dtype=torch.bfloat16, device=dev)
asc = torch.zeros(M, k_scale, dtype=torch.uint8, device=dev)
bsc = torch.zeros(N, k_scale, dtype=torch.uint8, device=dev)
grid = lambda META: (1,)
_gemm_fp4_kernel[grid](
aq, bt, y, asc, bsc, M, N, K_packed,
aq.stride(0), aq.stride(1), bt.stride(0), bt.stride(1),
0, y.stride(0), y.stride(1), asc.stride(0), asc.stride(1), 0, 0, k_scale,
SPLITK_BLOCK_SIZE=2*K_packed, BLOCK_SIZE_M=32, BLOCK_SIZE_N=32,
BLOCK_SIZE_K=1024, GROUP_SIZE_M=2, NUM_KSPLIT=1,
num_warps=4, num_stages=2, waves_per_eu=4, matrix_instr_nonkdim=16,
)
torch.cuda.synchronize()
except Exception:
pass
_prewarm_triton()
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
global _warmed
ext = _get_ext()
if not _warmed:
ext.warmup_asm(); _warmed = True
M, K, N = A.shape[0], A.shape[1], B.shape[0]
key = (M, K, N); dev = A.device
if key not in _ws_cache:
mp=((M+31)//32)*32; sm=((mp+255)//256)*256
_ws_cache[key]=(
torch.zeros(M,N,dtype=torch.float32,device=dev),
torch.empty(M,N,dtype=torch.bfloat16,device=dev),
torch.empty(M,K//2,dtype=torch.uint8,device=dev),
torch.empty(sm,K//32,dtype=torch.uint8,device=dev),
torch.empty(mp,N,dtype=torch.bfloat16,device=dev),
)
fp32_ws,bf16_out,aq_buf,asc_buf,asm_out=_ws_cache[key]
# S5: M=64 — quant_raw_v2 + Triton BN=32, BK=1024, warps=4 (AITer config)
if 64<=M<128:
k_scale=K//SCALE_GROUP_SIZE
asc_key=(M,K,'raw')
if asc_key not in _ws_cache:
_ws_cache[asc_key]=torch.empty(M,k_scale,dtype=torch.uint8,device=dev)
asc_raw=_ws_cache[asc_key]
ext.quant_raw_v2(A,aq_buf,asc_raw)
B_q_u8=B_q.view(torch.uint8) if B_q.dtype!=torch.uint8 else B_q
B_t=B_q_u8.T
B_sc_u8=B_scale_sh.view(torch.uint8) if B_scale_sh.dtype!=torch.uint8 else B_scale_sh
sn_pad=B_sc_u8.shape[1]; K_packed=K//2; y=bf16_out
grid=lambda META:(META["NUM_KSPLIT"]*triton.cdiv(M,META["BLOCK_SIZE_M"])*triton.cdiv(N,META["BLOCK_SIZE_N"]),)
_gemm_fp4_kernel[grid](
aq_buf,B_t,y,asc_raw,B_sc_u8,M,N,K_packed,
aq_buf.stride(0),aq_buf.stride(1),B_t.stride(0),B_t.stride(1),
0,y.stride(0),y.stride(1),asc_raw.stride(0),asc_raw.stride(1),0,0,sn_pad,
SPLITK_BLOCK_SIZE=2*K_packed,BLOCK_SIZE_M=32,BLOCK_SIZE_N=32,BLOCK_SIZE_K=1024,
GROUP_SIZE_M=2,NUM_KSPLIT=1,num_warps=4,num_stages=2,waves_per_eu=4,matrix_instr_nonkdim=16,
)
return y
# S6: M>=128 — quant_a + ASM GEMM (v517 path, best for large M)
if M>=128:
return ext.quant_and_asm_gemm(A,B_shuffle,B_scale_sh,N,aq_buf,asc_buf,asm_out)
# S2: M<=16, large K — ext_splitk (v517 path)
if M<=16 and K>1024:
return ext.gemm_ext_splitk_16x16_bsh_cached(A,B_shuffle,B_scale_sh,N,16,fp32_ws,bf16_out)
B_q_u8=B_q.view(torch.uint8) if B_q.dtype!=torch.uint8 else B_q
if M<=16 and N % 128 == 0:
return ext.gemm_fused_16x128(A,B_q_u8,B_scale_sh,N,bf16_out)
else:
return ext.gemm_fused_16x16(A,B_q_u8,B_scale_sh,N,bf16_out)
scrolls · 539 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