submission 732851
allen98. · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1394 lines, June 9 Researcher Reciprocity License v1.0.
kernel.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-732851?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:6b27ba8f981fdf998b3ef4949babc42900a77b15d5eae7912c8a7fd5efe1b855
license declaredunknown
license concludedunknown
authorsallen98.
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
FP4 GEMM v31 — Hybrid: HIP fused for M≤32 (ANY K) + Triton for M=33-127 + ASM for M≥128.num-warps = 4
num_warps=4,num_stages=1,waves_per_eu=4,matrix_instr_nonkdim=16,NUM_KSPLIT=1),shared-memory
__shared__ float red[4][64*4+16];split-k
int64_t sm, const std::string& knl_name, const std::string& co_path, int64_t splitK);stages = 1
num_warps=4,num_stages=1,waves_per_eu=4,matrix_instr_nonkdim=16,NUM_KSPLIT=1),tile-k = 512
const int nk = (K + 511) / 512; // BK=512 outer looptile-m = 16
(4, 2880, 512): dict(BLOCK_SIZE_M=16,BLOCK_SIZE_N=64,BLOCK_SIZE_K=512,GROUP_SIZE_M=1,tile-n = 64
(4, 2880, 512): dict(BLOCK_SIZE_M=16,BLOCK_SIZE_N=64,BLOCK_SIZE_K=512,GROUP_SIZE_M=1,vector-width = uint4
uint4 a0,a1,a2,a3;Kernel source
kernel.py1394 lines
"""
FP4 GEMM v31 — Hybrid: HIP fused for M≤32 (ANY K) + Triton for M=33-127 + ASM for M≥128.
Key change from v30:
- v31: Removed K threshold for HIP path. gemm_fused_16x16_kern now handles ALL K values.
For shape 2 (M=16 K=7168): single HIP kernel replaces 2 Triton launches.
Profiling showed 83% launch overhead for shape 2 — this eliminates it.
- ASM path for M≥128: hand-tuned 186% MFMA util beats Triton for M=256
"""
import functools, hashlib, os, sys, subprocess, shutil, tempfile, csv, glob
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
os.environ.setdefault("CXX", "clang++")
_IS_PROFILE_SUB = os.environ.get("PROFILE_SUBPROCESS") == "1"
_DO_PROFILE = os.environ.get("DO_PROFILE", "0") == "1" # Enable with DO_PROFILE=1 env var
_ORIG_DIR = os.environ.get("ORIG_DIR", os.path.dirname(os.path.abspath(__file__)))
if _ORIG_DIR not in sys.path:
sys.path.insert(0, _ORIG_DIR)
import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
SCALE_GROUP_SIZE = 32
# ═══════════════════════════════════════════════════════════════════
# HIP C++ kernels: 16x16 fused MFMA (M≤32 K≤1024) + ASM path (M≥128)
# ═══════════════════════════════════════════════════════════════════
CPP_SRC = r"""
#include <torch/extension.h>
void quantize_a_hip(torch::Tensor a, torch::Tensor a_q, torch::Tensor a_scale);
torch::Tensor gemm_fused_16x16_hip(torch::Tensor A, torch::Tensor Bq,
torch::Tensor Bsc, int64_t N, torch::Tensor out);
void quant_a_shuffled(torch::Tensor A, torch::Tensor Aq, torch::Tensor Asc, int64_t sm);
torch::Tensor quant_and_asm(torch::Tensor A, torch::Tensor Bsh,
torch::Tensor Bsc, int64_t N,
torch::Tensor Aq, torch::Tensor Asc, torch::Tensor out,
int64_t sm, const std::string& knl_name, const std::string& co_path, int64_t splitK);
void warmup_asm(const std::string& co_path, const std::string& knl_name);
void fast_hip16(int64_t A, int64_t Bq, int64_t Bsc, int64_t C,
int64_t M, int64_t N, int64_t K, int64_t strA, int64_t strBq, int64_t BscSN);
void fast_ext_splitk(int64_t A, int64_t Bq, int64_t Bsc, int64_t Cfp, int64_t Cbf,
int64_t M, int64_t N, int64_t K, int64_t strA, int64_t strBq, int64_t BscSN, int64_t ks);
void dispatch_hip(int64_t A, int64_t Bq, int64_t Bsc, int64_t C, int64_t Cfp,
int64_t M, int64_t N, int64_t K, int64_t strA, int64_t strBq, int64_t BscSN, int64_t ks);
void fast_reduce_splitk(int64_t src, int64_t dst, int64_t M, int64_t N, int64_t ksplit);
void launch_fused_npar(int64_t A, int64_t Bsh, int64_t Bsc, int64_t C,
int64_t M, int64_t N, int64_t K, int64_t strA, int64_t strBsh, int64_t BscSN);
void load_triton_hsaco(const std::string& path, const std::string& name);
void launch_triton_kernel(const std::string& key,
int64_t p0, int64_t p1, int64_t p2, int64_t p3,
int64_t a0, int64_t a1, int64_t a2, int64_t a3, int64_t a4,
int64_t a5, int64_t a6, int64_t a7, int64_t a8, int64_t a9, int64_t a10, int64_t a11,
int64_t gridX, int64_t gridY, int64_t blockX, int64_t shared_mem);
void launch_triton_fp4gemm(const std::string& key,
int64_t Aq, int64_t Bt, int64_t C, int64_t Asc, int64_t Bsc,
int64_t M, int64_t N, int64_t K, int64_t Ks,
int64_t sam, int64_t sbn, int64_t sck, int64_t scm, int64_t sasm,
int64_t gridX, int64_t shared_mem);
void launch_nofuse_triton(int64_t Aq, int64_t Bt, int64_t C,
int64_t Asc, int64_t Bsc,
int64_t M, int64_t N, int64_t K, int64_t Ks,
int64_t sam, int64_t sbn, int64_t scm, int64_t sasm);
"""
HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <stdint.h>
namespace {
constexpr int QUANT_GROUP = 32;
constexpr int GROUPS_PER_BLOCK = 4;
constexpr int BLOCK_THREADS = QUANT_GROUP * GROUPS_PER_BLOCK;
__device__ inline uint32_t float_to_bits(float x) {
union { float f; uint32_t u; } v; v.f = x; return v.u;
}
__device__ inline float bits_to_float(uint32_t x) {
union { uint32_t u; float f; } v; v.u = x; return v.f;
}
__device__ inline uint8_t encode_mxfp4(float x) {
constexpr float MAX_NORMAL = 6.0f, MIN_NORMAL = 1.0f;
constexpr int32_t DENORM_EXP = (127-1)+(23-1)+1;
constexpr int32_t DENORM_MASK_INT = DENORM_EXP << 23;
constexpr uint32_t NORMAL_ROUND_BIAS_BITS = 0xC11FFFFFu;
const uint32_t bits = float_to_bits(x);
const uint32_t sign = bits & 0x80000000u;
const float abs_x = bits_to_float(bits ^ sign);
bool saturate = abs_x >= MAX_NORMAL;
bool denormal = (!saturate) && (abs_x < MIN_NORMAL);
uint8_t value = 0x7;
if (denormal) {
float df = abs_x + bits_to_float((uint32_t)DENORM_MASK_INT);
value = (uint8_t)(float_to_bits(df) - (uint32_t)DENORM_MASK_INT);
} else if (!saturate) {
uint32_t ab = bits ^ sign;
uint32_t mo = (ab >> 22) & 1u;
int32_t r = (int32_t)ab + (int32_t)NORMAL_ROUND_BIAS_BITS + (int32_t)mo;
value = (uint8_t)((uint32_t)r >> 22);
}
return value | (sign ? 0x8u : 0u);
}
__global__ __launch_bounds__(BLOCK_THREADS)
void quantize_a_kernel(const __hip_bfloat16* __restrict__ a, uint8_t* __restrict__ a_q,
uint8_t* __restrict__ a_scale, int m, int k) {
const int row = blockIdx.y;
const int gib = threadIdx.x / QUANT_GROUP;
const int lane = threadIdx.x % QUANT_GROUP;
const int k_block = blockIdx.x * GROUPS_PER_BLOCK + gib;
if (k_block >= k / QUANT_GROUP) return;
const int col = k_block * QUANT_GROUP + lane;
float value = (row < m && col < k) ? (float)a[row * k + col] : 0.0f;
float abs_val = fabsf(value);
for (int off = 16; off > 0; off >>= 1) {
float o = __shfl_xor(abs_val, off); abs_val = fmaxf(abs_val, o);
}
uint32_t ab = float_to_bits(abs_val);
ab = (ab + 0x00200000u) & 0xFF800000u;
int ef = (ab >> 23u) & 0xFFu;
int se = (ab == 0u) ? -127 : max(-127, min(127, ef - 127 - 2));
float qs = (se <= 126) ? bits_to_float((uint32_t)(127 - se) << 23) : 0.0f;
uint8_t sb = (uint8_t)(se + 127);
if (lane == 0) a_scale[row * (k / QUANT_GROUP) + k_block] = sb;
uint8_t nib = encode_mxfp4(value * qs);
uint8_t pn = __shfl_xor(nib, 1);
if ((lane & 1) == 0 && row < m)
a_q[row * (k / 2) + k_block * 16 + (lane / 2)] = nib | (pn << 4);
}
typedef int __attribute__((ext_vector_type(8))) i32x8;
typedef float __attribute__((ext_vector_type(4))) f32x4;
typedef __bf16 bf16v2_t __attribute__((ext_vector_type(2)));
__device__ __forceinline__ void quant_32_hw(const uint16_t* ap, int out[4], int32_t& spk) {
uint16_t mx = 0;
#pragma unroll
for (int j = 0; j < 32; j++) mx = max(mx, (uint16_t)(ap[j] & 0x7FFF));
uint32_t au = (((uint32_t)mx << 16) + 0x200000u) & 0xFF800000u;
int ef = (au >> 23u) & 0xFFu;
int su = (au == 0u) ? -127 : max(-127, min(127, ef - 127 - 2));
float hs = (su >= -126) ? __uint_as_float((uint32_t)(su + 127) << 23) : 0.0f;
#pragma unroll
for (int j = 0; j < 4; j++) {
out[j] = 0;
out[j] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(out[j], *reinterpret_cast<const bf16v2_t*>(&ap[j*8+0]), hs, 0);
out[j] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(out[j], *reinterpret_cast<const bf16v2_t*>(&ap[j*8+2]), hs, 1);
out[j] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(out[j], *reinterpret_cast<const bf16v2_t*>(&ap[j*8+4]), hs, 2);
out[j] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(out[j], *reinterpret_cast<const bf16v2_t*>(&ap[j*8+6]), hs, 3);
}
spk = (int32_t)(uint8_t)(su + 127);
}
__device__ __forceinline__ int32_t load_b_scale_16(const uint8_t* Bsc,
int ng, int bb, int SNG, int N) {
if (ng >= N) return 127;
return (int32_t)Bsc[(ng/32)*(SNG*256)+(bb/8)*256+(bb%4)*64+(ng%16)*4+((bb%8)/4)*2+(ng%32)/16];
}
__global__ void __launch_bounds__(256, 3)
gemm_fused_16x16_kern(const uint16_t* __restrict__ A, const uint8_t* __restrict__ Bq,
const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C,
int M, int N, int K, int strA, int strBq, int BscSN) {
const int mt = blockIdx.x, nt = blockIdx.y;
const int wid = threadIdx.x >> 6, lid = threadIdx.x & 63;
const int t_row = lid & 15, kpart = lid >> 4;
const int SNG = BscSN >> 3;
const int m_row = mt * 16 + t_row, a_off = m_row * strA, b_ng = nt * 16 + t_row;
int ts = K >> 7, spw = (ts + 3) >> 2;
int ks = wid * spw * 128, ke = min(ks + spw * 128, K);
f32x4 c = {0,0,0,0};
for (int kb = ks; kb < ke; kb += 128) {
int ko = kb + kpart * 32, bk = (kb >> 1) + kpart * 16;
uint4 a0,a1,a2,a3;
if (m_row<M&&ko+31<K){const uint4*s=reinterpret_cast<const uint4*>(&A[a_off+ko]);a0=s[0];a1=s[1];a2=s[2];a3=s[3];}
else{a0={0,0,0,0};a1={0,0,0,0};a2={0,0,0,0};a3={0,0,0,0};}
int bi[4];
if(b_ng<N&&bk+15<(K>>1)){uint4 bd=*reinterpret_cast<const uint4*>(&Bq[b_ng*strBq+bk]);bi[0]=((int*)&bd)[0];bi[1]=((int*)&bd)[1];bi[2]=((int*)&bd)[2];bi[3]=((int*)&bd)[3];}
else{bi[0]=0;bi[1]=0;bi[2]=0;bi[3]=0;}
int32_t bs=load_b_scale_16(Bsc,b_ng,(kb>>5)+kpart,SNG,N);
uint16_t al[32];uint4*d=reinterpret_cast<uint4*>(al);d[0]=a0;d[1]=a1;d[2]=a2;d[3]=a3;
int ai[4];int32_t as2;quant_32_hw(al,ai,as2);
i32x8 am={ai[0],ai[1],ai[2],ai[3],0,0,0,0};
i32x8 bm={bi[0],bi[1],bi[2],bi[3],0,0,0,0};
c=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(am,bm,c,4,4,0,as2,0,bs);
}
__shared__ float red[4][64*4+16];
#pragma unroll
for(int i=0;i<4;i++) red[wid][lid*4+i]=c[i];
__syncthreads();
if(wid==0){for(int j=0;j<4;j++){
float s=red[0][lid*4+j]+red[1][lid*4+j]+red[2][lid*4+j]+red[3][lid*4+j];
int mo=mt*16+kpart*4+j,no=nt*16+t_row;
if(mo<M&&no<N){uint32_t fp=__float_as_uint(s);fp+=0x7FFFu+((fp>>16)&1u);C[mo*N+no]=(uint16_t)(fp>>16u);}
}}
}
// v78: N-parallel kernel matching Triton's warp layout {1,4} from translation doc
// 32 M-rows × 64 N-cols per CTA. 4 warps × 16 N-cols each.
// 2 M-reps × 4 K-reps = 8 MFMAs per BK=512 step. B loaded once, reused across M-reps.
// B_shuffle layout: [N/16][K/32][16][16] — coalesced 256-byte tiles.
__global__ void __launch_bounds__(256, 2)
gemm_fused_npar_kern(const uint16_t* __restrict__ A, const uint8_t* __restrict__ Bsh,
const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C,
int M, int N, int K, int strA, int strBsh, int BscSN) {
const int mt = blockIdx.x, nt = blockIdx.y;
const int wid = threadIdx.x >> 6, lid = threadIdx.x & 63;
const int t_row = lid & 15, kpart = lid >> 4;
const int SNG = BscSN >> 3;
const int b_tile_n = nt * 4 + wid;
const int b_lane = t_row;
const int b_ng = nt * 64 + wid * 16 + t_row;
const int nkt = K / 32;
// 2 M-rep base rows (32-row M-tile)
const int m_row0 = mt * 32 + t_row;
const int m_row1 = mt * 32 + 16 + t_row;
const int a_off0 = m_row0 * strA;
const int a_off1 = m_row1 * strA;
f32x4 c0 = {0,0,0,0};
f32x4 c1 = {0,0,0,0};
const int nk = (K + 511) / 512; // BK=512 outer loop
for (int ki = 0; ki < nk; ki++) {
#pragma unroll
for (int krep = 0; krep < 4; krep++) { // 4 K-reps per BK=512 step
int kbase = ki * 512 + krep * 128;
int ko = kbase + kpart * 32;
// Load B (shared between M-reps)
int btk = kbase / 32 + kpart;
int boff = b_tile_n * nkt * 256 + btk * 256 + b_lane * 16;
int bi[4];
if(b_ng<N && btk<nkt){uint4 bd=*reinterpret_cast<const uint4*>(&Bsh[boff]);
bi[0]=((int*)&bd)[0];bi[1]=((int*)&bd)[1];bi[2]=((int*)&bd)[2];bi[3]=((int*)&bd)[3];}
else{bi[0]=0;bi[1]=0;bi[2]=0;bi[3]=0;}
int32_t bs = load_b_scale_16(Bsc, b_ng, kbase/32+kpart, SNG, N);
i32x8 bm = {bi[0],bi[1],bi[2],bi[3],0,0,0,0};
// M-rep 0: load A, quant, MFMA
{
uint16_t al[32]; int ai[4]; int32_t as2;
if(m_row0<M&&ko+31<K){const uint4*s=reinterpret_cast<const uint4*>(&A[a_off0+ko]);
uint4*d=reinterpret_cast<uint4*>(al);d[0]=s[0];d[1]=s[1];d[2]=s[2];d[3]=s[3];}
else{for(int j=0;j<32;j++)al[j]=(m_row0<M&&ko+j<K)?A[a_off0+ko+j]:0;}
quant_32_hw(al,ai,as2);
i32x8 am={ai[0],ai[1],ai[2],ai[3],0,0,0,0};
c0=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(am,bm,c0,4,4,0,as2,0,bs);
}
// M-rep 1: reuse B, load different A rows
{
uint16_t al[32]; int ai[4]; int32_t as2;
if(m_row1<M&&ko+31<K){const uint4*s=reinterpret_cast<const uint4*>(&A[a_off1+ko]);
uint4*d=reinterpret_cast<uint4*>(al);d[0]=s[0];d[1]=s[1];d[2]=s[2];d[3]=s[3];}
else{for(int j=0;j<32;j++)al[j]=(m_row1<M&&ko+j<K)?A[a_off1+ko+j]:0;}
quant_32_hw(al,ai,as2);
i32x8 am={ai[0],ai[1],ai[2],ai[3],0,0,0,0};
c1=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(am,bm,c1,4,4,0,as2,0,bs);
}
}
}
// Output M-rep 0
int no = nt*64 + wid*16 + t_row;
int mo0 = mt*32 + kpart*4;
#pragma unroll
for(int j=0;j<4;j++){
if(mo0+j<M&&no<N){uint32_t fp=__float_as_uint(c0[j]);fp+=0x7FFFu+((fp>>16)&1u);
C[(mo0+j)*N+no]=(uint16_t)(fp>>16u);}
}
// Output M-rep 1
int mo1 = mt*32 + 16 + kpart*4;
#pragma unroll
for(int j=0;j<4;j++){
if(mo1+j<M&&no<N){uint32_t fp=__float_as_uint(c1[j]);fp+=0x7FFFu+((fp>>16)&1u);
C[(mo1+j)*N+no]=(uint16_t)(fp>>16u);}
}
}
// v80: Non-fused FP4 GEMM reproducing Triton's exact ISA structure
// Key findings from ISA dump:
// - isTransposed=true: MFMA src0=B_data, src1=A_data (SWAPPED)
// - 5th operand = B_scale, 6th = A_scale
// - opSel 0,1,2,3 via op_sel + op_sel_hi for 4 K-reps
// - 4 K-reps per BK=512, 2 M-reps, 8 MFMAs per K-step
// - Bt layout: [K/2, N] with stride(0)=1 (K contiguous), stride(1)=K/2
__global__ void __launch_bounds__(256, 1)
gemm_nofuse_triton(const uint8_t* __restrict__ Aq, const uint8_t* __restrict__ Bt,
uint16_t* __restrict__ C,
const uint8_t* __restrict__ Asc, const uint8_t* __restrict__ Bsc,
int M, int N, int K, int Ks, int sam, int sbn, int scm, int sasm) {
// K = Kp = original_K / 2 (packed FP4 byte count per row)
const int tid = threadIdx.x;
const int wid = tid >> 6;
const int lid = tid & 63;
const int t_row = lid & 15;
const int kpart = lid >> 4;
// Grid mapping (same as Triton)
int npn = (N + 63) / 64;
int pm = blockIdx.x / npn;
int pn = blockIdx.x % npn;
// M/N tile offsets with wrapping
int m0 = (pm * 32 + t_row) % M; // M-rep 0 row
int m1 = (pm * 32 + 16 + t_row) % M; // M-rep 1 row
int n_col = (pn * 64 + wid * 16 + t_row) % N; // N column for this thread
// Accumulators (matches Triton a[0:3], a[4:7])
f32x4 acc0 = {0,0,0,0};
f32x4 acc1 = {0,0,0,0};
// BK=512 → 4 K-reps of 128 FP4 each. K is packed byte count.
// nk = K / 256 (each BK step = 256 packed bytes = 512 FP4 elements)
int nk = K / 256; // K here = Kp = original_K / 2
for (int ki = 0; ki < nk; ki++) {
int k_byte_base = ki * 256; // byte offset for this BK step
// Pack 4 A scales for M-rep 0 (bytes for k-rep 0,1,2,3)
int32_t packed_asc0 = 0, packed_asc1 = 0;
// Pack 4 B scales
int32_t packed_bsc = 0;
{
uint8_t a0s[4], a1s[4], bs[4];
#pragma unroll
for (int kr = 0; kr < 4; kr++) {
int k_group = ki * 16 + kr * 4 + kpart; // scale group index
a0s[kr] = Asc[m0 * sasm + k_group];
a1s[kr] = Asc[m1 * sasm + k_group];
// B scale uses shuffled layout: _bsc_off(n_col, k_group, Ks)
int bsc_off = (n_col/32)*(Ks*32) + (k_group/8)*256 + (k_group%4)*64
+ (n_col%16)*4 + ((k_group/4)%2)*2 + ((n_col/16)%2);
bs[kr] = Bsc[bsc_off];
}
__builtin_memcpy(&packed_asc0, a0s, 4);
__builtin_memcpy(&packed_asc1, a1s, 4);
__builtin_memcpy(&packed_bsc, bs, 4);
}
// 4 K-reps manually unrolled (opSel must be compile-time constant)
#define KREP(KR) { \
int k_byte = k_byte_base + (KR) * 64 + kpart * 16; \
int b_off = k_byte + n_col * sbn; \
uint4 bd = *reinterpret_cast<const uint4*>(&Bt[b_off]); \
i32x8 bm = {((int*)&bd)[0],((int*)&bd)[1],((int*)&bd)[2],((int*)&bd)[3],0,0,0,0}; \
uint4 ad0 = *reinterpret_cast<const uint4*>(&Aq[m0 * sam + k_byte]); \
i32x8 am0 = {((int*)&ad0)[0],((int*)&ad0)[1],((int*)&ad0)[2],((int*)&ad0)[3],0,0,0,0}; \
uint4 ad1 = *reinterpret_cast<const uint4*>(&Aq[m1 * sam + k_byte]); \
i32x8 am1 = {((int*)&ad1)[0],((int*)&ad1)[1],((int*)&ad1)[2],((int*)&ad1)[3],0,0,0,0}; \
acc0 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4( \
bm, am0, acc0, 4, 4, (KR), packed_bsc, (KR), packed_asc0); \
acc1 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4( \
bm, am1, acc1, 4, 4, (KR), packed_bsc, (KR), packed_asc1); \
}
KREP(0) KREP(1) KREP(2) KREP(3)
#undef KREP
}
// Output with isTransposed=true: a[j] = C[N=kpart*4+j, M=t_row]
// M-row comes from t_row (src1=A dim), N-col from kpart*4+j (src0=B dim)
int out_m0 = pm * 32 + t_row; // M-rep 0
int out_m1 = pm * 32 + 16 + t_row; // M-rep 1
int out_n_base = pn * 64 + wid * 16 + kpart * 4;
if (out_m0 < M) {
#pragma unroll
for (int j = 0; j < 4; j++) {
if (out_n_base + j < N) {
uint32_t fp = __float_as_uint(acc0[j]); fp += 0x7FFFu + ((fp >> 16) & 1u);
C[out_m0 * scm + out_n_base + j] = (uint16_t)(fp >> 16u);
}
}
}
if (out_m1 < M) {
#pragma unroll
for (int j = 0; j < 4; j++) {
if (out_n_base + j < N) {
uint32_t fp = __float_as_uint(acc1[j]); fp += 0x7FFFu + ((fp >> 16) & 1u);
C[out_m1 * scm + out_n_base + j] = (uint16_t)(fp >> 16u);
}
}
}
}
// ── ASM GEMM launcher ──
static std::unordered_map<std::string, std::pair<hipModule_t, hipFunction_t>> _asm_cache;
struct __attribute__((packed)) AsmArgs {
void* D; char _0[8]; void* C; char _1[8];
void* A; char _2[8]; void* B; char _3[8];
float alpha; char _4[12]; float beta; char _5[12];
unsigned sD0; char _6[12]; unsigned sD1; char _7[12];
unsigned sC0; char _8[12]; unsigned sC1; char _9[12];
unsigned sA0; char _10[12]; unsigned sA1; char _11[12];
unsigned sB0; char _12[12]; unsigned sB1; char _13[12];
unsigned M; char _14[12]; unsigned N; char _15[12];
unsigned K; char _16[12];
void* SA; char _17[8]; void* SB; char _18[8];
unsigned sSA0; char _19[12]; unsigned sSA1; char _20[12];
unsigned sSB0; char _21[12]; unsigned sSB1; char _22[12];
int log2ks;
};
hipFunction_t get_asm_fn(const std::string& co_path, const std::string& knl_name) {
auto key = co_path + "::" + knl_name;
auto it = _asm_cache.find(key);
if (it != _asm_cache.end()) return it->second.second;
hipModule_t mod; hipFunction_t fn;
hipModuleLoad(&mod, co_path.c_str());
hipModuleGetFunction(&fn, mod, knl_name.c_str());
_asm_cache[key] = {mod, fn};
return fn;
}
__global__ void __launch_bounds__(64, 16)
quant_a_kern(const uint16_t* __restrict__ A, uint8_t* __restrict__ Aq,
uint8_t* __restrict__ Asc, int M, int K, int strA, int sm) {
#if defined(__gfx950__)
int idx = blockIdx.x * blockDim.x + threadIdx.x;
int ns = K / 32, total = sm * ns;
if (idx >= total) return;
int row = idx / ns, blk = idx % ns, ko = blk * 32;
uint16_t v[32];
if (row < M) {
const uint4* s = reinterpret_cast<const uint4*>(&A[row * strA + ko]);
uint4* d = reinterpret_cast<uint4*>(v); d[0]=s[0];d[1]=s[1];d[2]=s[2];d[3]=s[3];
} else { for(int j=0;j<32;j++) v[j]=0; }
uint16_t mx=0;
for(int j=0;j<32;j++) mx=max(mx,(uint16_t)(v[j]&0x7FFF));
uint32_t au=(((uint32_t)mx<<16)+0x200000u)&0xFF800000u;
int ef=(au>>23u)&0xFFu;
int su=(au==0u)?-127:max(-127,min(127,ef-127-2));
float hs=(su>=-126)?__uint_as_float((uint32_t)(su+127)<<23):0.0f;
uint8_t sv=(uint8_t)(su+127);
if(row<M){
unsigned int p[4];
for(int i=0;i<4;i++){p[i]=0;
p[i]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p[i],*reinterpret_cast<bf16v2_t*>(&v[i*8+0]),hs,0);
p[i]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p[i],*reinterpret_cast<bf16v2_t*>(&v[i*8+2]),hs,1);
p[i]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p[i],*reinterpret_cast<bf16v2_t*>(&v[i*8+4]),hs,2);
p[i]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p[i],*reinterpret_cast<bf16v2_t*>(&v[i*8+6]),hs,3);
}
*reinterpret_cast<uint4*>(&Aq[row*(K/2)+blk*16])=*reinterpret_cast<uint4*>(p);
}
int sn=K/32;
int i0=row/32,i1=(row%32)/16,i2=row%16;
int i3=blk/8,i4=(blk%8)/4,i5=blk%4;
Asc[i0*(32*sn)+i3*256+i5*64+i2*4+i4*2+i1]=sv;
#endif
}
} // namespace
void quantize_a_hip(torch::Tensor a, torch::Tensor a_q, torch::Tensor a_scale) {
int m=(int)a.size(0), k=(int)a.size(1);
int nkg=k/QUANT_GROUP;
dim3 g((nkg+GROUPS_PER_BLOCK-1)/GROUPS_PER_BLOCK, m);
quantize_a_kernel<<<g, BLOCK_THREADS>>>(
reinterpret_cast<const __hip_bfloat16*>(a.data_ptr<at::BFloat16>()),
a_q.data_ptr<uint8_t>(), a_scale.data_ptr<uint8_t>(), m, k);
}
torch::Tensor gemm_fused_16x16_hip(torch::Tensor A, torch::Tensor Bq,
torch::Tensor Bsc, int64_t N, torch::Tensor out) {
int M=(int)A.size(0),K=(int)A.size(1);
gemm_fused_16x16_kern<<<dim3((M+15)/16,((int)N+15)/16),256>>>(
reinterpret_cast<const uint16_t*>(A.data_ptr()),
reinterpret_cast<const uint8_t*>(Bq.data_ptr()),
reinterpret_cast<const uint8_t*>(Bsc.data_ptr()),
reinterpret_cast<uint16_t*>(out.data_ptr()),
M,(int)N,K,(int)A.stride(0),(int)(Bq.stride(0)*Bq.element_size()),(int)Bsc.size(1));
return out;
}
void quant_a_shuffled(torch::Tensor A, torch::Tensor Aq, torch::Tensor Asc, int64_t sm) {
int M=(int)A.size(0),K=(int)A.size(1);
int ns=K/32, total=(int)sm*ns;
quant_a_kern<<<(total+63)/64, 64>>>(
reinterpret_cast<const uint16_t*>(A.data_ptr<at::BFloat16>()),
Aq.data_ptr<uint8_t>(), Asc.data_ptr<uint8_t>(), M, K, (int)A.stride(0), (int)sm);
}
torch::Tensor quant_and_asm(torch::Tensor A, torch::Tensor Bsh,
torch::Tensor Bsc, int64_t N,
torch::Tensor Aq, torch::Tensor Asc, torch::Tensor out,
int64_t sm, const std::string& knl_name, const std::string& co_path, int64_t splitK) {
int M=(int)A.size(0),K=(int)A.size(1);
int ns=K/32, total=(int)sm*ns;
quant_a_kern<<<(total+63)/64, 64>>>(
reinterpret_cast<const uint16_t*>(A.data_ptr<at::BFloat16>()),
Aq.data_ptr<uint8_t>(), Asc.data_ptr<uint8_t>(), M, K, (int)A.stride(0), (int)sm);
auto fn = get_asm_fn(co_path, knl_name);
AsmArgs a={};
a.D=out.data_ptr();a.C=out.data_ptr();
a.A=Aq.data_ptr();a.B=Bsh.data_ptr();
a.alpha=1.0f;a.beta=0.0f;
a.sC0=(unsigned)out.stride(0);a.sC1=1;
a.sA0=(unsigned)(Aq.stride(0)*2);a.sA1=1;
a.sB0=(unsigned)(Bsh.stride(0)*2);a.sB1=1;
a.M=(unsigned)M;a.N=(unsigned)N;a.K=(unsigned)K;
a.SA=Asc.data_ptr();a.SB=Bsc.data_ptr();
a.sSA0=(unsigned)Asc.stride(0);a.sSA1=1;
a.sSB0=(unsigned)Bsc.stride(0);a.sSB1=1;
a.log2ks=(int)splitK;
size_t asz=sizeof(a);
void*cfg[]={HIP_LAUNCH_PARAM_BUFFER_POINTER,&a,HIP_LAUNCH_PARAM_BUFFER_SIZE,&asz,HIP_LAUNCH_PARAM_END};
// Parse tile_m x tile_n from kernel name (e.g. "..._64x128E" -> tm=64,tn=128)
unsigned tm=32, tn=128;
auto pos=knl_name.rfind('x');
if(pos!=std::string::npos) {
tn=std::stoi(knl_name.substr(pos+1));
// Find tile_m: scan backwards from 'x' to find the start of the number
auto pos2=pos-1;
while(pos2>0 && knl_name[pos2-1]>='0' && knl_name[pos2-1]<='9') pos2--;
tm=std::stoi(knl_name.substr(pos2, pos-pos2));
}
unsigned mp=((unsigned)M+tm-1)/tm*tm;
hipModuleLaunchKernel(fn,(N+tn-1)/tn,mp/tm,1,256,1,1,0,0,nullptr,(void**)cfg);
return out;
}
void warmup_asm(const std::string& co_path, const std::string& knl_name) {
get_asm_fn(co_path, knl_name);
}
// Raw-pointer fast path: avoids pybind11 tensor marshalling
void fast_hip16(int64_t A, int64_t Bq, int64_t Bsc, int64_t C,
int64_t M, int64_t N, int64_t K, int64_t strA, int64_t strBq, int64_t BscSN) {
gemm_fused_16x16_kern<<<dim3(((int)M+15)/16,((int)N+15)/16),256>>>(
(const uint16_t*)A, (const uint8_t*)Bq, (const uint8_t*)Bsc, (uint16_t*)C,
(int)M, (int)N, (int)K, (int)strA, (int)strBq, (int)BscSN);
}
// ext_splitk: on-the-fly quant + atomicAdd split-K (for M<=16 K>1024)
__device__ __forceinline__ void quant_32_inline(const uint16_t* ap, int out[4], int32_t& spk) {
uint16_t mx=0;
#pragma unroll
for(int j=0;j<32;j++) mx=max(mx,(uint16_t)(ap[j]&0x7FFF));
uint32_t au=(((uint32_t)mx<<16)+0x200000u)&0xFF800000u;
int ef=(au>>23u)&0xFFu;
int su=(au==0u)?-127:max(-127,min(127,ef-127-2));
float hs=(su>=-126)?__uint_as_float((uint32_t)(su+127)<<23):0.0f;
#pragma unroll
for(int j=0;j<4;j++){out[j]=0;
out[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(out[j],*reinterpret_cast<const bf16v2_t*>(&ap[j*8+0]),hs,0);
out[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(out[j],*reinterpret_cast<const bf16v2_t*>(&ap[j*8+2]),hs,1);
out[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(out[j],*reinterpret_cast<const bf16v2_t*>(&ap[j*8+4]),hs,2);
out[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(out[j],*reinterpret_cast<const bf16v2_t*>(&ap[j*8+6]),hs,3);
}
spk=(int32_t)(uint8_t)(su+127);
}
__global__ void __launch_bounds__(64, 8)
ext_splitk_16x16_kern(const uint16_t* __restrict__ A, const uint8_t* __restrict__ Bq,
const uint8_t* __restrict__ Bsc, float* __restrict__ Cfp,
int M, int N, int K, int strA, int strBq, int BscSN, int kper) {
#if defined(__gfx950__)
const int mt=blockIdx.x,nt=blockIdx.y,ks=blockIdx.z;
const int lid=threadIdx.x,t_row=lid&15,kpart=lid>>4;
const int SNG=BscSN>>3,m_row=mt*16+t_row,b_ng=nt*16+t_row;
int kst=(ks*kper/128)*128,ken=min(((ks+1)*kper+127)/128*128,K);
if(kst>=ken) return;
f32x4 c_acc={0,0,0,0};
for(int kb=kst;kb<ken;kb+=128){
int k_off=kb+kpart*32;
uint16_t al[32]; int a_i[4]; int32_t a_spk;
if(m_row<M&&k_off+31<K){const uint4*s=reinterpret_cast<const uint4*>(&A[m_row*strA+k_off]);
uint4*d=reinterpret_cast<uint4*>(al);d[0]=s[0];d[1]=s[1];d[2]=s[2];d[3]=s[3];}
else{for(int j=0;j<32;j++)al[j]=(m_row<M&&k_off+j<K)?A[m_row*strA+k_off+j]:0;}
quant_32_inline(al,a_i,a_spk);
int bk=kb/2+kpart*16; int b_i[4]; int32_t b_spk;
if(b_ng<N&&bk+15<K/2){uint4 bd=*reinterpret_cast<const uint4*>(&Bq[b_ng*strBq+bk]);
b_i[0]=((int*)&bd)[0];b_i[1]=((int*)&bd)[1];b_i[2]=((int*)&bd)[2];b_i[3]=((int*)&bd)[3];}
else{b_i[0]=0;b_i[1]=0;b_i[2]=0;b_i[3]=0;}
b_spk=load_b_scale_16(Bsc,b_ng,kb/32+kpart,SNG,N);
i32x8 av={a_i[0],a_i[1],a_i[2],a_i[3],0,0,0,0};
i32x8 bv={b_i[0],b_i[1],b_i[2],b_i[3],0,0,0,0};
c_acc=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(av,bv,c_acc,4,4,0,a_spk,0,b_spk);
}
for(int j=0;j<4;j++){int mo=mt*16+kpart*4+j,no=nt*16+t_row;
if(mo<M&&no<N) atomicAdd(&Cfp[mo*N+no],c_acc[j]);}
#endif
}
__global__ void __launch_bounds__(1024)
fp32_to_bf16_zero(float*src,uint16_t*dst,int n){
int i=blockIdx.x*1024+threadIdx.x;
if(i<n){float v=src[i];src[i]=0.0f;
uint32_t fp=__float_as_uint(v);fp+=0x7FFFu+((fp>>16)&1u);dst[i]=(uint16_t)(fp>>16u);}
}
// Raw-pointer ext_splitk launch (2 kernels, ~3µs total vs Triton's ~9.6µs)
void fast_ext_splitk(int64_t A, int64_t Bq, int64_t Bsc, int64_t Cfp, int64_t Cbf,
int64_t M, int64_t N, int64_t K, int64_t strA, int64_t strBq, int64_t BscSN, int64_t ks) {
int kper=((K/128+(int)ks-1)/(int)ks)*128;
int aks=(K+kper-1)/kper;
ext_splitk_16x16_kern<<<dim3(((int)M+15)/16,((int)N+15)/16,aks),64>>>(
(const uint16_t*)A,(const uint8_t*)Bq,(const uint8_t*)Bsc,(float*)Cfp,
(int)M,(int)N,(int)K,(int)strA,(int)strBq,(int)BscSN,kper);
int t=(int)M*(int)N;
fp32_to_bf16_zero<<<(t+1023)/1024,1024>>>((float*)Cfp,(uint16_t*)Cbf,t);
}
// C++ dispatch for M<=32 K<=1024 (single kernel, minimal overhead)
void dispatch_hip(int64_t A, int64_t Bq, int64_t Bsc, int64_t C, int64_t Cfp,
int64_t M, int64_t N, int64_t K, int64_t strA, int64_t strBq, int64_t BscSN, int64_t ks) {
gemm_fused_16x16_kern<<<dim3(((int)M+15)/16,((int)N+15)/16),256>>>(
(const uint16_t*)A,(const uint8_t*)Bq,(const uint8_t*)Bsc,(uint16_t*)C,
(int)M,(int)N,(int)K,(int)strA,(int)strBq,(int)BscSN);
}
// v78: N-parallel with 32-row M-tiles matching Triton's {1,4} warp layout
void launch_fused_npar(int64_t A, int64_t Bsh, int64_t Bsc, int64_t C,
int64_t M, int64_t N, int64_t K, int64_t strA, int64_t strBsh, int64_t BscSN) {
gemm_fused_npar_kern<<<dim3(((int)M+31)/32,((int)N+63)/64),256>>>(
(const uint16_t*)A,(const uint8_t*)Bsh,(const uint8_t*)Bsc,(uint16_t*)C,
(int)M,(int)N,(int)K,(int)strA,(int)strBsh,(int)BscSN);
}
// v80: Non-fused launcher reproducing Triton's ISA structure
void launch_nofuse_triton(int64_t Aq, int64_t Bt, int64_t C,
int64_t Asc, int64_t Bsc,
int64_t M, int64_t N, int64_t K, int64_t Ks,
int64_t sam, int64_t sbn, int64_t scm, int64_t sasm) {
int npn = ((int)N + 63) / 64;
int npm = ((int)M + 31) / 32;
int grid = npm * npn;
gemm_nofuse_triton<<<grid, 256>>>(
(const uint8_t*)Aq, (const uint8_t*)Bt, (uint16_t*)C,
(const uint8_t*)Asc, (const uint8_t*)Bsc,
(int)M, (int)N, (int)K, (int)Ks, (int)sam, (int)sbn, (int)scm, (int)sasm);
}
// ── HIP reduce kernel: replaces Triton _reduce_kernel (saves 11µs launch) ──
// For split-K=8: sum 8 fp32 partials → bf16 output
__global__ void __launch_bounds__(256)
hip_reduce_splitk(float* __restrict__ src, uint16_t* __restrict__ dst,
int M, int N, int ksplit) {
int idx = blockIdx.x * blockDim.x + threadIdx.x;
int total = M * N;
if (idx >= total) return;
float sum = 0.0f;
for (int k = 0; k < ksplit; k++)
sum += src[k * total + idx];
// Convert to bf16 and ZERO source for next atomic_add call
src[idx] = 0.0f;
uint32_t fp = __float_as_uint(sum);
fp += 0x7FFFu + ((fp >> 16) & 1u);
dst[idx] = (uint16_t)(fp >> 16u);
}
void fast_reduce_splitk(int64_t src, int64_t dst, int64_t M, int64_t N, int64_t ksplit) {
int total = (int)M * (int)N;
hip_reduce_splitk<<<(total+255)/256, 256>>>((float*)src, (uint16_t*)dst,
(int)M, (int)N, (int)ksplit);
}
// ── Triton HSACO loader: bypass 11µs Python dispatch → 2.5µs HIP launch ──
static std::unordered_map<std::string, std::pair<hipModule_t, hipFunction_t>> _triton_cache;
void load_triton_hsaco(const std::string& path, const std::string& name) {
auto key = path + "::" + name;
if (_triton_cache.count(key)) return;
hipModule_t mod; hipFunction_t fn;
if (hipModuleLoad(&mod, path.c_str()) != hipSuccess) return;
if (hipModuleGetFunction(&fn, mod, name.c_str()) != hipSuccess) return;
_triton_cache[key] = {mod, fn};
}
// Generic Triton kernel launcher: 4 ptrs + up to 12 int args
void launch_triton_kernel(const std::string& key,
int64_t p0, int64_t p1, int64_t p2, int64_t p3,
int64_t a0, int64_t a1, int64_t a2, int64_t a3, int64_t a4,
int64_t a5, int64_t a6, int64_t a7, int64_t a8, int64_t a9, int64_t a10, int64_t a11,
int64_t gridX, int64_t gridY, int64_t blockX, int64_t shared_mem) {
auto it = _triton_cache.find(key);
if (it == _triton_cache.end()) return;
void* vp0=(void*)p0; void* vp1=(void*)p1; void* vp2=(void*)p2; void* vp3=(void*)p3;
int32_t i0=(int32_t)a0,i1=(int32_t)a1,i2=(int32_t)a2,i3=(int32_t)a3,i4=(int32_t)a4;
int64_t i5=a5,i6=a6,i7=a7,i8=a8,i9=a9,i10=a10,i11=a11;
void* args[] = {&vp0,&vp1,&vp2,&vp3,&i0,&i1,&i2,&i3,&i4,&i5,&i6,&i7,&i8,&i9,&i10,&i11};
hipModuleLaunchKernel(it->second.second,
(unsigned)gridX,(unsigned)gridY,1,(unsigned)blockX,1,1,
(unsigned)shared_mem,0,args,nullptr);
}
// Launch Triton _gemm_fp4_kernel via cached HSACO with exact 96-byte kernarg layout
// Derived from ISA metadata dump: kernarg_segment_size=96, 5 ptrs + 9 scalars + padding + 2 unused ptrs
void launch_triton_fp4gemm(const std::string& key,
int64_t Aq, int64_t Bt, int64_t C, int64_t Asc, int64_t Bsc,
int64_t M, int64_t N, int64_t K, int64_t Ks,
int64_t sam, int64_t sbn, int64_t sck, int64_t scm, int64_t sasm,
int64_t gridX, int64_t shared_mem) {
auto it = _triton_cache.find(key);
if (it == _triton_cache.end()) return;
// Pack 96-byte kernarg buffer matching Triton's compiled layout
struct __attribute__((packed)) {
void* a_ptr; // 0
void* b_ptr; // 8
void* c_ptr; // 16
void* asc_ptr; // 24
void* bsc_ptr; // 32
int32_t M; // 40
int32_t N; // 44
int32_t K; // 48
int32_t Ks; // 52
int32_t sam; // 56
int32_t sbn; // 60
int32_t sck; // 64
int32_t scm; // 68
int32_t sasm; // 72
int32_t _pad; // 76
void* _ptr0; // 80 (unused by kernel, but kernarg_size=96)
void* _ptr1; // 88
} args;
args.a_ptr = (void*)Aq;
args.b_ptr = (void*)Bt;
args.c_ptr = (void*)C;
args.asc_ptr = (void*)Asc;
args.bsc_ptr = (void*)Bsc;
args.M = (int32_t)M;
args.N = (int32_t)N;
args.K = (int32_t)K;
args.Ks = (int32_t)Ks;
args.sam = (int32_t)sam;
args.sbn = (int32_t)sbn;
args.sck = (int32_t)sck;
args.scm = (int32_t)scm;
args.sasm = (int32_t)sasm;
args._pad = 0;
args._ptr0 = nullptr;
args._ptr1 = nullptr;
size_t sz = sizeof(args);
void* cfg[] = {HIP_LAUNCH_PARAM_BUFFER_POINTER, &args,
HIP_LAUNCH_PARAM_BUFFER_SIZE, &sz,
HIP_LAUNCH_PARAM_END};
hipModuleLaunchKernel(it->second.second,
(unsigned)gridX, 1, 1, 256, 1, 1,
(unsigned)shared_mem, 0, nullptr, (void**)cfg);
}
"""
@functools.lru_cache(maxsize=1)
def _hip_module():
d = hashlib.sha1(HIP_SRC.encode()).hexdigest()[:16]
return load_inline(name=f"hip30_{d}", cpp_sources=[CPP_SRC], cuda_sources=[HIP_SRC],
functions=["quantize_a_hip","gemm_fused_16x16_hip","fast_hip16","fast_ext_splitk",
"dispatch_hip","fast_reduce_splitk","load_triton_hsaco","launch_triton_kernel",
"launch_triton_fp4gemm","launch_nofuse_triton",
"quant_a_shuffled","quant_and_asm","warmup_asm","launch_fused_npar"],
extra_cflags=["-O3"],
extra_cuda_cflags=["-O3","--offload-arch=gfx950","-std=c++17"],
with_cuda=True, verbose=False)
# ═══════════════════════════════════════════════════════════════════
# Triton kernels: fused quant+GEMM (M≤64) + standard GEMM (fallback)
# ═══════════════════════════════════════════════════════════════════
@triton.jit
def _mxfp4_quant_inline(x, BLOCK_SIZE_K: tl.constexpr, BLOCK_SIZE_M: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr):
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_K // MXFP4_QUANT_BLOCK_SIZE
x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
se = tl.log2(amax).floor() - 2
se = tl.clamp(se, min=-127, max=127)
bs = se.to(tl.uint8) + 127
qs = tl.exp2(-se)
qx = x * qs
qx = qx.to(tl.uint32, bitcast=True)
s = qx & 0x80000000
qx = qx ^ s
qf = qx.to(tl.float32, bitcast=True)
sat = qf >= 6.0
den = (not sat) & (qf < 1.0)
nor = not (sat | den)
dexp: tl.constexpr = (127-1)+(23-1)+1
dmi: tl.constexpr = dexp << 23
dmf: tl.constexpr = tl.cast(dmi, tl.float32, bitcast=True)
dx = qf + dmf
dx = dx.to(tl.uint32, bitcast=True)
dx -= dmi
dx = dx.to(tl.uint8)
nx = qx
mo = (nx >> 22) & 1
nx += 0xC11FFFFF # ((1-127)<<23) + (1<<21) - 1 as uint32
nx += mo
nx = nx >> 22
nx = nx.to(tl.uint8)
v = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
v = tl.where(nor, nx, v)
v = tl.where(den, dx, v)
sl = s >> (23+8-1-2)
sl = sl.to(tl.uint8)
v = v | sl
v = tl.reshape(v, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2])
evens, odds = tl.split(v)
xfp4 = evens | (odds << 4)
return xfp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // 2), bs.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
@triton.jit
def remap_xcd(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):
ppx = (GRID_MN + NUM_XCDS - 1) // NUM_XCDS
tall = GRID_MN % NUM_XCDS
tall = NUM_XCDS if tall == 0 else tall
xcd = pid % NUM_XCDS
lp = pid // NUM_XCDS
if xcd < tall: pid = xcd * ppx + lp
else: pid = tall * ppx + (xcd - tall) * (ppx - 1) + lp
return pid
@triton.jit
def pid_grid(pid, npm, npn, GROUP_SIZE_M: tl.constexpr = 1):
if GROUP_SIZE_M == 1:
return pid // npn, pid % npn
else:
nig = GROUP_SIZE_M * npn
gid = pid // nig
fm = gid * GROUP_SIZE_M
gsm = min(npm - fm, GROUP_SIZE_M)
tl.assume(gsm >= 0)
return fm + (pid % gsm), (pid % nig) // gsm
@triton.jit
def _bsc_off(row, col, Ks):
return ((row//32)*(Ks*32) + (col//8)*256 + (col%4)*64
+ (row%16)*4 + ((col//4)%2)*2 + ((row//16)%2))
@triton.heuristics({"EVEN_K": lambda a: (a["K"]%(a["BLOCK_SIZE_K"]//2)==0)
and (a["SPLITK_BLOCK_SIZE"]%a["BLOCK_SIZE_K"]==0)
and (a["K"]%(a["SPLITK_BLOCK_SIZE"]//2)==0)})
@triton.jit
def _fused_gemm_fp4_kernel(
a_ptr, b_ptr, c_ptr, bsc_ptr,
M, N, K, actual_K, Ks,
sam, sak, sbk, sbn, sck, scm, scn,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr,
SPLITK_BLOCK_SIZE: tl.constexpr, EVEN_K: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr,
waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr):
tl.assume(sam>0);tl.assume(sak>0);tl.assume(sbk>0);tl.assume(sbn>0);tl.assume(scm>0);tl.assume(scn>0)
GMN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
pu = tl.program_id(0)
pu = remap_xcd(pu, GMN * NUM_KSPLIT, NUM_XCDS=8)
pk = pu % NUM_KSPLIT
pid = pu // NUM_KSPLIT
npm = tl.cdiv(M, BLOCK_SIZE_M); npn = tl.cdiv(N, BLOCK_SIZE_N)
if NUM_KSPLIT == 1: pm, pn = pid_grid(pid, npm, npn, GROUP_SIZE_M=GROUP_SIZE_M)
else: pm = pid // npn; pn = pid % npn
tl.assume(pm>=0);tl.assume(pn>=0)
SPK: tl.constexpr = BLOCK_SIZE_K // 32
if (pk * SPLITK_BLOCK_SIZE // 2) < K:
nki = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
okp = tl.arange(0, BLOCK_SIZE_K // 2)
okps = pk * (SPLITK_BLOCK_SIZE // 2) + okp
oka = tl.arange(0, BLOCK_SIZE_K)
okas = pk * SPLITK_BLOCK_SIZE + oka
oam = (pm * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
obn = (pn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
ap = a_ptr + (oam[:, None] * sam + okas[None, :] * sak)
bp = b_ptr + (okps[:, None] * sbk + obn[None, :] * sbn)
ksb = pk * (SPLITK_BLOCK_SIZE // 32)
oksl = tl.arange(0, SPK)
acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for ki in range(pk * nki, (pk + 1) * nki):
if EVEN_K: abf = tl.load(ap)
else: abf = tl.load(ap, mask=oka[None, :] < actual_K - ki * BLOCK_SIZE_K, other=0.0)
af32 = abf.to(tl.float32)
afp4, asc = _mxfp4_quant_inline(af32, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
coks = ksb + oksl
bso = _bsc_off(obn[:, None], coks[None, :], Ks)
bscv = tl.load(bsc_ptr + bso, cache_modifier=".cg")
if EVEN_K: b = tl.load(bp, cache_modifier=".cg")
else: b = tl.load(bp, mask=okp[:, None] < K - ki * (BLOCK_SIZE_K // 2), other=0, cache_modifier=".cg")
acc = tl.dot_scaled(afp4, asc, "e2m1", b, bscv, "e2m1", acc)
ap += BLOCK_SIZE_K * sak
bp += (BLOCK_SIZE_K // 2) * sbk
ksb += SPK
ocm = pm * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
ocn = pn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
c_mask = (ocm[:, None] < M) & (ocn[None, :] < N)
c = acc.to(c_ptr.type.element_ty)
cp = c_ptr + scm * ocm[:, None] + scn * ocn[None, :] + pk * sck
tl.store(cp, c, mask=c_mask)
@triton.heuristics({"EVEN_K": lambda a: (a["K"]%(a["BLOCK_SIZE_K"]//2)==0)
and (a["SPLITK_BLOCK_SIZE"]%a["BLOCK_SIZE_K"]==0)
and (a["K"]%(a["SPLITK_BLOCK_SIZE"]//2)==0)})
@triton.jit
def _gemm_fp4_kernel(
a_ptr, b_ptr, c_ptr, asc_ptr, bsc_ptr,
M, N, K, Ks,
sam, sak, sbk, sbn, sck, scm, scn, sasm, sask,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr,
SPLITK_BLOCK_SIZE: tl.constexpr, EVEN_K: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr,
waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr):
tl.assume(sam>0);tl.assume(sak>0);tl.assume(sbk>0);tl.assume(sbn>0);tl.assume(scm>0);tl.assume(scn>0)
tl.assume(sasm>0);tl.assume(sask>0)
GMN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
pu = tl.program_id(0)
pu = remap_xcd(pu, GMN * NUM_KSPLIT, NUM_XCDS=8)
pk = pu % NUM_KSPLIT
pid = pu // NUM_KSPLIT
npm = tl.cdiv(M, BLOCK_SIZE_M); npn = tl.cdiv(N, BLOCK_SIZE_N)
if NUM_KSPLIT == 1: pm, pn = pid_grid(pid, npm, npn, GROUP_SIZE_M=GROUP_SIZE_M)
else: pm = pid // npn; pn = pid % npn
tl.assume(pm>=0);tl.assume(pn>=0)
SPK: tl.constexpr = BLOCK_SIZE_K // 32
if (pk * SPLITK_BLOCK_SIZE // 2) < K:
nki = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
ok = tl.arange(0, BLOCK_SIZE_K // 2)
oks = pk * (SPLITK_BLOCK_SIZE // 2) + ok
oam = (pm * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
obn = (pn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
ap = a_ptr + (oam[:, None] * sam + oks[None, :] * sak)
bp = b_ptr + (oks[:, None] * sbk + obn[None, :] * sbn)
oaks = (pk * (SPLITK_BLOCK_SIZE // 32)) + tl.arange(0, SPK)
ascp = asc_ptr + oam[:, None] * sasm + oaks[None, :] * sask
ksb = pk * (SPLITK_BLOCK_SIZE // 32)
oksl = tl.arange(0, SPK)
acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k in range(pk * nki, (pk + 1) * nki):
ascv = tl.load(ascp)
coks = ksb + oksl
bso = _bsc_off(obn[:, None], coks[None, :], Ks)
bscv = tl.load(bsc_ptr + bso, cache_modifier=".cg")
if EVEN_K: a = tl.load(ap); b = tl.load(bp, cache_modifier=".cg")
else:
a = tl.load(ap, mask=ok[None, :] < K - k * (BLOCK_SIZE_K // 2), other=0)
b = tl.load(bp, mask=ok[:, None] < K - k * (BLOCK_SIZE_K // 2), other=0, cache_modifier=".cg")
acc = tl.dot_scaled(a, ascv, "e2m1", b, bscv, "e2m1", acc)
ap += (BLOCK_SIZE_K // 2) * sak
bp += (BLOCK_SIZE_K // 2) * sbk
ascp += SPK * sask
ksb += SPK
c = acc.to(c_ptr.type.element_ty)
ocm = pm * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
ocn = pn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
cp = c_ptr + scm * ocm[:, None] + scn * ocn[None, :] + pk * sck
tl.store(cp, c, mask=(ocm[:, None] < M) & (ocn[None, :] < N))
@triton.jit
def _reduce_kernel(ci, co, M, N, sick, sicm, sicn, socm, socn,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
ACTUAL_KSPLIT: tl.constexpr, MAX_KSPLIT: tl.constexpr):
pm = tl.program_id(0); pn = tl.program_id(1)
om = (pm*BLOCK_SIZE_M+tl.arange(0,BLOCK_SIZE_M))%M
on = (pn*BLOCK_SIZE_N+tl.arange(0,BLOCK_SIZE_N))%N
ok = tl.arange(0,MAX_KSPLIT)
ip = ci + ok[:,None,None]*sick + om[None,:,None]*sicm + on[None,None,:]*sicn
if ACTUAL_KSPLIT==MAX_KSPLIT: c = tl.load(ip)
else: c = tl.load(ip, mask=ok[:,None,None]<ACTUAL_KSPLIT)
c = tl.sum(c, axis=0).to(co.type.element_ty)
op = co + om[:,None]*socm + on[None,:]*socn
tl.store(op, c)
# ═══════════════════════════════════════════════════════════════════
# Per-shape configs
# ═══════════════════════════════════════════════════════════════════
SHAPE_CONFIGS = {
(4, 2880, 512): dict(BLOCK_SIZE_M=16,BLOCK_SIZE_N=64,BLOCK_SIZE_K=512,GROUP_SIZE_M=1,
num_warps=4,num_stages=1,waves_per_eu=4,matrix_instr_nonkdim=16,NUM_KSPLIT=1),
# v125: Shape 2 num_warps=8 — matches CDNA4 blog 8-wave scheduling
(16, 2112, 7168): dict(BLOCK_SIZE_M=16,BLOCK_SIZE_N=128,BLOCK_SIZE_K=256,GROUP_SIZE_M=1,
num_warps=8,num_stages=2,waves_per_eu=1,matrix_instr_nonkdim=16,NUM_KSPLIT=16),
(32, 4096, 512): dict(BLOCK_SIZE_M=32,BLOCK_SIZE_N=64,BLOCK_SIZE_K=512,GROUP_SIZE_M=1,
num_warps=4,num_stages=1,waves_per_eu=4,matrix_instr_nonkdim=16,NUM_KSPLIT=1),
(32, 2880, 512): dict(BLOCK_SIZE_M=32,BLOCK_SIZE_N=64,BLOCK_SIZE_K=512,GROUP_SIZE_M=1,
num_warps=4,num_stages=1,waves_per_eu=4,matrix_instr_nonkdim=16,NUM_KSPLIT=1),
# Revert shape 5 to baseline (8 warps failed correctness)
(64, 7168, 2048): dict(BLOCK_SIZE_M=32,BLOCK_SIZE_N=64,BLOCK_SIZE_K=512,GROUP_SIZE_M=4,
num_warps=4,num_stages=5,waves_per_eu=1,matrix_instr_nonkdim=16,NUM_KSPLIT=1),
}
CONFIGS = {
4: dict(BLOCK_SIZE_M=16,BLOCK_SIZE_N=128,BLOCK_SIZE_K=256,GROUP_SIZE_M=1,
num_warps=4,num_stages=2,waves_per_eu=3,matrix_instr_nonkdim=16,NUM_KSPLIT=16),
16: dict(BLOCK_SIZE_M=16,BLOCK_SIZE_N=128,BLOCK_SIZE_K=256,GROUP_SIZE_M=1,
num_warps=4,num_stages=2,waves_per_eu=3,matrix_instr_nonkdim=16,NUM_KSPLIT=16),
32: dict(BLOCK_SIZE_M=32,BLOCK_SIZE_N=128,BLOCK_SIZE_K=256,GROUP_SIZE_M=1,
num_warps=4,num_stages=2,waves_per_eu=3,matrix_instr_nonkdim=16,NUM_KSPLIT=1),
64: dict(BLOCK_SIZE_M=64,BLOCK_SIZE_N=256,BLOCK_SIZE_K=256,GROUP_SIZE_M=1,
num_warps=4,num_stages=3,waves_per_eu=2,matrix_instr_nonkdim=32,NUM_KSPLIT=1),
256: dict(BLOCK_SIZE_M=128,BLOCK_SIZE_N=256,BLOCK_SIZE_K=256,GROUP_SIZE_M=2,
num_warps=4,num_stages=3,waves_per_eu=2,matrix_instr_nonkdim=32,NUM_KSPLIT=1),
}
def get_config(M, N=None, K=None):
if N is not None and K is not None:
key = (M, N, K)
if key in SHAPE_CONFIGS: return SHAPE_CONFIGS[key].copy()
for t in sorted(CONFIGS.keys()):
if M <= t: return CONFIGS[t].copy()
return CONFIGS[256].copy()
def get_splitk(K, BSK, NSK):
SPBS = triton.cdiv(2 * triton.cdiv(K, NSK), BSK) * BSK
while NSK > 1 and BSK > 16:
if K % (SPBS // 2) == 0 and SPBS % BSK == 0 and K % (BSK // 2) == 0: break
elif K % (SPBS // 2) != 0 and NSK > 1: NSK //= 2
elif SPBS % BSK != 0:
if NSK > 1: NSK //= 2
elif BSK > 16: BSK //= 2
elif K % (BSK // 2) != 0 and BSK > 16: BSK //= 2
else: break
SPBS = triton.cdiv(2 * triton.cdiv(K, NSK), BSK) * BSK
NSK = triton.cdiv(K, SPBS // 2)
return SPBS, BSK, NSK
# ═══════════════════════════════════════════════════════════════════
# Hybrid dispatch
# ═══════════════════════════════════════════════════════════════════
ASM_BASE = "/home/runner/aiter/hsa/gfx950/f4gemm/"
ASM_CONFIGS_SHAPE = {
# Per-shape best ASM tile (from discovery: 26 tiles available)
# 64x128 for M=64: perfect M-tile fit, no padding. 128x128 for M=256 was worse (15.1 vs 12.5)
(64, 7168, 2048): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",
"f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128.co", 0),
}
ASM_DEFAULT = ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
"f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co", 0)
FUSED_M_THRESHOLD = 32 # Fused quant M≤32 ONLY. M=64 fused causes 20µs regression even with BLOCK_M=16.
# v31: Removed K threshold. HIP gemm_fused_16x16_kern handles ANY K (single kernel, LDS reduction).
# Key insight: For M=16 K=7168, Triton uses 2 launches (8µs overhead) vs HIP 1 launch (2µs overhead).
# GPU self-time is only 2.2µs — overhead dominates. Single HIP kernel: ~6µs vs Triton's ~12µs.
_OUT_CACHE = {}
_AQ_CACHE = {}
_HIP16_OUT = {}
_ASM_WS = {}
_warmed = False
_hip = None
_triton_hsaco_keys = {} # {shape: (key, grid, block, shared_mem, num_args)}
def _warm_triton_hsaco(hip_mod):
"""Trigger Triton JIT for shape 5, find cached .hsaco, load via hipModule."""
import glob as _glob
# Shape 5: M=64, N=7168, K=2048 — trigger compilation with dummy data
dev = torch.device('cuda')
m, n, k = 64, 7168, 2048
Kp = k >> 1
Aq = torch.zeros(m, k//2, dtype=torch.uint8, device=dev)
Asc = torch.zeros(m, k//32, dtype=torch.uint8, device=dev)
Bq8 = torch.zeros(n, k//2, dtype=torch.uint8, device=dev)
Bt = Bq8.T
Bsc8 = torch.zeros(n, k//32, dtype=torch.uint8, device=dev) # approximate
out = torch.zeros(m, n, dtype=torch.bfloat16, device=dev)
config = SHAPE_CONFIGS.get((64, 7168, 2048), get_config(64, 7168, 2048))
config["SPLITK_BLOCK_SIZE"] = 2 * Kp
grid_size = ((m + config["BLOCK_SIZE_M"] - 1) // config["BLOCK_SIZE_M"]) * \
((n + config["BLOCK_SIZE_N"] - 1) // config["BLOCK_SIZE_N"])
try:
_gemm_fp4_kernel[(grid_size,)](
Aq, Bt, out, Asc, Bsc8,
m, n, Kp, Bsc8.shape[1],
Aq.stride(0), Aq.stride(1), Bt.stride(0), Bt.stride(1),
0, out.stride(0), out.stride(1), Asc.stride(0), Asc.stride(1),
SPLITK_BLOCK_SIZE=config["SPLITK_BLOCK_SIZE"],
**{k2: config[k2] for k2 in ["BLOCK_SIZE_M","BLOCK_SIZE_N","BLOCK_SIZE_K",
"GROUP_SIZE_M","NUM_KSPLIT","num_warps","num_stages","waves_per_eu","matrix_instr_nonkdim"]})
torch.cuda.synchronize()
except Exception as e:
print(f"[HSACO] Triton warmup failed: {e}")
return
# Scan Triton cache for the compiled .hsaco
cache_base = os.path.expanduser("~/.triton/cache")
hsaco_path = None
for root, _, files in os.walk(cache_base):
for f in files:
if f == '_gemm_fp4_kernel.hsaco':
hsaco_path = os.path.join(root, f)
break
if hsaco_path:
break
if hsaco_path:
hsaco_key = hsaco_path + "::_gemm_fp4_kernel"
hip_mod.load_triton_hsaco(hsaco_path, "_gemm_fp4_kernel")
custom_kernel._hsaco_key = hsaco_key
print(f"[HSACO] Loaded: {hsaco_path}")
else:
print("[HSACO] .hsaco not found in cache")
# Also compile .amdgcn → .co (ISA reproduction)
amdgcn_path = None
for root2, _, files2 in os.walk(cache_base):
for f2 in files2:
if f2 == '_gemm_fp4_kernel.amdgcn':
amdgcn_path = os.path.join(root2, f2)
break
if amdgcn_path: break
if amdgcn_path:
import subprocess as _sp
co_path = "/tmp/_gemm_fp4_isa.co"
obj_path = "/tmp/_gemm_fp4_isa.o"
# Find llvm-mc (try common ROCm paths)
llvm_mc = None
for p in ["/opt/rocm/llvm/bin/llvm-mc", "/usr/bin/llvm-mc",
"/opt/rocm/lib/llvm/bin/llvm-mc"]:
if os.path.exists(p):
llvm_mc = p; break
lld = None
for p in ["/opt/rocm/llvm/bin/ld.lld", "/usr/bin/ld.lld",
"/opt/rocm/lib/llvm/bin/ld.lld"]:
if os.path.exists(p):
lld = p; break
if llvm_mc and lld:
r1 = _sp.run([llvm_mc, "-triple=amdgcn-amd-amdhsa", "-mcpu=gfx950",
"-filetype=obj", amdgcn_path, "-o", obj_path],
capture_output=True, text=True)
if r1.returncode == 0:
r2 = _sp.run([lld, "-shared", obj_path, "-o", co_path],
capture_output=True, text=True)
if r2.returncode == 0:
co_key = co_path + "::_gemm_fp4_kernel"
hip_mod.load_triton_hsaco(co_path, "_gemm_fp4_kernel")
custom_kernel._isa_co_key = co_key
print(f"[ISA-CO] Compiled & loaded: {co_path}")
else:
print(f"[ISA-CO] ld.lld failed: {r2.stderr[:200]}")
else:
print(f"[ISA-CO] llvm-mc failed: {r1.stderr[:200]}")
else:
print(f"[ISA-CO] llvm-mc={llvm_mc} lld={lld} — tools not found")
# Free dummy tensors
del Aq, Asc, Bq8, Bt, Bsc8, out
def _get_aq(m, k, dev):
key = (m, k, dev)
if key not in _AQ_CACHE:
_AQ_CACHE[key] = (torch.empty((m, k//2), dtype=torch.uint8, device=dev),
torch.empty((m, k//SCALE_GROUP_SIZE), dtype=torch.uint8, device=dev))
return _AQ_CACHE[key]
def custom_kernel(data: input_t) -> output_t:
global _warmed, _hip
if not _warmed:
_hip = _hip_module()
knl, co, _ = ASM_DEFAULT
_hip.warmup_asm(ASM_BASE + co, knl)
for _, v in ASM_CONFIGS_SHAPE.items():
_hip.warmup_asm(ASM_BASE + v[1], v[0])
_warm_triton_hsaco(_hip)
_warmed = True
A = data[0]
m, k = A.shape
n = data[2].shape[0] # B_q.shape[0]
Kp = k >> 1
hip = _hip
B_q = data[2]
B_shuffle = data[3]
B_scale_sh = data[4]
# ── Path 1: HIP fused 16x16 for M≤32 K≤1024 (single kernel, LDS K-reduction) ──
# v31 REVERTED: HIP fused for large K regressed 12→20.7µs (uncoalesced B access).
# Triton tl.dot_scaled uses coalesced tiled loads — must keep Triton for K>1024.
if m <= 32 and k <= 1024:
key_hip = (m, n)
if key_hip not in _HIP16_OUT:
_HIP16_OUT[key_hip] = (
torch.empty(m, n, dtype=torch.bfloat16, device=A.device),
int(B_q.stride(0) * B_q.element_size()), # cache strBq
int(B_scale_sh.size(1)), # cache BscSN
)
out, strBq, BscSN = _HIP16_OUT[key_hip]
hip.dispatch_hip(A.data_ptr(), B_q.data_ptr(), B_scale_sh.data_ptr(),
out.data_ptr(), 0, m, n, k, A.stride(0), strBq, BscSN, 0)
return out
# ── Path 1b: v78 N-parallel DISABLED — inline quant too slow for K>1024 (32.3µs vs 12.6µs Triton)
# Need non-fused version (pre-quantized A) to eliminate quant overhead. See v78 results.
if False: # 32 < m <= 127 and n % 64 == 0:
key_np = (m, n)
if key_np not in _HIP16_OUT:
out = torch.empty(m, n, dtype=torch.bfloat16, device=A.device)
_HIP16_OUT[key_np] = (out, int(B_shuffle.stride(0) * B_shuffle.element_size()), int(B_scale_sh.size(1)))
out, strBsh, BscSN = _HIP16_OUT[key_np]
# Pass B_shuffle (coalesced [N/16][K/32][16][16]) instead of B_q (strided [N, K/2])
hip.launch_fused_npar(A.data_ptr(), B_shuffle.data_ptr(), B_scale_sh.data_ptr(),
out.data_ptr(), m, n, k, A.stride(0), strBsh, BscSN)
return out
# ── Path 2: HIP ASM for M≥128 (hand-tuned 186% MFMA util, beats Triton) ──
if m >= 128:
key_asm = (m, n, k)
if key_asm not in _ASM_WS:
# Parse tile_m from kernel name for correct padding
knl_tmp, _, _ = ASM_CONFIGS_SHAPE.get((m, n, k), ASM_DEFAULT)
tile_m = 32 # default
xpos = knl_tmp.rfind('x')
if xpos > 0:
i = xpos - 1
while i > 0 and knl_tmp[i-1].isdigit(): i -= 1
tile_m = int(knl_tmp[i:xpos])
mp = ((m + tile_m - 1) // tile_m) * tile_m
sm = mp
_ASM_WS[key_asm] = {
'aq': torch.empty(m, k // 2, dtype=torch.uint8, device=A.device),
'asc': torch.empty(sm, k // 32, dtype=torch.uint8, device=A.device),
'out': torch.empty(mp, n, dtype=torch.bfloat16, device=A.device),
'sm': sm, 'mp': mp,
}
w = _ASM_WS[key_asm]
knl, co, sk = ASM_CONFIGS_SHAPE.get((m, n, k), ASM_DEFAULT)
hip.quant_and_asm(A, B_shuffle, B_scale_sh, n,
w['aq'], w['asc'], w['out'], w['sm'],
knl, ASM_BASE + co, sk)
return w['out'][:m] if w['mp'] > m else w['out']
# ── Path 2b: Launch Triton's ISA compiled to .co (shape 5 ONLY) ──
if m == 64 and n == 7168 and k == 2048 and hasattr(custom_kernel, '_isa_co_key'):
Aq, Asc = _get_aq(m, k, A.device)
hip.quantize_a_hip(A, Aq, Asc)
Bq8 = B_q.view(torch.uint8) if B_q.dtype != torch.uint8 else B_q
Bt = Bq8.T
Bsc8 = B_scale_sh.view(torch.uint8) if B_scale_sh.dtype != torch.uint8 else B_scale_sh
key_out = ('isa_co', m, n)
if key_out not in _OUT_CACHE:
_OUT_CACHE[key_out] = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
out = _OUT_CACHE[key_out]
Kp = k >> 1
grid_size = ((m + 31) // 32) * ((n + 63) // 64)
hip.launch_triton_fp4gemm(
custom_kernel._isa_co_key,
Aq.data_ptr(), Bt.data_ptr(), out.data_ptr(), Asc.data_ptr(), Bsc8.data_ptr(),
m, n, Kp, Bsc8.shape[1],
Aq.stride(0), Bt.stride(1), 0, out.stride(0), Asc.stride(0),
grid_size, 104448)
return out
# ── Path 3: Triton GEMM (M=16-64 with K>1024, or M=33-127) ──
Bq8 = B_q.view(torch.uint8) if B_q.dtype != torch.uint8 else B_q
Bt = Bq8.T
Bsc8 = B_scale_sh.view(torch.uint8) if B_scale_sh.dtype != torch.uint8 else B_scale_sh
Ks_stride = Bsc8.shape[1]
use_fused = (m <= FUSED_M_THRESHOLD)
if not use_fused:
Aq, Asc = _get_aq(m, k, A.device)
hip.quantize_a_hip(A, Aq, Asc)
config = get_config(m, n, k)
if config["NUM_KSPLIT"] > 1:
SPBS, BSK, NSK = get_splitk(Kp, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"])
config["SPLITK_BLOCK_SIZE"] = SPBS
config["BLOCK_SIZE_K"] = BSK
config["NUM_KSPLIT"] = NSK
else:
config["SPLITK_BLOCK_SIZE"] = 2 * Kp
if config["BLOCK_SIZE_K"] >= 2 * Kp:
config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * Kp)
config["SPLITK_BLOCK_SIZE"] = 2 * Kp
config["NUM_KSPLIT"] = 1
config["BLOCK_SIZE_K"] = max(config["BLOCK_SIZE_K"], 128)
NSK = config["NUM_KSPLIT"]
out_key = (m, n, NSK, A.device)
cached = _OUT_CACHE.get(out_key)
if cached is not None:
if NSK > 1: y_pp, y = cached
else: y = cached; y_pp = None
elif NSK > 1:
y_pp = torch.empty((NSK, m, n), dtype=torch.float32, device=A.device)
y = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
_OUT_CACHE[out_key] = (y_pp, y)
else:
y = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
_OUT_CACHE[out_key] = y
y_pp = None
out_t = y if NSK == 1 else y_pp
grid = lambda META: (META["NUM_KSPLIT"] * triton.cdiv(m, META["BLOCK_SIZE_M"]) * triton.cdiv(n, META["BLOCK_SIZE_N"]),)
if use_fused:
_fused_gemm_fp4_kernel[grid](
A, Bt, out_t, Bsc8,
m, n, Kp, k, Ks_stride,
A.stride(0), A.stride(1), Bt.stride(0), Bt.stride(1),
0 if NSK == 1 else y_pp.stride(0), out_t.stride(-2), out_t.stride(-1),
SPLITK_BLOCK_SIZE=config["SPLITK_BLOCK_SIZE"],
**{k2: config[k2] for k2 in ["BLOCK_SIZE_M","BLOCK_SIZE_N","BLOCK_SIZE_K",
"GROUP_SIZE_M","NUM_KSPLIT","num_warps","num_stages","waves_per_eu","matrix_instr_nonkdim"]})
else:
_gemm_fp4_kernel[grid](
Aq, Bt, out_t, Asc, Bsc8,
m, n, Kp, Ks_stride,
Aq.stride(0), Aq.stride(1), Bt.stride(0), Bt.stride(1),
0 if NSK == 1 else y_pp.stride(0), out_t.stride(-2), out_t.stride(-1),
Asc.stride(0), Asc.stride(1),
SPLITK_BLOCK_SIZE=config["SPLITK_BLOCK_SIZE"],
**{k2: config[k2] for k2 in ["BLOCK_SIZE_M","BLOCK_SIZE_N","BLOCK_SIZE_K",
"GROUP_SIZE_M","NUM_KSPLIT","num_warps","num_stages","waves_per_eu","matrix_instr_nonkdim"]})
if NSK > 1:
RBLM, RBLN = 16, 16
ANSK = triton.cdiv(Kp, config["SPLITK_BLOCK_SIZE"] // 2)
_reduce_kernel[(triton.cdiv(m, RBLM), triton.cdiv(n, RBLN))](
y_pp, y, m, n,
y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
y.stride(0), y.stride(1),
BLOCK_SIZE_M=RBLM, BLOCK_SIZE_N=RBLN,
ACTUAL_KSPLIT=ANSK, MAX_KSPLIT=triton.next_power_of_2(NSK))
return y
# ═══════════════════════════════════════════════════════════════════
# Profiling (DO_PROFILE=0 by default, set to 1 to enable)
# ═══════════════════════════════════════════════════════════════════
_COUNTER_GROUPS = [
"TCC_HIT_sum TCC_MISS_sum TCC_EA0_RDREQ_sum TCC_EA0_WRREQ_sum",
"SQ_LDS_BANK_CONFLICT SQ_INSTS_MFMA SQ_INSTS_VALU SQ_BUSY_CYCLES",
"SQ_VALU_MFMA_BUSY_CYCLES SQ_INSTS_VMEM_RD SQ_WAIT_ANY SQ_ACTIVE_INST_VMEM",
]
_SHAPES = [(4,2880,512),(16,2112,7168),(32,4096,512),(32,2880,512),(64,7168,2048),(256,3072,1536)]
def _make_profile_inputs():
try:
from reference import generate_input
return [generate_input(m=M, n=N, k=K, seed=42) for M, N, K in _SHAPES]
except Exception as e:
print(f"[PROFILE] generate_input failed: {e}"); return []
def _run_profiling_kernels():
inputs = _make_profile_inputs()
if not inputs: return
for d in inputs: custom_kernel(d)
torch.cuda.synchronize()
for i, d in enumerate(inputs):
print(f"[PROFILE] Shape M={_SHAPES[i][0]} N={_SHAPES[i][1]} K={_SHAPES[i][2]}")
custom_kernel(d); torch.cuda.synchronize()
def _parse_results(outdir):
dispatches = {}
_OUR = ('splitk','gemm_fused_16x16','fp32_to_bf16','quant_a_kern','f4gemm',
'_fused_gemm_fp4','_gemm_fp4','_reduce_kernel','quantize_a_kernel')
for cf in sorted(glob.glob(os.path.join(outdir, "**/*counter_collection.csv"), recursive=True)):
with open(cf) as f:
for row in csv.DictReader(f):
did, kn, cn, cv = row.get("Dispatch_Id",""), row.get("Kernel_Name",""), row.get("Counter_Name",""), row.get("Counter_Value","0")
if not kn or not any(k in kn for k in _OUR): continue
if did not in dispatches:
dispatches[did] = {'_k':kn,'_g':row.get("Grid_Size",""),'_b':row.get("Workgroup_Size",""),
'_v':row.get("VGPR_Count",""),'_a':row.get("Accum_VGPR_Count",""),'_s':row.get("Scratch_Size",""),'_l':row.get("LDS_Block_Size","")}
dispatches[did][cn] = cv
print("\n" + "="*80); print("ROCPROFV3 HARDWARE COUNTER RESULTS"); print("="*80)
for did in sorted(dispatches.keys(), key=lambda x: int(x)):
d = dispatches[did]
short = d['_k'].split('::')[-1][:50] if '::' in d['_k'] else d['_k'][:50]
print(f"\n--- Dispatch {did}: {short}")
print(f" Grid={d['_g']} Block={d['_b']} VGPR={d['_v']} AGPR={d['_a']} Scratch={d['_s']} LDS={d['_l']}")
for k2 in sorted(d.keys()):
if not k2.startswith('_'): print(f" {k2:40s} = {d[k2]}")
sys.stdout.flush()
def _run_profiling():
print("[PROFILE] Starting rocprofv3...")
tmpdir = tempfile.mkdtemp(prefix="rp_")
cf = os.path.join(tmpdir, "counters.txt")
with open(cf, "w") as f:
for g in _COUNTER_GROUPS: f.write(f"pmc: {g}\n")
runner = os.path.join(tmpdir, "runner.py")
shutil.copy(os.path.abspath(__file__), runner)
outdir = os.path.join(tmpdir, "out"); os.makedirs(outdir)
env = os.environ.copy()
env.update(PROFILE_SUBPROCESS="1", DO_PROFILE="0", ORIG_DIR=_ORIG_DIR)
cmd = ["rocprofv3","--input",cf,"--output-directory",outdir,"--output-format","csv","--",sys.executable,runner]
try:
r = subprocess.run(cmd, env=env, capture_output=True, text=True, timeout=280)
if r.stdout:
for l in r.stdout.splitlines()[-30:]: print(f" sub: {l}")
if r.stderr:
for l in r.stderr.splitlines()[-10:]: print(f" err: {l}")
_parse_results(outdir)
except Exception as e:
print(f"[PROFILE] Error: {e}")
if _IS_PROFILE_SUB:
_run_profiling_kernels(); sys.exit(0)
elif _DO_PROFILE:
_run_profiling()
# torch.compile failed: can't trace pybind11 C++ extensions.
# Would need torch.library custom op registration to work.
scrolls · 1394 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