submission 596367
rjvkr2021 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 946 lines, June 9 Researcher Reciprocity License v1.0.
v3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-596367?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:1fbdfb50df4b5af57cc0b2ef8b57625b81eff0858d119b04805c9b195183208a
license declaredunknown
license concludedunknown
authorsrjvkr2021
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.shared-memory
__shared__ __align__(16) uint8_t Alds[BLOCK_M*LDS_ROW];__shared__ uint8_t Asclds[BLOCK_M*8];tile-m = 16
constexpr int MFMA_K=128, DOUBLE_K=256, LDS_ROW=128, HALF_K=64, SCALE_GROUP=32, BLOCK_M=16, BLOCK_M_LARGE=32;tile-n = 64
constexpr int BLOCK_N=64, BNT=128;vector-width = int4
const int4 r0 = reinterpret_cast<const int4*>(s)[0];Kernel source
v3.py946 lines
"""
FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
Quant logic follows aiter op_tests/test_gemm_a4w4.py (get_triton_quant(QuantType.per_1x32)).
v3: Double-buffered K-loop for ksps>=4 to overlap memory loads with MFMA compute.
"""
import os
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
CPP_SOURCE = r"""
#include <torch/extension.h>
torch::Tensor dispatch_gemm(torch::Tensor A, torch::Tensor Bq, torch::Tensor Bsc, int M, int N, int K);
"""
HIP_SOURCE = r"""
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <torch/extension.h>
constexpr int MFMA_K=128, DOUBLE_K=256, LDS_ROW=128, HALF_K=64, SCALE_GROUP=32, BLOCK_M=16, BLOCK_M_LARGE=32;
typedef int __attribute__((ext_vector_type(4))) int4_vec;
typedef float __attribute__((ext_vector_type(4))) float4_vec;
typedef int32_t __attribute__((ext_vector_type(4))) i32x4;
typedef __bf16 __attribute__((ext_vector_type(2))) bf16x2;
typedef uint32_t __attribute__((address_space(3)))* as3_uint32_ptr;
extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(i32x4 rsrc,as3_uint32_ptr lds_ptr,int size,int voffset,int soffset,int offset,int aux) __asm("llvm.amdgcn.raw.buffer.load.lds");
struct buffer_resource{uint64_t ptr;uint32_t range;uint32_t config;};
__device__ __forceinline__ i32x4 make_srsrc(const void*p,uint32_t r){buffer_resource s={reinterpret_cast<uint64_t>(p),r,0x110000};return *reinterpret_cast<const i32x4*>(&s);}
__device__ __forceinline__ float4_vec mfma_fp4(int4_vec A,int4_vec B,float4_vec C,int sA,int sB){float4_vec D;asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%3,%4,%5 cbsz:4 blgp:4":"=v"(D):"v"(A),"v"(B),"v"(C),"v"(sA),"v"(sB));return D;}
__device__ __forceinline__ int lds_swz(int o){return o^(((o&2047)>>8)<<4);}
__device__ __forceinline__ void compute_scale(float mx,uint8_t&sc,float&sf){if(mx>0.f){uint32_t b=__float_as_uint(mx);b=(b+0x200000u)&0xFF800000u;int su=((b>>23)&0xFF)-129;su=su<-127?-127:(su>127?127:su);sc=(uint8_t)(su+127);sf=__uint_as_float((uint32_t)(su+127)<<23);}else{sc=0;sf=0.f;}}
__device__ __forceinline__ bf16x2 as_bf16x2(uint32_t w){return __builtin_bit_cast(bf16x2,w);}
__device__ __forceinline__ void update_max_word(uint32_t w,float&mx){mx=fmaxf(mx,__uint_as_float((w&0x7fffu)<<16));mx=fmaxf(mx,__uint_as_float(w&0x7fff0000u));}
__device__ __forceinline__ uint32_t pack4_words(uint32_t w0,uint32_t w1,uint32_t w2,uint32_t w3,float sf){
uint32_t p=0;
p=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p,as_bf16x2(w0),sf,0);
p=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p,as_bf16x2(w1),sf,1);
p=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p,as_bf16x2(w2),sf,2);
p=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p,as_bf16x2(w3),sf,3);
return p;
}
// A-quant helper: loads bf16 from global, quantizes to FP4, stores to LDS
__device__ __forceinline__ void do_a_quant(
const __hip_bfloat16*__restrict__ A, int M, int K,
int tid, int bm, int ke,
uint8_t* Alds_buf, uint8_t* Asclds_buf)
{
const int qr = tid >> 3, qg = tid & 7;
const int gr = bm + qr, ko = ke + qg * 32;
uint32_t p0=0, p1=0, p2=0, p3=0; uint8_t asc = 0x7f;
if (gr < M) {
const __hip_bfloat16* s = A + gr * K + ko;
const int4 r0 = reinterpret_cast<const int4*>(s)[0];
const int4 r1 = reinterpret_cast<const int4*>(s)[1];
const int4 r2 = reinterpret_cast<const int4*>(s)[2];
const int4 r3 = reinterpret_cast<const int4*>(s)[3];
const uint32_t w0 = (uint32_t)r0.x, w1 = (uint32_t)r0.y, w2 = (uint32_t)r0.z, w3 = (uint32_t)r0.w;
const uint32_t w4 = (uint32_t)r1.x, w5 = (uint32_t)r1.y, w6 = (uint32_t)r1.z, w7 = (uint32_t)r1.w;
const uint32_t w8 = (uint32_t)r2.x, w9 = (uint32_t)r2.y, w10 = (uint32_t)r2.z, w11 = (uint32_t)r2.w;
const uint32_t w12 = (uint32_t)r3.x, w13 = (uint32_t)r3.y, w14 = (uint32_t)r3.z, w15 = (uint32_t)r3.w;
float l = 0;
update_max_word(w0,l); update_max_word(w1,l); update_max_word(w2,l); update_max_word(w3,l);
update_max_word(w4,l); update_max_word(w5,l); update_max_word(w6,l); update_max_word(w7,l);
update_max_word(w8,l); update_max_word(w9,l); update_max_word(w10,l); update_max_word(w11,l);
update_max_word(w12,l); update_max_word(w13,l); update_max_word(w14,l); update_max_word(w15,l);
float sf; compute_scale(l, asc, sf);
p0=pack4_words(w0,w1,w2,w3,sf);
p1=pack4_words(w4,w5,w6,w7,sf);
p2=pack4_words(w8,w9,w10,w11,sf);
p3=pack4_words(w12,w13,w14,w15,sf);
}
int4 wr; wr.x=(int)p0; wr.y=(int)p1; wr.z=(int)p2; wr.w=(int)p3;
*reinterpret_cast<int4*>(&Alds_buf[lds_swz(qr*LDS_ROW+qg*16)]) = wr;
Asclds_buf[qr*8+qg] = asc;
}
// B-load helper: issues buffer_load_lds for B tiles
__device__ __forceinline__ void do_b_load(
i32x4 srsrc, int btid, int bn, int N, int kb, int bstride,
uint8_t* Blds_buf)
{
constexpr int BLOCK_N=64, BNT=128;
constexpr int BL = (BLOCK_N * LDS_ROW / 16 + BNT - 1) / BNT;
for (int ld = 0; ld < BL; ld++) {
const int f = (ld * BNT + btid) << 4;
const int r = f >> 7; const int gn = bn + r;
if (r < BLOCK_N && gn < N) {
const int sf = lds_swz(f), sc = sf & 127, ac = kb + sc;
llvm_amdgcn_raw_buffer_load_lds(srsrc,
(as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Blds_buf)+f),
16,
(gn>>4)*(bstride<<4)+(ac>>5)*512+((ac>>4)&1)*256+(gn&15)*16,
0, 0, 0);
}
}
}
// B-scale load helper
__device__ __forceinline__ void do_b_scale_load(
const uint8_t*__restrict__ Bsc, int stid, int bn, int N, int scs, int so,
uint8_t* Bslds_buf)
{
constexpr int BLOCK_N=64;
const int row = stid & (BLOCK_N - 1), gb = (stid >> 6) << 2;
const int gn = bn + row;
if (gn < N) {
const int rb=(gn>>5)*32*scs+(gn&15)*4+((gn>>4)&1);
for(int g=0;g<4;g++){
const int grp=gb+g,ac=so+grp;
Bslds_buf[row*8+grp]=Bsc[rb+(ac&3)*64+((ac&7)>>2)*2+(ac>>3)*256];
}
} else {
for(int g=0;g<4;g++) Bslds_buf[row*8+gb+g]=0x7f;
}
}
// Original gemm_bn64 kernel (unchanged from baseline)
__global__ __launch_bounds__(256, 4)
void gemm_bn64(const __hip_bfloat16*__restrict__ A,const uint8_t*__restrict__ Bq,const uint8_t*__restrict__ Bsc,float*__restrict__ ws,__hip_bfloat16*__restrict__ C,int M,int N,int K,int ksps){
constexpr int BLOCK_N=64, NT=256;
const int wid=threadIdx.x>>6,lm=threadIdx.x&15,lk=(threadIdx.x&63)>>4,tid=threadIdx.x;
const int bn=blockIdx.x*BLOCK_N,bm=blockIdx.y*BLOCK_M,wn=bn+(wid<<4),sid=blockIdx.z;const int bstride=K>>1,scs=K>>5;
__shared__ __align__(16) uint8_t Alds[BLOCK_M*LDS_ROW];__shared__ uint8_t Asclds[BLOCK_M*8];
__shared__ __align__(16) uint8_t Blds[BLOCK_N*LDS_ROW];__shared__ uint8_t Bslds[BLOCK_N*8];
const i32x4 srsrc=make_srsrc(Bq,N*bstride);float4_vec acc={0,0,0,0};
for(int ks=sid*ksps;ks<sid*ksps+ksps;ks++){const int ke=ks*DOUBLE_K,kb=ks*LDS_ROW;
if (tid < 128) {
const int qr = tid >> 3, qg = tid & 7;
const int gr = bm + qr, ko = ke + qg * 32;
uint32_t p0=0, p1=0, p2=0, p3=0; uint8_t asc = 0x7f;
if (gr < M) {
const __hip_bfloat16* s = A + gr * K + ko;
const int4 r0 = reinterpret_cast<const int4*>(s)[0];
const int4 r1 = reinterpret_cast<const int4*>(s)[1];
const int4 r2 = reinterpret_cast<const int4*>(s)[2];
const int4 r3 = reinterpret_cast<const int4*>(s)[3];
const uint32_t w0 = (uint32_t)r0.x, w1 = (uint32_t)r0.y, w2 = (uint32_t)r0.z, w3 = (uint32_t)r0.w;
const uint32_t w4 = (uint32_t)r1.x, w5 = (uint32_t)r1.y, w6 = (uint32_t)r1.z, w7 = (uint32_t)r1.w;
const uint32_t w8 = (uint32_t)r2.x, w9 = (uint32_t)r2.y, w10 = (uint32_t)r2.z, w11 = (uint32_t)r2.w;
const uint32_t w12 = (uint32_t)r3.x, w13 = (uint32_t)r3.y, w14 = (uint32_t)r3.z, w15 = (uint32_t)r3.w;
float l = 0;
update_max_word(w0,l); update_max_word(w1,l); update_max_word(w2,l); update_max_word(w3,l);
update_max_word(w4,l); update_max_word(w5,l); update_max_word(w6,l); update_max_word(w7,l);
update_max_word(w8,l); update_max_word(w9,l); update_max_word(w10,l); update_max_word(w11,l);
update_max_word(w12,l); update_max_word(w13,l); update_max_word(w14,l); update_max_word(w15,l);
float sf; compute_scale(l, asc, sf);
p0=pack4_words(w0,w1,w2,w3,sf);
p1=pack4_words(w4,w5,w6,w7,sf);
p2=pack4_words(w8,w9,w10,w11,sf);
p3=pack4_words(w12,w13,w14,w15,sf);
}
int4 wr; wr.x=(int)p0; wr.y=(int)p1; wr.z=(int)p2; wr.w=(int)p3;
*reinterpret_cast<int4*>(&Alds[lds_swz(qr*LDS_ROW+qg*16)]) = wr;
Asclds[qr*8+qg] = asc;
}
if (tid >= 128) {
const int btid = tid - 128;
constexpr int BNT = 128;
constexpr int BL = (BLOCK_N * LDS_ROW / 16 + BNT - 1) / BNT;
for (int ld = 0; ld < BL; ld++) {
const int f = (ld * BNT + btid) << 4;
const int r = f >> 7; const int gn = bn + r;
if (r < BLOCK_N && gn < N) {
const int sf = lds_swz(f), sc = sf & 127, ac = kb + sc;
llvm_amdgcn_raw_buffer_load_lds(srsrc,(as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Blds)+f),16,(gn>>4)*(bstride<<4)+(ac>>5)*512+((ac>>4)&1)*256+(gn&15)*16,0,0,0);
}
}
}
{const int so=ks<<3;
if (tid >= 128) {
const int stid = tid - 128;
const int row = stid & (BLOCK_N - 1), gb = (stid >> 6) << 2;
const int gn = bn + row;
if (gn < N) { const int rb=(gn>>5)*32*scs+(gn&15)*4+((gn>>4)&1);
for(int g=0;g<4;g++){const int grp=gb+g,ac=so+grp;Bslds[row*8+grp]=Bsc[rb+(ac&3)*64+((ac&7)>>2)*2+(ac>>3)*256];}
} else { for(int g=0;g<4;g++)Bslds[row*8+gb+g]=0x7f; }
}}
asm volatile("s_waitcnt vmcnt(0)");__syncthreads();
int4_vec A0,A1;int as0,as1;
{const int o=lds_swz(lm*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[o]);A0={t.x,t.y,t.z,t.w};}
as0=(int)Asclds[(lm<<3)+lk];
{const int o=lds_swz(lm*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[o]);A1={t.x,t.y,t.z,t.w};}
as1=(int)Asclds[(lm<<3)+4+lk];
const int br=(wid<<4)+lm;
int4_vec B0,B1;int bs0,bs1;
{const int o=lds_swz(br*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[o]);B0={t.x,t.y,t.z,t.w};}
bs0=(int)Bslds[br*8+lk];
{const int o=lds_swz(br*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[o]);B1={t.x,t.y,t.z,t.w};}
bs1=(int)Bslds[br*8+4+lk];
acc=mfma_fp4(A0,B0,acc,as0,bs0);
acc=mfma_fp4(A1,B1,acc,as1,bs1);
}
const int or_=bm+(lk<<2),oc=wn+lm;
if(oc<N){const float*ap=reinterpret_cast<const float*>(&acc);
if(C){for(int r=0;r<4;r++){int g=or_+r;if(g<M)C[g*N+oc]=__float2bfloat16(ap[r]);}}
else{float*w=ws+sid*M*N;for(int r=0;r<4;r++){int g=or_+r;if(g<M)w[g*N+oc]=ap[r];}}}
}
template<int KSPS>
__global__ __launch_bounds__(256, 4)
void gemm_bn64_k(const __hip_bfloat16*__restrict__ A,const uint8_t*__restrict__ Bq,const uint8_t*__restrict__ Bsc,float*__restrict__ ws,__hip_bfloat16*__restrict__ C,int M,int N,int K){
constexpr int BLOCK_N=64, NT=256;
const int wid=threadIdx.x>>6,lm=threadIdx.x&15,lk=(threadIdx.x&63)>>4,tid=threadIdx.x;
const int bn=blockIdx.x*BLOCK_N,bm=blockIdx.y*BLOCK_M,wn=bn+(wid<<4),sid=blockIdx.z;const int bstride=K>>1,scs=K>>5;
__shared__ __align__(16) uint8_t Alds[BLOCK_M*LDS_ROW];__shared__ uint8_t Asclds[BLOCK_M*8];
__shared__ __align__(16) uint8_t Blds[BLOCK_N*LDS_ROW];__shared__ uint8_t Bslds[BLOCK_N*8];
const i32x4 srsrc=make_srsrc(Bq,N*bstride);float4_vec acc={0,0,0,0};
#pragma unroll
for(int kstep=0;kstep<KSPS;kstep++){const int ks=sid*KSPS+kstep,ke=ks*DOUBLE_K,kb=ks*LDS_ROW;
if (tid < 128) {
const int qr = tid >> 3, qg = tid & 7;
const int gr = bm + qr, ko = ke + qg * 32;
uint32_t p0=0, p1=0, p2=0, p3=0; uint8_t asc = 0x7f;
if (gr < M) {
const __hip_bfloat16* s = A + gr * K + ko;
const int4 r0 = reinterpret_cast<const int4*>(s)[0];
const int4 r1 = reinterpret_cast<const int4*>(s)[1];
const int4 r2 = reinterpret_cast<const int4*>(s)[2];
const int4 r3 = reinterpret_cast<const int4*>(s)[3];
const uint32_t w0 = (uint32_t)r0.x, w1 = (uint32_t)r0.y, w2 = (uint32_t)r0.z, w3 = (uint32_t)r0.w;
const uint32_t w4 = (uint32_t)r1.x, w5 = (uint32_t)r1.y, w6 = (uint32_t)r1.z, w7 = (uint32_t)r1.w;
const uint32_t w8 = (uint32_t)r2.x, w9 = (uint32_t)r2.y, w10 = (uint32_t)r2.z, w11 = (uint32_t)r2.w;
const uint32_t w12 = (uint32_t)r3.x, w13 = (uint32_t)r3.y, w14 = (uint32_t)r3.z, w15 = (uint32_t)r3.w;
float l = 0;
update_max_word(w0,l); update_max_word(w1,l); update_max_word(w2,l); update_max_word(w3,l);
update_max_word(w4,l); update_max_word(w5,l); update_max_word(w6,l); update_max_word(w7,l);
update_max_word(w8,l); update_max_word(w9,l); update_max_word(w10,l); update_max_word(w11,l);
update_max_word(w12,l); update_max_word(w13,l); update_max_word(w14,l); update_max_word(w15,l);
float sf; compute_scale(l, asc, sf);
p0=pack4_words(w0,w1,w2,w3,sf);
p1=pack4_words(w4,w5,w6,w7,sf);
p2=pack4_words(w8,w9,w10,w11,sf);
p3=pack4_words(w12,w13,w14,w15,sf);
}
int4 wr; wr.x=(int)p0; wr.y=(int)p1; wr.z=(int)p2; wr.w=(int)p3;
*reinterpret_cast<int4*>(&Alds[lds_swz(qr*LDS_ROW+qg*16)]) = wr;
Asclds[qr*8+qg] = asc;
}
if (tid >= 128) {
const int btid = tid - 128;
constexpr int BNT = 128;
constexpr int BL = (BLOCK_N * LDS_ROW / 16 + BNT - 1) / BNT;
for (int ld = 0; ld < BL; ld++) {
const int f = (ld * BNT + btid) << 4;
const int r = f >> 7; const int gn = bn + r;
if (r < BLOCK_N && gn < N) {
const int sf = lds_swz(f), sc = sf & 127, ac = kb + sc;
llvm_amdgcn_raw_buffer_load_lds(srsrc,(as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Blds)+f),16,(gn>>4)*(bstride<<4)+(ac>>5)*512+((ac>>4)&1)*256+(gn&15)*16,0,0,0);
}
}
}
{const int so=ks<<3;
if (tid >= 128) {
const int stid = tid - 128;
const int row = stid & (BLOCK_N - 1), gb = (stid >> 6) << 2;
const int gn = bn + row;
if (gn < N) { const int rb=(gn>>5)*32*scs+(gn&15)*4+((gn>>4)&1);
for(int g=0;g<4;g++){const int grp=gb+g,ac=so+grp;Bslds[row*8+grp]=Bsc[rb+(ac&3)*64+((ac&7)>>2)*2+(ac>>3)*256];}
} else { for(int g=0;g<4;g++)Bslds[row*8+gb+g]=0x7f; }
}}
asm volatile("s_waitcnt vmcnt(0)");__syncthreads();
int4_vec A0,A1;int as0,as1;
{const int o=lds_swz(lm*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[o]);A0={t.x,t.y,t.z,t.w};}
as0=(int)Asclds[(lm<<3)+lk];
{const int o=lds_swz(lm*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[o]);A1={t.x,t.y,t.z,t.w};}
as1=(int)Asclds[(lm<<3)+4+lk];
const int br=(wid<<4)+lm;
int4_vec B0,B1;int bs0,bs1;
{const int o=lds_swz(br*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[o]);B0={t.x,t.y,t.z,t.w};}
bs0=(int)Bslds[br*8+lk];
{const int o=lds_swz(br*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[o]);B1={t.x,t.y,t.z,t.w};}
bs1=(int)Bslds[br*8+4+lk];
acc=mfma_fp4(A0,B0,acc,as0,bs0);
acc=mfma_fp4(A1,B1,acc,as1,bs1);
}
const int or_=bm+(lk<<2),oc=wn+lm;
if(oc<N){const float*ap=reinterpret_cast<const float*>(&acc);
if(C){for(int r=0;r<4;r++){int g=or_+r;if(g<M)C[g*N+oc]=__float2bfloat16(ap[r]);}}
else{float*w=ws+sid*M*N;for(int r=0;r<4;r++){int g=or_+r;if(g<M)w[g*N+oc]=ap[r];}}}
}
// Double-buffered version for KSPS >= 4
// Overlaps memory loads for step N+1 with MFMA compute for step N.
template<int KSPS>
__global__ __launch_bounds__(256, 2)
void gemm_bn64_db(const __hip_bfloat16*__restrict__ A,const uint8_t*__restrict__ Bq,const uint8_t*__restrict__ Bsc,float*__restrict__ ws,__hip_bfloat16*__restrict__ C,int M,int N,int K){
constexpr int BLOCK_N=64;
const int wid=threadIdx.x>>6,lm=threadIdx.x&15,lk=(threadIdx.x&63)>>4,tid=threadIdx.x;
const int bn=blockIdx.x*BLOCK_N,bm=blockIdx.y*BLOCK_M,wn=bn+(wid<<4),sid=blockIdx.z;
const int bstride=K>>1,scs=K>>5;
// Double-buffered LDS
__shared__ __align__(16) uint8_t Alds[2][BLOCK_M*LDS_ROW];
__shared__ uint8_t Asclds[2][BLOCK_M*8];
__shared__ __align__(16) uint8_t Blds[2][BLOCK_N*LDS_ROW];
__shared__ uint8_t Bslds[2][BLOCK_N*8];
const i32x4 srsrc=make_srsrc(Bq,N*bstride);
float4_vec acc={0,0,0,0};
const int ks_base = sid * KSPS;
// === PROLOGUE: load step 0 into buffer 0 ===
{
const int ks = ks_base;
const int ke = ks * DOUBLE_K, kb = ks * LDS_ROW;
if (tid < 128) {
do_a_quant(A, M, K, tid, bm, ke, Alds[0], Asclds[0]);
}
if (tid >= 128) {
do_b_load(srsrc, tid - 128, bn, N, kb, bstride, Blds[0]);
}
{
const int so = ks << 3;
if (tid >= 128) {
do_b_scale_load(Bsc, tid - 128, bn, N, scs, so, Bslds[0]);
}
}
asm volatile("s_waitcnt vmcnt(0)");
__syncthreads();
}
// === MAIN LOOP: for steps 1..KSPS-1, load next step while computing current ===
#pragma unroll
for (int kstep = 1; kstep < KSPS; kstep++) {
const int cur = (kstep - 1) & 1; // buffer with ready data
const int nxt = kstep & 1; // buffer to load into
const int ks_next = ks_base + kstep;
const int ke_next = ks_next * DOUBLE_K, kb_next = ks_next * LDS_ROW;
// Read current step's data from LDS into registers FIRST (before issuing new loads)
int4_vec A0,A1; int as0,as1;
{const int o=lds_swz(lm*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[cur][o]);A0={t.x,t.y,t.z,t.w};}
as0=(int)Asclds[cur][(lm<<3)+lk];
{const int o=lds_swz(lm*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[cur][o]);A1={t.x,t.y,t.z,t.w};}
as1=(int)Asclds[cur][(lm<<3)+4+lk];
const int br=(wid<<4)+lm;
int4_vec B0,B1; int bs0,bs1;
{const int o=lds_swz(br*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[cur][o]);B0={t.x,t.y,t.z,t.w};}
bs0=(int)Bslds[cur][br*8+lk];
{const int o=lds_swz(br*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[cur][o]);B1={t.x,t.y,t.z,t.w};}
bs1=(int)Bslds[cur][br*8+4+lk];
// Issue loads for next step into nxt buffer
// These overlap with the MFMA below since they target different LDS regions
if (tid < 128) {
do_a_quant(A, M, K, tid, bm, ke_next, Alds[nxt], Asclds[nxt]);
}
if (tid >= 128) {
do_b_load(srsrc, tid - 128, bn, N, kb_next, bstride, Blds[nxt]);
}
{
const int so = ks_next << 3;
if (tid >= 128) {
do_b_scale_load(Bsc, tid - 128, bn, N, scs, so, Bslds[nxt]);
}
}
// MFMA on current step's data (from registers)
acc=mfma_fp4(A0,B0,acc,as0,bs0);
acc=mfma_fp4(A1,B1,acc,as1,bs1);
// Wait for next step's loads to complete
asm volatile("s_waitcnt vmcnt(0)");
__syncthreads();
}
// === EPILOGUE: compute last step from the last loaded buffer ===
{
const int cur = (KSPS - 1) & 1;
int4_vec A0,A1; int as0,as1;
{const int o=lds_swz(lm*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[cur][o]);A0={t.x,t.y,t.z,t.w};}
as0=(int)Asclds[cur][(lm<<3)+lk];
{const int o=lds_swz(lm*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[cur][o]);A1={t.x,t.y,t.z,t.w};}
as1=(int)Asclds[cur][(lm<<3)+4+lk];
const int br=(wid<<4)+lm;
int4_vec B0,B1; int bs0,bs1;
{const int o=lds_swz(br*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[cur][o]);B0={t.x,t.y,t.z,t.w};}
bs0=(int)Bslds[cur][br*8+lk];
{const int o=lds_swz(br*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[cur][o]);B1={t.x,t.y,t.z,t.w};}
bs1=(int)Bslds[cur][br*8+4+lk];
acc=mfma_fp4(A0,B0,acc,as0,bs0);
acc=mfma_fp4(A1,B1,acc,as1,bs1);
}
const int or_=bm+(lk<<2),oc=wn+lm;
if(oc<N){const float*ap=reinterpret_cast<const float*>(&acc);
if(C){for(int r=0;r<4;r++){int g=or_+r;if(g<M)C[g*N+oc]=__float2bfloat16(ap[r]);}}
else{float*w=ws+sid*M*N;for(int r=0;r<4;r++){int g=or_+r;if(g<M)w[g*N+oc]=ap[r];}}}
}
// Double-buffered version for large-M (BLOCK_M=32, 512 threads)
template<int KSPS>
__global__ __launch_bounds__(512, 1)
void gemm_bn64_m32_db(const __hip_bfloat16*__restrict__ A,const uint8_t*__restrict__ Bq,const uint8_t*__restrict__ Bsc,float*__restrict__ ws,__hip_bfloat16*__restrict__ C,int M,int N,int K){
constexpr int BLOCK_N=64;
const int wid=threadIdx.x>>6,lm=threadIdx.x&15,lk=(threadIdx.x&63)>>4,tid=threadIdx.x;
const int wg=wid>>2,wc=wid&3;
const int bn=blockIdx.x*BLOCK_N,bm=blockIdx.y*BLOCK_M_LARGE,col=bn+(wc<<4),sid=blockIdx.z;
const int bstride=K>>1,scs=K>>5;
__shared__ __align__(16) uint8_t Alds[2][BLOCK_M_LARGE*LDS_ROW];
__shared__ uint8_t Asclds[2][BLOCK_M_LARGE*8];
__shared__ __align__(16) uint8_t Blds[2][BLOCK_N*LDS_ROW];
__shared__ uint8_t Bslds[2][BLOCK_N*8];
const i32x4 srsrc=make_srsrc(Bq,N*bstride);
float4_vec acc={0,0,0,0};
const int ks_base = sid * KSPS;
// === PROLOGUE: load step 0 into buffer 0 ===
{
const int ks = ks_base;
const int ke = ks * DOUBLE_K, kb = ks * LDS_ROW;
if (tid < 256) {
const int qr = tid >> 3, qg = tid & 7;
const int gr = bm + qr, ko = ke + qg * 32;
uint32_t p0=0, p1=0, p2=0, p3=0; uint8_t asc = 0x7f;
if (gr < M) {
const __hip_bfloat16* s = A + gr * K + ko;
const int4 r0 = reinterpret_cast<const int4*>(s)[0];
const int4 r1 = reinterpret_cast<const int4*>(s)[1];
const int4 r2 = reinterpret_cast<const int4*>(s)[2];
const int4 r3 = reinterpret_cast<const int4*>(s)[3];
const uint32_t w0 = (uint32_t)r0.x, w1 = (uint32_t)r0.y, w2 = (uint32_t)r0.z, w3 = (uint32_t)r0.w;
const uint32_t w4 = (uint32_t)r1.x, w5 = (uint32_t)r1.y, w6 = (uint32_t)r1.z, w7 = (uint32_t)r1.w;
const uint32_t w8 = (uint32_t)r2.x, w9 = (uint32_t)r2.y, w10 = (uint32_t)r2.z, w11 = (uint32_t)r2.w;
const uint32_t w12 = (uint32_t)r3.x, w13 = (uint32_t)r3.y, w14 = (uint32_t)r3.z, w15 = (uint32_t)r3.w;
float l = 0;
update_max_word(w0,l); update_max_word(w1,l); update_max_word(w2,l); update_max_word(w3,l);
update_max_word(w4,l); update_max_word(w5,l); update_max_word(w6,l); update_max_word(w7,l);
update_max_word(w8,l); update_max_word(w9,l); update_max_word(w10,l); update_max_word(w11,l);
update_max_word(w12,l); update_max_word(w13,l); update_max_word(w14,l); update_max_word(w15,l);
float sf; compute_scale(l, asc, sf);
p0=pack4_words(w0,w1,w2,w3,sf);
p1=pack4_words(w4,w5,w6,w7,sf);
p2=pack4_words(w8,w9,w10,w11,sf);
p3=pack4_words(w12,w13,w14,w15,sf);
}
int4 wr; wr.x=(int)p0; wr.y=(int)p1; wr.z=(int)p2; wr.w=(int)p3;
*reinterpret_cast<int4*>(&Alds[0][lds_swz(qr*LDS_ROW+qg*16)]) = wr;
Asclds[0][qr*8+qg] = asc;
}
if (tid >= 256) {
const int btid = tid - 256;
constexpr int BNT = 256;
constexpr int BL = (BLOCK_N * LDS_ROW / 16 + BNT - 1) / BNT;
for (int ld = 0; ld < BL; ld++) {
const int f = (ld * BNT + btid) << 4;
const int r = f >> 7; const int gn = bn + r;
if (r < BLOCK_N && gn < N) {
const int sf = lds_swz(f), sc = sf & 127, ac = kb + sc;
llvm_amdgcn_raw_buffer_load_lds(srsrc,(as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Blds[0])+f),16,(gn>>4)*(bstride<<4)+(ac>>5)*512+((ac>>4)&1)*256+(gn&15)*16,0,0,0);
}
}
}
{
const int so = ks << 3;
if (tid >= 256 && tid < 384) {
const int stid = tid - 256;
const int row = stid & (BLOCK_N - 1), gb = (stid >> 6) << 2;
const int gn = bn + row;
if (gn < N) { const int rb=(gn>>5)*32*scs+(gn&15)*4+((gn>>4)&1);
for(int g=0;g<4;g++){const int grp=gb+g,ac=so+grp;Bslds[0][row*8+grp]=Bsc[rb+(ac&3)*64+((ac&7)>>2)*2+(ac>>3)*256];}
} else { for(int g=0;g<4;g++)Bslds[0][row*8+gb+g]=0x7f; }
}
}
asm volatile("s_waitcnt vmcnt(0)");
__syncthreads();
}
// === MAIN LOOP ===
#pragma unroll
for (int kstep = 1; kstep < KSPS; kstep++) {
const int cur = (kstep - 1) & 1;
const int nxt = kstep & 1;
const int ks_next = ks_base + kstep;
const int ke_next = ks_next * DOUBLE_K, kb_next = ks_next * LDS_ROW;
// Read current data from LDS
const int ar=(wg<<4)+lm;
int4_vec A0,A1; int as0,as1;
{const int o=lds_swz(ar*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[cur][o]);A0={t.x,t.y,t.z,t.w};}
as0=(int)Asclds[cur][ar*8+lk];
{const int o=lds_swz(ar*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[cur][o]);A1={t.x,t.y,t.z,t.w};}
as1=(int)Asclds[cur][ar*8+4+lk];
const int br=(wc<<4)+lm;
int4_vec B0,B1; int bs0,bs1;
{const int o=lds_swz(br*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[cur][o]);B0={t.x,t.y,t.z,t.w};}
bs0=(int)Bslds[cur][br*8+lk];
{const int o=lds_swz(br*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[cur][o]);B1={t.x,t.y,t.z,t.w};}
bs1=(int)Bslds[cur][br*8+4+lk];
// Issue loads for next step
if (tid < 256) {
const int qr = tid >> 3, qg = tid & 7;
const int gr = bm + qr, ko = ke_next + qg * 32;
uint32_t p0=0, p1=0, p2=0, p3=0; uint8_t asc = 0x7f;
if (gr < M) {
const __hip_bfloat16* s = A + gr * K + ko;
const int4 r0 = reinterpret_cast<const int4*>(s)[0];
const int4 r1 = reinterpret_cast<const int4*>(s)[1];
const int4 r2 = reinterpret_cast<const int4*>(s)[2];
const int4 r3 = reinterpret_cast<const int4*>(s)[3];
const uint32_t ww0 = (uint32_t)r0.x, ww1 = (uint32_t)r0.y, ww2 = (uint32_t)r0.z, ww3 = (uint32_t)r0.w;
const uint32_t ww4 = (uint32_t)r1.x, ww5 = (uint32_t)r1.y, ww6 = (uint32_t)r1.z, ww7 = (uint32_t)r1.w;
const uint32_t ww8 = (uint32_t)r2.x, ww9 = (uint32_t)r2.y, ww10 = (uint32_t)r2.z, ww11 = (uint32_t)r2.w;
const uint32_t ww12 = (uint32_t)r3.x, ww13 = (uint32_t)r3.y, ww14 = (uint32_t)r3.z, ww15 = (uint32_t)r3.w;
float l = 0;
update_max_word(ww0,l); update_max_word(ww1,l); update_max_word(ww2,l); update_max_word(ww3,l);
update_max_word(ww4,l); update_max_word(ww5,l); update_max_word(ww6,l); update_max_word(ww7,l);
update_max_word(ww8,l); update_max_word(ww9,l); update_max_word(ww10,l); update_max_word(ww11,l);
update_max_word(ww12,l); update_max_word(ww13,l); update_max_word(ww14,l); update_max_word(ww15,l);
float sf; compute_scale(l, asc, sf);
p0=pack4_words(ww0,ww1,ww2,ww3,sf);
p1=pack4_words(ww4,ww5,ww6,ww7,sf);
p2=pack4_words(ww8,ww9,ww10,ww11,sf);
p3=pack4_words(ww12,ww13,ww14,ww15,sf);
}
int4 wr; wr.x=(int)p0; wr.y=(int)p1; wr.z=(int)p2; wr.w=(int)p3;
*reinterpret_cast<int4*>(&Alds[nxt][lds_swz(qr*LDS_ROW+qg*16)]) = wr;
Asclds[nxt][qr*8+qg] = asc;
}
if (tid >= 256) {
const int btid = tid - 256;
constexpr int BNT = 256;
constexpr int BL = (BLOCK_N * LDS_ROW / 16 + BNT - 1) / BNT;
for (int ld = 0; ld < BL; ld++) {
const int f = (ld * BNT + btid) << 4;
const int r = f >> 7; const int gn = bn + r;
if (r < BLOCK_N && gn < N) {
const int sf = lds_swz(f), sc = sf & 127, ac = kb_next + sc;
llvm_amdgcn_raw_buffer_load_lds(srsrc,(as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Blds[nxt])+f),16,(gn>>4)*(bstride<<4)+(ac>>5)*512+((ac>>4)&1)*256+(gn&15)*16,0,0,0);
}
}
}
{
const int so = ks_next << 3;
if (tid >= 256 && tid < 384) {
const int stid = tid - 256;
const int row = stid & (BLOCK_N - 1), gb = (stid >> 6) << 2;
const int gn = bn + row;
if (gn < N) { const int rb=(gn>>5)*32*scs+(gn&15)*4+((gn>>4)&1);
for(int g=0;g<4;g++){const int grp=gb+g,ac=so+grp;Bslds[nxt][row*8+grp]=Bsc[rb+(ac&3)*64+((ac&7)>>2)*2+(ac>>3)*256];}
} else { for(int g=0;g<4;g++)Bslds[nxt][row*8+gb+g]=0x7f; }
}
}
// MFMA on current step's data
acc=mfma_fp4(A0,B0,acc,as0,bs0);
acc=mfma_fp4(A1,B1,acc,as1,bs1);
asm volatile("s_waitcnt vmcnt(0)");
__syncthreads();
}
// === EPILOGUE ===
{
const int cur = (KSPS - 1) & 1;
const int ar=(wg<<4)+lm;
int4_vec A0,A1; int as0,as1;
{const int o=lds_swz(ar*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[cur][o]);A0={t.x,t.y,t.z,t.w};}
as0=(int)Asclds[cur][ar*8+lk];
{const int o=lds_swz(ar*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[cur][o]);A1={t.x,t.y,t.z,t.w};}
as1=(int)Asclds[cur][ar*8+4+lk];
const int br=(wc<<4)+lm;
int4_vec B0,B1; int bs0,bs1;
{const int o=lds_swz(br*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[cur][o]);B0={t.x,t.y,t.z,t.w};}
bs0=(int)Bslds[cur][br*8+lk];
{const int o=lds_swz(br*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[cur][o]);B1={t.x,t.y,t.z,t.w};}
bs1=(int)Bslds[cur][br*8+4+lk];
acc=mfma_fp4(A0,B0,acc,as0,bs0);
acc=mfma_fp4(A1,B1,acc,as1,bs1);
}
const int or_=bm+(wg<<4)+(lk<<2),oc=col+lm;
if(oc<N){const float*ap=reinterpret_cast<const float*>(&acc);
if(C){for(int r=0;r<4;r++){int g=or_+r;if(g<M)C[g*N+oc]=__float2bfloat16(ap[r]);}}
else{float*w=ws+sid*M*N;for(int r=0;r<4;r++){int g=or_+r;if(g<M)w[g*N+oc]=ap[r];}}}
}
// Original gemm_bn64_m32 kernel (unchanged)
__global__ void gemm_bn64_m32(const __hip_bfloat16*__restrict__ A,const uint8_t*__restrict__ Bq,const uint8_t*__restrict__ Bsc,float*__restrict__ ws,__hip_bfloat16*__restrict__ C,int M,int N,int K,int ksps){
constexpr int BLOCK_N=64;
const int wid=threadIdx.x>>6,lm=threadIdx.x&15,lk=(threadIdx.x&63)>>4,tid=threadIdx.x;
const int wg=wid>>2,wc=wid&3;
const int bn=blockIdx.x*BLOCK_N,bm=blockIdx.y*BLOCK_M_LARGE,col=bn+(wc<<4),sid=blockIdx.z;const int bstride=K>>1,scs=K>>5;
__shared__ __align__(16) uint8_t Alds[BLOCK_M_LARGE*LDS_ROW];__shared__ uint8_t Asclds[BLOCK_M_LARGE*8];
__shared__ __align__(16) uint8_t Blds[BLOCK_N*LDS_ROW];__shared__ uint8_t Bslds[BLOCK_N*8];
const i32x4 srsrc=make_srsrc(Bq,N*bstride);float4_vec acc={0,0,0,0};
for(int ks=sid*ksps;ks<sid*ksps+ksps;ks++){const int ke=ks*DOUBLE_K,kb=ks*LDS_ROW;
if (tid < 256) {
const int qr = tid >> 3, qg = tid & 7;
const int gr = bm + qr, ko = ke + qg * 32;
uint32_t p0=0, p1=0, p2=0, p3=0; uint8_t asc = 0x7f;
if (gr < M) {
const __hip_bfloat16* s = A + gr * K + ko;
const int4 r0 = reinterpret_cast<const int4*>(s)[0];
const int4 r1 = reinterpret_cast<const int4*>(s)[1];
const int4 r2 = reinterpret_cast<const int4*>(s)[2];
const int4 r3 = reinterpret_cast<const int4*>(s)[3];
const uint32_t w0 = (uint32_t)r0.x, w1 = (uint32_t)r0.y, w2 = (uint32_t)r0.z, w3 = (uint32_t)r0.w;
const uint32_t w4 = (uint32_t)r1.x, w5 = (uint32_t)r1.y, w6 = (uint32_t)r1.z, w7 = (uint32_t)r1.w;
const uint32_t w8 = (uint32_t)r2.x, w9 = (uint32_t)r2.y, w10 = (uint32_t)r2.z, w11 = (uint32_t)r2.w;
const uint32_t w12 = (uint32_t)r3.x, w13 = (uint32_t)r3.y, w14 = (uint32_t)r3.z, w15 = (uint32_t)r3.w;
float l = 0;
update_max_word(w0,l); update_max_word(w1,l); update_max_word(w2,l); update_max_word(w3,l);
update_max_word(w4,l); update_max_word(w5,l); update_max_word(w6,l); update_max_word(w7,l);
update_max_word(w8,l); update_max_word(w9,l); update_max_word(w10,l); update_max_word(w11,l);
update_max_word(w12,l); update_max_word(w13,l); update_max_word(w14,l); update_max_word(w15,l);
float sf; compute_scale(l, asc, sf);
p0=pack4_words(w0,w1,w2,w3,sf);
p1=pack4_words(w4,w5,w6,w7,sf);
p2=pack4_words(w8,w9,w10,w11,sf);
p3=pack4_words(w12,w13,w14,w15,sf);
}
int4 wr; wr.x=(int)p0; wr.y=(int)p1; wr.z=(int)p2; wr.w=(int)p3;
*reinterpret_cast<int4*>(&Alds[lds_swz(qr*LDS_ROW+qg*16)]) = wr;
Asclds[qr*8+qg] = asc;
}
if (tid >= 256) {
const int btid = tid - 256;
constexpr int BNT = 256;
constexpr int BL = (BLOCK_N * LDS_ROW / 16 + BNT - 1) / BNT;
for (int ld = 0; ld < BL; ld++) {
const int f = (ld * BNT + btid) << 4;
const int r = f >> 7; const int gn = bn + r;
if (r < BLOCK_N && gn < N) {
const int sf = lds_swz(f), sc = sf & 127, ac = kb + sc;
llvm_amdgcn_raw_buffer_load_lds(srsrc,(as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Blds)+f),16,(gn>>4)*(bstride<<4)+(ac>>5)*512+((ac>>4)&1)*256+(gn&15)*16,0,0,0);
}
}
}
{const int so=ks<<3;
if (tid >= 256 && tid < 384) {
const int stid = tid - 256;
const int row = stid & (BLOCK_N - 1), gb = (stid >> 6) << 2;
const int gn = bn + row;
if (gn < N) { const int rb=(gn>>5)*32*scs+(gn&15)*4+((gn>>4)&1);
for(int g=0;g<4;g++){const int grp=gb+g,ac=so+grp;Bslds[row*8+grp]=Bsc[rb+(ac&3)*64+((ac&7)>>2)*2+(ac>>3)*256];}
} else { for(int g=0;g<4;g++)Bslds[row*8+gb+g]=0x7f; }
}}
asm volatile("s_waitcnt vmcnt(0)");__syncthreads();
const int ar=(wg<<4)+lm;
int4_vec A0,A1;int as0,as1;
{const int o=lds_swz(ar*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[o]);A0={t.x,t.y,t.z,t.w};}
as0=(int)Asclds[ar*8+lk];
{const int o=lds_swz(ar*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[o]);A1={t.x,t.y,t.z,t.w};}
as1=(int)Asclds[ar*8+4+lk];
const int br=(wc<<4)+lm;
int4_vec B0,B1;int bs0,bs1;
{const int o=lds_swz(br*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[o]);B0={t.x,t.y,t.z,t.w};}
bs0=(int)Bslds[br*8+lk];
{const int o=lds_swz(br*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[o]);B1={t.x,t.y,t.z,t.w};}
bs1=(int)Bslds[br*8+4+lk];
acc=mfma_fp4(A0,B0,acc,as0,bs0);
acc=mfma_fp4(A1,B1,acc,as1,bs1);
}
const int or_=bm+(wg<<4)+(lk<<2),oc=col+lm;
if(oc<N){const float*ap=reinterpret_cast<const float*>(&acc);
if(C){for(int r=0;r<4;r++){int g=or_+r;if(g<M)C[g*N+oc]=__float2bfloat16(ap[r]);}}
else{float*w=ws+sid*M*N;for(int r=0;r<4;r++){int g=or_+r;if(g<M)w[g*N+oc]=ap[r];}}}
}
template<int KSPS>
__global__ void gemm_bn64_m32_k(const __hip_bfloat16*__restrict__ A,const uint8_t*__restrict__ Bq,const uint8_t*__restrict__ Bsc,float*__restrict__ ws,__hip_bfloat16*__restrict__ C,int M,int N,int K){
constexpr int BLOCK_N=64;
const int wid=threadIdx.x>>6,lm=threadIdx.x&15,lk=(threadIdx.x&63)>>4,tid=threadIdx.x;
const int wg=wid>>2,wc=wid&3;
const int bn=blockIdx.x*BLOCK_N,bm=blockIdx.y*BLOCK_M_LARGE,col=bn+(wc<<4),sid=blockIdx.z;const int bstride=K>>1,scs=K>>5;
__shared__ __align__(16) uint8_t Alds[BLOCK_M_LARGE*LDS_ROW];__shared__ uint8_t Asclds[BLOCK_M_LARGE*8];
__shared__ __align__(16) uint8_t Blds[BLOCK_N*LDS_ROW];__shared__ uint8_t Bslds[BLOCK_N*8];
const i32x4 srsrc=make_srsrc(Bq,N*bstride);float4_vec acc={0,0,0,0};
#pragma unroll
for(int kstep=0;kstep<KSPS;kstep++){const int ks=sid*KSPS+kstep,ke=ks*DOUBLE_K,kb=ks*LDS_ROW;
if (tid < 256) {
const int qr = tid >> 3, qg = tid & 7;
const int gr = bm + qr, ko = ke + qg * 32;
uint32_t p0=0, p1=0, p2=0, p3=0; uint8_t asc = 0x7f;
if (gr < M) {
const __hip_bfloat16* s = A + gr * K + ko;
const int4 r0 = reinterpret_cast<const int4*>(s)[0];
const int4 r1 = reinterpret_cast<const int4*>(s)[1];
const int4 r2 = reinterpret_cast<const int4*>(s)[2];
const int4 r3 = reinterpret_cast<const int4*>(s)[3];
const uint32_t w0 = (uint32_t)r0.x, w1 = (uint32_t)r0.y, w2 = (uint32_t)r0.z, w3 = (uint32_t)r0.w;
const uint32_t w4 = (uint32_t)r1.x, w5 = (uint32_t)r1.y, w6 = (uint32_t)r1.z, w7 = (uint32_t)r1.w;
const uint32_t w8 = (uint32_t)r2.x, w9 = (uint32_t)r2.y, w10 = (uint32_t)r2.z, w11 = (uint32_t)r2.w;
const uint32_t w12 = (uint32_t)r3.x, w13 = (uint32_t)r3.y, w14 = (uint32_t)r3.z, w15 = (uint32_t)r3.w;
float l = 0;
update_max_word(w0,l); update_max_word(w1,l); update_max_word(w2,l); update_max_word(w3,l);
update_max_word(w4,l); update_max_word(w5,l); update_max_word(w6,l); update_max_word(w7,l);
update_max_word(w8,l); update_max_word(w9,l); update_max_word(w10,l); update_max_word(w11,l);
update_max_word(w12,l); update_max_word(w13,l); update_max_word(w14,l); update_max_word(w15,l);
float sf; compute_scale(l, asc, sf);
p0=pack4_words(w0,w1,w2,w3,sf);
p1=pack4_words(w4,w5,w6,w7,sf);
p2=pack4_words(w8,w9,w10,w11,sf);
p3=pack4_words(w12,w13,w14,w15,sf);
}
int4 wr; wr.x=(int)p0; wr.y=(int)p1; wr.z=(int)p2; wr.w=(int)p3;
*reinterpret_cast<int4*>(&Alds[lds_swz(qr*LDS_ROW+qg*16)]) = wr;
Asclds[qr*8+qg] = asc;
}
if (tid >= 256) {
const int btid = tid - 256;
constexpr int BNT = 256;
constexpr int BL = (BLOCK_N * LDS_ROW / 16 + BNT - 1) / BNT;
for (int ld = 0; ld < BL; ld++) {
const int f = (ld * BNT + btid) << 4;
const int r = f >> 7; const int gn = bn + r;
if (r < BLOCK_N && gn < N) {
const int sf = lds_swz(f), sc = sf & 127, ac = kb + sc;
llvm_amdgcn_raw_buffer_load_lds(srsrc,(as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Blds)+f),16,(gn>>4)*(bstride<<4)+(ac>>5)*512+((ac>>4)&1)*256+(gn&15)*16,0,0,0);
}
}
}
{const int so=ks<<3;
if (tid >= 256 && tid < 384) {
const int stid = tid - 256;
const int row = stid & (BLOCK_N - 1), gb = (stid >> 6) << 2;
const int gn = bn + row;
if (gn < N) { const int rb=(gn>>5)*32*scs+(gn&15)*4+((gn>>4)&1);
for(int g=0;g<4;g++){const int grp=gb+g,ac=so+grp;Bslds[row*8+grp]=Bsc[rb+(ac&3)*64+((ac&7)>>2)*2+(ac>>3)*256];}
} else { for(int g=0;g<4;g++)Bslds[row*8+gb+g]=0x7f; }
}}
asm volatile("s_waitcnt vmcnt(0)");__syncthreads();
const int ar=(wg<<4)+lm;
int4_vec A0,A1;int as0,as1;
{const int o=lds_swz(ar*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[o]);A0={t.x,t.y,t.z,t.w};}
as0=(int)Asclds[ar*8+lk];
{const int o=lds_swz(ar*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[o]);A1={t.x,t.y,t.z,t.w};}
as1=(int)Asclds[ar*8+4+lk];
const int br=(wc<<4)+lm;
int4_vec B0,B1;int bs0,bs1;
{const int o=lds_swz(br*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[o]);B0={t.x,t.y,t.z,t.w};}
bs0=(int)Bslds[br*8+lk];
{const int o=lds_swz(br*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[o]);B1={t.x,t.y,t.z,t.w};}
bs1=(int)Bslds[br*8+4+lk];
acc=mfma_fp4(A0,B0,acc,as0,bs0);
acc=mfma_fp4(A1,B1,acc,as1,bs1);
}
const int or_=bm+(wg<<4)+(lk<<2),oc=col+lm;
if(oc<N){const float*ap=reinterpret_cast<const float*>(&acc);
if(C){for(int r=0;r<4;r++){int g=or_+r;if(g<M)C[g*N+oc]=__float2bfloat16(ap[r]);}}
else{float*w=ws+sid*M*N;for(int r=0;r<4;r++){int g=or_+r;if(g<M)w[g*N+oc]=ap[r];}}}
}
template<int SK>
__global__ void reduce_k_t(const float*__restrict__ ws,__hip_bfloat16*__restrict__ C,int M,int N){
const int i=blockIdx.x*blockDim.x+threadIdx.x;if(i>=M*N)return;
const int stride=M*N;
float vals[SK];
#pragma unroll
for(int k=0;k<SK;k++) vals[k]=ws[k*stride+i];
float s=0;
#pragma unroll
for(int k=0;k<SK;k++) s+=vals[k];
C[i]=__float2bfloat16(s);
}
__global__ void reduce_k_gen(const float*__restrict__ ws,__hip_bfloat16*__restrict__ C,int M,int N,int sk){
const int i=blockIdx.x*blockDim.x+threadIdx.x;if(i>=M*N)return;
const int stride=M*N;const float*p=ws+i;
float s0=0,s1=0,s2=0,s3=0;
int k=0;
for(;k+3<sk;k+=4){s0+=p[k*stride];s1+=p[(k+1)*stride];s2+=p[(k+2)*stride];s3+=p[(k+3)*stride];}
for(;k<sk;k++)s0+=p[k*stride];
C[i]=__float2bfloat16(s0+s1+s2+s3);}
__global__ void reduce_k_14(const float*__restrict__ ws,__hip_bfloat16*__restrict__ C,int M,int N){
const int i=blockIdx.x*blockDim.x+threadIdx.x;if(i>=M*N)return;
const int stride=M*N;const float*p=ws+i;
float s0=p[0];
s0+=p[4*stride];
s0+=p[8*stride];
s0+=p[12*stride];
float s1=p[1*stride];
s1+=p[5*stride];
s1+=p[9*stride];
s1+=p[13*stride];
float s2=p[2*stride];
s2+=p[6*stride];
s2+=p[10*stride];
float s3=p[3*stride];
s3+=p[7*stride];
s3+=p[11*stride];
C[i]=__float2bfloat16(s0+s1+s2+s3);
}
torch::Tensor dispatch_gemm(torch::Tensor A, torch::Tensor Bq, torch::Tensor Bsc, int M, int N, int K) {
auto C = torch::empty({M, N}, torch::dtype(torch::kBFloat16).device(torch::kCUDA));
const auto* a_ptr = reinterpret_cast<const __hip_bfloat16*>(A.data_ptr());
const auto* bq_ptr = reinterpret_cast<const uint8_t*>(Bq.data_ptr());
const auto* bsc_ptr = reinterpret_cast<const uint8_t*>(Bsc.data_ptr());
auto* c_ptr = reinterpret_cast<__hip_bfloat16*>(C.data_ptr());
int ks = K / 256;
const int bn = 64;
const bool large_m = M >= 128;
const int block_m = large_m ? BLOCK_M_LARGE : BLOCK_M;
int bmn = ((N + bn - 1) / bn) * ((M + block_m - 1) / block_m);
int sk = 1;
if (!large_m && bmn < 304) {
if (ks > 8) {
int target = (462 + bmn - 1) / bmn;
int best = 1;
for (int s = 1; s <= ks; s++) {
if (ks % s == 0 && s <= target) best = s;
}
sk = best;
} else {
int target = (304 + bmn - 1) / bmn;
if (target < 1) target = 1;
int best = 1;
for (int s = 1; s <= ks; s++) {
if (ks % s == 0 && s <= target) best = s;
}
while (best > 1 && ks / best < 2) best /= 2;
sk = best > 0 ? best : 1;
}
}
if (sk == 1) {
if (large_m) {
dim3 b(512), g((N + bn - 1) / bn, (M + BLOCK_M_LARGE - 1) / BLOCK_M_LARGE, 1);
// Use double-buffered kernel for ks >= 4 (significant K-loop to overlap)
switch(ks) {
case 6: hipLaunchKernelGGL(gemm_bn64_m32_db<6>, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, (float*)nullptr, c_ptr, M, N, K); break;
default: hipLaunchKernelGGL(gemm_bn64_m32, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, (float*)nullptr, c_ptr, M, N, K, ks); break;
}
} else {
dim3 b(256), g((N + bn - 1) / bn, (M + BLOCK_M - 1) / BLOCK_M, 1);
// Use double-buffered kernel for ks >= 4
switch(ks) {
case 2: hipLaunchKernelGGL(gemm_bn64_k<2>, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, (float*)nullptr, c_ptr, M, N, K); break;
case 4: hipLaunchKernelGGL(gemm_bn64_db<4>, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, (float*)nullptr, c_ptr, M, N, K); break;
case 6: hipLaunchKernelGGL(gemm_bn64_db<6>, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, (float*)nullptr, c_ptr, M, N, K); break;
case 8: hipLaunchKernelGGL(gemm_bn64_db<8>, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, (float*)nullptr, c_ptr, M, N, K); break;
default: hipLaunchKernelGGL(gemm_bn64, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, (float*)nullptr, c_ptr, M, N, K, ks); break;
}
}
} else {
auto ws = torch::empty({sk, M, N}, torch::dtype(torch::kFloat32).device(torch::kCUDA));
auto* ws_ptr = reinterpret_cast<float*>(ws.data_ptr());
int ksps = ks / sk;
if (large_m) {
dim3 b(512), g((N + bn - 1) / bn, (M + BLOCK_M_LARGE - 1) / BLOCK_M_LARGE, sk);
if (ksps == 6) {
hipLaunchKernelGGL(gemm_bn64_m32_k<6>, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, ws_ptr, (__hip_bfloat16*)nullptr, M, N, K);
} else {
hipLaunchKernelGGL(gemm_bn64_m32, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, ws_ptr, (__hip_bfloat16*)nullptr, M, N, K, ksps);
}
} else {
dim3 b(256), g((N + bn - 1) / bn, (M + BLOCK_M - 1) / BLOCK_M, sk);
if (ksps == 2) {
hipLaunchKernelGGL(gemm_bn64_k<2>, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, ws_ptr, (__hip_bfloat16*)nullptr, M, N, K);
} else {
hipLaunchKernelGGL(gemm_bn64, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, ws_ptr, (__hip_bfloat16*)nullptr, M, N, K, ksps);
}
}
int n = M * N;
dim3 rb(256), rg((n+255)/256);
switch(sk) {
case 2: hipLaunchKernelGGL(reduce_k_t<2>, rg, rb, 0, 0, ws_ptr, c_ptr, M, N); break;
case 4: hipLaunchKernelGGL(reduce_k_t<4>, rg, rb, 0, 0, ws_ptr, c_ptr, M, N); break;
case 7: hipLaunchKernelGGL(reduce_k_t<7>, rg, rb, 0, 0, ws_ptr, c_ptr, M, N); break;
case 14: hipLaunchKernelGGL(reduce_k_14, rg, rb, 0, 0, ws_ptr, c_ptr, M, N); break;
default: hipLaunchKernelGGL(reduce_k_gen, rg, rb, 0, 0, ws_ptr, c_ptr, M, N, sk); break;
}
}
return C;
}
"""
_module = None
def _get_module():
global _module
if _module is None:
_module = load_inline(
name="gemm_v3",
cpp_sources=CPP_SOURCE,
cuda_sources=HIP_SOURCE,
functions=["dispatch_gemm"],
verbose=False,
extra_cuda_cflags=["-O3", "-fno-gpu-rdc", "-ffp-contract=fast"],
)
return _module
@torch.inference_mode()
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]
return _get_module().dispatch_gemm(A, B_shuffle.view(torch.uint8), B_scale_sh.view(torch.uint8), m, n, k)
scrolls · 946 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