submission 729970
aswinkumar · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1862 lines, June 9 Researcher Reciprocity License v1.0.
submission_v90.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-729970?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:cfc8ea814dcc3a8216a43bcd0c6a0d0c8b6b32d0750d529c859694580861459b
license declaredunknown
license concludedunknown
authorsaswinkumar
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
__device__ __forceinline__ void quant_a(uint32_t ap[16], v4i32& fp4, uint32_t& e8m0) {shared-memory
__shared__ float lds[4 * 256];split-k
__global__ void kern_splitk(vector-width = uint4
__device__ __forceinline__ uint4 lds_load_u128(lds_u32_ptr ptr) {Kernel source
submission_v90.py1862 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
v99_afirst: v98 + A-first load reorder. Issue A loads → GLL+Bsc → quant A.
Pre-barrier: pairs of (buffer_load B[s], global_load Bsc[s]) for predictable vmcnt drain.
Landmarks at key vmcnt points for later binary patching.
Pure C++ via <<<>>>.
"""
import os
os.environ.setdefault('PYTORCH_ROCM_ARCH', 'gfx950')
import torch
from torch.utils.cpp_extension import load_inline
input_t = tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]
output_t = torch.Tensor
HIP_SRC = r"""
#include <hip/hip_runtime.h>
#include <torch/extension.h>
namespace py = pybind11;
typedef int v4i32 __attribute__((__vector_size__(16)));
typedef float v4f32 __attribute__((__vector_size__(16)));
typedef const __attribute__((address_space(3))) uint32_t* lds_u32_ptr;
__device__ __forceinline__ float bf16_to_f32(uint16_t x) {
return __uint_as_float((uint32_t)x << 16);
}
__device__ __forceinline__ uint16_t f32_to_bf16(float x) {
uint32_t u = __float_as_uint(x);
u += 0x7FFFu + ((u >> 16) & 1u);
return (uint16_t)(u >> 16);
}
__device__ __forceinline__ uint4 lds_load_u128(lds_u32_ptr ptr) {
uint4 value;
asm volatile("ds_read_b128 %0, %1"
: "=v"(value)
: "v"(ptr)
: "memory");
return value;
}
__device__ __forceinline__ int bsc_idx(int bn, int bkg, int pk) {
return (bn>>5)*(pk*32)+(bkg>>3)*256+(bkg&3)*64+(bn&15)*4+((bkg&7)>>2)*2+((bn&31)>>4);
}
__device__ __forceinline__ void store_tile(
uint16_t* C, v4f32& a, int mt, int nt, int gid, int lid, int M, int N
) {
int on=nt*16+lid, om=mt*16+gid*4;
if(on<N){ for(int r=0;r<4;r++) if(om+r<M) C[(om+r)*N+on]=f32_to_bf16(a[r]); }
}
__device__ __forceinline__ void store_tile_full(
uint16_t* C, v4f32& a, int mt, int nt, int gid, int lid, int N
) {
int on=nt*16+lid, om=mt*16+gid*4;
C[(om+0)*N+on]=f32_to_bf16(a[0]);
C[(om+1)*N+on]=f32_to_bf16(a[1]);
C[(om+2)*N+on]=f32_to_bf16(a[2]);
C[(om+3)*N+on]=f32_to_bf16(a[3]);
}
// --- Original quant (still used by kern_fused for non-LDS path) ---
__device__ __forceinline__ void quant_a(uint32_t ap[16], v4i32& fp4, uint32_t& e8m0) {
uint32_t mx[16];
for(int i=0;i<16;i++) mx[i]=ap[i]&0x7FFF7FFFu;
for(int i=0;i<8;i++) asm volatile("v_pk_max_u16 %0,%1,%2":"=v"(mx[i]):"v"(mx[i]),"v"(mx[i+8]));
for(int i=0;i<4;i++) asm volatile("v_pk_max_u16 %0,%1,%2":"=v"(mx[i]):"v"(mx[i]),"v"(mx[i+4]));
for(int i=0;i<2;i++) asm volatile("v_pk_max_u16 %0,%1,%2":"=v"(mx[i]):"v"(mx[i]),"v"(mx[i+2]));
asm volatile("v_pk_max_u16 %0,%1,%2":"=v"(mx[0]):"v"(mx[0]),"v"(mx[1]));
uint32_t h=mx[0]>>16,l=mx[0]&0xFFFFu; uint16_t am=(uint16_t)(h>l?h:l);
float sf;
if(am==0u){e8m0=0u;sf=0.0f;}
else{unsigned ai=__float_as_uint(bf16_to_f32(am));ai=(ai+0x200000u)&0xFF800000u;
int su=(int)((ai>>23)&0xFFu)-129;su=max(-127,min(127,su));
e8m0=(uint32_t)(su+127);sf=__uint_as_float(e8m0<<23);}
for(int v=0;v<4;v++){uint32_t r=0;
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2":"+v"(r):"v"(ap[v*4+0]),"v"(sf));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,0]":"+v"(r):"v"(ap[v*4+1]),"v"(sf));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,0,1]":"+v"(r):"v"(ap[v*4+2]),"v"(sf));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,1]":"+v"(r):"v"(ap[v*4+3]),"v"(sf));
fp4[v]=(int)r;}
}
// --- Optimized quant: takes 4 uint4 directly, no intermediate array ---
__device__ __forceinline__ void quant_a_direct(
uint4 c0, uint4 c1, uint4 c2, uint4 c3,
v4i32& fp4, uint32_t& e8m0
) {
// Absolute value (clear sign bits of packed bf16)
const uint32_t ABSMASK = 0x7FFF7FFFu;
uint32_t m0=c0.x&ABSMASK, m1=c0.y&ABSMASK, m2=c0.z&ABSMASK, m3=c0.w&ABSMASK;
uint32_t m4=c1.x&ABSMASK, m5=c1.y&ABSMASK, m6=c1.z&ABSMASK, m7=c1.w&ABSMASK;
uint32_t m8=c2.x&ABSMASK, m9=c2.y&ABSMASK, mA=c2.z&ABSMASK, mB=c2.w&ABSMASK;
uint32_t mC=c3.x&ABSMASK, mD=c3.y&ABSMASK, mE=c3.z&ABSMASK, mF=c3.w&ABSMASK;
// Tree reduction: 16→8→4→2→1 using pk_max_u16
// NOTE: "+v" (read-write) because %0 is both source and destination
asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m0):"v"(m8));
asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m1):"v"(m9));
asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m2):"v"(mA));
asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m3):"v"(mB));
asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m4):"v"(mC));
asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m5):"v"(mD));
asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m6):"v"(mE));
asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m7):"v"(mF));
asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m0):"v"(m4));
asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m1):"v"(m5));
asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m2):"v"(m6));
asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m3):"v"(m7));
asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m0):"v"(m2));
asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m1):"v"(m3));
asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m0):"v"(m1));
// Extract scalar max from packed result
uint32_t h=m0>>16, l=m0&0xFFFFu;
uint16_t am=(uint16_t)(h>l?h:l);
// Compute E8M0 scale
float sf;
if(am==0u){e8m0=0u;sf=0.0f;}
else{unsigned ai=__float_as_uint(bf16_to_f32(am));ai=(ai+0x200000u)&0xFF800000u;
int su=(int)((ai>>23)&0xFFu)-129;su=max(-127,min(127,su));
e8m0=(uint32_t)(su+127);sf=__uint_as_float(e8m0<<23);}
// Convert bf16 pairs to FP4 — operate directly on uint4 fields
uint32_t r0=0, r1=0, r2=0, r3=0;
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2":"+v"(r0):"v"(c0.x),"v"(sf));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,0]":"+v"(r0):"v"(c0.y),"v"(sf));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,0,1]":"+v"(r0):"v"(c0.z),"v"(sf));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,1]":"+v"(r0):"v"(c0.w),"v"(sf));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2":"+v"(r1):"v"(c1.x),"v"(sf));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,0]":"+v"(r1):"v"(c1.y),"v"(sf));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,0,1]":"+v"(r1):"v"(c1.z),"v"(sf));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,1]":"+v"(r1):"v"(c1.w),"v"(sf));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2":"+v"(r2):"v"(c2.x),"v"(sf));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,0]":"+v"(r2):"v"(c2.y),"v"(sf));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,0,1]":"+v"(r2):"v"(c2.z),"v"(sf));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,1]":"+v"(r2):"v"(c2.w),"v"(sf));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2":"+v"(r3):"v"(c3.x),"v"(sf));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,0]":"+v"(r3):"v"(c3.y),"v"(sf));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,0,1]":"+v"(r3):"v"(c3.z),"v"(sf));
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,1]":"+v"(r3):"v"(c3.w),"v"(sf));
fp4[0]=(int)r0; fp4[1]=(int)r1; fp4[2]=(int)r2; fp4[3]=(int)r3;
}
// ==========================================================================
// KERNEL 1: Fused quant+GEMM (single wave, pipelined, for K<=1024)
// ==========================================================================
__launch_bounds__(64)
__global__ void kern_fused(
const uint32_t* __restrict__ A, const uint8_t* __restrict__ Bs,
const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C,
const int M, const int N, const int K, const int kt, const int pk
) {
const int lane=threadIdx.x, mt=blockIdx.x, nt=blockIdx.y;
const int lid=lane&15, gid=lane>>4, mg=mt*16+lid, Kh=K>>1;
const int bn=nt*16+lid, bb=nt*(kt*512);
v4f32 acc={0,0,0,0};
uint32_t ap[16];
if(mg<M){const uint4* av=reinterpret_cast<const uint4*>(A+mg*Kh+(gid*32>>1));
uint4 c0=av[0],c1=av[1],c2=av[2],c3=av[3];
ap[0]=c0.x;ap[1]=c0.y;ap[2]=c0.z;ap[3]=c0.w;ap[4]=c1.x;ap[5]=c1.y;ap[6]=c1.z;ap[7]=c1.w;
ap[8]=c2.x;ap[9]=c2.y;ap[10]=c2.z;ap[11]=c2.w;ap[12]=c3.x;ap[13]=c3.y;ap[14]=c3.z;ap[15]=c3.w;
} else { for(int i=0;i<16;i++) ap[i]=0u; }
v4i32 af; uint32_t ae; quant_a(ap,af,ae);
uint4 bc=*reinterpret_cast<const uint4*>(Bs+bb+(gid/2)*512+(gid&1)*256+lid*16);
v4i32 bf; bf[0]=(int)bc.x;bf[1]=(int)bc.y;bf[2]=(int)bc.z;bf[3]=(int)bc.w;
uint32_t be=(uint32_t)Bsc[bsc_idx(bn,gid,pk)];
for(int ki=1,Ki=K>>7;ki<Ki;ki++){int kb=ki<<7;
uint32_t apn[16];
if(mg<M){const uint4* avn=reinterpret_cast<const uint4*>(A+mg*Kh+((kb+gid*32)>>1));
uint4 d0=avn[0],d1=avn[1],d2=avn[2],d3=avn[3];
apn[0]=d0.x;apn[1]=d0.y;apn[2]=d0.z;apn[3]=d0.w;apn[4]=d1.x;apn[5]=d1.y;apn[6]=d1.z;apn[7]=d1.w;
apn[8]=d2.x;apn[9]=d2.y;apn[10]=d2.z;apn[11]=d2.w;apn[12]=d3.x;apn[13]=d3.y;apn[14]=d3.z;apn[15]=d3.w;
} else { for(int i=0;i<16;i++) apn[i]=0u; }
uint4 bcn=*reinterpret_cast<const uint4*>(Bs+bb+(kb/64+gid/2)*512+(gid&1)*256+lid*16);
v4i32 bfn; bfn[0]=(int)bcn.x;bfn[1]=(int)bcn.y;bfn[2]=(int)bcn.z;bfn[3]=(int)bcn.w;
uint32_t ben=(uint32_t)Bsc[bsc_idx(bn,kb/32+gid,pk)];
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
:"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(be));
quant_a(apn,af,ae); bf=bfn; be=ben;
}
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
:"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(be));
store_tile(C,acc,mt,nt,gid,lid,M,N);
}
// ==========================================================================
// KERNEL 1b: Template-specialized fused quant+GEMM
// Template params: cM, cN, cK are compile-time shape constants.
// Compiler can constant-fold Ki, Kh, kt, pk, address math, eliminate branches.
// For K<=1024: 1 K-tile per wave. For K>1024: loop over Ki/4 tiles per wave.
// ==========================================================================
template<int cM, int cN, int cK>
__launch_bounds__(256)
__global__ void kern_fused_spec(
const uint32_t* __restrict__ A, const uint8_t* __restrict__ Bs,
const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C
) {
// All shape-dependent values are compile-time constants
constexpr int cKi = cK >> 7; // K-tiles
constexpr int cKh = cK >> 1; // K/2
constexpr int ckt = cK / 64;
constexpr int cpk = ((cK/32 + 7) / 8) * 8;
const int tid = threadIdx.x;
const int wave_id = tid >> 6;
const int lane = tid & 63;
const int mt = blockIdx.x, nt = blockIdx.y;
const int lid = lane & 15, gid = lane >> 4;
const int mg = mt*16 + lid;
// Compile-time K-tile assignment per wave
constexpr int tiles_per_wave = (cKi + 3) / 4; // ceil(Ki/4)
// Compile-time: which K-tiles this wave handles
const int ki_start = wave_id * tiles_per_wave;
const int ki_end_raw = ki_start + tiles_per_wave;
const int ki_end = (ki_end_raw < cKi) ? ki_end_raw : cKi;
const int bb = nt * (ckt * 512);
const int bn = nt*16 + lid;
v4f32 acc = {0,0,0,0};
if (ki_start < cKi) { // compile-time: waves 0-3 all active when Ki>=4
// Prologue
int kb = ki_start << 7;
uint32_t ap[16];
if (mg < cM) { // compile-time M enables branch elimination
const uint4* av = reinterpret_cast<const uint4*>(A + mg*cKh + ((kb+(gid<<5))>>1));
uint4 c0=av[0],c1=av[1],c2=av[2],c3=av[3];
ap[0]=c0.x;ap[1]=c0.y;ap[2]=c0.z;ap[3]=c0.w;
ap[4]=c1.x;ap[5]=c1.y;ap[6]=c1.z;ap[7]=c1.w;
ap[8]=c2.x;ap[9]=c2.y;ap[10]=c2.z;ap[11]=c2.w;
ap[12]=c3.x;ap[13]=c3.y;ap[14]=c3.z;ap[15]=c3.w;
} else { for(int i=0;i<16;i++) ap[i]=0u; }
v4i32 af; uint32_t ae;
quant_a(ap, af, ae);
int b_off = bb + ((kb>>6)+(gid>>1))*512 + (gid&1)*256 + (lid<<4);
uint4 bc = *reinterpret_cast<const uint4*>(Bs + b_off);
v4i32 bf; bf[0]=(int)bc.x;bf[1]=(int)bc.y;bf[2]=(int)bc.z;bf[3]=(int)bc.w;
uint32_t be = (uint32_t)Bsc[bsc_idx(bn, (kb>>5)+gid, cpk)];
// Pipelined K-loop (compiler unrolls when tiles_per_wave is small constant)
for (int ki = ki_start+1; ki < ki_end; ki++) {
int kbn = ki << 7;
uint32_t apn[16];
if (mg < cM) {
const uint4* avn = reinterpret_cast<const uint4*>(A + mg*cKh + ((kbn+(gid<<5))>>1));
uint4 d0=avn[0],d1=avn[1],d2=avn[2],d3=avn[3];
apn[0]=d0.x;apn[1]=d0.y;apn[2]=d0.z;apn[3]=d0.w;
apn[4]=d1.x;apn[5]=d1.y;apn[6]=d1.z;apn[7]=d1.w;
apn[8]=d2.x;apn[9]=d2.y;apn[10]=d2.z;apn[11]=d2.w;
apn[12]=d3.x;apn[13]=d3.y;apn[14]=d3.z;apn[15]=d3.w;
} else { for(int i=0;i<16;i++) apn[i]=0u; }
int b_offn = bb + ((kbn>>6)+(gid>>1))*512 + (gid&1)*256 + (lid<<4);
uint4 bcn = *reinterpret_cast<const uint4*>(Bs + b_offn);
v4i32 bfn; bfn[0]=(int)bcn.x;bfn[1]=(int)bcn.y;bfn[2]=(int)bcn.z;bfn[3]=(int)bcn.w;
uint32_t ben = (uint32_t)Bsc[bsc_idx(bn, (kbn>>5)+gid, cpk)];
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
:"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(be));
quant_a(apn, af, ae); bf=bfn; be=ben;
}
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
:"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(be));
}
// LDS reduction — transposed layout
__shared__ float lds[4 * 256];
int lds_w = wave_id * 256 + lane;
lds[lds_w]=acc[0]; lds[lds_w+64]=acc[1]; lds[lds_w+128]=acc[2]; lds[lds_w+192]=acc[3];
__syncthreads();
if (wave_id == 0) {
v4f32 sum;
sum[0]=lds[lane]+lds[lane+256]+lds[lane+512]+lds[lane+768];
sum[1]=lds[lane+64]+lds[lane+320]+lds[lane+576]+lds[lane+832];
sum[2]=lds[lane+128]+lds[lane+384]+lds[lane+640]+lds[lane+896];
sum[3]=lds[lane+192]+lds[lane+448]+lds[lane+704]+lds[lane+960];
store_tile(C, sum, mt, nt, gid, lid, cM, cN);
}
}
// ==========================================================================
// KERNEL 2: Warp-parallel A quantization
// ==========================================================================
#define QPG 16
__global__ void kern_quant(
const uint32_t* __restrict__ A, unsigned char* __restrict__ fp4,
unsigned char* __restrict__ sc,
const int M, const int K, const int Ksc, const int Mpad, const int Kscpad
) {
const int tid=blockIdx.x*blockDim.x+threadIdx.x;
const int gid=tid/QPG, lig=tid%QPG;
if(gid>=M*Ksc)return;
const int row=gid/Ksc, grp=gid%Ksc;
uint32_t p=*(A+row*(K/2)+grp*16+lig);
uint32_t mx=p&0x7FFF7FFFu;
for(int o=8;o>=1;o>>=1){uint32_t ot=__shfl_xor(mx,o);
asm volatile("v_pk_max_u16 %0,%1,%2":"=v"(mx):"v"(mx),"v"(ot));}
uint32_t h=mx>>16,l=mx&0xFFFFu;uint16_t am=(uint16_t)(h>l?h:l);
float sf;unsigned char e8m0;
if(am==0u){e8m0=0;sf=0.0f;}
else{unsigned ai=__float_as_uint(bf16_to_f32(am));ai=(ai+0x200000u)&0xFF800000u;
int su=(int)((ai>>23)&0xFFu)-129;su=max(-127,min(127,su));
e8m0=(unsigned char)(su+127);sf=__uint_as_float((uint32_t)e8m0<<23);}
uint32_t fp4b; asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2":"=v"(fp4b):"v"(p),"v"(sf));
fp4[row*(K/2)+grp*16+lig]=(unsigned char)(fp4b&0xFFu);
if(lig==0){int sn=Kscpad;
int si=(row/32)*(32*sn)+(grp/8)*256+(grp%4)*64+(row%16)*4+((grp%8)/4)*2+((row%32)/16);
if(si<Mpad*Kscpad)sc[si]=e8m0;}
}
// ==========================================================================
// KERNEL 3: Lean GEMM (pre-quantized A)
// ==========================================================================
__launch_bounds__(64)
__global__ void kern_gemm(
const uint8_t* __restrict__ Afp4, const uint8_t* __restrict__ Asc,
const uint8_t* __restrict__ Bs, const uint8_t* __restrict__ Bsc,
uint16_t* __restrict__ C,
const int M, const int N, const int K, const int kt, const int pkA, const int pkB
) {
const int lane=threadIdx.x,mt=blockIdx.x,nt=blockIdx.y;
const int lid=lane&15,gid=lane>>4,mg=mt*16+lid,Kh=K>>1;
const int bn=nt*16+lid,bb=nt*(kt*512);
v4f32 acc={0,0,0,0};
v4i32 af,bf;uint32_t ae,be;
if(mg<M){uint4 ac=*reinterpret_cast<const uint4*>(Afp4+mg*Kh+(gid*32)/2);
af[0]=(int)ac.x;af[1]=(int)ac.y;af[2]=(int)ac.z;af[3]=(int)ac.w;}
else{af[0]=af[1]=af[2]=af[3]=0;}
ae=mg<M?(uint32_t)Asc[bsc_idx(mg,gid,pkA)]:0u;
{uint4 bc=*reinterpret_cast<const uint4*>(Bs+bb+(gid/2)*512+(gid&1)*256+lid*16);
bf[0]=(int)bc.x;bf[1]=(int)bc.y;bf[2]=(int)bc.z;bf[3]=(int)bc.w;
be=(uint32_t)Bsc[bsc_idx(bn,gid,pkB)];}
for(int ki=1,Ki=K>>7;ki<Ki;ki++){int kb=ki<<7;
v4i32 afn,bfn;uint32_t aen,ben;
if(mg<M){uint4 acn=*reinterpret_cast<const uint4*>(Afp4+mg*Kh+(kb+gid*32)/2);
afn[0]=(int)acn.x;afn[1]=(int)acn.y;afn[2]=(int)acn.z;afn[3]=(int)acn.w;}
else{afn[0]=afn[1]=afn[2]=afn[3]=0;}
aen=mg<M?(uint32_t)Asc[bsc_idx(mg,kb/32+gid,pkA)]:0u;
{uint4 bcn=*reinterpret_cast<const uint4*>(Bs+bb+(kb/64+gid/2)*512+(gid&1)*256+lid*16);
bfn[0]=(int)bcn.x;bfn[1]=(int)bcn.y;bfn[2]=(int)bcn.z;bfn[3]=(int)bcn.w;
ben=(uint32_t)Bsc[bsc_idx(bn,kb/32+gid,pkB)];}
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
:"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(be));
af=afn;ae=aen;bf=bfn;be=ben;}
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
:"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(be));
store_tile(C,acc,mt,nt,gid,lid,M,N);
}
// ==========================================================================
// KERNEL 3b: Lean GEMM with 32x32x64 MFMA (pre-quantized A)
// B-sharing across 32 M-rows happens WITHIN the wave — no LDS, no barriers!
// For M>=32 with K>1024. 4x fewer tiles, 2x less B bandwidth per output element.
// ==========================================================================
typedef float v16f32 __attribute__((__vector_size__(64)));
// ==========================================================================
// KERNEL 4: split-K GEMM with atomicAdd
// ==========================================================================
__launch_bounds__(64)
__global__ void kern_splitk(
const uint8_t* __restrict__ Afp4, const uint8_t* __restrict__ Asc,
const uint8_t* __restrict__ Bs, const uint8_t* __restrict__ Bsc,
float* __restrict__ temp,
const int M, const int N, const int K,
const int kt, const int pkA, const int pkB,
const int n_tiles, const int k_split
) {
int bid=blockIdx.x,tid_t=bid/k_split,sid=bid%k_split;
int mt=tid_t/n_tiles,nt=tid_t%n_tiles;
const int lane=threadIdx.x,lid=lane&15,gid=lane>>4,mg=mt*16+lid,Kh=K>>1;
const int bn=nt*16+lid,bb=nt*(kt*512);
int Ki=K>>7,kip=Ki/k_split,kis=sid*kip,kie=(sid==k_split-1)?Ki:kis+kip;
v4f32 acc={0,0,0,0};
if(kis<kie){int kb0=kis<<7;
v4i32 af,bf;uint32_t ae,be;
if(mg<M){uint4 ac=*reinterpret_cast<const uint4*>(Afp4+mg*Kh+(kb0+gid*32)/2);
af[0]=(int)ac.x;af[1]=(int)ac.y;af[2]=(int)ac.z;af[3]=(int)ac.w;}
else{af[0]=af[1]=af[2]=af[3]=0;}
ae=mg<M?(uint32_t)Asc[bsc_idx(mg,kb0/32+gid,pkA)]:0u;
{uint4 bc=*reinterpret_cast<const uint4*>(Bs+bb+(kb0/64+gid/2)*512+(gid&1)*256+lid*16);
bf[0]=(int)bc.x;bf[1]=(int)bc.y;bf[2]=(int)bc.z;bf[3]=(int)bc.w;
be=(uint32_t)Bsc[bsc_idx(bn,kb0/32+gid,pkB)];}
for(int ki=kis+1;ki<kie;ki++){int kb=ki<<7;
v4i32 afn,bfn;uint32_t aen,ben;
if(mg<M){uint4 acn=*reinterpret_cast<const uint4*>(Afp4+mg*Kh+(kb+gid*32)/2);
afn[0]=(int)acn.x;afn[1]=(int)acn.y;afn[2]=(int)acn.z;afn[3]=(int)acn.w;}
else{afn[0]=afn[1]=afn[2]=afn[3]=0;}
aen=mg<M?(uint32_t)Asc[bsc_idx(mg,kb/32+gid,pkA)]:0u;
{uint4 bcn=*reinterpret_cast<const uint4*>(Bs+bb+(kb/64+gid/2)*512+(gid&1)*256+lid*16);
bfn[0]=(int)bcn.x;bfn[1]=(int)bcn.y;bfn[2]=(int)bcn.z;bfn[3]=(int)bcn.w;
ben=(uint32_t)Bsc[bsc_idx(bn,kb/32+gid,pkB)];}
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
:"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(be));
af=afn;ae=aen;bf=bfn;be=ben;}
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
:"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(be));}
int on=nt*16+lid,om=mt*16+gid*4;
if(on<N){for(int r=0;r<4;r++)if(om+r<M) atomicAdd(&temp[(om+r)*N+on],acc[r]);}
}
__global__ void kern_convert(const float* src, uint16_t* dst, int n) {
int i=blockIdx.x*blockDim.x+threadIdx.x; if(i>=n)return;
uint32_t u=__float_as_uint(src[i]); u+=0x7FFFu+((u>>16)&1u); dst[i]=(uint16_t)(u>>16);
}
// ==========================================================================
// KERNEL 6b: Hybrid inline-A-quant + GLL-B for Shape 1.
// Each wave owns one K chunk, preloads its B chunk into private LDS, then
// quantizes A inline while feeding the MFMA loop. The only barrier is the
// final LDS reduction across split waves.
// ==========================================================================
template<int BATCH, int NUM_SPLITS>
__launch_bounds__(512, 2)
__global__ void kern_hybrid_s1(
const uint32_t* __restrict__ A, const uint8_t* __restrict__ Bs,
const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C,
const int M, const int N, const int K, const int kt, const int pk
) {
static_assert(BATCH <= 14, "supports up to 14 staged B tiles");
constexpr int N_WAVES = NUM_SPLITS;
const int tid = threadIdx.x;
const int wave_id = tid >> 6;
const int lane = tid & 63;
const int lid = lane & 15;
const int gid = lane >> 4;
// 2D grid: blockIdx.x = mt (M-tile), blockIdx.y = nt (N-tile)
// For Shape 1 (M=16): grid(1, 132) → mt always 0
// For Shape 4 (M=64): grid(4, 448) → mt = 0..3
const int mt = blockIdx.x;
const int nt = blockIdx.y;
const int Kh = K >> 1;
const int Ki = K >> 7;
const int Ki_chunk = Ki / NUM_SPLITS;
const int k_start = wave_id * Ki_chunk * 128;
extern __shared__ uint32_t lds_hybrid[];
const int MY_B_DW = wave_id * BATCH * 256;
const int my_b_byte = MY_B_DW * 4;
const int B_LDS_DW_TOTAL = N_WAVES * BATCH * 256;
float* red_buf = reinterpret_cast<float*>(&lds_hybrid[B_LDS_DW_TOTAL]);
const int bn = nt * 16 + lid;
const int bb = nt * (kt * 512);
const int nt_total = (N + 15) / 16;
const bool valid_wave = (nt < nt_total);
uint32_t be0=0,be1=0,be2=0,be3=0,be4=0,be5=0,be6=0,be7=0;
uint32_t be8=0,be9=0,be10=0,be11=0,be12=0,be13=0;
if (valid_wave) {
for (int s = 0; s < BATCH && s < Ki_chunk; s++) {
int kb = k_start + (s << 7);
const uint8_t* baddr = Bs + bb + (kb/64+gid/2)*512 + (gid&1)*256 + lid*16;
uint32_t m0v = (uint32_t)(my_b_byte + s * 1024);
asm volatile("s_mov_b32 m0, %0\n\t"
"s_nop 4\n\t"
"global_load_lds_dwordx4 %1, off nt"
:: "s"(__builtin_amdgcn_readfirstlane(m0v)), "v"(baddr) : "memory");
}
#define LOAD_BE_H(I) do { int _kb=k_start+((I)<<7); be##I=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; } while(0)
LOAD_BE_H(0); if(Ki_chunk>1) LOAD_BE_H(1); if(Ki_chunk>2) LOAD_BE_H(2); if(Ki_chunk>3) LOAD_BE_H(3);
if(Ki_chunk>4) LOAD_BE_H(4); if(Ki_chunk>5) LOAD_BE_H(5); if(Ki_chunk>6) LOAD_BE_H(6); if(Ki_chunk>7) LOAD_BE_H(7);
if(Ki_chunk>8) LOAD_BE_H(8); if(Ki_chunk>9) LOAD_BE_H(9); if(Ki_chunk>10) LOAD_BE_H(10); if(Ki_chunk>11) LOAD_BE_H(11);
if(Ki_chunk>12) LOAD_BE_H(12); if(Ki_chunk>13) LOAD_BE_H(13);
#undef LOAD_BE_H
}
if (!valid_wave) return;
asm volatile("s_waitcnt vmcnt(%c0)" :: "n"(BATCH) : "memory");
const int mg = mt * 16 + lid;
auto load_a_tile = [&](int ki_local, uint4& c0, uint4& c1, uint4& c2, uint4& c3) {
int kb = k_start + (ki_local << 7);
if (mg < M) {
const uint4* av = reinterpret_cast<const uint4*>(A + mg*Kh + ((kb + (gid<<5)) >> 1));
c0 = av[0]; c1 = av[1]; c2 = av[2]; c3 = av[3];
} else { c0=c1=c2=c3=make_uint4(0,0,0,0); }
};
auto load_b_tile = [&](int slot, v4i32& bf_out) {
int boff = MY_B_DW + slot * 256 + lane * 4;
uint4 bc = lds_load_u128((lds_u32_ptr)(&lds_hybrid[boff]));
bf_out[0]=(int)bc.x; bf_out[1]=(int)bc.y; bf_out[2]=(int)bc.z; bf_out[3]=(int)bc.w;
};
auto get_be = [&](int idx) -> uint32_t {
switch (idx) {
case 0: return be0; case 1: return be1; case 2: return be2; case 3: return be3;
case 4: return be4; case 5: return be5; case 6: return be6; case 7: return be7;
case 8: return be8; case 9: return be9; case 10: return be10; case 11: return be11;
case 12: return be12; case 13: return be13;
default: return 0u;
}
};
v4f32 acc = {0,0,0,0};
uint4 a0c0, a0c1, a0c2, a0c3;
load_a_tile(0, a0c0, a0c1, a0c2, a0c3);
v4i32 af; uint32_t ae;
quant_a_direct(a0c0, a0c1, a0c2, a0c3, af, ae);
v4i32 bf;
load_b_tile(0, bf);
for (int ki_local = 1; ki_local < Ki_chunk; ++ki_local) {
int far_ki = ki_local + BATCH - 1;
if (far_ki < Ki_chunk) {
int kb = k_start + (far_ki << 7);
const uint8_t* baddr = Bs + bb + (kb/64+gid/2)*512 + (gid&1)*256 + lid*16;
uint32_t m0v = (uint32_t)(my_b_byte + ((ki_local - 1) % BATCH) * 1024);
asm volatile("s_mov_b32 m0, %0\n\t"
"s_nop 4\n\t"
"global_load_lds_dwordx4 %1, off nt"
:: "s"(__builtin_amdgcn_readfirstlane(m0v)), "v"(baddr) : "memory");
}
uint4 next_c0, next_c1, next_c2, next_c3;
load_a_tile(ki_local, next_c0, next_c1, next_c2, next_c3);
v4i32 bfn;
load_b_tile(ki_local % BATCH, bfn);
uint32_t be_cur = get_be(ki_local - 1);
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
:"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(be_cur));
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
quant_a_direct(next_c0, next_c1, next_c2, next_c3, af, ae);
bf = bfn;
}
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
:"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(get_be(Ki_chunk - 1)));
int red_off = wave_id * 256 + lane * 4;
red_buf[red_off+0] = acc[0]; red_buf[red_off+1] = acc[1];
red_buf[red_off+2] = acc[2]; red_buf[red_off+3] = acc[3];
__syncthreads();
if (wave_id == 0) {
float s0=0, s1=0, s2=0, s3=0;
for (int w = 0; w < NUM_SPLITS; w++) {
int off = w * 256 + lane * 4;
s0 += red_buf[off+0]; s1 += red_buf[off+1];
s2 += red_buf[off+2]; s3 += red_buf[off+3];
}
v4f32 result; result[0]=s0; result[1]=s1; result[2]=s2; result[3]=s3;
store_tile(C, result, mt, nt, gid, lid, M, N);
}
}
// ==========================================================================
// Dedicated kernel: M=64 N=7168 K=2048 Ki=16 N_N=8 BATCH=10 SPLIT=10
__launch_bounds__(512)
__global__ void kern_s4(
const uint32_t* __restrict__ A, const uint8_t* __restrict__ Bs,
const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C,
const int M, const int N, const int K, const int kt, const int pk
) {
constexpr int N_M = 1, N_N = 8, BATCH = 10;
constexpr int N_WAVES = 8;
constexpr int A_TILE = 384;
constexpr int A_FP4_STRIDE = 5;
constexpr int Ki = 16;
const int tid = threadIdx.x;
const int wave_id = tid >> 6;
const int lane = tid & 63;
const int lid = lane & 15;
const int gid = lane >> 4;
const int mt = blockIdx.x;
const int nt = blockIdx.y * N_N + wave_id;
const int Kh = K >> 1;
extern __shared__ uint32_t lds_gll[];
const int B_LDS_DW = Ki * A_TILE;
const int MY_B_DW = B_LDS_DW + wave_id * BATCH * 256;
const int my_b_byte = MY_B_DW * 4;
const int bn = nt * 16 + lid;
const int bb = nt * (kt * 512);
const int nt_total = (N + 15) / 16;
const bool valid_wave = (nt < nt_total && mt < (M + 15) / 16);
// === Quant A → LDS (2 tile(s) per wave) ===
{
const int mg_q = mt * 16 + lid;
// Quant tile 0
{
const int ki_t = wave_id * 2 + 0;
int kb = ki_t << 7;
uint4 c0,c1,c2,c3;
if (mg_q < M) {
const uint4* av = reinterpret_cast<const uint4*>(A + mg_q*Kh + ((kb+(gid<<5))>>1));
c0=av[0];c1=av[1];c2=av[2];c3=av[3];
} else { c0=c1=c2=c3=make_uint4(0,0,0,0); }
v4i32 fp4; uint32_t e8m0;
quant_a_direct(c0, c1, c2, c3, fp4, e8m0);
int base = ki_t * A_TILE;
lds_gll[base+lane*A_FP4_STRIDE+0]=(uint32_t)fp4[0];
lds_gll[base+lane*A_FP4_STRIDE+1]=(uint32_t)fp4[1];
lds_gll[base+lane*A_FP4_STRIDE+2]=(uint32_t)fp4[2];
lds_gll[base+lane*A_FP4_STRIDE+3]=(uint32_t)fp4[3];
lds_gll[base+320+lane]=e8m0;
}
// Quant tile 1
{
const int ki_t = wave_id * 2 + 1;
int kb = ki_t << 7;
uint4 c0,c1,c2,c3;
if (mg_q < M) {
const uint4* av = reinterpret_cast<const uint4*>(A + mg_q*Kh + ((kb+(gid<<5))>>1));
c0=av[0];c1=av[1];c2=av[2];c3=av[3];
} else { c0=c1=c2=c3=make_uint4(0,0,0,0); }
v4i32 fp4; uint32_t e8m0;
quant_a_direct(c0, c1, c2, c3, fp4, e8m0);
int base = ki_t * A_TILE;
lds_gll[base+lane*A_FP4_STRIDE+0]=(uint32_t)fp4[0];
lds_gll[base+lane*A_FP4_STRIDE+1]=(uint32_t)fp4[1];
lds_gll[base+lane*A_FP4_STRIDE+2]=(uint32_t)fp4[2];
lds_gll[base+lane*A_FP4_STRIDE+3]=(uint32_t)fp4[3];
lds_gll[base+320+lane]=e8m0;
}
}
// === Load B+Bsc: SPLIT=10 pre-barrier, 0 post-barrier ===
uint32_t be0,be1,be2,be3,be4,be5,be6,be7,be8,be9,be10,be11,be12,be13,be14,be15;
if (valid_wave) {
const uint8_t* b_base = Bs + bb;
uint64_t b_addr = (uint64_t)b_base;
asm volatile(
"s_mov_b32 s[16], %0\n\t"
"s_mov_b32 s[17], %1\n\t"
"s_mov_b32 s[18], 0xFFFFFFFF\n\t"
"s_mov_b32 s[19], 0x20000"
:: "s"(__builtin_amdgcn_readfirstlane((uint32_t)b_addr)),
"s"(__builtin_amdgcn_readfirstlane((uint32_t)(b_addr >> 32)))
: "s16", "s17", "s18", "s19"
);
// Pre-barrier: 10 interleaved B+Bsc pairs
{ uint32_t _m0v = (uint32_t)(my_b_byte + 0 * 1024);
uint32_t _buf_off = (0*2+gid/2)*512 + (gid&1)*256 + lid*16;
asm volatile("s_mov_b32 m0, %0\n\t"
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
: "memory", "s16", "s17", "s18", "s19");
{ int _kb = 0 << 7; be0 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
{ uint32_t _m0v = (uint32_t)(my_b_byte + 1 * 1024);
uint32_t _buf_off = (1*2+gid/2)*512 + (gid&1)*256 + lid*16;
asm volatile("s_mov_b32 m0, %0\n\t"
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
: "memory", "s16", "s17", "s18", "s19");
{ int _kb = 1 << 7; be1 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
{ uint32_t _m0v = (uint32_t)(my_b_byte + 2 * 1024);
uint32_t _buf_off = (2*2+gid/2)*512 + (gid&1)*256 + lid*16;
asm volatile("s_mov_b32 m0, %0\n\t"
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
: "memory", "s16", "s17", "s18", "s19");
{ int _kb = 2 << 7; be2 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
{ uint32_t _m0v = (uint32_t)(my_b_byte + 3 * 1024);
uint32_t _buf_off = (3*2+gid/2)*512 + (gid&1)*256 + lid*16;
asm volatile("s_mov_b32 m0, %0\n\t"
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
: "memory", "s16", "s17", "s18", "s19");
{ int _kb = 3 << 7; be3 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
{ uint32_t _m0v = (uint32_t)(my_b_byte + 4 * 1024);
uint32_t _buf_off = (4*2+gid/2)*512 + (gid&1)*256 + lid*16;
asm volatile("s_mov_b32 m0, %0\n\t"
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
: "memory", "s16", "s17", "s18", "s19");
{ int _kb = 4 << 7; be4 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
{ uint32_t _m0v = (uint32_t)(my_b_byte + 5 * 1024);
uint32_t _buf_off = (5*2+gid/2)*512 + (gid&1)*256 + lid*16;
asm volatile("s_mov_b32 m0, %0\n\t"
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
: "memory", "s16", "s17", "s18", "s19");
{ int _kb = 5 << 7; be5 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
{ uint32_t _m0v = (uint32_t)(my_b_byte + 6 * 1024);
uint32_t _buf_off = (6*2+gid/2)*512 + (gid&1)*256 + lid*16;
asm volatile("s_mov_b32 m0, %0\n\t"
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
: "memory", "s16", "s17", "s18", "s19");
{ int _kb = 6 << 7; be6 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
{ uint32_t _m0v = (uint32_t)(my_b_byte + 7 * 1024);
uint32_t _buf_off = (7*2+gid/2)*512 + (gid&1)*256 + lid*16;
asm volatile("s_mov_b32 m0, %0\n\t"
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
: "memory", "s16", "s17", "s18", "s19");
{ int _kb = 7 << 7; be7 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
{ uint32_t _m0v = (uint32_t)(my_b_byte + 8 * 1024);
uint32_t _buf_off = (8*2+gid/2)*512 + (gid&1)*256 + lid*16;
asm volatile("s_mov_b32 m0, %0\n\t"
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
: "memory", "s16", "s17", "s18", "s19");
{ int _kb = 8 << 7; be8 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
{ uint32_t _m0v = (uint32_t)(my_b_byte + 9 * 1024);
uint32_t _buf_off = (9*2+gid/2)*512 + (gid&1)*256 + lid*16;
asm volatile("s_mov_b32 m0, %0\n\t"
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
: "memory", "s16", "s17", "s18", "s19");
{ int _kb = 9 << 7; be9 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
{ int _kb=10<<7; be10=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
{ int _kb=11<<7; be11=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
{ int _kb=12<<7; be12=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
{ int _kb=13<<7; be13=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
{ int _kb=14<<7; be14=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
{ int _kb=15<<7; be15=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
} else {
be0=0; be1=0; be2=0; be3=0; be4=0; be5=0; be6=0; be7=0; be8=0; be9=0; be10=0; be11=0; be12=0; be13=0; be14=0; be15=0;
}
__syncthreads();
if (!valid_wave) return;
const int a_offset = 0;
v4f32 acc = {0,0,0,0};
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
#define DB_READ_B_K(...) DB_READ_B_IMPL(__VA_ARGS__)
#define DB_READ_A01_K(...) DB_READ_A01_IMPL(__VA_ARGS__)
#define DB_READ_A23_K(...) DB_READ_A23_IMPL(__VA_ARGS__)
#define DB_READ_AE_K(...) DB_READ_AE_IMPL(__VA_ARGS__)
#define DB_ISSUE_K(...) DB_ISSUE_IMPL(__VA_ARGS__)
#define DB_PACK_K(...) DB_PACK_IMPL(__VA_ARGS__)
#define DB_MFMA_K(...) DB_MFMA_IMPL(__VA_ARGS__)
#define DB_TILE_A_K(...) DB_TILE_A_IMPL(__VA_ARGS__)
#define DB_TILE_B_K(...) DB_TILE_B_IMPL(__VA_ARGS__)
#define DO_MFMA_RT_K(...) DO_MFMA_RT_IMPL(__VA_ARGS__)
#define ISSUE_GLL_K(...) ISSUE_GLL_IMPL(__VA_ARGS__)
#define DB_READ_B_IMPL(dst, addr) \
asm volatile("ds_read_b128 %0, %1" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
#define DB_READ_A01_IMPL(dst, addr) \
asm volatile("ds_read2_b32 %0, %1 offset1:1" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
#define DB_READ_A23_IMPL(dst, addr) \
asm volatile("ds_read2_b32 %0, %1 offset0:2 offset1:3" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
#define DB_READ_AE_IMPL(dst, addr) \
asm volatile("ds_read_b32 %0, %1" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
#define DB_ISSUE_IMPL(KI, SLOT, bb, aa01, aa23, aae) do { \
int _boff = MY_B_DW + (SLOT) * 256 + lane * 4; \
DB_READ_B_IMPL(bb, _boff); \
int _abase = a_offset + (KI) * A_TILE + lane * A_FP4_STRIDE; \
DB_READ_A01_IMPL(aa01, _abase); \
DB_READ_A23_IMPL(aa23, _abase); \
DB_READ_AE_IMPL(aae, a_offset + (KI) * A_TILE + 320 + lane); \
} while(0)
#define DB_PACK_IMPL(bb, aa01, aa23, aae, bf_v, af_v, ae_v) do { \
bf_v[0]=(int)bb.x; bf_v[1]=(int)bb.y; bf_v[2]=(int)bb.z; bf_v[3]=(int)bb.w; \
af_v[0]=(int)(uint32_t)aa01; af_v[1]=(int)(uint32_t)(aa01>>32); \
af_v[2]=(int)(uint32_t)aa23; af_v[3]=(int)(uint32_t)(aa23>>32); \
ae_v = (uint32_t)aae; \
} while(0)
#define DB_MFMA_IMPL(af_v, bf_v, ae_v, be_v) \
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4" \
:"+v"(acc):"v"(af_v),"v"(bf_v),"v"(ae_v),"v"(be_v))
#define DB_TILE_A_IMPL(KI, SLOT, BE_VAR, NEXT_KI, NEXT_SLOT) do { \
DB_PACK_IMPL(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae); \
DB_MFMA_IMPL(maf, mbf, mae, BE_VAR); \
DB_ISSUE_IMPL(NEXT_KI, NEXT_SLOT, bb_a, a01_a, a23_a, ae_a); \
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory"); \
} while(0)
#define DB_TILE_B_IMPL(KI, SLOT, BE_VAR, NEXT_KI, NEXT_SLOT) do { \
DB_PACK_IMPL(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae); \
DB_MFMA_IMPL(maf, mbf, mae, BE_VAR); \
DB_ISSUE_IMPL(NEXT_KI, NEXT_SLOT, bb_b, a01_b, a23_b, ae_b); \
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory"); \
} while(0)
#define DO_MFMA_RT_IMPL(KI_RT, SLOT, BE_VAR) do { \
v4i32 _bf; \
{ int _boff = MY_B_DW + (SLOT) * 256 + lane * 4; \
uint4 _bc = lds_load_u128((lds_u32_ptr)(&lds_gll[_boff])); \
_bf[0]=(int)_bc.x; _bf[1]=(int)_bc.y; _bf[2]=(int)_bc.z; _bf[3]=(int)_bc.w; } \
v4i32 _af; uint32_t _ae; \
{ int _base = a_offset + (KI_RT) * A_TILE; \
_af[0]=(int)lds_gll[_base+lane*A_FP4_STRIDE+0]; _af[1]=(int)lds_gll[_base+lane*A_FP4_STRIDE+1]; \
_af[2]=(int)lds_gll[_base+lane*A_FP4_STRIDE+2]; _af[3]=(int)lds_gll[_base+lane*A_FP4_STRIDE+3]; \
_ae=lds_gll[_base+320+lane]; } \
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4" \
:"+v"(acc):"v"(_af),"v"(_bf),"v"(_ae),"v"(BE_VAR)); \
} while(0)
#define ISSUE_GLL_IMPL(KI_RT, SLOT) do { \
uint32_t _buf_off = ((KI_RT)*2+gid/2)*512 + (gid&1)*256 + lid*16; \
uint32_t _m0 = (uint32_t)(my_b_byte + (SLOT) * 1024); \
const uint8_t* _bbase = Bs + bb; \
uint64_t _baddr = (uint64_t)_bbase; \
asm volatile( \
"s_mov_b32 s[16], %2\n\t" \
"s_mov_b32 s[17], %3\n\t" \
"s_mov_b32 s[18], 0xFFFFFFFF\n\t" \
"s_mov_b32 s[19], 0x20000\n\t" \
"s_mov_b32 m0, %0\n\t" \
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds" \
:: "s"(__builtin_amdgcn_readfirstlane(_m0)), "v"(_buf_off), \
"s"(__builtin_amdgcn_readfirstlane((uint32_t)_baddr)), \
"s"(__builtin_amdgcn_readfirstlane((uint32_t)(_baddr >> 32))) \
: "memory", "s16", "s17", "s18", "s19" \
); \
} while(0)
{
uint4 bb_a, bb_b; uint64_t a01_a, a23_a, a01_b, a23_b; uint32_t ae_a, ae_b;
v4i32 maf, mbf; uint32_t mae;
DB_ISSUE_IMPL(0, 0, bb_a, a01_a, a23_a, ae_a);
DB_ISSUE_IMPL(1, 1, bb_b, a01_b, a23_b, ae_b);
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");
// Tile 0: MFMA + GLL(10)→slot0 + prefetch 2
DB_PACK_IMPL(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae);
DB_MFMA_IMPL(maf, mbf, mae, be0);
ISSUE_GLL_IMPL(10, 0);
DB_ISSUE_IMPL(2, 2, bb_a, a01_a, a23_a, ae_a);
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");
// Tile 1: MFMA + GLL(11)→slot1 + prefetch 3
DB_PACK_IMPL(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae);
DB_MFMA_IMPL(maf, mbf, mae, be1);
ISSUE_GLL_IMPL(11, 1);
DB_ISSUE_IMPL(3, 3, bb_b, a01_b, a23_b, ae_b);
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");
// Tile 2: MFMA + GLL(12)→slot2 + prefetch 4
DB_PACK_IMPL(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae);
DB_MFMA_IMPL(maf, mbf, mae, be2);
ISSUE_GLL_IMPL(12, 2);
DB_ISSUE_IMPL(4, 4, bb_a, a01_a, a23_a, ae_a);
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");
// Tile 3: MFMA + GLL(13)→slot3 + prefetch 5
DB_PACK_IMPL(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae);
DB_MFMA_IMPL(maf, mbf, mae, be3);
ISSUE_GLL_IMPL(13, 3);
DB_ISSUE_IMPL(5, 5, bb_b, a01_b, a23_b, ae_b);
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");
// Tile 4: MFMA + GLL(14)→slot4 + prefetch 6
DB_PACK_IMPL(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae);
DB_MFMA_IMPL(maf, mbf, mae, be4);
ISSUE_GLL_IMPL(14, 4);
DB_ISSUE_IMPL(6, 6, bb_a, a01_a, a23_a, ae_a);
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");
// Tile 5: MFMA + GLL(15)→slot5 + prefetch 7
DB_PACK_IMPL(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae);
DB_MFMA_IMPL(maf, mbf, mae, be5);
ISSUE_GLL_IMPL(15, 5);
DB_ISSUE_IMPL(7, 7, bb_b, a01_b, a23_b, ae_b);
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");
DB_TILE_A_IMPL(6, 6, be6, 8, 8);
DB_TILE_B_IMPL(7, 7, be7, 9, 9);
// Tile 8: no more pre-loaded to prefetch
DB_PACK_IMPL(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae);
DB_MFMA_IMPL(maf, mbf, mae, be8);
asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");
// Tile 9: last pre-loaded
DB_PACK_IMPL(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae);
DB_MFMA_IMPL(maf, mbf, mae, be9);
// Tiles 10..15: GLL loaded
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
DO_MFMA_RT_IMPL(10, 0, be10);
DO_MFMA_RT_IMPL(11, 1, be11);
DO_MFMA_RT_IMPL(12, 2, be12);
DO_MFMA_RT_IMPL(13, 3, be13);
DO_MFMA_RT_IMPL(14, 4, be14);
DO_MFMA_RT_IMPL(15, 5, be15);
}
#undef DB_READ_B_IMPL
#undef DB_READ_B_K
#undef DB_READ_A01_IMPL
#undef DB_READ_A01_K
#undef DB_READ_A23_IMPL
#undef DB_READ_A23_K
#undef DB_READ_AE_IMPL
#undef DB_READ_AE_K
#undef DB_ISSUE_IMPL
#undef DB_ISSUE_K
#undef DB_PACK_IMPL
#undef DB_PACK_K
#undef DB_MFMA_IMPL
#undef DB_MFMA_K
#undef DB_TILE_A_IMPL
#undef DB_TILE_A_K
#undef DB_TILE_B_IMPL
#undef DB_TILE_B_K
#undef DO_MFMA_RT_IMPL
#undef DO_MFMA_RT_K
#undef ISSUE_GLL_IMPL
#undef ISSUE_GLL_K
store_tile(C, acc, mt, nt, gid, lid, M, N);
}
// Dedicated kernel: M=256 N=3072 K=1536 Ki=12 N_N=12 BATCH=8 SPLIT=8
__launch_bounds__(768)
__global__ void kern_s5(
const uint32_t* __restrict__ A, const uint8_t* __restrict__ Bs,
const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C,
const int M, const int N, const int K, const int kt, const int pk
) {
constexpr int N_M = 1, N_N = 12, BATCH = 8;
constexpr int N_WAVES = 12;
constexpr int A_TILE = 384;
constexpr int A_FP4_STRIDE = 5;
constexpr int Ki = 12;
const int tid = threadIdx.x;
const int wave_id = tid >> 6;
const int lane = tid & 63;
const int lid = lane & 15;
const int gid = lane >> 4;
const int mt = blockIdx.x;
const int nt = blockIdx.y * N_N + wave_id;
const int Kh = K >> 1;
extern __shared__ uint32_t lds_gll[];
const int B_LDS_DW = Ki * A_TILE;
const int MY_B_DW = B_LDS_DW + wave_id * BATCH * 256;
const int my_b_byte = MY_B_DW * 4;
const int bn = nt * 16 + lid;
const int bb = nt * (kt * 512);
const int nt_total = (N + 15) / 16;
const bool valid_wave = (nt < nt_total && mt < (M + 15) / 16);
// === Quant A → LDS (1 tile(s) per wave) ===
{
const int mg_q = mt * 16 + lid;
// Quant tile 0
{
const int ki_t = wave_id * 1 + 0;
int kb = ki_t << 7;
uint4 c0,c1,c2,c3;
if (mg_q < M) {
const uint4* av = reinterpret_cast<const uint4*>(A + mg_q*Kh + ((kb+(gid<<5))>>1));
c0=av[0];c1=av[1];c2=av[2];c3=av[3];
} else { c0=c1=c2=c3=make_uint4(0,0,0,0); }
v4i32 fp4; uint32_t e8m0;
quant_a_direct(c0, c1, c2, c3, fp4, e8m0);
int base = ki_t * A_TILE;
lds_gll[base+lane*A_FP4_STRIDE+0]=(uint32_t)fp4[0];
lds_gll[base+lane*A_FP4_STRIDE+1]=(uint32_t)fp4[1];
lds_gll[base+lane*A_FP4_STRIDE+2]=(uint32_t)fp4[2];
lds_gll[base+lane*A_FP4_STRIDE+3]=(uint32_t)fp4[3];
lds_gll[base+320+lane]=e8m0;
}
}
// === Load B+Bsc: SPLIT=8 pre-barrier, 0 post-barrier ===
uint32_t be0,be1,be2,be3,be4,be5,be6,be7,be8,be9,be10,be11;
if (valid_wave) {
const uint8_t* b_base = Bs + bb;
uint64_t b_addr = (uint64_t)b_base;
asm volatile(
"s_mov_b32 s[16], %0\n\t"
"s_mov_b32 s[17], %1\n\t"
"s_mov_b32 s[18], 0xFFFFFFFF\n\t"
"s_mov_b32 s[19], 0x20000"
:: "s"(__builtin_amdgcn_readfirstlane((uint32_t)b_addr)),
"s"(__builtin_amdgcn_readfirstlane((uint32_t)(b_addr >> 32)))
: "s16", "s17", "s18", "s19"
);
// Pre-barrier: 8 interleaved B+Bsc pairs
{ uint32_t _m0v = (uint32_t)(my_b_byte + 0 * 1024);
uint32_t _buf_off = (0*2+gid/2)*512 + (gid&1)*256 + lid*16;
asm volatile("s_mov_b32 m0, %0\n\t"
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
: "memory", "s16", "s17", "s18", "s19");
{ int _kb = 0 << 7; be0 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
{ uint32_t _m0v = (uint32_t)(my_b_byte + 1 * 1024);
uint32_t _buf_off = (1*2+gid/2)*512 + (gid&1)*256 + lid*16;
asm volatile("s_mov_b32 m0, %0\n\t"
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
: "memory", "s16", "s17", "s18", "s19");
{ int _kb = 1 << 7; be1 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
{ uint32_t _m0v = (uint32_t)(my_b_byte + 2 * 1024);
uint32_t _buf_off = (2*2+gid/2)*512 + (gid&1)*256 + lid*16;
asm volatile("s_mov_b32 m0, %0\n\t"
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
: "memory", "s16", "s17", "s18", "s19");
{ int _kb = 2 << 7; be2 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
{ uint32_t _m0v = (uint32_t)(my_b_byte + 3 * 1024);
uint32_t _buf_off = (3*2+gid/2)*512 + (gid&1)*256 + lid*16;
asm volatile("s_mov_b32 m0, %0\n\t"
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
: "memory", "s16", "s17", "s18", "s19");
{ int _kb = 3 << 7; be3 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
{ uint32_t _m0v = (uint32_t)(my_b_byte + 4 * 1024);
uint32_t _buf_off = (4*2+gid/2)*512 + (gid&1)*256 + lid*16;
asm volatile("s_mov_b32 m0, %0\n\t"
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
: "memory", "s16", "s17", "s18", "s19");
{ int _kb = 4 << 7; be4 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
{ uint32_t _m0v = (uint32_t)(my_b_byte + 5 * 1024);
uint32_t _buf_off = (5*2+gid/2)*512 + (gid&1)*256 + lid*16;
asm volatile("s_mov_b32 m0, %0\n\t"
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
: "memory", "s16", "s17", "s18", "s19");
{ int _kb = 5 << 7; be5 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
{ uint32_t _m0v = (uint32_t)(my_b_byte + 6 * 1024);
uint32_t _buf_off = (6*2+gid/2)*512 + (gid&1)*256 + lid*16;
asm volatile("s_mov_b32 m0, %0\n\t"
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
: "memory", "s16", "s17", "s18", "s19");
{ int _kb = 6 << 7; be6 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
{ uint32_t _m0v = (uint32_t)(my_b_byte + 7 * 1024);
uint32_t _buf_off = (7*2+gid/2)*512 + (gid&1)*256 + lid*16;
asm volatile("s_mov_b32 m0, %0\n\t"
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
: "memory", "s16", "s17", "s18", "s19");
{ int _kb = 7 << 7; be7 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
{ int _kb=8<<7; be8=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
{ int _kb=9<<7; be9=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
{ int _kb=10<<7; be10=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
{ int _kb=11<<7; be11=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
} else {
be0=0; be1=0; be2=0; be3=0; be4=0; be5=0; be6=0; be7=0; be8=0; be9=0; be10=0; be11=0;
}
__syncthreads();
if (!valid_wave) return;
const int a_offset = 0;
v4f32 acc = {0,0,0,0};
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
#define DB_READ_B_K(...) DB_READ_B_IMPL(__VA_ARGS__)
#define DB_READ_A01_K(...) DB_READ_A01_IMPL(__VA_ARGS__)
#define DB_READ_A23_K(...) DB_READ_A23_IMPL(__VA_ARGS__)
#define DB_READ_AE_K(...) DB_READ_AE_IMPL(__VA_ARGS__)
#define DB_ISSUE_K(...) DB_ISSUE_IMPL(__VA_ARGS__)
#define DB_PACK_K(...) DB_PACK_IMPL(__VA_ARGS__)
#define DB_MFMA_K(...) DB_MFMA_IMPL(__VA_ARGS__)
#define DB_TILE_A_K(...) DB_TILE_A_IMPL(__VA_ARGS__)
#define DB_TILE_B_K(...) DB_TILE_B_IMPL(__VA_ARGS__)
#define DO_MFMA_RT_K(...) DO_MFMA_RT_IMPL(__VA_ARGS__)
#define ISSUE_GLL_K(...) ISSUE_GLL_IMPL(__VA_ARGS__)
#define DB_READ_B_IMPL(dst, addr) \
asm volatile("ds_read_b128 %0, %1" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
#define DB_READ_A01_IMPL(dst, addr) \
asm volatile("ds_read2_b32 %0, %1 offset1:1" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
#define DB_READ_A23_IMPL(dst, addr) \
asm volatile("ds_read2_b32 %0, %1 offset0:2 offset1:3" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
#define DB_READ_AE_IMPL(dst, addr) \
asm volatile("ds_read_b32 %0, %1" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
#define DB_ISSUE_IMPL(KI, SLOT, bb, aa01, aa23, aae) do { \
int _boff = MY_B_DW + (SLOT) * 256 + lane * 4; \
DB_READ_B_IMPL(bb, _boff); \
int _abase = a_offset + (KI) * A_TILE + lane * A_FP4_STRIDE; \
DB_READ_A01_IMPL(aa01, _abase); \
DB_READ_A23_IMPL(aa23, _abase); \
DB_READ_AE_IMPL(aae, a_offset + (KI) * A_TILE + 320 + lane); \
} while(0)
#define DB_PACK_IMPL(bb, aa01, aa23, aae, bf_v, af_v, ae_v) do { \
bf_v[0]=(int)bb.x; bf_v[1]=(int)bb.y; bf_v[2]=(int)bb.z; bf_v[3]=(int)bb.w; \
af_v[0]=(int)(uint32_t)aa01; af_v[1]=(int)(uint32_t)(aa01>>32); \
af_v[2]=(int)(uint32_t)aa23; af_v[3]=(int)(uint32_t)(aa23>>32); \
ae_v = (uint32_t)aae; \
} while(0)
#define DB_MFMA_IMPL(af_v, bf_v, ae_v, be_v) \
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4" \
:"+v"(acc):"v"(af_v),"v"(bf_v),"v"(ae_v),"v"(be_v))
#define DB_TILE_A_IMPL(KI, SLOT, BE_VAR, NEXT_KI, NEXT_SLOT) do { \
DB_PACK_IMPL(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae); \
DB_MFMA_IMPL(maf, mbf, mae, BE_VAR); \
DB_ISSUE_IMPL(NEXT_KI, NEXT_SLOT, bb_a, a01_a, a23_a, ae_a); \
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory"); \
} while(0)
#define DB_TILE_B_IMPL(KI, SLOT, BE_VAR, NEXT_KI, NEXT_SLOT) do { \
DB_PACK_IMPL(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae); \
DB_MFMA_IMPL(maf, mbf, mae, BE_VAR); \
DB_ISSUE_IMPL(NEXT_KI, NEXT_SLOT, bb_b, a01_b, a23_b, ae_b); \
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory"); \
} while(0)
#define DO_MFMA_RT_IMPL(KI_RT, SLOT, BE_VAR) do { \
v4i32 _bf; \
{ int _boff = MY_B_DW + (SLOT) * 256 + lane * 4; \
uint4 _bc = lds_load_u128((lds_u32_ptr)(&lds_gll[_boff])); \
_bf[0]=(int)_bc.x; _bf[1]=(int)_bc.y; _bf[2]=(int)_bc.z; _bf[3]=(int)_bc.w; } \
v4i32 _af; uint32_t _ae; \
{ int _base = a_offset + (KI_RT) * A_TILE; \
_af[0]=(int)lds_gll[_base+lane*A_FP4_STRIDE+0]; _af[1]=(int)lds_gll[_base+lane*A_FP4_STRIDE+1]; \
_af[2]=(int)lds_gll[_base+lane*A_FP4_STRIDE+2]; _af[3]=(int)lds_gll[_base+lane*A_FP4_STRIDE+3]; \
_ae=lds_gll[_base+320+lane]; } \
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4" \
:"+v"(acc):"v"(_af),"v"(_bf),"v"(_ae),"v"(BE_VAR)); \
} while(0)
#define ISSUE_GLL_IMPL(KI_RT, SLOT) do { \
uint32_t _buf_off = ((KI_RT)*2+gid/2)*512 + (gid&1)*256 + lid*16; \
uint32_t _m0 = (uint32_t)(my_b_byte + (SLOT) * 1024); \
const uint8_t* _bbase = Bs + bb; \
uint64_t _baddr = (uint64_t)_bbase; \
asm volatile( \
"s_mov_b32 s[16], %2\n\t" \
"s_mov_b32 s[17], %3\n\t" \
"s_mov_b32 s[18], 0xFFFFFFFF\n\t" \
"s_mov_b32 s[19], 0x20000\n\t" \
"s_mov_b32 m0, %0\n\t" \
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds" \
:: "s"(__builtin_amdgcn_readfirstlane(_m0)), "v"(_buf_off), \
"s"(__builtin_amdgcn_readfirstlane((uint32_t)_baddr)), \
"s"(__builtin_amdgcn_readfirstlane((uint32_t)(_baddr >> 32))) \
: "memory", "s16", "s17", "s18", "s19" \
); \
} while(0)
{
uint4 bb_a, bb_b; uint64_t a01_a, a23_a, a01_b, a23_b; uint32_t ae_a, ae_b;
v4i32 maf, mbf; uint32_t mae;
DB_ISSUE_IMPL(0, 0, bb_a, a01_a, a23_a, ae_a);
DB_ISSUE_IMPL(1, 1, bb_b, a01_b, a23_b, ae_b);
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");
// Tile 0: MFMA + GLL(8)→slot0 + prefetch 2
DB_PACK_IMPL(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae);
DB_MFMA_IMPL(maf, mbf, mae, be0);
ISSUE_GLL_IMPL(8, 0);
DB_ISSUE_IMPL(2, 2, bb_a, a01_a, a23_a, ae_a);
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");
// Tile 1: MFMA + GLL(9)→slot1 + prefetch 3
DB_PACK_IMPL(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae);
DB_MFMA_IMPL(maf, mbf, mae, be1);
ISSUE_GLL_IMPL(9, 1);
DB_ISSUE_IMPL(3, 3, bb_b, a01_b, a23_b, ae_b);
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");
// Tile 2: MFMA + GLL(10)→slot2 + prefetch 4
DB_PACK_IMPL(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae);
DB_MFMA_IMPL(maf, mbf, mae, be2);
ISSUE_GLL_IMPL(10, 2);
DB_ISSUE_IMPL(4, 4, bb_a, a01_a, a23_a, ae_a);
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");
// Tile 3: MFMA + GLL(11)→slot3 + prefetch 5
DB_PACK_IMPL(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae);
DB_MFMA_IMPL(maf, mbf, mae, be3);
ISSUE_GLL_IMPL(11, 3);
DB_ISSUE_IMPL(5, 5, bb_b, a01_b, a23_b, ae_b);
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");
DB_TILE_A_IMPL(4, 4, be4, 6, 6);
DB_TILE_B_IMPL(5, 5, be5, 7, 7);
// Tile 6: no more pre-loaded to prefetch
DB_PACK_IMPL(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae);
DB_MFMA_IMPL(maf, mbf, mae, be6);
asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");
// Tile 7: last pre-loaded
DB_PACK_IMPL(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae);
DB_MFMA_IMPL(maf, mbf, mae, be7);
// Tiles 8..11: GLL loaded
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
DO_MFMA_RT_IMPL(8, 0, be8);
DO_MFMA_RT_IMPL(9, 1, be9);
DO_MFMA_RT_IMPL(10, 2, be10);
DO_MFMA_RT_IMPL(11, 3, be11);
}
#undef DB_READ_B_IMPL
#undef DB_READ_B_K
#undef DB_READ_A01_IMPL
#undef DB_READ_A01_K
#undef DB_READ_A23_IMPL
#undef DB_READ_A23_K
#undef DB_READ_AE_IMPL
#undef DB_READ_AE_K
#undef DB_ISSUE_IMPL
#undef DB_ISSUE_K
#undef DB_PACK_IMPL
#undef DB_PACK_K
#undef DB_MFMA_IMPL
#undef DB_MFMA_K
#undef DB_TILE_A_IMPL
#undef DB_TILE_A_K
#undef DB_TILE_B_IMPL
#undef DB_TILE_B_K
#undef DO_MFMA_RT_IMPL
#undef DO_MFMA_RT_K
#undef ISSUE_GLL_IMPL
#undef ISSUE_GLL_K
store_tile(C, acc, mt, nt, gid, lid, M, N);
}
// KERNEL 5b: GLL batch — B via GLOBAL_LOAD_LDS, batch processing
// Issue D GLL loads to LDS, process D MFMAs while next batch's GLL flies.
// No B VGPRs in pipeline — only bf (4 VGPRs) read from LDS per MFMA.
// ==========================================================================
template<int N_M, int N_N, int BATCH>
__launch_bounds__(N_M * N_N * 64)
__global__ void kern_fused_gll(
const uint32_t* __restrict__ A, const uint8_t* __restrict__ Bs,
const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C,
const int M, const int N, const int K, const int kt, const int pk
) {
constexpr int N_WAVES = N_M * N_N;
constexpr int A_TILE = 384;
constexpr int A_FP4_STRIDE = 5;
const int tid = threadIdx.x;
const int wave_id = tid >> 6;
const int lane = tid & 63;
const int lid = lane & 15;
const int gid = lane >> 4;
const int mt_local = wave_id / N_N;
const int nt_local = wave_id % N_N;
const int mt = blockIdx.x * N_M + mt_local;
const int nt = blockIdx.y * N_N + nt_local;
const int mg = mt * 16 + lid;
const int Kh = K >> 1;
const int Ki = K >> 7;
extern __shared__ uint32_t lds_gll[];
// B LDS layout: after A data. Each wave gets BATCH slots of 256 dwords.
const int B_LDS_DW = N_M * Ki * A_TILE;
const int MY_B_DW = B_LDS_DW + wave_id * BATCH * 256;
const int my_b_byte = MY_B_DW * 4;
const int bn = nt * 16 + lid;
const int bb = nt * (kt * 512);
const int nt_total = (N + 15) / 16;
const bool valid_wave = (nt < nt_total && mt < (M + 15) / 16);
// REORDERED Phase 1+2: load A regs → issue GLL+Bsc → quant A → store LDS
// A loads go to vmcnt first. Then GLL+Bsc go to vmcnt.
// Compiler's vmcnt(0) before quant drains A (oldest). GLL+Bsc keep flying.
// After landmark patching: vmcnt(0) → vmcnt(N) to drain only A loads.
int total_work = N_M * Ki;
int work_per_wave = (total_work + N_WAVES - 1) / N_WAVES;
int w_start = wave_id * work_per_wave;
int w_end = w_start + work_per_wave;
if (w_end > total_work) w_end = total_work;
int num_work = w_end - w_start;
// Step 1: Issue A loads into registers (compiler tracks these as vmcnt)
uint4 aq_c0_0={},aq_c1_0={},aq_c2_0={},aq_c3_0={};
uint4 aq_c0_1={},aq_c1_1={},aq_c2_1={},aq_c3_1={};
int aq_m_idx_0=0, aq_ki_0=0, aq_m_idx_1=0, aq_ki_1=0;
if (num_work > 0) {
int w = w_start;
aq_m_idx_0 = w / Ki; aq_ki_0 = w % Ki;
int mt_q = blockIdx.x * N_M + aq_m_idx_0;
int mg_q = mt_q * 16 + lid;
int kb = aq_ki_0 << 7;
if (mg_q < M) {
const uint4* av = reinterpret_cast<const uint4*>(A + mg_q*Kh + ((kb+(gid<<5))>>1));
aq_c0_0=av[0]; aq_c1_0=av[1]; aq_c2_0=av[2]; aq_c3_0=av[3];
}
}
if (num_work > 1) {
int w = w_start + 1;
aq_m_idx_1 = w / Ki; aq_ki_1 = w % Ki;
int mt_q = blockIdx.x * N_M + aq_m_idx_1;
int mg_q = mt_q * 16 + lid;
int kb = aq_ki_1 << 7;
if (mg_q < M) {
const uint4* av = reinterpret_cast<const uint4*>(A + mg_q*Kh + ((kb+(gid<<5))>>1));
aq_c0_1=av[0]; aq_c1_1=av[1]; aq_c2_1=av[2]; aq_c3_1=av[3];
}
}
// vmcnt now has 4*num_work A loads (up to 8 for 2 tiles)
// Step 2: Issue GLL+Bsc (interleaved pairs) — BEFORE quant!
// These go to vmcnt AFTER A loads. GLL+Bsc fly during quant ALU.
// Pre-barrier: INTERLEAVED (buffer_load B[s], global_load Bsc[s]) pairs
// vmcnt queue: [B0,Bsc0, B1,Bsc1, ..., B11,Bsc11] = 24 entries for BATCH=12
// Each pair = 2 vmcnt entries. vmcnt(N) drains oldest, keeping N newest.
uint32_t be0=0,be1=0,be2=0,be3=0,be4=0,be5=0,be6=0,be7=0,be8=0,be9=0,be10=0,be11=0,be12=0,be13=0,be14=0,be15=0;
if (valid_wave) {
// Setup buffer descriptor for B in s[16:19]
const uint8_t* b_base = Bs + bb;
uint64_t b_addr = (uint64_t)b_base;
asm volatile(
"s_mov_b32 s[16], %0\n\t"
"s_mov_b32 s[17], %1\n\t"
"s_mov_b32 s[18], 0xFFFFFFFF\n\t"
"s_mov_b32 s[19], 0x20000"
:: "s"(__builtin_amdgcn_readfirstlane((uint32_t)b_addr)),
"s"(__builtin_amdgcn_readfirstlane((uint32_t)(b_addr >> 32)))
: "s16", "s17", "s18", "s19"
);
// Interleaved: issue B[s] then Bsc[s] as a pair, repeat for each slot
#define ISSUE_PAIR(S) do { \
int _kb = (S) << 7; \
uint32_t _m0v = (uint32_t)(my_b_byte + (S) * 1024); \
uint32_t _buf_off = (_kb/64+gid/2)*512 + (gid&1)*256 + lid*16; \
asm volatile( \
"s_mov_b32 m0, %0\n\t" \
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds" \
:: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off) \
: "memory", "s16", "s17", "s18", "s19" \
); \
be##S = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; \
} while(0)
// Issue interleaved pairs for pre-loaded BATCH slots only
if constexpr (BATCH > 0) ISSUE_PAIR(0);
if constexpr (BATCH > 1) { if (Ki>1) ISSUE_PAIR(1); }
if constexpr (BATCH > 2) { if (Ki>2) ISSUE_PAIR(2); }
if constexpr (BATCH > 3) { if (Ki>3) ISSUE_PAIR(3); }
if constexpr (BATCH > 4) { if (Ki>4) ISSUE_PAIR(4); }
if constexpr (BATCH > 5) { if (Ki>5) ISSUE_PAIR(5); }
if constexpr (BATCH > 6) { if (Ki>6) ISSUE_PAIR(6); }
if constexpr (BATCH > 7) { if (Ki>7) ISSUE_PAIR(7); }
if constexpr (BATCH > 8) { if (Ki>8) ISSUE_PAIR(8); }
if constexpr (BATCH > 9) { if (Ki>9) ISSUE_PAIR(9); }
if constexpr (BATCH > 10) { if (Ki>10) ISSUE_PAIR(10); }
if constexpr (BATCH > 11) { if (Ki>11) ISSUE_PAIR(11); }
// Bsc-only for tiles beyond BATCH (no GLL — loaded mid-loop or not at all)
#define LOAD_BE_ONLY(I) do { int _kb=(I)<<7; be##I=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; } while(0)
if constexpr (BATCH <= 6) { if (Ki>6) LOAD_BE_ONLY(6); }
if constexpr (BATCH <= 7) { if (Ki>7) LOAD_BE_ONLY(7); }
if constexpr (BATCH <= 8) { if (Ki>8) LOAD_BE_ONLY(8); }
if constexpr (BATCH <= 9) { if (Ki>9) LOAD_BE_ONLY(9); }
if constexpr (BATCH <= 10) { if (Ki>10) LOAD_BE_ONLY(10); }
if constexpr (BATCH <= 11) { if (Ki>11) LOAD_BE_ONLY(11); }
if (Ki>12) LOAD_BE_ONLY(12);
if (Ki>13) LOAD_BE_ONLY(13);
if (Ki>14) LOAD_BE_ONLY(14);
if (Ki>15) LOAD_BE_ONLY(15);
#undef LOAD_BE_ONLY
#undef ISSUE_PAIR
}
// vmcnt queue: [A_0(×4), A_1(×4), GLL_0,Bsc_0, ..., GLL_11,Bsc_11, Bsc_12..15]
// For S4: 8 (A) + 24 (12 GLL+Bsc pairs) + 4 (Bsc only) = 36
// Step 3: Quant A from preloaded registers
// Compiler will insert vmcnt(0) before using aq_c0_0 etc.
// LANDMARK 3: marks the vmcnt before quant tile 0 for patching
// After patch: vmcnt(0) → vmcnt(28) = drain 8 A loads, keep 28 GLL+Bsc
asm volatile("s_nop 3" ::: "memory"); // LANDMARK 3: before quant tile 0
if (num_work > 0) {
v4i32 fp4; uint32_t e8m0;
quant_a_direct(aq_c0_0, aq_c1_0, aq_c2_0, aq_c3_0, fp4, e8m0);
int base = aq_m_idx_0 * Ki * A_TILE + aq_ki_0 * A_TILE;
lds_gll[base+lane*A_FP4_STRIDE+0]=(uint32_t)fp4[0];
lds_gll[base+lane*A_FP4_STRIDE+1]=(uint32_t)fp4[1];
lds_gll[base+lane*A_FP4_STRIDE+2]=(uint32_t)fp4[2];
lds_gll[base+lane*A_FP4_STRIDE+3]=(uint32_t)fp4[3];
lds_gll[base+320+lane]=e8m0;
}
asm volatile("s_nop 4" ::: "memory"); // LANDMARK 4: before quant tile 1
if (num_work > 1) {
v4i32 fp4; uint32_t e8m0;
quant_a_direct(aq_c0_1, aq_c1_1, aq_c2_1, aq_c3_1, fp4, e8m0);
int base = aq_m_idx_1 * Ki * A_TILE + aq_ki_1 * A_TILE;
lds_gll[base+lane*A_FP4_STRIDE+0]=(uint32_t)fp4[0];
lds_gll[base+lane*A_FP4_STRIDE+1]=(uint32_t)fp4[1];
lds_gll[base+lane*A_FP4_STRIDE+2]=(uint32_t)fp4[2];
lds_gll[base+lane*A_FP4_STRIDE+3]=(uint32_t)fp4[3];
lds_gll[base+320+lane]=e8m0;
}
__syncthreads();
if (!valid_wave) return;
const int a_offset = mt_local * Ki * A_TILE;
v4f32 acc = {0,0,0,0};
// Runtime MFMA iteration: ki and be are runtime values
#define DO_MFMA_RT(KI_RT, SLOT, BE_VAR) do { \
v4i32 _bf; \
{ int _boff = MY_B_DW + (SLOT) * 256 + lane * 4; \
uint4 _bc = lds_load_u128((lds_u32_ptr)(&lds_gll[_boff])); \
_bf[0]=(int)_bc.x; _bf[1]=(int)_bc.y; _bf[2]=(int)_bc.z; _bf[3]=(int)_bc.w; } \
v4i32 _af; uint32_t _ae; \
{ int _base = a_offset + (KI_RT) * A_TILE; \
_af[0]=(int)lds_gll[_base+lane*A_FP4_STRIDE+0]; _af[1]=(int)lds_gll[_base+lane*A_FP4_STRIDE+1]; \
_af[2]=(int)lds_gll[_base+lane*A_FP4_STRIDE+2]; _af[3]=(int)lds_gll[_base+lane*A_FP4_STRIDE+3]; \
_ae=lds_gll[_base+320+lane]; } \
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4" \
:"+v"(acc):"v"(_af),"v"(_bf),"v"(_ae),"v"(BE_VAR)); \
} while(0)
#define ISSUE_GLL_RT(KI_RT, SLOT) do { \
int _kb = (KI_RT) << 7; \
uint32_t _buf_off = (_kb/64+gid/2)*512 + (gid&1)*256 + lid*16; \
uint32_t _m0 = (uint32_t)(my_b_byte + (SLOT) * 1024); \
const uint8_t* _bbase = Bs + bb; \
uint64_t _baddr = (uint64_t)_bbase; \
asm volatile( \
"s_mov_b32 s[16], %2\n\t" \
"s_mov_b32 s[17], %3\n\t" \
"s_mov_b32 s[18], 0xFFFFFFFF\n\t" \
"s_mov_b32 s[19], 0x20000\n\t" \
"s_mov_b32 m0, %0\n\t" \
"buffer_load_dwordx4 %1, s[16:19], 0 offen lds" \
:: "s"(__builtin_amdgcn_readfirstlane(_m0)), "v"(_buf_off), \
"s"(__builtin_amdgcn_readfirstlane((uint32_t)_baddr)), \
"s"(__builtin_amdgcn_readfirstlane((uint32_t)(_baddr >> 32))) \
: "memory", "s16", "s17", "s18", "s19" \
); \
} while(0)
#define LOAD_BSC_RT(VAR, KI_RT) do { int _kb=(KI_RT)<<7; VAR=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; } while(0)
if (Ki > 16) {
// === BATCHED LOOP for large Ki (e.g., Ki=56) ===
// Only for Ki > 16 where we can't use the unrolled paths.
// Uses BATCH GLL slots + be0..be7 (8 Bsc), reloaded per batch.
// First batch already pre-loaded (BATCH GLL + 8 Bsc).
// vmcnt(8): drain BATCH GLL, keep 8 Bsc.
// Hardware: BATCH + 8. Drain to 8 → removes BATCH GLL entries.
asm volatile("s_waitcnt vmcnt(8)" ::: "memory");
for (int batch = 0; batch < Ki; batch += BATCH) {
int remaining = Ki - batch;
int count = remaining < BATCH ? remaining : BATCH;
// 8 unrolled MFMA iterations using be0..be7, slots 0..7
if (count > 0) DO_MFMA_RT(batch+0, 0, be0);
if (count > 1) DO_MFMA_RT(batch+1, 1, be1);
if (count > 2) DO_MFMA_RT(batch+2, 2, be2);
if (count > 3) DO_MFMA_RT(batch+3, 3, be3);
if (count > 4) DO_MFMA_RT(batch+4, 4, be4);
if (count > 5) DO_MFMA_RT(batch+5, 5, be5);
if (count > 6) DO_MFMA_RT(batch+6, 6, be6);
if (count > 7) DO_MFMA_RT(batch+7, 7, be7);
int next_batch = batch + BATCH;
if (next_batch < Ki) {
int next_count = Ki - next_batch;
if (next_count > BATCH) next_count = BATCH;
// Issue next batch's GLL → slots 0..7 (already consumed)
for (int s = 0; s < next_count; s++) {
ISSUE_GLL_RT(next_batch + s, s);
}
// Reload Bsc for next batch
if (next_count > 0) LOAD_BSC_RT(be0, next_batch+0);
if (next_count > 1) LOAD_BSC_RT(be1, next_batch+1);
if (next_count > 2) LOAD_BSC_RT(be2, next_batch+2);
if (next_count > 3) LOAD_BSC_RT(be3, next_batch+3);
if (next_count > 4) LOAD_BSC_RT(be4, next_batch+4);
if (next_count > 5) LOAD_BSC_RT(be5, next_batch+5);
if (next_count > 6) LOAD_BSC_RT(be6, next_batch+6);
if (next_count > 7) LOAD_BSC_RT(be7, next_batch+7);
// Drain GLL, keep 8 Bsc. Hardware: next_count GLL + <=8 Bsc.
// vmcnt(8): drain GLL entries, keep Bsc.
asm volatile("s_waitcnt vmcnt(8)" ::: "memory");
}
}
} else if constexpr (BATCH >= 16) {
// vmcnt: drain BATCH GLL, keep Ki Bsc (Ki <= BATCH)
if (Ki == 12) { asm volatile("s_waitcnt vmcnt(12)" ::: "memory"); }
else if (Ki == 16) { asm volatile("s_waitcnt vmcnt(16)" ::: "memory"); }
else { asm volatile("s_waitcnt vmcnt(0)" ::: "memory"); }
// ALL Ki tiles pre-loaded — straight-through
if (Ki > 0) DO_MFMA_RT(0, 0, be0); if (Ki > 1) DO_MFMA_RT(1, 1, be1);
if (Ki > 2) DO_MFMA_RT(2, 2, be2); if (Ki > 3) DO_MFMA_RT(3, 3, be3);
if (Ki > 4) DO_MFMA_RT(4, 4, be4); if (Ki > 5) DO_MFMA_RT(5, 5, be5);
if (Ki > 6) DO_MFMA_RT(6, 6, be6); if (Ki > 7) DO_MFMA_RT(7, 7, be7);
if (Ki > 8) DO_MFMA_RT(8, 8, be8); if (Ki > 9) DO_MFMA_RT(9, 9, be9);
if (Ki > 10) DO_MFMA_RT(10, 10, be10); if (Ki > 11) DO_MFMA_RT(11, 11, be11);
if (Ki > 12) DO_MFMA_RT(12, 12, be12); if (Ki > 13) DO_MFMA_RT(13, 13, be13);
if (Ki > 14) DO_MFMA_RT(14, 14, be14); if (Ki > 15) DO_MFMA_RT(15, 15, be15);
} else if constexpr (BATCH == 12) {
// BATCH=12: DOUBLE-BUFFERED lgkmcnt(4) ds_reads
// Two register sets. All ds_reads via inline asm for lgkmcnt control.
if (Ki == 12) { asm volatile("s_waitcnt vmcnt(12)" ::: "memory"); }
else { asm volatile("s_waitcnt vmcnt(0)" ::: "memory"); }
#define DB_READ_B(dst, addr) \
asm volatile("ds_read_b128 %0, %1" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
#define DB_READ_A01(dst, addr) \
asm volatile("ds_read2_b32 %0, %1 offset1:1" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
#define DB_READ_A23(dst, addr) \
asm volatile("ds_read2_b32 %0, %1 offset0:2 offset1:3" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
#define DB_READ_AE(dst, addr) \
asm volatile("ds_read_b32 %0, %1" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
#define DB_ISSUE(KI, SLOT, bb, aa01, aa23, aae) do { \
int _boff = MY_B_DW + (SLOT) * 256 + lane * 4; \
DB_READ_B(bb, _boff); \
int _abase = a_offset + (KI) * A_TILE + lane * A_FP4_STRIDE; \
DB_READ_A01(aa01, _abase); \
DB_READ_A23(aa23, _abase); \
DB_READ_AE(aae, a_offset + (KI) * A_TILE + 320 + lane); \
} while(0)
#define DB_PACK(bb, aa01, aa23, aae, bf_v, af_v, ae_v) do { \
bf_v[0]=(int)bb.x; bf_v[1]=(int)bb.y; bf_v[2]=(int)bb.z; bf_v[3]=(int)bb.w; \
af_v[0]=(int)(uint32_t)aa01; af_v[1]=(int)(uint32_t)(aa01>>32); \
af_v[2]=(int)(uint32_t)aa23; af_v[3]=(int)(uint32_t)(aa23>>32); \
ae_v = (uint32_t)aae; \
} while(0)
#define DB_MFMA(af_v, bf_v, ae_v, be_v) \
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4" \
:"+v"(acc):"v"(af_v),"v"(bf_v),"v"(ae_v),"v"(be_v))
#define DB_TILE_A(KI, SLOT, BE_VAR, NEXT_KI, NEXT_SLOT) do { \
DB_PACK(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae); \
DB_MFMA(maf, mbf, mae, BE_VAR); \
DB_ISSUE(NEXT_KI, NEXT_SLOT, bb_a, a01_a, a23_a, ae_a); \
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory"); \
} while(0)
#define DB_TILE_B(KI, SLOT, BE_VAR, NEXT_KI, NEXT_SLOT) do { \
DB_PACK(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae); \
DB_MFMA(maf, mbf, mae, BE_VAR); \
DB_ISSUE(NEXT_KI, NEXT_SLOT, bb_b, a01_b, a23_b, ae_b); \
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory"); \
} while(0)
uint4 bb_a, bb_b; uint64_t a01_a, a23_a, a01_b, a23_b; uint32_t ae_a, ae_b;
v4i32 maf, mbf; uint32_t mae;
asm volatile("s_nop 1" ::: "memory"); // LANDMARK 1: after vmcnt drain, before MFMA
// Prologue: prefetch tiles 0 and 1
DB_ISSUE(0, 0, bb_a, a01_a, a23_a, ae_a);
DB_ISSUE(1, 1, bb_b, a01_b, a23_b, ae_b);
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");
DB_TILE_A(0, 0, be0, 2, 2);
DB_TILE_B(1, 1, be1, 3, 3);
// Tile 2: MFMA + issue GLL[12..15] → slots 0..3 (consumed)
DB_PACK(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae);
DB_MFMA(maf, mbf, mae, be2);
if (Ki > 12) ISSUE_GLL_RT(12, 0);
if (Ki > 13) ISSUE_GLL_RT(13, 1);
if (Ki > 14) ISSUE_GLL_RT(14, 2);
if (Ki > 15) ISSUE_GLL_RT(15, 3);
DB_ISSUE(4, 4, bb_a, a01_a, a23_a, ae_a);
asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");
DB_TILE_B(3, 3, be3, 5, 5);
DB_TILE_A(4, 4, be4, 6, 6);
DB_TILE_B(5, 5, be5, 7, 7);
DB_TILE_A(6, 6, be6, 8, 8);
DB_TILE_B(7, 7, be7, 9, 9);
DB_TILE_A(8, 8, be8, 10, 10);
DB_TILE_B(9, 9, be9, 11, 11);
// Tile 10: no more pre-loaded tiles
DB_PACK(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae);
DB_MFMA(maf, mbf, mae, be10);
asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");
// Tile 11
DB_PACK(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae);
DB_MFMA(maf, mbf, mae, be11);
// Tiles 12-15: GLL loaded, use DO_MFMA_RT
asm volatile("s_nop 2" ::: "memory"); // LANDMARK 2: before mid-loop vmcnt drain
if (Ki > 12) {
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
DO_MFMA_RT(12, 0, be12);
}
if (Ki > 13) DO_MFMA_RT(13, 1, be13);
if (Ki > 14) DO_MFMA_RT(14, 2, be14);
if (Ki > 15) DO_MFMA_RT(15, 3, be15);
#undef DB_READ_B
#undef DB_READ_A01
#undef DB_READ_A23
#undef DB_READ_AE
#undef DB_ISSUE
#undef DB_PACK
#undef DB_MFMA
#undef DB_TILE_A
#undef DB_TILE_B
} else if constexpr (BATCH == 6) {
// BATCH=6: 3+3 double-buffer for Ki<=12
// Pre-loaded: 6 GLL (slots 0..5) + Ki Bsc
// vmcnt: drain 6 GLL, keep Ki Bsc
if (Ki == 12) { asm volatile("s_waitcnt vmcnt(12)" ::: "memory"); }
else { asm volatile("s_waitcnt vmcnt(0)" ::: "memory"); }
// First 3 MFMAs: ki=0..2 (slots 0..2)
if (Ki > 0) DO_MFMA_RT(0, 0, be0);
if (Ki > 1) DO_MFMA_RT(1, 1, be1);
if (Ki > 2) DO_MFMA_RT(2, 2, be2);
// Burst GLL[6..8] → slots 0..2 (consumed)
if (Ki > 6) ISSUE_GLL_RT(6, 0);
if (Ki > 7) ISSUE_GLL_RT(7, 1);
if (Ki > 8) ISSUE_GLL_RT(8, 2);
// Next 3 MFMAs: ki=3..5 (slots 3..5)
if (Ki > 3) DO_MFMA_RT(3, 3, be3);
if (Ki > 4) DO_MFMA_RT(4, 4, be4);
if (Ki > 5) DO_MFMA_RT(5, 5, be5);
// Burst GLL[9..11] → slots 3..5 (consumed)
if (Ki > 9) ISSUE_GLL_RT(9, 3);
if (Ki > 10) ISSUE_GLL_RT(10, 4);
if (Ki > 11) ISSUE_GLL_RT(11, 5);
// Drain all GLL before consuming batch 1
if (Ki > 6) {
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
DO_MFMA_RT(6, 0, be6);
}
if (Ki > 7) DO_MFMA_RT(7, 1, be7);
if (Ki > 8) DO_MFMA_RT(8, 2, be8);
if (Ki > 9) DO_MFMA_RT(9, 3, be9);
if (Ki > 10) DO_MFMA_RT(10, 4, be10);
if (Ki > 11) DO_MFMA_RT(11, 5, be11);
} else {
// BATCH=8: batch processing for Ki=8..16
// vmcnt: drain BATCH GLL, keep Ki Bsc
if (Ki == 12) { asm volatile("s_waitcnt vmcnt(12)" ::: "memory"); }
else if (Ki == 16) { asm volatile("s_waitcnt vmcnt(16)" ::: "memory"); }
else { asm volatile("s_waitcnt vmcnt(8)" ::: "memory"); }
// First half: ki=0..3 (slots 0..3)
if (Ki > 0) DO_MFMA_RT(0, 0, be0); if (Ki > 1) DO_MFMA_RT(1, 1, be1);
if (Ki > 2) DO_MFMA_RT(2, 2, be2); if (Ki > 3) DO_MFMA_RT(3, 3, be3);
// Burst GLL[8..11] → slots 0..3
if (Ki > 8) ISSUE_GLL_RT(8, 0); if (Ki > 9) ISSUE_GLL_RT(9, 1);
if (Ki > 10) ISSUE_GLL_RT(10, 2); if (Ki > 11) ISSUE_GLL_RT(11, 3);
// Second half: ki=4..7 (slots 4..7)
if (Ki > 4) DO_MFMA_RT(4, 4, be4); if (Ki > 5) DO_MFMA_RT(5, 5, be5);
if (Ki > 6) DO_MFMA_RT(6, 6, be6); if (Ki > 7) DO_MFMA_RT(7, 7, be7);
// Burst GLL[12..15] → slots 4..7
if (Ki > 12) ISSUE_GLL_RT(12, 4); if (Ki > 13) ISSUE_GLL_RT(13, 5);
if (Ki > 14) ISSUE_GLL_RT(14, 6); if (Ki > 15) ISSUE_GLL_RT(15, 7);
// Drain GLL, consume batch 1
if (Ki > 8) {
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
DO_MFMA_RT(8, 0, be8);
}
if (Ki > 9) DO_MFMA_RT(9, 1, be9);
if (Ki > 10) DO_MFMA_RT(10, 2, be10); if (Ki > 11) DO_MFMA_RT(11, 3, be11);
if (Ki > 12) DO_MFMA_RT(12, 4, be12); if (Ki > 13) DO_MFMA_RT(13, 5, be13);
if (Ki > 14) DO_MFMA_RT(14, 6, be14); if (Ki > 15) DO_MFMA_RT(15, 7, be15);
}
#undef LOAD_BSC_RT
#undef ISSUE_GLL_RT
#undef DO_MFMA_RT
store_tile(C, acc, mt, nt, gid, lid, M, N);
}
// ==========================================================================
// MASTER C++ DISPATCH: Raw pointers, zero ATen overhead in hot path.
// All buffers pre-allocated once in init(), reused via raw pointers.
// ==========================================================================
static uint8_t* s_fp4_ptr = nullptr;
static uint8_t* s_sc_ptr = nullptr;
static float* s_temp_ptr = nullptr;
static uint16_t* s_C_ptr = nullptr;
static int s_fp4_cap=0, s_sc_cap=0, s_temp_cap=0, s_C_cap=0;
// One-time allocation of max-size buffers for all benchmark shapes
void preallocate(int64_t max_M, int64_t max_N, int64_t max_K) {
int max_MN = (int)(max_M * max_N);
int max_Kh = (int)(max_K / 2);
int max_Ksc = (int)(max_K / 32);
int max_Mpad = ((int)max_M + 255) / 256 * 256;
int max_Kscpad = (max_Ksc + 7) / 8 * 8;
int fp4_need = (int)(max_M * max_Kh);
int sc_need = max_Mpad * max_Kscpad;
int temp_need = max_MN;
int c_need = max_MN;
if (fp4_need > s_fp4_cap) {
if (s_fp4_ptr) hipFree(s_fp4_ptr);
hipMalloc(&s_fp4_ptr, fp4_need); s_fp4_cap = fp4_need;
}
if (sc_need > s_sc_cap) {
if (s_sc_ptr) hipFree(s_sc_ptr);
hipMalloc(&s_sc_ptr, sc_need); s_sc_cap = sc_need;
}
if (temp_need * 4 > s_temp_cap) {
if (s_temp_ptr) hipFree(s_temp_ptr);
hipMalloc(&s_temp_ptr, temp_need * sizeof(float)); s_temp_cap = temp_need * 4;
}
if (c_need * 2 > s_C_cap) {
if (s_C_ptr) hipFree(s_C_ptr);
hipMalloc(&s_C_ptr, c_need * sizeof(uint16_t)); s_C_cap = c_need * 2;
}
}
torch::Tensor dispatch(
torch::Tensor A, torch::Tensor Bs, torch::Tensor Bsc,
int64_t M, int64_t N, int64_t K
) {
const int NUM_CUS = 304;
int mt = (int)((M+15)/16), nt = (int)((N+15)/16);
int tiles = mt * nt;
int Ki = (int)(K >> 7);
int kt = (int)(K / 64);
int pk = (((int)(K/32)+7)/8)*8;
// Output: cached tensor, only recreated if shape changes
static torch::Tensor s_C_tensor;
static int64_t s_C_M = 0, s_C_N = 0;
if (M != s_C_M || N != s_C_N) {
s_C_tensor = torch::from_blob(s_C_ptr, {M, N},
torch::TensorOptions().dtype(torch::kBFloat16).device(A.device()));
s_C_M = M; s_C_N = N;
}
uint16_t* C_ptr = (uint16_t*)s_C_tensor.data_ptr();
// Helper macro: launch specialized fused kernel
#define LAUNCH_FUSED(cM, cN, cK) do { \
dim3 grid((unsigned)((cM+15)/16), (unsigned)((cN+15)/16)); \
kern_fused_spec<cM, cN, cK><<<grid, 256>>>( \
(const uint32_t*)A.data_ptr(), (const uint8_t*)Bs.data_ptr(), \
(const uint8_t*)Bsc.data_ptr(), C_ptr); \
} while(0)
// Dispatch to specialized kernel if shape matches benchmark shapes
// Specialized fused: K<=1024 shapes + M=16 K=7168 (underutilized split-k)
if (M==4 && N==2880 && K==512) { LAUNCH_FUSED(4, 2880, 512); }
else if (M==16 && N==2112 && K==7168) {
// Hybrid inline-A + GLL-B + nt — sweep winner: sp8_b2
constexpr int BATCH_S1 = 2;
constexpr int SPLITS = 8;
dim3 grid(1u, 132u);
size_t lds_b = (size_t)SPLITS * BATCH_S1 * 1024;
size_t lds_red = (size_t)SPLITS * 256 * sizeof(float);
kern_hybrid_s1<BATCH_S1, SPLITS><<<grid, SPLITS*64, lds_b + lds_red>>>(
(const uint32_t*)A.data_ptr(), (const uint8_t*)Bs.data_ptr(),
(const uint8_t*)Bsc.data_ptr(), C_ptr,
16, 2112, 7168, kt, pk);
}
else if (M==32 && N==4096 && K==512) { LAUNCH_FUSED(32, 4096, 512); }
else if (M==32 && N==2880 && K==512) { LAUNCH_FUSED(32, 2880, 512); }
// v64: GLL BATCH=12 for Shape 4 — best config
// 12 GLL pre-loaded, burst 4 mid-loop, 8 MFMAs overlap for burst
// LDS: A(24576) + B(8 waves × 12 slots × 1024 = 98304) = 122880 = 120KB
else if (M==64 && N==7168 && K==2048) {
// Dedicated kern_s4: all optimizations, hardcoded constants
dim3 grid(4u, 56u);
size_t lds_a = (size_t)1 * 16 * 384 * sizeof(uint32_t);
size_t lds_b = (size_t)8 * 10 * 1024;
kern_s4<<<grid, 512, lds_a + lds_b>>>(
(const uint32_t*)A.data_ptr(), (const uint8_t*)Bs.data_ptr(),
(const uint8_t*)Bsc.data_ptr(), C_ptr,
64, 7168, 2048, kt, pk);
}
else if (M==256 && N==3072 && K==1536) {
// Dedicated kern_s5: Ki=12, N_N=12, BATCH=6
dim3 grid(16u, 16u);
size_t lds_a = (size_t)1 * 12 * 384 * sizeof(uint32_t);
size_t lds_b = (size_t)12 * 8 * 1024;
kern_s5<<<grid, 768, lds_a + lds_b>>>(
(const uint32_t*)A.data_ptr(), (const uint8_t*)Bs.data_ptr(),
(const uint8_t*)Bsc.data_ptr(), C_ptr,
256, 3072, 1536, kt, pk);
}
else if (K <= 1024) {
// Fallback for unknown K<=1024 shapes: use original kern_fused
dim3 grid((unsigned)mt, (unsigned)nt);
kern_fused<<<grid, 64>>>(
(const uint32_t*)A.data_ptr(), (const uint8_t*)Bs.data_ptr(),
(const uint8_t*)Bsc.data_ptr(), C_ptr,
(int)M, (int)N, (int)K, kt, pk);
} else {
int Ksc = (int)(K / 32);
int Mpad = ((int)M + 255) / 256 * 256;
int Kscpad = (Ksc + 7) / 8 * 8;
int qtot = (int)M * Ksc * QPG, qthr = 256, qblk = (qtot+qthr-1)/qthr;
kern_quant<<<qblk, qthr>>>(
(const uint32_t*)A.data_ptr(),
s_fp4_ptr, s_sc_ptr,
(int)M, (int)K, Ksc, Mpad, Kscpad);
if (tiles < NUM_CUS && Ki >= 16) {
int ksplit = Ki / 4;
if (ksplit < 2) ksplit = 2;
if (ksplit > 4) ksplit = 4;
hipMemsetAsync(s_temp_ptr, 0, (int)(M*N) * sizeof(float), 0);
int total_blocks = mt * nt * ksplit;
kern_splitk<<<total_blocks, 64>>>(
(const uint8_t*)s_fp4_ptr, (const uint8_t*)s_sc_ptr,
(const uint8_t*)Bs.data_ptr(), (const uint8_t*)Bsc.data_ptr(),
s_temp_ptr,
(int)M, (int)N, (int)K, kt, pk, pk, nt, ksplit);
int MN = (int)(M * N);
kern_convert<<<(MN+255)/256, 256>>>(s_temp_ptr, C_ptr, MN);
} else {
dim3 grid((unsigned)mt, (unsigned)nt);
kern_gemm<<<grid, 64>>>(
(const uint8_t*)s_fp4_ptr, (const uint8_t*)s_sc_ptr,
(const uint8_t*)Bs.data_ptr(), (const uint8_t*)Bsc.data_ptr(),
C_ptr,
(int)M, (int)N, (int)K, kt, pk, pk);
}
}
return s_C_tensor;
}
// Zero-Python-overhead dispatch: takes the data tuple directly in C++
// Avoids: Python tuple unpack, .shape access, .contiguous(), pybind11 6-arg conversion
// data = (A, B, B_q, B_shuffle, B_scale_sh)
torch::Tensor dispatch_fast(py::tuple data) {
// Direct PyTuple access — no vector construction, no per-element type check
torch::Tensor A = py::cast<torch::Tensor>(data[0]);
torch::Tensor Bs = py::cast<torch::Tensor>(data[3]); // B_shuffle
torch::Tensor Bsc = py::cast<torch::Tensor>(data[4]); // B_scale_sh
int64_t M = A.size(0);
int64_t K = A.size(1);
int64_t N = py::cast<torch::Tensor>(data[1]).size(0); // B.shape[0]
return dispatch(A, Bs, Bsc, M, N, K);
}
"""
CPP_SRC = """
namespace py = pybind11;
void preallocate(int64_t max_M, int64_t max_N, int64_t max_K);
torch::Tensor dispatch(torch::Tensor A, torch::Tensor Bs, torch::Tensor Bsc,
int64_t M, int64_t N, int64_t K);
torch::Tensor dispatch_fast(py::tuple data);
"""
print("[v102] Compiling all-C++ dispatch kernel...")
hip_module = load_inline(
name='fused_mxfp4_v102',
cpp_sources=[CPP_SRC],
cuda_sources=[HIP_SRC],
functions=['preallocate', 'dispatch', 'dispatch_fast'],
verbose=False,
extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3",
"-mllvm", "-amdgpu-kernarg-preload-count=16"],
)
# Pre-allocate ALL buffers for all benchmark shapes (max M=256, N=7168, K=7168)
hip_module.preallocate(256, 7168, 7168)
print("[v102] Done.")
def custom_kernel(data: input_t) -> output_t:
return hip_module.dispatch_fast(data)
scrolls · 1862 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