submission 740115
LunNova · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 913 lines, June 9 Researcher Reciprocity License v1.0.
submission_standalone_best_v3b.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-740115?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:c00ca907b60877ff80e87721a29bdbf7c916c8e465cf22214c9008197d1b16a0
license declaredunknown
license concludedunknown
authorsLunNova
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ char lds_bytes[3136];split-k
constexpr int kSplitK=7, kKPerSplit=1024, kNumWaves=4, kItersPerWave=2;tile-n = 64
constexpr int kN=2112, kK=7168, kKHalf=3584, kTileN=64, kMfmaCols=4;Kernel source
submission_standalone_best_v3b.py913 lines
"""Standalone best-of GEMM: 6 shapes, load_inline C++ dispatch."""
from __future__ import annotations
import os, sys, tempfile
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
# ── M=64 hand-unrolled iteration generator ──────────────────────
def _gen_m64_iters():
table = [
("a_buf0", "a_buf2", True, True),
("a_buf1", "a_buf0", True, True),
("a_buf2", "a_buf1", True, True),
("a_buf0", "a_buf2", True, True),
("a_buf1", "a_buf0", True, True),
("a_buf2", "a_buf1", True, True),
("a_buf0", "a_buf2", False, True),
("a_buf1", "a_buf0", False, False),
]
lines = []
for k, (consume, far, load_a, load_b) in enumerate(table):
lines.append(f" {{ // iter {k}")
lines.append(f" const int kg_ = (ki_base + {k}) * 2 + half;")
lines.append( " int bsc0_ = (int)lds_bscale[g::soff<64>(b_col0, kg_) - bsc_base];")
lines.append( " int bsc1_ = (int)lds_bscale[g::soff<64>(b_col1, kg_) - bsc_base];")
lines.append(f" g::v4i32 ar_; int as_; g::quant_a({consume}, ar_, as_);")
if load_b:
lines.append(f" {{ const uint32_t* bn0_ = reinterpret_cast<const uint32_t*>(bb0 + (ki_base+{k+1})*512 + half*256);")
lines.append( " b_nxt0 = {(int)bn0_[0],(int)bn0_[1],(int)bn0_[2],(int)bn0_[3]};")
lines.append(f" const uint32_t* bn1_ = reinterpret_cast<const uint32_t*>(bb1 + (ki_base+{k+1})*512 + half*256);")
lines.append( " b_nxt1 = {(int)bn1_[0],(int)bn1_[1],(int)bn1_[2],(int)bn1_[3]}; }")
lines.append( " g::mfma_pair(ar_, b_cur0, b_cur1, acc0, acc1, as_, bsc0_, bsc1_);")
if load_a:
lines.append(f" g::load_a(a_row + ((ki_base+{k+2})*2 + half)*16, {far});")
if load_b:
lines.append( " b_cur0 = b_nxt0; b_cur1 = b_nxt1;")
lines.append( " }")
return "\n".join(lines)
# ═════════════════════════════════════════════════════════════════
# HIP source sections
# ═════════════════════════════════════════════════════════════════
_HIP_HEADER = r"""
#include <torch/extension.h>
#include <torch/csrc/autograd/python_variable.h>
#include <pybind11/pybind11.h>
namespace py = pybind11;
#include <hip/hip_runtime.h>
#include <cstdint>
#include <cstdio>
#include <string>
#define CVT_PK_FP4_BF16_B0(dst, src, scale) \
dst = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(dst, src, scale, 0)
#define CVT_PK_FP4_BF16_B1(dst, src, scale) \
dst = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(dst, src, scale, 1)
#define CVT_PK_FP4_BF16_B2(dst, src, scale) \
dst = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(dst, src, scale, 2)
#define CVT_PK_FP4_BF16_B3(dst, src, scale) \
dst = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(dst, src, scale, 3)
// ═══════════════════════════════════════════════════════
// Shared utilities
// ═══════════════════════════════════════════════════════
namespace g {
using v4i32 = int __attribute__((ext_vector_type(4)));
using v4f32 = float __attribute__((ext_vector_type(4)));
using v16f32 = float __attribute__((ext_vector_type(16)));
typedef __bf16 v2bf16 __attribute__((ext_vector_type(2)));
__device__ __forceinline__ uint32_t bcu32(float v) { union{float f;uint32_t u;}x; x.f=v; return x.u; }
__device__ __forceinline__ float bcf32(uint32_t v) { union{uint32_t u;float f;}x; x.u=v; return x.f; }
__device__ __forceinline__ uint16_t f2bf16(float v) {
uint32_t b=bcu32(v); b+=((b>>16)&1u)+0x7FFFu; return (uint16_t)(b>>16);
}
__device__ __forceinline__ uint8_t e8m0sc(uint16_t m) {
uint16_t r=(uint16_t)(m+0x20u); int e=(int)((r>>7)&0xFFu);
return (e<=2)?0u:(e>=255)?254u:(uint8_t)(e-2);
}
template<int PG>
__device__ __forceinline__ int soff(int row, int group) {
return (row>>5)*(32*PG)+(group>>3)*256+(group&3)*64+(row&15)*4+((group>>2)&1)*2+((row>>4)&1);
}
__device__ __forceinline__ void load_a(const uint32_t* __restrict__ p, uint32_t* d) {
*reinterpret_cast<uint4*>(&d[0]) = *(reinterpret_cast<const uint4*>(p)+0);
*reinterpret_cast<uint4*>(&d[4]) = *(reinterpret_cast<const uint4*>(p)+1);
*reinterpret_cast<uint4*>(&d[8]) = *(reinterpret_cast<const uint4*>(p)+2);
*reinterpret_cast<uint4*>(&d[12]) = *(reinterpret_cast<const uint4*>(p)+3);
}
__device__ __forceinline__ void quant_a(const uint32_t* dw, v4i32& out, int& sc) {
uint32_t mx=0u;
#pragma unroll
for(int i=0;i<16;++i){uint32_t w=dw[i],hi=(w>>16)&0x7FFFu,lo=w&0x7FFFu;mx=(hi>mx)?hi:mx;mx=(lo>mx)?lo:mx;}
uint8_t s=e8m0sc((uint16_t)mx); sc=(int)s;
float fwd=(s==0u)?bcf32(0x00400000u):bcf32((uint32_t)s<<23);
const v2bf16*p=reinterpret_cast<const v2bf16*>(dw);
uint32_t pk[4]={0,0,0,0};
CVT_PK_FP4_BF16_B0(pk[0],p[0],fwd); CVT_PK_FP4_BF16_B0(pk[1],p[4],fwd);
CVT_PK_FP4_BF16_B0(pk[2],p[8],fwd); CVT_PK_FP4_BF16_B0(pk[3],p[12],fwd);
CVT_PK_FP4_BF16_B1(pk[0],p[1],fwd); CVT_PK_FP4_BF16_B1(pk[1],p[5],fwd);
CVT_PK_FP4_BF16_B1(pk[2],p[9],fwd); CVT_PK_FP4_BF16_B1(pk[3],p[13],fwd);
CVT_PK_FP4_BF16_B2(pk[0],p[2],fwd); CVT_PK_FP4_BF16_B2(pk[1],p[6],fwd);
CVT_PK_FP4_BF16_B2(pk[2],p[10],fwd); CVT_PK_FP4_BF16_B2(pk[3],p[14],fwd);
CVT_PK_FP4_BF16_B3(pk[0],p[3],fwd); CVT_PK_FP4_BF16_B3(pk[1],p[7],fwd);
CVT_PK_FP4_BF16_B3(pk[2],p[11],fwd); CVT_PK_FP4_BF16_B3(pk[3],p[15],fwd);
out = {(int)pk[0],(int)pk[1],(int)pk[2],(int)pk[3]};
}
// v_pk_max_u16 tree reduction — fewer instructions than scalar loop
__device__ __forceinline__ void quant_a_pk(const uint32_t* dw, v4i32& out, int& sc) {
const uint32_t mask = 0x7FFF7FFFu;
uint32_t p0=dw[0]&mask,p1=dw[1]&mask,p2=dw[2]&mask,p3=dw[3]&mask;
uint32_t p4=dw[4]&mask,p5=dw[5]&mask,p6=dw[6]&mask,p7=dw[7]&mask;
uint32_t p8=dw[8]&mask,p9=dw[9]&mask,pA=dw[10]&mask,pB=dw[11]&mask;
uint32_t pC=dw[12]&mask,pD=dw[13]&mask,pE=dw[14]&mask,pF=dw[15]&mask;
uint32_t m01,m23,m45,m67,m89,mAB,mCD,mEF;
asm("v_pk_max_u16 %0,%1,%2":"=v"(m01):"v"(p0),"v"(p1));
asm("v_pk_max_u16 %0,%1,%2":"=v"(m23):"v"(p2),"v"(p3));
asm("v_pk_max_u16 %0,%1,%2":"=v"(m45):"v"(p4),"v"(p5));
asm("v_pk_max_u16 %0,%1,%2":"=v"(m67):"v"(p6),"v"(p7));
asm("v_pk_max_u16 %0,%1,%2":"=v"(m89):"v"(p8),"v"(p9));
asm("v_pk_max_u16 %0,%1,%2":"=v"(mAB):"v"(pA),"v"(pB));
asm("v_pk_max_u16 %0,%1,%2":"=v"(mCD):"v"(pC),"v"(pD));
asm("v_pk_max_u16 %0,%1,%2":"=v"(mEF):"v"(pE),"v"(pF));
uint32_t t0,t1,t2,t3;
asm("v_pk_max_u16 %0,%1,%2":"=v"(t0):"v"(m01),"v"(m23));
asm("v_pk_max_u16 %0,%1,%2":"=v"(t1):"v"(m45),"v"(m67));
asm("v_pk_max_u16 %0,%1,%2":"=v"(t2):"v"(m89),"v"(mAB));
asm("v_pk_max_u16 %0,%1,%2":"=v"(t3):"v"(mCD),"v"(mEF));
uint32_t u0,u1;
asm("v_pk_max_u16 %0,%1,%2":"=v"(u0):"v"(t0),"v"(t1));
asm("v_pk_max_u16 %0,%1,%2":"=v"(u1):"v"(t2),"v"(t3));
uint32_t fpk;
asm("v_pk_max_u16 %0,%1,%2":"=v"(fpk):"v"(u0),"v"(u1));
uint32_t mx = (fpk>>16) > (fpk&0xFFFFu) ? (fpk>>16) : (fpk&0xFFFFu);
uint8_t s=e8m0sc((uint16_t)mx); sc=(int)s;
float fwd=(s==0u)?bcf32(0x00400000u):bcf32((uint32_t)s<<23);
const v2bf16*p=reinterpret_cast<const v2bf16*>(dw);
uint32_t pk[4]={0,0,0,0};
CVT_PK_FP4_BF16_B0(pk[0],p[0],fwd); CVT_PK_FP4_BF16_B0(pk[1],p[4],fwd);
CVT_PK_FP4_BF16_B0(pk[2],p[8],fwd); CVT_PK_FP4_BF16_B0(pk[3],p[12],fwd);
CVT_PK_FP4_BF16_B1(pk[0],p[1],fwd); CVT_PK_FP4_BF16_B1(pk[1],p[5],fwd);
CVT_PK_FP4_BF16_B1(pk[2],p[9],fwd); CVT_PK_FP4_BF16_B1(pk[3],p[13],fwd);
CVT_PK_FP4_BF16_B2(pk[0],p[2],fwd); CVT_PK_FP4_BF16_B2(pk[1],p[6],fwd);
CVT_PK_FP4_BF16_B2(pk[2],p[10],fwd); CVT_PK_FP4_BF16_B2(pk[3],p[14],fwd);
CVT_PK_FP4_BF16_B3(pk[0],p[3],fwd); CVT_PK_FP4_BF16_B3(pk[1],p[7],fwd);
CVT_PK_FP4_BF16_B3(pk[2],p[11],fwd); CVT_PK_FP4_BF16_B3(pk[3],p[15],fwd);
out = {(int)pk[0],(int)pk[1],(int)pk[2],(int)pk[3]};
}
__device__ __forceinline__ v4f32 mfma16v(v4i32 a, v4i32 b, v4f32 c, int as, int bs) {
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
:"+v"(c):"v"(a),"v"(b),"v"(as),"v"(bs)); return c;
}
__device__ __forceinline__ v4f32 mfma16a(v4i32 a, v4i32 b, v4f32 c, int as, int bs) {
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
:"+a"(c):"v"(a),"v"(b),"v"(as),"v"(bs)); return c;
}
__device__ __forceinline__ void mfma_pair(v4i32 a, v4i32 b0, v4i32 b1,
v16f32& c0, v16f32& c1, int as, int bs0, int bs1) {
__builtin_amdgcn_sched_barrier(0); __builtin_amdgcn_s_setprio(1);
asm volatile(
"v_mfma_scale_f32_32x32x64_f8f6f4 %0,%2,%3,%0,%4,%5 cbsz:4 blgp:4\n"
"v_mfma_scale_f32_32x32x64_f8f6f4 %1,%2,%6,%1,%4,%7 cbsz:4 blgp:4\n"
:"+a"(c0),"+a"(c1):"v"(a),"v"(b0),"v"(as),"v"(bs0),"v"(b1),"v"(bs1));
__builtin_amdgcn_s_setprio(0);
}
} // namespace g
"""
# ═══════════════════════════════════════════════════════
# K1: M=4, N=2880, K=512 (4-wave cooperative A-quant, 32x32x64)
# grid=90, block=256, LDS=3136
# From sub_f1b3a2ab493a/4x2880x512.hip
# ═══════════════════════════════════════════════════════
_HIP_M4 = r"""
namespace fused_m4_common {
using v4i32 = int __attribute__((ext_vector_type(4)));
using v16f32 = float __attribute__((ext_vector_type(16)));
__device__ __forceinline__ uint32_t m4_bcu32(float v) { union{float f;uint32_t u;}x; x.f=v; return x.u; }
__device__ __forceinline__ float m4_bcf32(uint32_t v) { union{uint32_t u;float f;}x; x.u=v; return x.f; }
__device__ __forceinline__ uint16_t m4_f2bf16(float v) {
uint32_t bits=m4_bcu32(v); bits+=((bits>>16)&1u)+0x7FFFu; return (uint16_t)(bits>>16);
}
__device__ __forceinline__ uint8_t m4_e8m0sc(uint16_t max_abs) {
uint16_t r=(uint16_t)(max_abs+0x20u); int e=(int)((r>>7)&0xFFu);
return (e<=2)?0u:(e>=255)?254u:(uint8_t)(e-2);
}
__device__ __forceinline__ int m4_soff(int row, int group) {
return (row>>5)*(32*16)+(group>>3)*256+(group&3)*64+(row&15)*4+((group>>2)&1)*2+((row>>4)&1);
}
__device__ __forceinline__ float m4_e8m0_inv(uint8_t biased) {
if (biased==0u) return m4_bcf32(0x7F000000u);
return m4_bcf32((uint32_t)(254u-biased)<<23);
}
__device__ __forceinline__ uint8_t m4_pack_fp4(float scaled) {
uint32_t bits=m4_bcu32(scaled);
uint32_t sign=bits&0x80000000u; bits^=sign;
float mag=m4_bcf32(bits);
uint8_t encoded=0x7u;
if (mag<6.0f) {
if (mag<1.0f) {
constexpr uint32_t kDM=(uint32_t)(149)<<23;
uint32_t d=m4_bcu32(mag+m4_bcf32(kDM)); d-=kDM;
encoded=(uint8_t)d;
} else {
uint32_t normal=bits;
uint32_t mo=(normal>>22)&0x1u;
normal=(uint32_t)((int32_t)normal+(-126*(1<<23)+(1<<21)-1));
normal+=mo; normal>>=22;
encoded=(uint8_t)normal;
}
}
return (uint8_t)((encoded|(uint8_t)(sign>>28))&0x0Fu);
}
__device__ __forceinline__ v16f32 m4_mfma(v4i32 a, v4i32 b, v16f32 acc, int as, int bs) {
asm volatile("v_mfma_scale_f32_32x32x64_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
:"+v"(acc):"v"(a),"v"(b),"v"(as),"v"(bs)); return acc;
}
} // namespace fused_m4_common
extern "C" __global__ __launch_bounds__(256)
void fused_m4_m4_n2880_k512(
const uint16_t* __restrict__ A,
const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh,
uint16_t* __restrict__ C)
{
using namespace fused_m4_common;
constexpr int M=4, N=2880, K=512, K_HALF=256, K_tiles=8, padded_groups=16;
const int col_base = blockIdx.x * 32;
if (col_base >= N) return;
const int tid=threadIdx.x;
const int wave_id=tid>>6;
const int tid_in_wave=tid&63;
const int lane=tid&31;
const int half=(tid>>5)&1;
__shared__ char lds_bytes[3136];
uint8_t* lds_a_fp4 = reinterpret_cast<uint8_t*>(lds_bytes+0);
uint8_t* lds_a_scale = reinterpret_cast<uint8_t*>(lds_bytes+1024);
float* lds_reduce = reinterpret_cast<float*>(lds_bytes+1088);
const int b_col = col_base + lane;
const bool b_valid = (b_col < N);
const int n_in_tile = lane & 15;
const int n_tile = (col_base>>4) + (lane>>4);
const int ki0 = wave_id*2;
const int ki1 = ki0+1;
v4i32 br0={0,0,0,0}, br1={0,0,0,0};
int bs0=127, bs1=127;
if (b_valid) {
const int b_flat0 = ((n_tile*K_tiles+ki0)*2+half)*256 + n_in_tile*16;
const uint32_t* s0 = reinterpret_cast<const uint32_t*>(B_shuffle+b_flat0);
br0 = {(int)s0[0],(int)s0[1],(int)s0[2],(int)s0[3]};
bs0 = (int)B_scale_sh[m4_soff(b_col, ki0*2+half)];
const int b_flat1 = ((n_tile*K_tiles+ki1)*2+half)*256 + n_in_tile*16;
const uint32_t* s1 = reinterpret_cast<const uint32_t*>(B_shuffle+b_flat1);
br1 = {(int)s1[0],(int)s1[1],(int)s1[2],(int)s1[3]};
bs1 = (int)B_scale_sh[m4_soff(b_col, ki1*2+half)];
}
// Cooperative A quantization into LDS
{
const int row = wave_id;
const int group = tid_in_wave>>2;
const int tib = tid_in_wave&3;
const uint16_t* ap = A + row*K + group*32 + tib*8;
uint16_t ar[8]; uint16_t mx=0u;
#pragma unroll
for (int j=0;j<8;++j) { ar[j]=ap[j]; uint16_t ab=ar[j]&0x7FFFu; mx=(ab>mx)?ab:mx; }
mx=(uint16_t)max((int)mx, __shfl_xor((int)mx,1,64));
mx=(uint16_t)max((int)mx, __shfl_xor((int)mx,2,64));
uint8_t sb = m4_e8m0sc(mx);
float is = m4_e8m0_inv(sb);
uint32_t pd=0u;
#pragma unroll
for (int j=0;j<4;++j) {
uint8_t lo = m4_pack_fp4(m4_bcf32((uint32_t)(ar[j*2])<<16)*is);
uint8_t hi = m4_pack_fp4(m4_bcf32((uint32_t)(ar[j*2+1])<<16)*is);
pd |= ((uint32_t)((lo&0xFu)|(hi<<4))<<(j*8));
}
*reinterpret_cast<uint32_t*>(lds_a_fp4 + row*K_HALF + group*16 + tib*4) = pd;
if (tib==0) lds_a_scale[row*16+group] = sb;
}
__syncthreads();
const bool real = (lane < M);
v16f32 acc;
#pragma unroll
for (int i=0;i<16;++i) acc[i]=0.0f;
{ // MFMA iteration 0
v4i32 ar={0,0,0,0}; int as=127;
if (real) {
const uint32_t* s=reinterpret_cast<const uint32_t*>(lds_a_fp4+lane*K_HALF+ki0*32+half*16);
ar={(int)s[0],(int)s[1],(int)s[2],(int)s[3]};
as=(int)lds_a_scale[lane*16+ki0*2+half];
}
acc = m4_mfma(ar, br0, acc, as, bs0);
}
{ // MFMA iteration 1
v4i32 ar={0,0,0,0}; int as=127;
if (real) {
const uint32_t* s=reinterpret_cast<const uint32_t*>(lds_a_fp4+lane*K_HALF+ki1*32+half*16);
ar={(int)s[0],(int)s[1],(int)s[2],(int)s[3]};
as=(int)lds_a_scale[lane*16+ki1*2+half];
}
acc = m4_mfma(ar, br1, acc, as, bs1);
}
// Reduction
if (half==0) {
#pragma unroll
for (int r=0;r<M;++r) lds_reduce[wave_id*32*M+lane*M+r] = acc[r];
}
__syncthreads();
if (wave_id==0 && half==0 && (col_base+lane)<N) {
#pragma unroll
for (int r=0;r<M;++r) {
float sum=0.0f;
#pragma unroll
for (int w=0;w<4;++w) sum += lds_reduce[w*32*M+lane*M+r];
C[r*N+col_base+lane] = m4_f2bf16(sum);
}
}
}
"""
# ═══════════════════════════════════════════════════════
# K2: M=16, N=2112, K=7168 (two-pass split-K=7, workspace + reduce)
# From gen_m16_varI.py — no atomics, no device-scope staging
# GEMM: grid=231, block=256, shared=16384 (extern)
# Reduce: grid=132, block=256
# ═══════════════════════════════════════════════════════
_HIP_M16 = r"""
// ── PASS 1: GEMM kernel — stores f32 partials to workspace[7][16][2112] ──
extern "C" __global__ __attribute__((amdgpu_flat_work_group_size(256,256)))
void m16_varI_gemm(
const uint16_t* __restrict__ A,
const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh,
float* __restrict__ workspace)
{
using namespace g;
constexpr int kN=2112, kK=7168, kKHalf=3584, kTileN=64, kMfmaCols=4;
constexpr int kSplitK=7, kKPerSplit=1024, kNumWaves=4, kItersPerWave=2;
constexpr int kNumNTiles=33, kBPanelStride=kKHalf*16;
constexpr int kLdsFloatsPerCol=16*kNumWaves*16; // 1024
constexpr int kNXCD=8, kC=4, kBPC=32, kLimit=224;
int xy = blockIdx.x;
if (xy < kLimit) {
int xcd=xy%kNXCD, local_=xy/kNXCD;
int chunk=local_/kC, pos=local_%kC;
xy = chunk*kBPC + xcd*kC + pos;
}
const int k_split = xy / kNumNTiles;
const int n_tile = xy % kNumNTiles;
const int col_base = n_tile * kTileN;
const int tid=threadIdx.x, wave_id=tid>>6, lane=tid&63;
const int row=lane&15, k_quarter=lane>>4;
extern __shared__ char lds_raw[];
float* lds = reinterpret_cast<float*>(lds_raw);
const uint32_t* a_row = reinterpret_cast<const uint32_t*>(A + row * kK);
const uint8_t* b_bases[kMfmaCols];
#pragma unroll
for (int c=0;c<kMfmaCols;++c) {
int b_col = col_base + c*16 + row;
b_bases[c] = B_shuffle + (b_col>>4)*kBPanelStride + (b_col&15)*16;
}
const int split_iter_base = k_split * (kKPerSplit/128);
const int ki_base = split_iter_base + wave_id * kItersPerWave;
const int kg0 = ki_base*4+k_quarter;
const int kg1 = (ki_base+1)*4+k_quarter;
// All loads upfront
asm volatile("" ::: "memory");
uint32_t a_buf0[16], a_buf1[16];
{ const uint32_t* s = a_row + kg0*16;
#pragma unroll
for (int i=0;i<16;++i) a_buf0[i]=s[i]; }
{ const uint32_t* s = a_row + kg1*16;
#pragma unroll
for (int i=0;i<16;++i) a_buf1[i]=s[i]; }
v4i32 b0[kMfmaCols], b1[kMfmaCols];
int bs0[kMfmaCols], bs1[kMfmaCols];
#pragma unroll
for (int c=0;c<kMfmaCols;++c) {
const uint32_t* s0=reinterpret_cast<const uint32_t*>(b_bases[c]+ki_base*1024+k_quarter*256);
b0[c]={(int)s0[0],(int)s0[1],(int)s0[2],(int)s0[3]};
const uint32_t* s1=reinterpret_cast<const uint32_t*>(b_bases[c]+(ki_base+1)*1024+k_quarter*256);
b1[c]={(int)s1[0],(int)s1[1],(int)s1[2],(int)s1[3]};
int b_col = col_base+c*16+row;
bs0[c]=(int)B_scale_sh[soff<224>(b_col,kg0)];
bs1[c]=(int)B_scale_sh[soff<224>(b_col,kg1)];
}
asm volatile("":"+v"(bs0[0]),"+v"(bs0[1]),"+v"(bs0[2]),"+v"(bs0[3]),
"+v"(bs1[0]),"+v"(bs1[1]),"+v"(bs1[2]),"+v"(bs1[3])::"memory");
// Quant + MFMA
v4f32 acc[kMfmaCols];
#pragma unroll
for (int c=0;c<kMfmaCols;++c) acc[c]={0,0,0,0};
{ v4i32 ar; int as; quant_a(a_buf0, ar, as);
#pragma unroll
for (int c=0;c<kMfmaCols;++c) acc[c]=mfma16v(ar, b0[c], acc[c], as, bs0[c]); }
{ v4i32 ar; int as; quant_a(a_buf1, ar, as);
#pragma unroll
for (int c=0;c<kMfmaCols;++c) acc[c]=mfma16v(ar, b1[c], acc[c], as, bs1[c]); }
// LDS reduce 4 waves
const int oc16=lane&15, rb4=(lane>>4)*4;
#pragma unroll
for (int c=0;c<kMfmaCols;++c)
#pragma unroll
for (int r=0;r<4;++r)
lds[c*kLdsFloatsPerCol + (rb4+r)*kNumWaves*16 + wave_id*16 + oc16] = acc[c][r];
__syncthreads();
// Reduce and store to workspace
const int ws_base = k_split * 16 * kN;
{ int c=wave_id;
#pragma unroll
for (int r=0;r<4;++r) {
int out_row=rb4+r;
if (out_row < 16) {
int lb = c*kLdsFloatsPerCol + out_row*kNumWaves*16 + oc16;
float sum=0.0f;
#pragma unroll
for (int w=0;w<kNumWaves;++w) sum += lds[lb+w*16];
int gc = col_base + c*16 + oc16;
workspace[ws_base + out_row*kN + gc] = sum;
}
}
}
}
// ── PASS 2: Reduce kernel — sum 7 partials, convert to bf16 ──
extern "C" __global__ __attribute__((amdgpu_flat_work_group_size(256,256)))
void m16_varI_reduce(
const float* __restrict__ workspace,
uint16_t* __restrict__ C)
{
using namespace g;
constexpr int kM=16, kN=2112, kSplitK=7;
const int idx = blockIdx.x * 256 + threadIdx.x;
if (idx >= kM*kN) return;
float sum=0.0f;
#pragma unroll
for (int s=0;s<kSplitK;++s) sum += workspace[s*kM*kN + idx];
C[idx] = f2bf16(sum);
}
"""
# ═══════════════════════════════════════════════════════
# K3: M=32 kernels (template body, 2 shapes)
# 32x2880x512: grid=180, 32x4096x512: grid=256
# block=256, shared=8192 (static)
# ═══════════════════════════════════════════════════════
_HIP_M32 = r"""
template <int kN>
__device__ void m32_body(const uint16_t* __restrict__ A,
const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh,
uint16_t* __restrict__ C)
{
using namespace g;
const int n_tile=blockIdx.x/2, m_tile=blockIdx.x%2;
const int col_base=n_tile*32, row_base=m_tile*16;
const int tid=threadIdx.x, wave_id=tid>>6, lane=tid&63;
const int row16=lane&15, kq=lane>>4;
__shared__ char lds_raw[8192];
float* lds = reinterpret_cast<float*>(lds_raw);
const int ki=wave_id, kg=ki*4+kq;
const int a_row=row_base+row16;
int bsc[2];
#pragma unroll
for (int c=0;c<2;++c) {
int bc=col_base+c*16+row16;
bsc[c]=(int)B_scale_sh[soff<16>(bc,kg)];
}
asm volatile("":"+v"(bsc[0]),"+v"(bsc[1])::"memory");
v4i32 b_regs[2];
#pragma unroll
for (int c=0;c<2;++c) {
int bc=col_base+c*16+row16;
const uint8_t* bb=B_shuffle+(bc>>4)*(256*16)+(bc&15)*16;
const uint32_t* bs=reinterpret_cast<const uint32_t*>(bb+ki*1024+kq*256);
b_regs[c]={(int)bs[0],(int)bs[1],(int)bs[2],(int)bs[3]};
}
uint32_t a_buf[16];
{ const uint32_t* as=reinterpret_cast<const uint32_t*>(A+a_row*512+kg*32);
#pragma unroll
for(int i=0;i<16;++i) a_buf[i]=as[i]; }
v4f32 acc0={0,0,0,0}, acc1={0,0,0,0};
{ v4i32 ar; int asc; quant_a(a_buf,ar,asc);
acc0=mfma16a(ar,b_regs[0],acc0,asc,bsc[0]);
acc1=mfma16a(ar,b_regs[1],acc1,asc,bsc[1]); }
const int oc=lane&15, rb4=(lane>>4)*4;
#pragma unroll
for (int r=0;r<4;++r) {
int or_=rb4+r;
lds[or_*128+wave_id*32+oc]=acc0[r];
lds[or_*128+wave_id*32+oc+16]=acc1[r];
}
__syncthreads();
const int e0=tid, e1=tid+256;
const int lr0=e0/32,lc0=e0%32, lr1=e1/32,lc1=e1%32;
const int lb0=lr0*128+lc0, lb1=lr1*128+lc1;
float s0=lds[lb0]+lds[lb0+32]+lds[lb0+64]+lds[lb0+96];
float s1=lds[lb1]+lds[lb1+32]+lds[lb1+64]+lds[lb1+96];
C[(row_base+lr0)*kN+col_base+lc0]=f2bf16(s0);
C[(row_base+lr1)*kN+col_base+lc1]=f2bf16(s1);
}
extern "C" __global__ __attribute__((amdgpu_flat_work_group_size(256,256)))
void m32_mk4l_m32_n2880_k512(const uint16_t* __restrict__ A,
const uint8_t* __restrict__ B_shuffle, const uint8_t* __restrict__ B_scale_sh,
uint16_t* __restrict__ C) { m32_body<2880>(A,B_shuffle,B_scale_sh,C); }
extern "C" __global__ __attribute__((amdgpu_flat_work_group_size(256,256)))
void m32_mk4l_m32_n4096_k512(const uint16_t* __restrict__ A,
const uint8_t* __restrict__ B_shuffle, const uint8_t* __restrict__ B_scale_sh,
uint16_t* __restrict__ C) { m32_body<4096>(A,B_shuffle,B_scale_sh,C); }
"""
# ═══════════════════════════════════════════════════════
# K4: M=64, N=7168, K=2048 (32x32x64, fused MFMA pair)
# grid=224, block=256, shared=32768 (static)
# ═══════════════════════════════════════════════════════
_HIP_M64_PRE = r"""
extern "C" __global__ __attribute__((amdgpu_flat_work_group_size(256,256)))
void m64_mk3a_m64_n7168_k2048(
const uint16_t* __restrict__ A,
const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh,
uint16_t* __restrict__ C)
{
using namespace g;
// XCD: W=2, C=14, BPC=112, LIMIT=224=TOTAL
int xy = blockIdx.x;
{ int xcd=xy%8, loc=xy/8, chunk=loc/14, pos=loc%14;
xy = chunk*112 + xcd*14 + pos; }
const int l=xy%224, m_tile=(xy/224)*2+(l%2), n_tile=l/2;
const int col_base=n_tile*64, row_base=m_tile*32;
const int tid=threadIdx.x, wave_id=tid>>6, lane=tid&31, half=(tid>>5)&1;
__shared__ char lds_bytes[32768];
constexpr int lds_half = 4096;
const int b_col0=col_base+lane, b_col1=col_base+32+lane;
const uint8_t* bb0 = B_shuffle + (b_col0>>4)*(1024*16) + (b_col0&15)*16;
const uint8_t* bb1 = B_shuffle + (b_col1>>4)*(1024*16) + (b_col1&15)*16;
// Cooperative B scale preload (4096 bytes, 256 threads x 16 bytes)
{ int bsc_gbase = (col_base>>5)*(32*64);
const uint8_t* bsc_src = B_scale_sh + bsc_gbase;
int my_off = tid * 16;
#pragma unroll
for (int bi=0;bi<16;++bi)
reinterpret_cast<uint8_t*>(lds_bytes)[my_off+bi] = bsc_src[my_off+bi];
}
__syncthreads();
const uint8_t* lds_bscale = reinterpret_cast<const uint8_t*>(lds_bytes);
const int bsc_base = (col_base>>5)*(32*64);
const int my_row = row_base + lane;
const uint32_t* a_row = reinterpret_cast<const uint32_t*>(A + my_row * 2048);
const int ki_base = wave_id * 8;
v16f32 acc0, acc1;
#pragma unroll
for (int i=0;i<16;++i) { acc0[i]=0.0f; acc1[i]=0.0f; }
// Prologue B data
v4i32 b_cur0, b_cur1;
{ const uint32_t* bs0=reinterpret_cast<const uint32_t*>(bb0+ki_base*512+half*256);
b_cur0={(int)bs0[0],(int)bs0[1],(int)bs0[2],(int)bs0[3]};
const uint32_t* bs1=reinterpret_cast<const uint32_t*>(bb1+ki_base*512+half*256);
b_cur1={(int)bs1[0],(int)bs1[1],(int)bs1[2],(int)bs1[3]}; }
uint32_t a_buf0[16], a_buf1[16], a_buf2[16];
load_a(a_row + (ki_base*2+half)*16, a_buf0);
load_a(a_row + ((ki_base+1)*2+half)*16, a_buf1);
v4i32 b_nxt0, b_nxt1;
// 8 hand-unrolled iterations
"""
_HIP_M64_POST = r"""
// Barrier + LDS double reduction
__syncthreads();
float* lds_r = reinterpret_cast<float*>(lds_bytes);
#pragma unroll
for (int i=0;i<4;++i) {
#pragma unroll
for (int j=0;j<4;++j) {
int row=half*4+j+i*8;
int base=row*128+wave_id*32+lane;
lds_r[base]=acc0[i*4+j];
lds_r[base+lds_half]=acc1[i*4+j];
}
}
__syncthreads();
{ int i=wave_id;
#pragma unroll
for (int j=0;j<4;++j) {
int row=half*4+j+i*8;
int gr=row_base+row;
float s0=0.0f, s1=0.0f;
#pragma unroll
for (int w=0;w<4;++w) {
s0 += lds_r[row*128+w*32+lane];
s1 += lds_r[lds_half+row*128+w*32+lane];
}
C[gr*7168+col_base+lane]=f2bf16(s0);
C[gr*7168+col_base+32+lane]=f2bf16(s1);
}
}
}
"""
# ═══════════════════════════════════════════════════════
# K5: M=256, N=3072, K=1536 (32x32x64, fused MFMA pair)
# grid=384, block=256, shared=32768 (static)
# ═══════════════════════════════════════════════════════
_HIP_M256 = r"""
extern "C" __global__ __attribute__((amdgpu_flat_work_group_size(256,256)))
void m256_v6a_kernel(
const uint16_t* __restrict__ A,
const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh,
uint16_t* __restrict__ C)
{
using namespace g;
// XCD: W=8, C=6, BPC=48, LIMIT=384=TOTAL
int xy = blockIdx.x;
{ int xcd=xy%8, loc=xy/8, chunk=loc/6, pos=loc%6;
xy = chunk*48 + xcd*6 + pos; }
const int l=xy%384, m_tile=(xy/384)*8+(l%8), n_tile=l/8;
const int col_base=n_tile*64, row_base=m_tile*32;
const int tid=threadIdx.x, wave_id=tid>>6, lane=tid&31, half=(tid>>5)&1;
__shared__ char lds_bytes[32768];
constexpr int lds_half = 4096;
const int bc0=col_base+lane, bc1=col_base+32+lane;
const uint8_t* bb0 = B_shuffle + (bc0>>4)*(768*16) + (bc0&15)*16;
const uint8_t* bb1 = B_shuffle + (bc1>>4)*(768*16) + (bc1&15)*16;
const int my_row = row_base + lane;
const uint32_t* a_row = reinterpret_cast<const uint32_t*>(A + my_row * 1536);
const int ki_base = wave_id * 6;
v16f32 acc0, acc1;
#pragma unroll
for (int i=0;i<16;++i) { acc0[i]=0.0f; acc1[i]=0.0f; }
// Prologue A/B loads (fly during LDS preload)
uint32_t a_triple[3][16];
load_a(a_row + (ki_base*2+half)*16, a_triple[0]);
load_a(a_row + ((ki_base+1)*2+half)*16, a_triple[1]);
v4i32 b_cur0, b_cur1;
{ const uint32_t* bs0=reinterpret_cast<const uint32_t*>(bb0+ki_base*512+half*256);
b_cur0={(int)bs0[0],(int)bs0[1],(int)bs0[2],(int)bs0[3]};
const uint32_t* bs1=reinterpret_cast<const uint32_t*>(bb1+ki_base*512+half*256);
b_cur1={(int)bs1[0],(int)bs1[1],(int)bs1[2],(int)bs1[3]}; }
// Anti-stagger
if (blockIdx.x & 4) asm volatile("s_sleep 1" :::);
if (blockIdx.x & 8) asm volatile("s_sleep 1" :::);
// Cooperative B scale preload (3072 bytes, 12 bytes/thread)
{ int bsc_gbase = (col_base>>5)*(32*48);
const uint8_t* bsc_src = B_scale_sh + bsc_gbase;
constexpr int bpt = 12;
int my_off = tid * bpt;
#pragma unroll
for (int bi=0;bi<bpt;++bi)
reinterpret_cast<uint8_t*>(lds_bytes)[my_off+bi] = bsc_src[my_off+bi]; }
__syncthreads();
const uint8_t* lds_bscale = reinterpret_cast<const uint8_t*>(lds_bytes);
const int bsc_base = (col_base>>5)*(32*48);
// Main K loop: 6 iters
#pragma unroll
for (int k=0;k<6;++k) {
const int ki=ki_base+k;
int kg=ki*2+half;
int bsc0_=(int)lds_bscale[soff<48>(bc0,kg)-bsc_base];
int bsc1_=(int)lds_bscale[soff<48>(bc1,kg)-bsc_base];
v4i32 ar; int asc;
quant_a(a_triple[k%3], ar, asc);
v4i32 b_nxt0={0,0,0,0}, b_nxt1={0,0,0,0};
if (k+1<6) {
const uint32_t* bs0=reinterpret_cast<const uint32_t*>(bb0+(ki+1)*512+half*256);
b_nxt0={(int)bs0[0],(int)bs0[1],(int)bs0[2],(int)bs0[3]};
const uint32_t* bs1=reinterpret_cast<const uint32_t*>(bb1+(ki+1)*512+half*256);
b_nxt1={(int)bs1[0],(int)bs1[1],(int)bs1[2],(int)bs1[3]};
}
mfma_pair(ar, b_cur0, b_cur1, acc0, acc1, asc, bsc0_, bsc1_);
if (k+2<6) load_a(a_row+((ki+2)*2+half)*16, a_triple[(k+2)%3]);
b_cur0=b_nxt0; b_cur1=b_nxt1;
}
__syncthreads();
float* lds_r = reinterpret_cast<float*>(lds_bytes);
#pragma unroll
for (int i=0;i<4;++i) {
#pragma unroll
for (int j=0;j<4;++j) {
int row=half*4+j+i*8;
int base=row*128+wave_id*32+lane;
lds_r[base]=acc0[i*4+j];
lds_r[base+lds_half]=acc1[i*4+j];
}
}
__syncthreads();
{ int i=wave_id;
#pragma unroll
for (int j=0;j<4;++j) {
int row=half*4+j+i*8;
int gr=row_base+row;
float s0=0.0f, s1=0.0f;
#pragma unroll
for (int w=0;w<4;++w) {
s0 += lds_r[row*128+w*32+lane];
s1 += lds_r[lds_half+row*128+w*32+lane];
}
C[gr*3072+col_base+lane]=f2bf16(s0);
C[gr*3072+col_base+32+lane]=f2bf16(s1);
}
}
}
"""
# ═══════════════════════════════════════════════════════
# C++ Dispatch (hip_module_v2 style)
# ═══════════════════════════════════════════════════════
_HIP_DISPATCH = r"""
static py::function g_fallback_fn;
static bool g_has_fallback = false;
#define SK(m,n,k) ((uint64_t)(m)|((uint64_t)(n)<<16)|((uint64_t)(k)<<32))
torch::Tensor dispatch(py::tuple data) {
const auto& A = THPVariable_Unpack(data[0].ptr());
const auto& B = THPVariable_Unpack(data[1].ptr());
const auto& B_shuffle = THPVariable_Unpack(data[3].ptr());
const auto& B_scale_sh = THPVariable_Unpack(data[4].ptr());
int64_t m=A.size(0), k=A.size(1), n=B.size(0);
uint64_t sk = SK(m,n,k);
auto opts = torch::TensorOptions().dtype(torch::kBFloat16).device(torch::kCUDA);
const uint16_t* a = (const uint16_t*)A.data_ptr();
const uint8_t* bs = (const uint8_t*)B_shuffle.data_ptr();
const uint8_t* bsc= (const uint8_t*)B_scale_sh.data_ptr();
torch::Tensor output;
switch (sk) {
case SK(4,2880,512):
output = torch::empty({m,n}, opts);
hipLaunchKernelGGL(fused_m4_m4_n2880_k512, dim3(90),dim3(256),0,0,
a,bs,bsc,(uint16_t*)output.data_ptr());
break;
case SK(16,2112,7168): {
output = torch::empty({m,n}, opts);
auto ws_opts = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
auto workspace = torch::empty({7*16*2112}, ws_opts);
hipLaunchKernelGGL(m16_varI_gemm, dim3(231),dim3(256),16384,0,
a,bs,bsc,(float*)workspace.data_ptr());
hipLaunchKernelGGL(m16_varI_reduce, dim3(132),dim3(256),0,0,
(const float*)workspace.data_ptr(),(uint16_t*)output.data_ptr());
break;
}
case SK(32,2880,512):
output = torch::empty({m,n}, opts);
hipLaunchKernelGGL(m32_mk4l_m32_n2880_k512, dim3(180),dim3(256),0,0,
a,bs,bsc,(uint16_t*)output.data_ptr());
break;
case SK(32,4096,512):
output = torch::empty({m,n}, opts);
hipLaunchKernelGGL(m32_mk4l_m32_n4096_k512, dim3(256),dim3(256),0,0,
a,bs,bsc,(uint16_t*)output.data_ptr());
break;
case SK(64,7168,2048):
output = torch::empty({m,n}, opts);
hipLaunchKernelGGL(m64_mk3a_m64_n7168_k2048, dim3(224),dim3(256),0,0,
a,bs,bsc,(uint16_t*)output.data_ptr());
break;
case SK(256,3072,1536):
output = torch::empty({m,n}, opts);
hipLaunchKernelGGL(m256_v6a_kernel, dim3(384),dim3(256),0,0,
a,bs,bsc,(uint16_t*)output.data_ptr());
break;
default:
if (g_has_fallback)
return g_fallback_fn(data).cast<torch::Tensor>();
throw std::runtime_error("[dispatch] unsupported shape "
+ std::to_string(m)+"x"+std::to_string(n)+"x"+std::to_string(k));
}
#undef SK
return output;
}
void set_fallback(py::function fn) {
g_fallback_fn = std::move(fn);
g_has_fallback = true;
}
"""
# ═══════════════════════════════════════════════════════
# Assemble full HIP source
# ═══════════════════════════════════════════════════════
_FULL_HIP = (
_HIP_HEADER
+ _HIP_M4
+ _HIP_M16
+ _HIP_M32
+ _HIP_M64_PRE + "\n" + _gen_m64_iters() + "\n" + _HIP_M64_POST
+ _HIP_M256
+ _HIP_DISPATCH
)
_CPP_DECL = r"""
#include <torch/extension.h>
#include <pybind11/pybind11.h>
namespace py = pybind11;
torch::Tensor dispatch(py::tuple data);
void set_fallback(py::function fn);
"""
# ═══════════════════════════════════════════════════════
# Build
# ═══════════════════════════════════════════════════════
print("[standalone_best_v3b] Compiling...", file=sys.stderr)
_name = "gemm_standalone_best_v3b_ext"
_bdir = os.path.join(tempfile.gettempdir(), _name)
os.makedirs(_bdir, exist_ok=True)
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
_mod = load_inline(
name=_name,
cpp_sources=_CPP_DECL,
cuda_sources=_FULL_HIP,
functions=["dispatch", "set_fallback"],
with_cuda=True,
extra_cflags=["-O3"],
extra_cuda_cflags=[
"-O3", "-std=c++17",
"-ffast-math",
"-fgpu-flush-denormals-to-zero",
"-mllvm", "-amdgpu-kernarg-preload-count=10",
],
build_directory=_bdir,
verbose=bool(int(os.environ.get("INLINE_HIP_VERBOSE", "0"))),
)
# Register aiter fallback for unsupported shapes
def _aiter_fallback(data):
from aiter import dtypes
import aiter
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
A_fp4, A_scale = dynamic_mxfp4_quant(A)
A_scale_sh = e8m0_shuffle(A_scale)
A_q = A_fp4.view(dtypes.fp4x2)
A_scale_sh = A_scale_sh.view(dtypes.fp8_e8m0)
return aiter.gemm_a4w4(
A_q, B_shuffle, A_scale_sh, B_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True,
)
_mod.set_fallback(_aiter_fallback)
print("[standalone_best_v3b] Ready.", file=sys.stderr)
custom_kernel = _mod.dispatch
scrolls · 913 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