submission 754569
.jonnss · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 404 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754569?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:eebcabe56408df7e34392d93f6fcd0eca3946d285e669d25b1629ed001cefefc
license declaredunknown
license concludedunknown
authors.jonnss
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
vector-width = float4
if(i+3<n){float4 a=*reinterpret_cast<const float4*>(p+i);Kernel source
submission.py404 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
# Microscaled blocked matmul for CDNA4 wave processors.
# Activation tensor is narrowed to 4-bit on the fly using
# the paired nibble converter in the gfx950 datapath.
# Weight operand arrives in a hardware-friendly permuted
# layout and is unpacked inside the compute loop.
# K-dimension can be striped across wave groups when
# the tile count under-saturates the shader array.
from task import input_t, output_t
import torch, triton, triton.language as tl
import sys as _J, time as _K, gc as _H
torch.set_grad_enabled(False); _H.disable(); _J.setswitchinterval(1.0)
_pr = lambda m: print(m, file=_J.stderr, flush=True)
# total shader engines on this part
_SE = 256
# ==============================================================
# Native HIP lane merger — 128-bit vectorized K-stripe reducer
# Compiles at import time; Triton fallback if unavailable
# ==============================================================
_HSRC = r"""
#include <hip/hip_runtime.h>
__device__ __forceinline__ unsigned short narrow(float x){
unsigned int w; __builtin_memcpy(&w,&x,sizeof(w));
return (unsigned short)((w+((w>>16)&1)+0x7FFFu)>>16);}
template<int S> __global__ void wide_fold(const float*__restrict__ p,
unsigned short*__restrict__ d, int n){
int i=(blockIdx.x*blockDim.x+threadIdx.x)*4;
if(i+3<n){float4 a=*reinterpret_cast<const float4*>(p+i);
#pragma unroll
for(int s=1;s<S;s++){float4 v=*reinterpret_cast<const float4*>(p+s*n+i);
a.x+=v.x;a.y+=v.y;a.z+=v.z;a.w+=v.w;}
*reinterpret_cast<unsigned long long*>(d+i)=
(unsigned long long)narrow(a.x)|((unsigned long long)narrow(a.y)<<16)|
((unsigned long long)narrow(a.z)<<32)|((unsigned long long)narrow(a.w)<<48);
}else{for(int j=i;j<n&&j<i+4;j++){float a=p[j];
#pragma unroll
for(int s=1;s<S;s++)a+=p[s*n+j]; d[j]=narrow(a);}}}
__global__ void thin_fold(const float*__restrict__ p,unsigned short*__restrict__ d,int n,int s){
int g=blockIdx.x*blockDim.x+threadIdx.x;
if(g<n){float a=p[g]; for(int k=1;k<s;k++)a+=p[k*n+g]; d[g]=narrow(a);}}
void run_fold(torch::Tensor p,torch::Tensor d,int r,int c,int s){
int n=r*c; auto*sp=p.data_ptr<float>();
auto*dp=reinterpret_cast<unsigned short*>(d.data_ptr());
int t=64,bl=(n+t*4-1)/(t*4);
switch(s){
case 2:wide_fold<2><<<bl,t>>>(sp,dp,n);break;
case 3:wide_fold<3><<<bl,t>>>(sp,dp,n);break;
case 4:wide_fold<4><<<bl,t>>>(sp,dp,n);break;
case 7:wide_fold<7><<<bl,t>>>(sp,dp,n);break;
case 8:wide_fold<8><<<bl,t>>>(sp,dp,n);break;
case 14:wide_fold<14><<<bl,t>>>(sp,dp,n);break;
default:{int t2=256; thin_fold<<<(n+t2-1)/t2,t2>>>(sp,dp,n,s);break;}}}
"""
_HH = "void run_fold(torch::Tensor p,torch::Tensor d,int r,int c,int s);"
_GOT_HIP = False
try:
from torch.utils.cpp_extension import load_inline as _cc
_c0 = _K.time()
_hip_fold = _cc(name="nf4m", cpp_sources=[_HH], cuda_sources=[_HSRC],
functions=["run_fold"], verbose=False,
extra_cuda_cflags=["--offload-arch=gfx950","-O3"])
_GOT_HIP = True; _pr(f"[hip] fold ready ({_K.time()-_c0:.1f}s)")
except Exception as _e: _pr(f"[hip] fold unavail: {_e}")
# ==============================================================
# Register-level BF16 -> packed FP4 narrowing
# Each group of 32 elements shares one E8M0 scale byte.
# The gfx950 ISA instruction converts a pair of BF16 values
# into a single byte holding two FP4 nibbles.
# ==============================================================
@triton.jit
def _narrow_to_fp4(raw, NR: tl.constexpr, NC: tl.constexpr):
GWIDTH: tl.constexpr = 32
NGRP: tl.constexpr = NC // GWIDTH
f32 = raw.to(tl.float32).reshape(NR, NGRP, GWIDTH)
# find per-group peak, snap to power-of-two via bit rounding
pk = tl.max(tl.abs(f32), axis=-1, keep_dims=True)
pk = pk.to(tl.int32, bitcast=True)
pk = (pk + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
# derive unbiased exponent for the E8M0 scale encoding
be = ((pk >> 23) & 0xFF).to(tl.int32) - 127
ue = be - 2
ue = tl.minimum(tl.maximum(ue, -127), 127)
sc_byte = ue.to(tl.uint8) + 127
# build the IEEE754 float that the hw divider expects
dv = (ue.to(tl.int32) + 127).to(tl.uint32) << 23
dv = dv.to(tl.float32, bitcast=True)
# broadcast divider to every element pair
dv = tl.broadcast_to(dv, (NR, NGRP, GWIDTH)).reshape(NR, NC)
dv = dv.reshape(NR, NC // 2, 2)
dv_l, _ = tl.split(dv)
dv_f = dv_l.reshape(NR, NC // 2)
# pack adjacent bf16 into u32 words for the converter
u16 = raw.to(tl.uint16, bitcast=True).reshape(NR, NC // 2, 2)
w0, w1 = tl.split(u16)
pw = w0.to(tl.uint32) | (w1.to(tl.uint32) << 16)
pw = pw.reshape(NR, NC // 2)
# one ISA op produces both nibbles per pair
nib = tl.inline_asm_elementwise(
"v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",
"=v, v, v", [pw, dv_f],
dtype=tl.uint32, is_pure=True, pack=1)
out = (nib & 0xFF).to(tl.uint8).reshape(NR, NC // 2)
return out, sc_byte.reshape(NR, NGRP)
# ==============================================================
# Fused narrowing + scaled dot product engine
# Memory policy: L2-pin activations, write-through results,
# cache-global for weight operands. Relaxed FP math enabled.
# ==============================================================
@triton.heuristics({
"KDIV_EXACT": lambda args: (args["K"] % (args["TC"] // 2) == 0)
and (args["KSTRIPE"] % args["TC"] == 0)
and (args["K"] % (args["KSTRIPE"] // 2) == 0),
})
@triton.jit
def _tiled_gemm_engine(
xp, wp, yp, sp,
M, N, K,
sx_r, sx_c, sw_r, sw_c,
sy_s, sy_r, sy_c, ss_r, ss_c,
TR: tl.constexpr, TN: tl.constexpr, TC: tl.constexpr,
GSWIZ: tl.constexpr, NSTRIPE: tl.constexpr, KSTRIPE: tl.constexpr,
KDIV_EXACT: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr,
waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
ld_policy: tl.constexpr,
):
# stride positivity hints for the LLVM backend
tl.assume(sx_r > 0); tl.assume(sx_c > 0)
tl.assume(sw_r > 0); tl.assume(sw_c > 0)
tl.assume(sy_r > 0); tl.assume(sy_c > 0)
tl.assume(ss_r > 0); tl.assume(ss_c > 0)
SG: tl.constexpr = 32 # microscale group width
nr = tl.cdiv(M, TR); nc = tl.cdiv(N, TN)
# decompose flat program id into stripe + tile coordinates
pid = tl.program_id(axis=0)
sid = pid % NSTRIPE
wid = pid // NSTRIPE
# single-stripe: grouped swizzle for L2 locality
# multi-stripe: simple linear mapping
if NSTRIPE == 1:
gw = GSWIZ * nc
gi = wid // gw
gr = gi * GSWIZ
gl = min(nr - gr, GSWIZ)
ri = gr + ((wid % gw) % gl)
ci = (wid % gw) // gl
else:
ri = wid // nc
ci = wid % nc
tl.assume(ri >= 0); tl.assume(ci >= 0); tl.assume(sid >= 0)
# guard: skip if this stripe starts past K boundary
if (sid * KSTRIPE // 2) < K:
nloop = tl.cdiv(KSTRIPE // 2, TC // 2)
# activation tile: BF16 input rows
ro = (ri * TR + tl.arange(0, TR)) % M
co = sid * KSTRIPE + tl.arange(0, TC)
ap = xp + (ro[:, None] * sx_r + co[None, :] * sx_c)
# weight tile: pre-permuted FP4x2 blocks
lr = tl.arange(0, (TC // 2) * 16)
lo = sid * (KSTRIPE // 2) * 16 + lr
wr = (ci * (TN // 16) + tl.arange(0, TN // 16)) % (N // 16)
bpp = wp + (wr[:, None] * sw_r + lo[None, :] * sw_c)
# weight scale tile: shuffled E8M0 blocks
sr_off = (ci * TN + tl.arange(0, TN // 32) * 32)
sc_off = (sid * (KSTRIPE // SG) * 32) + tl.arange(0, TC // SG * 32)
spp = sp + sr_off[:, None] * ss_r + sc_off[None, :] * ss_c
# initialize tile accumulator
dot = tl.zeros((TR, TN), dtype=tl.float32)
for step in range(sid * nloop, (sid + 1) * nloop):
# fire all loads first for maximum MLP
if KDIV_EXACT:
xblk = tl.load(ap, eviction_policy="evict_last")
sblk = tl.load(spp, cache_modifier=ld_policy)
wblk = tl.load(bpp, cache_modifier=ld_policy)
else:
koff = (step - sid * nloop) * TC
xblk = tl.load(ap,
mask=tl.arange(0, TC)[None, :] < (2*K - sid*KSTRIPE - koff),
other=0.0, eviction_policy="evict_last")
sblk = tl.load(spp, cache_modifier=ld_policy)
wblk = tl.load(bpp,
mask=lr[None, :] < ((K - (sid*(KSTRIPE//2) + (step-sid*nloop)*(TC//2)))*16),
other=0, cache_modifier=ld_policy)
# narrow activations in register file
xn, xs = _narrow_to_fp4(xblk, TR, TC)
# undo the permutation on weight scales
ws = (sblk
.reshape(TN//32, TC//SG//8, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(TN, TC // SG))
# undo the permutation on weight data
wt = (wblk
.reshape(1, TN//16, TC//64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(TN, TC // 2).trans(1, 0))
# scaled dot with relaxed precision for better scheduling
dot = tl.dot_scaled(xn, xs, "e2m1", wt, ws, "e2m1", dot, fast_math=True)
# advance pointers by one K-block
ap += TC * sx_c
bpp += (TC // 2) * 16 * sw_c
spp += TC * ss_c
# write-through to avoid polluting L2 with output data
res = dot.to(yp.type.element_ty)
yr = ri * TR + tl.arange(0, TR).to(tl.int64)
yc = ci * TN + tl.arange(0, TN).to(tl.int64)
ypp = yp + sy_r * yr[:, None] + sy_c * yc[None, :] + sid * sy_s
ymask = (yr[:, None] < M) & (yc[None, :] < N)
tl.store(ypp, res, mask=ymask, cache_modifier=".wt")
# ==============================================================
# Triton-based stripe reducer (fallback when HIP unavailable)
# ==============================================================
@triton.jit
def _fold_partials(
fp, dp, M, N, s_fs, s_fr, s_fc, s_dr, s_dc,
FR: tl.constexpr, FC: tl.constexpr,
REAL: tl.constexpr, PAD: tl.constexpr,
):
pr = tl.program_id(0); pc = tl.program_id(1)
ro = (pr * FR + tl.arange(0, FR)) % M
co = (pc * FC + tl.arange(0, FC)) % N
bp = fp + (ro[:, None] * s_fr) + (co[None, :] * s_fc)
acc = tl.load(bp).to(tl.float32)
for i in tl.static_range(1, PAD):
if i < REAL: acc += tl.load(bp + i * s_fs).to(tl.float32)
dp2 = dp + (ro[:, None] * s_dr) + (co[None, :] * s_dc)
tl.store(dp2, acc.to(dp.type.element_ty))
# ==============================================================
# K-stripe alignment fixer
# ==============================================================
def _fix_stripe(kh, tc, ns):
ks = triton.cdiv((2 * triton.cdiv(kh, ns)), tc) * tc
while ns > 1 and tc > 16:
if kh%(ks//2)==0 and ks%tc==0 and kh%(tc//2)==0: break
elif kh%(ks//2)!=0 and ns>1: ns //= 2
elif ks%tc!=0:
if ns>1: ns //= 2
elif tc>16: tc //= 2
elif kh%(tc//2)!=0 and tc>16: tc //= 2
else: break
ks = triton.cdiv((2 * triton.cdiv(kh, ns)), tc) * tc
ns = triton.cdiv(kh, ks // 2)
return ks, tc, ns
# ==============================================================
# Profiled tile geometries from offline sweep
# ==============================================================
_SWEEP = {
(4,2880,512): (4,128,256,1,4,2,1,16,None,1),
(16,2112,7168): (16,128,512,1,4,2,3,16,".cg",14),
(32,4096,512): (16,32,256,1,4,3,3,16,".cg",1),
(32,2880,512): (8,128,256,1,4,2,2,16,None,1),
(64,7168,2048): (16,128,256,1,4,2,2,16,".cg",1),
(256,3072,1536): (16,256,512,1,8,2,2,16,None,1),
}
# unpack: (TR, TN, TC, GSWIZ, nw, ns, wpe, mi, cm, nstripe)
def _from_sweep(t):
return {"TR":t[0],"TN":t[1],"TC":t[2],"GSWIZ":t[3],
"num_warps":t[4],"num_stages":t[5],"waves_per_eu":t[6],
"matrix_instr_nonkdim":t[7],"ld_policy":t[8],"NSTRIPE":t[9]}
# ==============================================================
# Occupancy model — analytical fallback for unknown shapes
# ==============================================================
def _model_cfg(m, n, k):
if m <= 32:
tr, tn = 8, 128
we = ((m+tr-1)//tr)*((n+127)//128); ns = 1
if k>=4096: ns=7
elif k>=2048: ns = 4 if we*2<(_SE*3)//4 else 2
elif k>=1536: ns = 3 if we*2<(_SE*3)//4 else 2
tc = 256 if k<=ns*512 else 512
if we*ns < (_SE*3)//4: tn = 64
else:
tr = 16
if m<=128 and ((m+15)//16)*((n+127)//128) < (_SE*3)//4: tr = 8
we = ((m+tr-1)//tr)*((n+127)//128); tn, ns = 128, 1
if _SE//2<=we<=_SE and (k>=7168 or (k>=2048 and tr==8)): ns=2
elif we<_SE//2 and k>512:
if k>=4096: ns = 2 if we*2>=_SE else 7
elif k>=2048: ns=2
elif k>=1536: ns=3
tc = 256 if k<=max(ns*4096,2048) else 512
if we*ns < (_SE*3)//4: tn = 64
return {"TR":tr,"TN":max(tn,32),"TC":tc,"GSWIZ":1,
"num_warps":4,"num_stages":2,"waves_per_eu":2,
"matrix_instr_nonkdim":16,"ld_policy":".cg","NSTRIPE":ns}
# ==============================================================
# Cached config resolver + launch parameter precomputation
# ==============================================================
_RC = {}; _YC = {}; _FC = {}; _WC = {}; _LC = {}
def _alloc_y(m, n, ns, dev):
t = (m,n,ns)
if t not in _YC:
_YC[t] = (torch.empty((m,n),dtype=torch.bfloat16,device=dev),
torch.empty((ns,m,n),dtype=torch.float32,device=dev) if ns>1 else None)
return _YC[t]
def _cfg(m, n, k):
t = (m,n,k)
if t in _RC: return _RC[t]
c = _from_sweep(_SWEEP[t]) if t in _SWEEP else _model_cfg(m, n, k)
kh = k // 2
if c["NSTRIPE"] > 1:
ks, tc2, ns2 = _fix_stripe(kh, c["TC"], c["NSTRIPE"])
c["KSTRIPE"]=ks; c["TC"]=tc2; c["NSTRIPE"]=ns2
else:
c["KSTRIPE"] = 2*kh; c["NSTRIPE"] = 1
if c["TC"] >= 2*kh:
c["TC"]=triton.next_power_of_2(2*kh); c["KSTRIPE"]=2*kh; c["NSTRIPE"]=1
c["TN"]=max(c["TN"],32)
_RC[t] = c
return c
def _get_wt(wd, ws, n, kh):
a = wd.data_ptr()
if a not in _WC:
_WC[a] = (wd.view(torch.uint8).reshape(n//16, kh*16), ws.view(torch.uint8))
return _WC[a]
def _launch(m, n, k, dev):
t = (m,n,k)
if t in _LC: return _LC[t]
c = _cfg(m,n,k); kh=k//2; ns=c["NSTRIPE"]
y,pp = _alloc_y(m,n,ns,dev)
g = (ns * triton.cdiv(m,c["TR"]) * triton.cdiv(n,c["TN"]),)
if ns==1: s_s,s_r,s_c = 0,y.stride(0),y.stride(1)
else: s_s,s_r,s_c = pp.stride(0),pp.stride(1),pp.stride(2)
b = {'c':c,'kh':kh,'g':g,'ns':ns,'ss':s_s,'sr':s_r,'sc':s_c}
if ns>1:
b['rg']=(triton.cdiv(m,16),triton.cdiv(n,64))
b['rn']=triton.cdiv(kh,c["KSTRIPE"]//2)
b['rp']=triton.next_power_of_2(ns)
_LC[t] = b
return b
# ==============================================================
# Entry point
# ==============================================================
def _go(x, wd, ws, m, n, k):
b = _launch(m, n, k, x.device)
y, pp = _alloc_y(m, n, b['ns'], x.device)
wf, sf = _get_wt(wd, ws, n, b['kh'])
_tiled_gemm_engine[b['g']](
x, wf, y if b['ns']==1 else pp, sf,
m, n, b['kh'],
x.stride(0), x.stride(1), wf.stride(0), wf.stride(1),
b['ss'], b['sr'], b['sc'], sf.stride(0), sf.stride(1),
**b['c'])
if b['ns'] > 1:
if _GOT_HIP:
_hip_fold.run_fold(pp, y, m, n, b['rn'])
else:
_fold_partials[b['rg']](pp, y, m, n,
pp.stride(0),pp.stride(1),pp.stride(2),
y.stride(0),y.stride(1), 16, 64, b['rn'], b['rp'])
return y
def custom_kernel(data: input_t) -> output_t:
A = data[0]
return _go(A, data[3], data[4], A.shape[0], data[1].shape[0], A.shape[1])
scrolls · 404 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 754371.
#!POPCORN leaderboard amd-mxfp4-mm#!POPCORN gpu MI355X- """- CDNA4 blocked FP4 matmul via aiter preshuffle path.- Runtime patches the quantization subroutine with ISA-level- paired conversion and adjusts the accumulator idiom.- Occupancy-aware tiling with phased JIT warmup.- """+ # Microscaled blocked matmul for CDNA4 wave processors.+ # Activation tensor is narrowed to 4-bit on the fly using+ # the paired nibble converter in the gfx950 datapath.+ # Weight operand arrives in a hardware-friendly permuted+ # layout and is unpacked inside the compute loop.+ # K-dimension can be striped across wave groups when+ # the tile count under-saturates the shader array.- import os as _env- _env.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")- _env.environ.setdefault("CXX", "clang++")- import uuid as _uid- _env.environ["TRITON_CACHE_DIR"] = f"/tmp/_tc_{_uid.uuid4().hex[:8]}"-- # Synthesize a minimal config CSV so aiter skips its heavy build paths- _ASM_LABEL = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"- _CSV_TMP = "/tmp/_fp4_cfg.csv"- _ENGINES = 256- _DIM_PAIRS = [- (2880, 512), (2112, 7168), (4096, 512), (7168, 2048), (3072, 1536),- (2880, 1536), (4096, 1536), (2112, 512), (2112, 2048),- (7168, 512), (7168, 1536), (7168, 7168), (3072, 512),- (3072, 7168), (3072, 2048), (4096, 2048), (4096, 7168),- (2880, 2048), (2880, 7168),- ]- _BATCH_DIMS = [1, 2, 4, 8, 16, 32, 64, 128, 256]- _rows = ["cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio"]- for _d1, _d2 in _DIM_PAIRS:- for _b in _BATCH_DIMS:- _wave_cnt = ((_b + 31) // 32) * ((_d1 + 127) // 128)- _ratio = _ENGINES / max(_wave_cnt, 1)- _lg = 0- while _ratio >= pow(2, _lg + 1) and (pow(2, _lg + 1) * 128) < 2 * _d2:- _lg += 1- _lg = min(_lg, 3)- _rows.append(f"{_ENGINES},{_b},{_d1},{_d2},21,{_lg},1.0,{_ASM_LABEL},0,0,0.0")- with open(_CSV_TMP, "w") as _fh:- _fh.write("\n".join(_rows))- _env.environ["AITER_CONFIG_GEMM_A4W4"] = (- _CSV_TMP + ":/home/runner/aiter/aiter/configs/a4w4_blockscale_tuned_gemm.csv"- )-- import torch- torch.set_grad_enabled(False)- import triton- import triton.language as tl- from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (- _gemm_a16wfp4_preshuffle_kernel,- )- from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (- _gemm_afp4wfp4_reduce_kernel,- )from task import input_t, output_t- import sys as _io- import time as _clk- import gc as _mem- _io.setswitchinterval(1.0)- _log = lambda s: print(s, file=_io.stderr, flush=True)+ import torch, triton, triton.language as tl+ import sys as _J, time as _K, gc as _H+ torch.set_grad_enabled(False); _H.disable(); _J.setswitchinterval(1.0)+ _pr = lambda m: print(m, file=_J.stderr, flush=True)+ # total shader engines on this part+ _SE = 256- # ---- Override heuristics to avoid dynamic tile selection ----- try:- _gemm_a16wfp4_preshuffle_kernel.values['GRID_MN'] = lambda args: 1- _gemm_a16wfp4_preshuffle_kernel.values['EVEN_K'] = lambda args: True- _log("[setup] heuristics locked")- except Exception as _exc:- _log(f"[setup] heuristic override failed: {_exc}")- _env.environ["HIP_FORCE_DEV_KERNARG"] = "1"--- # ---- Inject ISA-level BF16->FP4 quantization ----- _log("[setup] patching quantizer with hardware conversion...")+ # ==============================================================+ # Native HIP lane merger — 128-bit vectorized K-stripe reducer+ # Compiles at import time; Triton fallback if unavailable+ # ==============================================================+ _HSRC = r"""+ #include <hip/hip_runtime.h>+ __device__ __forceinline__ unsigned short narrow(float x){+ unsigned int w; __builtin_memcpy(&w,&x,sizeof(w));+ return (unsigned short)((w+((w>>16)&1)+0x7FFFu)>>16);}+ template<int S> __global__ void wide_fold(const float*__restrict__ p,+ unsigned short*__restrict__ d, int n){+ int i=(blockIdx.x*blockDim.x+threadIdx.x)*4;+ if(i+3<n){float4 a=*reinterpret_cast<const float4*>(p+i);+ #pragma unroll+ for(int s=1;s<S;s++){float4 v=*reinterpret_cast<const float4*>(p+s*n+i);+ a.x+=v.x;a.y+=v.y;a.z+=v.z;a.w+=v.w;}+ *reinterpret_cast<unsigned long long*>(d+i)=+ (unsigned long long)narrow(a.x)|((unsigned long long)narrow(a.y)<<16)|+ ((unsigned long long)narrow(a.z)<<32)|((unsigned long long)narrow(a.w)<<48);+ }else{for(int j=i;j<n&&j<i+4;j++){float a=p[j];+ #pragma unroll+ for(int s=1;s<S;s++)a+=p[s*n+j]; d[j]=narrow(a);}}}+ __global__ void thin_fold(const float*__restrict__ p,unsigned short*__restrict__ d,int n,int s){+ int g=blockIdx.x*blockDim.x+threadIdx.x;+ if(g<n){float a=p[g]; for(int k=1;k<s;k++)a+=p[k*n+g]; d[g]=narrow(a);}}+ void run_fold(torch::Tensor p,torch::Tensor d,int r,int c,int s){+ int n=r*c; auto*sp=p.data_ptr<float>();+ auto*dp=reinterpret_cast<unsigned short*>(d.data_ptr());+ int t=64,bl=(n+t*4-1)/(t*4);+ switch(s){+ case 2:wide_fold<2><<<bl,t>>>(sp,dp,n);break;+ case 3:wide_fold<3><<<bl,t>>>(sp,dp,n);break;+ case 4:wide_fold<4><<<bl,t>>>(sp,dp,n);break;+ case 7:wide_fold<7><<<bl,t>>>(sp,dp,n);break;+ case 8:wide_fold<8><<<bl,t>>>(sp,dp,n);break;+ case 14:wide_fold<14><<<bl,t>>>(sp,dp,n);break;+ default:{int t2=256; thin_fold<<<(n+t2-1)/t2,t2>>>(sp,dp,n,s);break;}}}+ """+ _HH = "void run_fold(torch::Tensor p,torch::Tensor d,int r,int c,int s);"+ _GOT_HIP = Falsetry:- _kern_obj = (- _gemm_a16wfp4_preshuffle_kernel.fn- if hasattr(_gemm_a16wfp4_preshuffle_kernel, 'fn')- else _gemm_a16wfp4_preshuffle_kernel- )- _quant_ref = _kern_obj.__globals__['_mxfp4_quant_op']+ from torch.utils.cpp_extension import load_inline as _cc+ _c0 = _K.time()+ _hip_fold = _cc(name="nf4m", cpp_sources=[_HH], cuda_sources=[_HSRC],+ functions=["run_fold"], verbose=False,+ extra_cuda_cflags=["--offload-arch=gfx950","-O3"])+ _GOT_HIP = True; _pr(f"[hip] fold ready ({_K.time()-_c0:.1f}s)")+ except Exception as _e: _pr(f"[hip] fold unavail: {_e}")- _patched_body = '''def _mxfp4_quant_op(- x,- BLOCK_SIZE_N,- BLOCK_SIZE_M,- MXFP4_QUANT_BLOCK_SIZE,- ):- """ISA-accelerated BF16 to packed FP4 via v_cvt_scalef32_pk_fp4_bf16."""- NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE- HALF_BLOCK: tl.constexpr = MXFP4_QUANT_BLOCK_SIZE // 2- x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)+ # ==============================================================+ # Register-level BF16 -> packed FP4 narrowing+ # Each group of 32 elements shares one E8M0 scale byte.+ # The gfx950 ISA instruction converts a pair of BF16 values+ # into a single byte holding two FP4 nibbles.+ # ==============================================================+ @triton.jit+ def _narrow_to_fp4(raw, NR: tl.constexpr, NC: tl.constexpr):+ GWIDTH: tl.constexpr = 32+ NGRP: tl.constexpr = NC // GWIDTH- 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+ f32 = raw.to(tl.float32).reshape(NR, NGRP, GWIDTH)- amax_exp = (amax >> 23) & 0xFF- scale_e8m0_unbiased = (amax_exp.to(tl.int32) - 129).to(tl.float32)- scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)+ # find per-group peak, snap to power-of-two via bit rounding+ pk = tl.max(tl.abs(f32), axis=-1, keep_dims=True)+ pk = pk.to(tl.int32, bitcast=True)+ pk = (pk + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000- bs_e8m0 = (scale_e8m0_unbiased + 127).to(tl.float32).to(tl.uint8)+ # derive unbiased exponent for the E8M0 scale encoding+ be = ((pk >> 23) & 0xFF).to(tl.int32) - 127+ ue = be - 2+ ue = tl.minimum(tl.maximum(ue, -127), 127)+ sc_byte = ue.to(tl.uint8) + 127- biased_exp_f = tl.maximum(scale_e8m0_unbiased + 127.0, 1.0)- hw_scale = (biased_exp_f.to(tl.int32).to(tl.uint32) << 23).to(tl.float32, bitcast=True)+ # build the IEEE754 float that the hw divider expects+ dv = (ue.to(tl.int32) + 127).to(tl.uint32) << 23+ dv = dv.to(tl.float32, bitcast=True)- x_bf16 = x.to(tl.bfloat16)- x_pairs = x_bf16.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, HALF_BLOCK, 2)- evens, odds = tl.split(x_pairs)- lo = evens.to(tl.uint16, bitcast=True).to(tl.uint32)- hi = odds.to(tl.uint16, bitcast=True).to(tl.uint32)- packed_bf16 = lo | (hi << 16)+ # broadcast divider to every element pair+ dv = tl.broadcast_to(dv, (NR, NGRP, GWIDTH)).reshape(NR, NC)+ dv = dv.reshape(NR, NC // 2, 2)+ dv_l, _ = tl.split(dv)+ dv_f = dv_l.reshape(NR, NC // 2)- result = tl.inline_asm_elementwise(+ # pack adjacent bf16 into u32 words for the converter+ u16 = raw.to(tl.uint16, bitcast=True).reshape(NR, NC // 2, 2)+ w0, w1 = tl.split(u16)+ pw = w0.to(tl.uint32) | (w1.to(tl.uint32) << 16)+ pw = pw.reshape(NR, NC // 2)++ # one ISA op produces both nibbles per pair+ nib = tl.inline_asm_elementwise("v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",- "=v,v,v",- [packed_bf16, hw_scale],- dtype=tl.uint32,- is_pure=True,- pack=1,- )+ "=v, v, v", [pw, dv_f],+ dtype=tl.uint32, is_pure=True, pack=1)+ out = (nib & 0xFF).to(tl.uint8).reshape(NR, NC // 2)+ return out, sc_byte.reshape(NR, NGRP)- x_fp4 = (result & 0xFF).to(tl.uint8)- x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)- return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)- '''+ # ==============================================================+ # Fused narrowing + scaled dot product engine+ # Memory policy: L2-pin activations, write-through results,+ # cache-global for weight operands. Relaxed FP math enabled.+ # ==============================================================+ @triton.heuristics({+ "KDIV_EXACT": lambda args: (args["K"] % (args["TC"] // 2) == 0)+ and (args["KSTRIPE"] % args["TC"] == 0)+ and (args["K"] % (args["KSTRIPE"] // 2) == 0),+ })+ @triton.jit+ def _tiled_gemm_engine(+ xp, wp, yp, sp,+ M, N, K,+ sx_r, sx_c, sw_r, sw_c,+ sy_s, sy_r, sy_c, ss_r, ss_c,+ TR: tl.constexpr, TN: tl.constexpr, TC: tl.constexpr,+ GSWIZ: tl.constexpr, NSTRIPE: tl.constexpr, KSTRIPE: tl.constexpr,+ KDIV_EXACT: tl.constexpr,+ num_warps: tl.constexpr, num_stages: tl.constexpr,+ waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,+ ld_policy: tl.constexpr,+ ):+ # stride positivity hints for the LLVM backend+ tl.assume(sx_r > 0); tl.assume(sx_c > 0)+ tl.assume(sw_r > 0); tl.assume(sw_c > 0)+ tl.assume(sy_r > 0); tl.assume(sy_c > 0)+ tl.assume(ss_r > 0); tl.assume(ss_c > 0)- if hasattr(_quant_ref, '_unsafe_update_src'):- _quant_ref._unsafe_update_src(_patched_body)- else:- _quant_ref._src = _patched_body- if hasattr(_quant_ref, 'src'):- _quant_ref.src = _patched_body- if hasattr(_quant_ref, 'hash'):- _quant_ref.hash = None+ SG: tl.constexpr = 32 # microscale group width+ nr = tl.cdiv(M, TR); nc = tl.cdiv(N, TN)- _orig_src = _kern_obj._src- _mod_src = _orig_src.replace(- 'accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")',- 'accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", acc=accumulator)'- )- if _mod_src != _orig_src:- _kern_obj._unsafe_update_src(_mod_src)- _log("[setup] quantizer + kernel patched OK")- else:- _log("[setup] quantizer patched, kernel string mismatch")+ # decompose flat program id into stripe + tile coordinates+ pid = tl.program_id(axis=0)+ sid = pid % NSTRIPE+ wid = pid // NSTRIPE- _check = _quant_ref._src if hasattr(_quant_ref, '_src') else ''- _log(f"[setup] has hw asm: {'inline_asm_elementwise' in _check}")- except Exception as _exc:- import traceback- _log(f"[setup] patch FAILED: {_exc}")- traceback.print_exc(file=_io.stderr)--- # ---- Native HIP accumulator merger for K-partitioned runs ----- _MERGER_HIP = r"""- #include <hip/hip_runtime.h>-- __device__ __forceinline__ unsigned short to_bf16(float val) {- unsigned int raw;- __builtin_memcpy(&raw, &val, sizeof(raw));- unsigned int bias = ((raw >> 16) & 1) + 0x7FFFu;- return (unsigned short)((raw + bias) >> 16);- }-- template <int NP>- __global__ void vec4_sum(const float* __restrict__ src,- unsigned short* __restrict__ dst, int len) {- int base = (blockIdx.x * blockDim.x + threadIdx.x) * 4;- if (base + 3 < len) {- float4 acc = *reinterpret_cast<const float4*>(src + base);- #pragma unroll- for (int p = 1; p < NP; p++) {- float4 part = *reinterpret_cast<const float4*>(src + p * len + base);- acc.x += part.x; acc.y += part.y; acc.z += part.z; acc.w += part.w;- }- unsigned short a = to_bf16(acc.x), b = to_bf16(acc.y);- unsigned short c = to_bf16(acc.z), d = to_bf16(acc.w);- *reinterpret_cast<unsigned long long*>(dst + base) =- (unsigned long long)a | ((unsigned long long)b << 16) |- ((unsigned long long)c << 32) | ((unsigned long long)d << 48);- } else {- for (int j = base; j < len && j < base + 4; j++) {- float acc = src[j];- #pragma unroll- for (int p = 1; p < NP; p++) acc += src[p * len + j];- dst[j] = to_bf16(acc);- }- }- }-- __global__ void scalar_sum(const float* __restrict__ src,- unsigned short* __restrict__ dst,- int len, int np) {- int gid = blockIdx.x * blockDim.x + threadIdx.x;- if (gid < len) {- float acc = src[gid];- for (int p = 1; p < np; p++) acc += src[p * len + gid];- dst[gid] = to_bf16(acc);- }- }-- void merge_partials(torch::Tensor src, torch::Tensor dst, int R, int C, int np) {- int len = R * C;- const float* sp = src.data_ptr<float>();- unsigned short* dp = reinterpret_cast<unsigned short*>(dst.data_ptr());- const int thr = 64, stride = thr * 4;- const int nblk = (len + stride - 1) / stride;- switch (np) {- case 2: vec4_sum<2><<<nblk, thr>>>(sp, dp, len); break;- case 3: vec4_sum<3><<<nblk, thr>>>(sp, dp, len); break;- case 4: vec4_sum<4><<<nblk, thr>>>(sp, dp, len); break;- case 7: vec4_sum<7><<<nblk, thr>>>(sp, dp, len); break;- case 8: vec4_sum<8><<<nblk, thr>>>(sp, dp, len); break;- default: {- const int t2 = 256, b2 = (len + t2 - 1) / t2;- scalar_sum<<<b2, t2>>>(sp, dp, len, np);- break;- }- }- }- """- _MERGER_HDR = "void merge_partials(torch::Tensor src, torch::Tensor dst, int R, int C, int np);"-- _HAS_HIP_MERGER = False- try:- from torch.utils.cpp_extension import load_inline as _jit- _jit_t0 = _clk.time()- _hip_merger = _jit(- name="fp4_kmerge",- cpp_sources=[_MERGER_HDR],- cuda_sources=[_MERGER_HIP],- functions=["merge_partials"],- verbose=False,- extra_cuda_cflags=["--offload-arch=gfx950", "-O3"],- )- _HAS_HIP_MERGER = True- _log(f"[setup] HIP merger ready ({_clk.time()-_jit_t0:.1f}s)")- except Exception as _exc:- _log(f"[setup] HIP merger unavailable: {_exc}")--- # ---- Partition alignment utility ----- def _align_parts(kh, bk, np):- span = triton.cdiv((2 * triton.cdiv(kh, np)), bk) * bk- while np > 1 and bk > 16:- ok = (kh % (span // 2) == 0 and span % bk == 0 and kh % (bk // 2) == 0)- if ok:- break- elif kh % (span // 2) != 0 and np > 1:- np //= 2- elif span % bk != 0:- np = np // 2 if np > 1 else np- if np <= 1 and bk > 16:- bk //= 2- elif kh % (bk // 2) != 0 and bk > 16:- bk //= 2- else:- break- span = triton.cdiv((2 * triton.cdiv(kh, np)), bk) * bk- return span, bk, np--- # ---- Occupancy-driven config resolver ----- _resolved = {}-- def _shape_config(batch, cols, depth):- tag = (batch, cols, depth)- if tag in _resolved:- return _resolved[tag]- kh = depth // 2-- if batch <= 32:- bm, bn = 8, 128- wave_est = ((batch + bm - 1) // bm) * ((cols + 127) // 128)- np = 1- if depth >= 4096:- np = 7- elif depth >= 2048:- np = 2 if (wave_est * 2 >= (_ENGINES * 3) // 4 and wave_est * 2 <= _ENGINES) else 4- elif depth >= 1536:- np = 2 if (wave_est * 2 >= (_ENGINES * 3) // 4 and wave_est * 2 <= _ENGINES) else 3- bk = 256 if depth <= np * 512 or (np == 2 and depth <= np * 1024) else 512- if wave_est * np < (_ENGINES * 3) // 4:- bn = 64- total_wg = ((batch + bm - 1) // bm) * ((cols + bn - 1) // bn) * np- wpe = 2 if total_wg > _ENGINES else 1+ # single-stripe: grouped swizzle for L2 locality+ # multi-stripe: simple linear mapping+ if NSTRIPE == 1:+ gw = GSWIZ * nc+ gi = wid // gw+ gr = gi * GSWIZ+ gl = min(nr - gr, GSWIZ)+ ri = gr + ((wid % gw) % gl)+ ci = (wid % gw) // glelse:- bm = 16- if batch <= 128:- est16 = ((batch + 15) // 16) * ((cols + 127) // 128)- if est16 < (_ENGINES * 3) // 4:- bm = 8- wave_est = ((batch + bm - 1) // bm) * ((cols + 127) // 128)- bn, np = 128, 1- if _ENGINES // 2 <= wave_est <= _ENGINES and (depth >= 7168 or (depth >= 2048 and bm == 8)):- np = 2- elif wave_est < _ENGINES // 2 and depth > 512:- if depth >= 4096:- np = 2 if wave_est * 2 >= _ENGINES else 7- elif depth >= 2048:- np = 2- elif depth >= 1536:- np = 3- bk = 256 if depth <= max(np * 4096, 2048) else 512- if wave_est * np < (_ENGINES * 3) // 4:- bn = 64- total_wg = ((batch + bm - 1) // bm) * ((cols + bn - 1) // bn) * np- wpe = 2 if total_wg > _ENGINES else 1+ ri = wid // nc+ ci = wid % nc+ tl.assume(ri >= 0); tl.assume(ci >= 0); tl.assume(sid >= 0)- params = {- "BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": max(bn, 32), "BLOCK_SIZE_K": bk,- "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,- "waves_per_eu": wpe, "matrix_instr_nonkdim": 16,- "cache_modifier": ".cg", "NUM_KSPLIT": np,- }+ # guard: skip if this stripe starts past K boundary+ if (sid * KSTRIPE // 2) < K:+ nloop = tl.cdiv(KSTRIPE // 2, TC // 2)- if params["NUM_KSPLIT"] > 1:- span, bk2, np2 = _align_parts(kh, params["BLOCK_SIZE_K"], params["NUM_KSPLIT"])- params["SPLITK_BLOCK_SIZE"] = span- params["BLOCK_SIZE_K"] = bk2- params["NUM_KSPLIT"] = np2+ # activation tile: BF16 input rows+ ro = (ri * TR + tl.arange(0, TR)) % M+ co = sid * KSTRIPE + tl.arange(0, TC)+ ap = xp + (ro[:, None] * sx_r + co[None, :] * sx_c)- if params["BLOCK_SIZE_K"] >= 2 * kh:- params["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * kh)- params["SPLITK_BLOCK_SIZE"] = 2 * kh- params["NUM_KSPLIT"] = 1- params["BLOCK_SIZE_N"] = max(params["BLOCK_SIZE_N"], 32)+ # weight tile: pre-permuted FP4x2 blocks+ lr = tl.arange(0, (TC // 2) * 16)+ lo = sid * (KSTRIPE // 2) * 16 + lr+ wr = (ci * (TN // 16) + tl.arange(0, TN // 16)) % (N // 16)+ bpp = wp + (wr[:, None] * sw_r + lo[None, :] * sw_c)- if params["NUM_KSPLIT"] == 1:- params["SPLITK_BLOCK_SIZE"] = 2 * kh+ # weight scale tile: shuffled E8M0 blocks+ sr_off = (ci * TN + tl.arange(0, TN // 32) * 32)+ sc_off = (sid * (KSTRIPE // SG) * 32) + tl.arange(0, TC // SG * 32)+ spp = sp + sr_off[:, None] * ss_r + sc_off[None, :] * ss_c- real_np, padded_np = None, None- if params["NUM_KSPLIT"] > 1:- real_np = triton.cdiv(kh, params["SPLITK_BLOCK_SIZE"] // 2)- padded_np = triton.next_power_of_2(params["NUM_KSPLIT"])+ # initialize tile accumulator+ dot = tl.zeros((TR, TN), dtype=tl.float32)- m_tiles = triton.cdiv(batch, params["BLOCK_SIZE_M"])- n_tiles = triton.cdiv(cols, params["BLOCK_SIZE_N"])- launch_grid = (params["NUM_KSPLIT"] * m_tiles * n_tiles,)- red_grid = None- if params["NUM_KSPLIT"] > 1:- red_grid = (triton.cdiv(batch, 16), triton.cdiv(cols, 16))+ for step in range(sid * nloop, (sid + 1) * nloop):+ # fire all loads first for maximum MLP+ if KDIV_EXACT:+ xblk = tl.load(ap, eviction_policy="evict_last")+ sblk = tl.load(spp, cache_modifier=ld_policy)+ wblk = tl.load(bpp, cache_modifier=ld_policy)+ else:+ koff = (step - sid * nloop) * TC+ xblk = tl.load(ap,+ mask=tl.arange(0, TC)[None, :] < (2*K - sid*KSTRIPE - koff),+ other=0.0, eviction_policy="evict_last")+ sblk = tl.load(spp, cache_modifier=ld_policy)+ wblk = tl.load(bpp,+ mask=lr[None, :] < ((K - (sid*(KSTRIPE//2) + (step-sid*nloop)*(TC//2)))*16),+ other=0, cache_modifier=ld_policy)- bundle = (- params, real_np, padded_np, launch_grid, red_grid,- kh, params["BLOCK_SIZE_M"], params["BLOCK_SIZE_N"],- params["BLOCK_SIZE_K"], params["NUM_KSPLIT"],- params["SPLITK_BLOCK_SIZE"], params["waves_per_eu"],- )- _resolved[tag] = bundle- return bundle+ # narrow activations in register file+ xn, xs = _narrow_to_fp4(xblk, TR, TC)+ # undo the permutation on weight scales+ ws = (sblk+ .reshape(TN//32, TC//SG//8, 4, 16, 2, 2, 1)+ .permute(0, 5, 3, 1, 4, 2, 6)+ .reshape(TN, TC // SG))- # ---- Phased JIT warmup ----- _warm_t0 = _clk.time()- _warmed = {}- _phase1 = {}- _phase2 = {}- _red_set = set()+ # undo the permutation on weight data+ wt = (wblk+ .reshape(1, TN//16, TC//64, 2, 16, 16)+ .permute(0, 1, 4, 2, 3, 5)+ .reshape(TN, TC // 2).trans(1, 0))- for _dn, _dk in _DIM_PAIRS:- for _db in _BATCH_DIMS:- _cb, _rn, _rp, _, _, _, _, _, _, _, _, _ = _shape_config(_db, _dn, _dk)- _sig = (- _cb["BLOCK_SIZE_M"], _cb["BLOCK_SIZE_N"], _cb["BLOCK_SIZE_K"],- _cb["NUM_KSPLIT"], _cb["SPLITK_BLOCK_SIZE"], _cb["waves_per_eu"],- )- if _db <= 32 and _dk >= 1536:- _phase1.setdefault(_sig, True)- else:- _phase2.setdefault(_sig, True)- if _rn is not None:- _red_set.add((_rn, _rp))+ # scaled dot with relaxed precision for better scheduling+ dot = tl.dot_scaled(xn, xs, "e2m1", wt, ws, "e2m1", dot, fast_math=True)- for _s in _phase1:- _phase2.pop(_s, None)+ # advance pointers by one K-block+ ap += TC * sx_c+ bpp += (TC // 2) * 16 * sw_c+ spp += TC * ss_c- _log(f"[warm] {len(_phase1)} phase1 + {len(_phase2)} phase2, {len(_red_set)} reducers")+ # write-through to avoid polluting L2 with output data+ res = dot.to(yp.type.element_ty)+ yr = ri * TR + tl.arange(0, TR).to(tl.int64)+ yc = ci * TN + tl.arange(0, TN).to(tl.int64)+ ypp = yp + sy_r * yr[:, None] + sy_c * yc[None, :] + sid * sy_s+ ymask = (yr[:, None] < M) & (yc[None, :] < N)+ tl.store(ypp, res, mask=ymask, cache_modifier=".wt")- _dummy_x = torch.zeros(32, 8192, dtype=torch.bfloat16, device="cuda")- _dummy_w = torch.zeros(16, 65536, dtype=torch.uint8, device="cuda")- _dummy_s = torch.zeros(16, 65536, dtype=torch.uint8, device="cuda")- _dummy_pp = torch.zeros(16, 32, 256, dtype=torch.float32, device="cuda")- _dummy_y = torch.zeros(32, 256, dtype=torch.bfloat16, device="cuda")- def _fire_warmup(bm, bn, bk, ks, spk, wpe):- cfg = {- "BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk,- "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,- "waves_per_eu": wpe, "matrix_instr_nonkdim": 16,- "cache_modifier": ".cg", "NUM_KSPLIT": ks, "SPLITK_BLOCK_SIZE": spk,- }- target = _dummy_pp if ks > 1 else _dummy_y- _gemm_a16wfp4_preshuffle_kernel[(max(ks, 1),)](- _dummy_x, _dummy_w, target, _dummy_s, bm, bn, spk // 2,- _dummy_x.stride(0), _dummy_x.stride(1),- _dummy_w.stride(0), _dummy_w.stride(1),- 0 if ks <= 1 else _dummy_pp.stride(0),- _dummy_y.stride(0) if ks <= 1 else _dummy_pp.stride(1),- _dummy_y.stride(1) if ks <= 1 else _dummy_pp.stride(2),- _dummy_s.stride(0), _dummy_s.stride(1),- PREQUANT=True, **cfg,- )+ # ==============================================================+ # Triton-based stripe reducer (fallback when HIP unavailable)+ # ==============================================================+ @triton.jit+ def _fold_partials(+ fp, dp, M, N, s_fs, s_fr, s_fc, s_dr, s_dc,+ FR: tl.constexpr, FC: tl.constexpr,+ REAL: tl.constexpr, PAD: tl.constexpr,+ ):+ pr = tl.program_id(0); pc = tl.program_id(1)+ ro = (pr * FR + tl.arange(0, FR)) % M+ co = (pc * FC + tl.arange(0, FC)) % N+ bp = fp + (ro[:, None] * s_fr) + (co[None, :] * s_fc)+ acc = tl.load(bp).to(tl.float32)+ for i in tl.static_range(1, PAD):+ if i < REAL: acc += tl.load(bp + i * s_fs).to(tl.float32)+ dp2 = dp + (ro[:, None] * s_dr) + (co[None, :] * s_dc)+ tl.store(dp2, acc.to(dp.type.element_ty))- _log("[warm] phase 1 (no lsr)...")- for _sig in sorted(_phase1):- try:- _fire_warmup(*_sig)- _warmed[_sig] = 1- _log(f" {_sig[0]}x{_sig[1]}x{_sig[2]} ks={_sig[3]} ({_clk.time()-_warm_t0:.0f}s)")- except Exception as _exc:- _log(f" {_sig}: ERR {_exc}")- _env.environ["DISABLE_LLVM_OPT"] = "disable-lsr"- _log(f"[warm] phase 2 (with lsr) @ {_clk.time()-_warm_t0:.0f}s")+ # ==============================================================+ # K-stripe alignment fixer+ # ==============================================================+ def _fix_stripe(kh, tc, ns):+ ks = triton.cdiv((2 * triton.cdiv(kh, ns)), tc) * tc+ while ns > 1 and tc > 16:+ if kh%(ks//2)==0 and ks%tc==0 and kh%(tc//2)==0: break+ elif kh%(ks//2)!=0 and ns>1: ns //= 2+ elif ks%tc!=0:+ if ns>1: ns //= 2+ elif tc>16: tc //= 2+ elif kh%(tc//2)!=0 and tc>16: tc //= 2+ else: break+ ks = triton.cdiv((2 * triton.cdiv(kh, ns)), tc) * tc+ ns = triton.cdiv(kh, ks // 2)+ return ks, tc, ns- for _idx, _sig in enumerate(sorted(_phase2)):- if _clk.time() - _warm_t0 > 200:- _log(f" timeout, {len(_phase2) - _idx} skipped")- break- try:- _fire_warmup(*_sig)- _warmed[_sig] = 2- _log(f" {_sig[0]}x{_sig[1]}x{_sig[2]} ks={_sig[3]} ({_clk.time()-_warm_t0:.0f}s)")- except Exception as _exc:- _log(f" {_sig}: ERR {_exc}")- _log(f"[warm] reducers...")- for _rn, _rp in sorted(_red_set):- if _clk.time() - _warm_t0 > 230:- _log(" timeout")- break- try:- _gemm_afp4wfp4_reduce_kernel[(1, 1)](- _dummy_pp, _dummy_y, 16, 16,- _dummy_pp.stride(0), _dummy_pp.stride(1), _dummy_pp.stride(2),- _dummy_y.stride(0), _dummy_y.stride(1), 16, 16, _rn, _rp,- )- except Exception:- pass+ # ==============================================================+ # Profiled tile geometries from offline sweep+ # ==============================================================+ _SWEEP = {+ (4,2880,512): (4,128,256,1,4,2,1,16,None,1),+ (16,2112,7168): (16,128,512,1,4,2,3,16,".cg",14),+ (32,4096,512): (16,32,256,1,4,3,3,16,".cg",1),+ (32,2880,512): (8,128,256,1,4,2,2,16,None,1),+ (64,7168,2048): (16,128,256,1,4,2,2,16,".cg",1),+ (256,3072,1536): (16,256,512,1,8,2,2,16,None,1),+ }+ # unpack: (TR, TN, TC, GSWIZ, nw, ns, wpe, mi, cm, nstripe)- del _dummy_x, _dummy_w, _dummy_s, _dummy_pp, _dummy_y, _fire_warmup- del _phase1, _phase2, _red_set- torch.cuda.empty_cache()- _log(f"[warm] done: {len(_warmed)} configs in {_clk.time()-_warm_t0:.0f}s")+ def _from_sweep(t):+ return {"TR":t[0],"TN":t[1],"TC":t[2],"GSWIZ":t[3],+ "num_warps":t[4],"num_stages":t[5],"waves_per_eu":t[6],+ "matrix_instr_nonkdim":t[7],"ld_policy":t[8],"NSTRIPE":t[9]}- _mem.disable()+ # ==============================================================+ # Occupancy model — analytical fallback for unknown shapes+ # ==============================================================+ def _model_cfg(m, n, k):+ if m <= 32:+ tr, tn = 8, 128+ we = ((m+tr-1)//tr)*((n+127)//128); ns = 1+ if k>=4096: ns=7+ elif k>=2048: ns = 4 if we*2<(_SE*3)//4 else 2+ elif k>=1536: ns = 3 if we*2<(_SE*3)//4 else 2+ tc = 256 if k<=ns*512 else 512+ if we*ns < (_SE*3)//4: tn = 64+ else:+ tr = 16+ if m<=128 and ((m+15)//16)*((n+127)//128) < (_SE*3)//4: tr = 8+ we = ((m+tr-1)//tr)*((n+127)//128); tn, ns = 128, 1+ if _SE//2<=we<=_SE and (k>=7168 or (k>=2048 and tr==8)): ns=2+ elif we<_SE//2 and k>512:+ if k>=4096: ns = 2 if we*2>=_SE else 7+ elif k>=2048: ns=2+ elif k>=1536: ns=3+ tc = 256 if k<=max(ns*4096,2048) else 512+ if we*ns < (_SE*3)//4: tn = 64+ return {"TR":tr,"TN":max(tn,32),"TC":tc,"GSWIZ":1,+ "num_warps":4,"num_stages":2,"waves_per_eu":2,+ "matrix_instr_nonkdim":16,"ld_policy":".cg","NSTRIPE":ns}- # ---- Runtime dispatch state ----- _wt_cache = {}- _dest_buf = {}- _frag_buf = {}- _seen = set()+ # ==============================================================+ # Cached config resolver + launch parameter precomputation+ # ==============================================================+ _RC = {}; _YC = {}; _FC = {}; _WC = {}; _LC = {}- def _prepare_wt(data):- addr = data[3].data_ptr()- if addr not in _wt_cache:- n_dim = data[3].shape[0]- k_bytes = data[3].shape[1]- sr, sc = data[4].shape- n_grp = n_dim // 32- w_view = data[3].view(torch.uint8).reshape(n_dim // 16, k_bytes * 16)- s_view = data[4].view(torch.uint8).reshape(sr // 32, sc * 32)[:n_grp].contiguous()- _wt_cache[addr] = (w_view, s_view, w_view.stride(0), s_view.stride(0))- return _wt_cache[addr]+ def _alloc_y(m, n, ns, dev):+ t = (m,n,ns)+ if t not in _YC:+ _YC[t] = (torch.empty((m,n),dtype=torch.bfloat16,device=dev),+ torch.empty((ns,m,n),dtype=torch.float32,device=dev) if ns>1 else None)+ return _YC[t]+ def _cfg(m, n, k):+ t = (m,n,k)+ if t in _RC: return _RC[t]+ c = _from_sweep(_SWEEP[t]) if t in _SWEEP else _model_cfg(m, n, k)+ kh = k // 2+ if c["NSTRIPE"] > 1:+ ks, tc2, ns2 = _fix_stripe(kh, c["TC"], c["NSTRIPE"])+ c["KSTRIPE"]=ks; c["TC"]=tc2; c["NSTRIPE"]=ns2+ else:+ c["KSTRIPE"] = 2*kh; c["NSTRIPE"] = 1+ if c["TC"] >= 2*kh:+ c["TC"]=triton.next_power_of_2(2*kh); c["KSTRIPE"]=2*kh; c["NSTRIPE"]=1+ c["TN"]=max(c["TN"],32)+ _RC[t] = c+ return c- def custom_kernel(data: input_t) -> output_t:- X = data[0]- if not X.is_contiguous():- X = X.contiguous()- ndims = X.ndim- X_flat = X if ndims == 2 else X.view(-1, X.shape[-1])- batch = X_flat.shape[0]- cols = data[3].shape[0]- depth = data[3].shape[1] * 2+ def _get_wt(wd, ws, n, kh):+ a = wd.data_ptr()+ if a not in _WC:+ _WC[a] = (wd.view(torch.uint8).reshape(n//16, kh*16), ws.view(torch.uint8))+ return _WC[a]- (params, real_np, padded_np, launch_grid, red_grid,- kh, bm, bn, bk, ks, spk, wpe) = _shape_config(batch, cols, depth)+ def _launch(m, n, k, dev):+ t = (m,n,k)+ if t in _LC: return _LC[t]+ c = _cfg(m,n,k); kh=k//2; ns=c["NSTRIPE"]+ y,pp = _alloc_y(m,n,ns,dev)+ g = (ns * triton.cdiv(m,c["TR"]) * triton.cdiv(n,c["TN"]),)+ if ns==1: s_s,s_r,s_c = 0,y.stride(0),y.stride(1)+ else: s_s,s_r,s_c = pp.stride(0),pp.stride(1),pp.stride(2)+ b = {'c':c,'kh':kh,'g':g,'ns':ns,'ss':s_s,'sr':s_r,'sc':s_c}+ if ns>1:+ b['rg']=(triton.cdiv(m,16),triton.cdiv(n,64))+ b['rn']=triton.cdiv(kh,c["KSTRIPE"]//2)+ b['rp']=triton.next_power_of_2(ns)+ _LC[t] = b+ return b- tag = (batch, cols, depth)- if tag not in _seen:- _seen.add(tag)- _log(f"[run] {batch}x{cols}x{depth} bm={bm} bn={bn} bk={bk} ks={ks} wpe={wpe}")- okey = (batch, cols)- if okey not in _dest_buf:- _dest_buf[okey] = torch.empty((batch, cols), dtype=torch.bfloat16, device="cuda")- dest = _dest_buf[okey]-- w_view, s_view, sw0, ss0 = _prepare_wt(data)-- if ks > 1:- fkey = (padded_np, batch, cols)- if fkey not in _frag_buf:- _frag_buf[fkey] = torch.empty(- (padded_np, batch, cols), dtype=torch.float32, device="cuda"- )- frags = _frag_buf[fkey]- stride_p, stride_r = batch * cols, cols- else:- frags = None- stride_p, stride_r = 0, cols-- _gemm_a16wfp4_preshuffle_kernel[launch_grid](- X_flat, w_view,- dest if frags is None else frags,- s_view, batch, cols, kh,- depth, 1, sw0, 1,- stride_p, stride_r, 1,- ss0, 1,- BLOCK_SIZE_M=bm, BLOCK_SIZE_N=bn, BLOCK_SIZE_K=bk,- GROUP_SIZE_M=1, NUM_KSPLIT=ks, SPLITK_BLOCK_SIZE=spk,- num_warps=4, num_stages=2, waves_per_eu=wpe,- matrix_instr_nonkdim=16, cache_modifier=".cg",- PREQUANT=True,- )-- if frags is not None:- if _HAS_HIP_MERGER:- _hip_merger.merge_partials(frags, dest, batch, cols, real_np)+ # ==============================================================+ # Entry point+ # ==============================================================+ def _go(x, wd, ws, m, n, k):+ b = _launch(m, n, k, x.device)+ y, pp = _alloc_y(m, n, b['ns'], x.device)+ wf, sf = _get_wt(wd, ws, n, b['kh'])+ _tiled_gemm_engine[b['g']](+ x, wf, y if b['ns']==1 else pp, sf,+ m, n, b['kh'],+ x.stride(0), x.stride(1), wf.stride(0), wf.stride(1),+ b['ss'], b['sr'], b['sc'], sf.stride(0), sf.stride(1),+ **b['c'])+ if b['ns'] > 1:+ if _GOT_HIP:+ _hip_fold.run_fold(pp, y, m, n, b['rn'])else:- _gemm_afp4wfp4_reduce_kernel[red_grid](- frags, dest, batch, cols,- batch * cols, cols, 1, cols, 1,- 16, 16, real_np, padded_np,- )+ _fold_partials[b['rg']](pp, y, m, n,+ pp.stride(0),pp.stride(1),pp.stride(2),+ y.stride(0),y.stride(1), 16, 64, b['rn'], b['rp'])+ return y- return dest if ndims == 2 else dest.view(*X.shape[:-1], cols)+ def custom_kernel(data: input_t) -> output_t:+ A = data[0]+ return _go(A, data[3], data[4], A.shape[0], data[1].shape[0], A.shape[1])
scrolls · 884 diff lines total
Best evidence level for this revision: reported
JSON