Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
8.10µs
#25 of 1143
2026-04-07

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 = float4if(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 = False
try:
- _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) // gl
else:
- 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