submission 730077
johnny.t.shi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 176 lines, June 9 Researcher Reciprocity License v1.0.
v911_nk14.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-730077?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:693da62df42d122e18b819d997a1ac4b9abc5038e53c71dcc9f72b37564a5491
license declaredunknown
license concludedunknown
authorsjohnny.t.shi
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
"""v911: NK=14 for split-K shape — maximize CU utilization.Kernel source
v911_nk14.py176 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""v911: NK=14 for split-K shape — maximize CU utilization.
NK=7 → 119 blocks (46% CU). NK=14 → 238 blocks (93% CU).
Each block does 1 K-iter instead of 2. More parallelism for memory-bound shape.
Trade: reduce sums 14 partials instead of 7."""
import os
os.environ['DISABLE_LLVM_OPT'] = 'disable-lsr,disable-machine-licm,disable-machine-sink'
os.environ['TRITON_HIP_USE_BLOCK_PINGPONG'] = '1'
import torch, triton, triton.language as tl
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid
from aiter.ops.triton.gluon.gemm_afp4wfp4 import _gemm_afp4wfp4_reduce_kernel as _gluon_reduce_kernel
import aiter.ops.triton.gemm_afp4wfp4 as _gm
_gs=_gm.get_splitk
from task import input_t, output_t
# ===== NX=2 XCD remap helper =====
@triton.jit
def _rx(p, n, NX: tl.constexpr):
cs = tl.cdiv(n, NX)
return (p % NX) * cs + p // NX
# ===== FORKED preshuffle kernel with NX=2 XCD remap =====
@triton.heuristics({
"EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0) and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0) and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
"GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"]) * triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
})
@triton.jit
def _ps_nx2(a_ptr, b_ptr, c_ptr, b_scales_ptr, M, N, K,
stride_am, stride_ak, stride_bn, stride_bk, stride_ck, stride_cm, stride_cn, stride_bsn, stride_bsk,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
EVEN_K: tl.constexpr, num_warps: tl.constexpr, num_stages: tl.constexpr,
waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
GRID_MN: tl.constexpr, PREQUANT: tl.constexpr, cache_modifier: tl.constexpr):
tl.assume(stride_am > 0); tl.assume(stride_ak > 0); tl.assume(stride_bk > 0)
tl.assume(stride_bn > 0); tl.assume(stride_cm > 0); tl.assume(stride_cn > 0)
tl.assume(stride_bsk > 0); tl.assume(stride_bsn > 0)
pid_unified = tl.program_id(axis=0)
if pid_unified < GRID_MN * NUM_KSPLIT:
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
if NUM_KSPLIT == 1:
pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
else:
pid_m = pid // num_pid_n; pid_n = pid % num_pid_n
tl.assume(pid_m >= 0); tl.assume(pid_n >= 0); tl.assume(pid_k >= 0)
SCALE_GROUP_SIZE: tl.constexpr = 32
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak)
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
b_ptrs = b_ptr + (offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk)
offs_bsn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))) % N
offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)
b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
b_scales = (tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE))
if EVEN_K:
a_bf16 = tl.load(a_ptrs)
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
b = (b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2).trans(1, 0))
if PREQUANT:
a, a_scales = _hw_cvt_quant(a_bf16.to(tl.float32), BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
b_scale_ptrs += BLOCK_SIZE_K * stride_bsk
c = accumulator.to(c_ptr.type.element_ty)
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + pid_k * stride_ck
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, c, mask=c_mask)
# ===== v710 kernel for M>32 =====
@triton.jit
def _hw_cvt_quant(x, BLOCK_K: tl.constexpr, BLOCK_M: tl.constexpr, QUANT_BLOCK: tl.constexpr):
NUM_BLOCKS: tl.constexpr = BLOCK_K // QUANT_BLOCK
HALF_QB: tl.constexpr = QUANT_BLOCK // 2
x_3d = x.reshape(BLOCK_M, NUM_BLOCKS, QUANT_BLOCK)
amax = tl.max(tl.abs(x_3d), axis=-1, keep_dims=True)
amax_i32 = amax.to(tl.int32, bitcast=True)
amax_rounded = ((amax_i32 + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000)
amax_exp = (amax_rounded >> 23).to(tl.int32)
su = amax_exp - 129;su = tl.where(su < -126, -126, su);su = tl.where(su > 127, 127, su)
bs_e8m0 = (su + 127).to(tl.uint8)
# Fused CVT scale: pass 2^(su) so hardware does srcExp -= su → FP4(value * 2^(-su))
cvt_exp = su + 127
cvt_scale = (cvt_exp << 23).to(tl.float32, bitcast=True)
cvt_scale_bc = tl.broadcast_to(cvt_scale, (BLOCK_M, NUM_BLOCKS, HALF_QB))
cvt_scale_flat = cvt_scale_bc.reshape(BLOCK_M, BLOCK_K // 2)
# Raw x pairs — NO pre-multiply! Hardware CVT handles scaling.
x_flat = x_3d.reshape(BLOCK_M, BLOCK_K)
x_pairs = x_flat.reshape(BLOCK_M, BLOCK_K // 2, 2)
evens, odds = tl.split(x_pairs)
fp4_packed = tl.inline_asm_elementwise("v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3","=v,v,v,v",args=[evens, odds, cvt_scale_flat],dtype=tl.int32,is_pure=True,pack=1)
x_fp4 = fp4_packed.to(tl.uint8)
return x_fp4.reshape(BLOCK_M, BLOCK_K // 2), bs_e8m0.reshape(BLOCK_M, NUM_BLOCKS)
# ===== Configs =====
def _get_ps_config(M, N, K):
if K > 4096:
# v911: NK=14 → 238 blocks (93% CU) with 1 K-iter each
return {"BLOCK_SIZE_M":16,"BLOCK_SIZE_N":128,"BLOCK_SIZE_K":512,"GROUP_SIZE_M":1,"num_warps":4,"num_stages":2,"waves_per_eu":2,"matrix_instr_nonkdim":16,"cache_modifier":".cg","NUM_KSPLIT":14}
if M <= 4:
return {"BLOCK_SIZE_M":4,"BLOCK_SIZE_N":128,"BLOCK_SIZE_K":256,"GROUP_SIZE_M":1,"num_warps":4,"num_stages":2,"waves_per_eu":0,"matrix_instr_nonkdim":16,"cache_modifier":".cg","NUM_KSPLIT":1}
elif M <= 8:
return {"BLOCK_SIZE_M":8,"BLOCK_SIZE_N":128,"BLOCK_SIZE_K":256,"GROUP_SIZE_M":1,"num_warps":4,"num_stages":2,"waves_per_eu":0,"matrix_instr_nonkdim":16,"cache_modifier":".cg","NUM_KSPLIT":1}
elif K <= 1024:
return {"BLOCK_SIZE_M":8,"BLOCK_SIZE_N":128,"BLOCK_SIZE_K":256,"GROUP_SIZE_M":1,"num_warps":4,"num_stages":2,"waves_per_eu":2,"matrix_instr_nonkdim":16,"cache_modifier":None,"NUM_KSPLIT":1}
elif M <= 32:
return {"BLOCK_SIZE_M":32,"BLOCK_SIZE_N":64,"BLOCK_SIZE_K":512,"GROUP_SIZE_M":1,"num_warps":8,"num_stages":1,"waves_per_eu":2,"matrix_instr_nonkdim":16,"cache_modifier":None,"NUM_KSPLIT":1}
elif M <= 64:
return {"BLOCK_SIZE_M":8,"BLOCK_SIZE_N":128,"BLOCK_SIZE_K":512,"GROUP_SIZE_M":1,"num_warps":4,"num_stages":2,"waves_per_eu":2,"matrix_instr_nonkdim":16,"cache_modifier":".cg","NUM_KSPLIT":1}
else:
# M>64: BSK=256 (v895 confirmed -0.4µs on M=256), BSM=16, nw=4/wpe=2
return {"BLOCK_SIZE_M":16,"BLOCK_SIZE_N":128,"BLOCK_SIZE_K":256,"GROUP_SIZE_M":1,"num_warps":4,"num_stages":2,"waves_per_eu":2,"matrix_instr_nonkdim":16,"cache_modifier":".cg","NUM_KSPLIT":1}
_cf={}
def custom_kernel(data:input_t)->output_t:
A,B,Bq,Bs,Bss=data;M,K=A.shape;N=Bs.shape[0]
key=(M,K,N);c=_cf.get(key)
if c is None:
config = _get_ps_config(M, N, K)
K_kernel = K // 2; BSK = config["BLOCK_SIZE_K"]; BSN = max(config["BLOCK_SIZE_N"], 32); BSM = config["BLOCK_SIZE_M"]
NK = config["NUM_KSPLIT"]
if NK > 1: SBS, BSK, NK = _gs(K_kernel, BSK, NK)
else: SBS = 2 * K_kernel
grid_mn = triton.cdiv(M, BSM) * triton.cdiv(N, BSN)
grid_size = grid_mn * NK
out = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
y_pp = torch.empty((NK, M, N), dtype=torch.float32, device=A.device) if NK > 1 else None
AK = triton.cdiv(K_kernel, (SBS // 2)) if NK > 1 else 1
c = ('ps', out, y_pp, None, None, grid_size, K_kernel, BSM, BSN, BSK,
config["GROUP_SIZE_M"], NK, SBS, config["num_warps"], config["num_stages"],
config["waves_per_eu"], config["matrix_instr_nonkdim"], config["cache_modifier"],
AK, triton.next_power_of_2(NK) if NK > 1 else 1,
(triton.cdiv(M,16), triton.cdiv(N,64)) if NK > 1 else None)
_cf[key] = c
_, out, y_pp, Bw, Bsc, grid_size, K_kernel, BSM, BSN, BSK, GSM, NK, SBS, nw, ns, wpe, mind, cm, AK, MK, rgr = c
b_ptr = Bs.data_ptr()
if Bw is None or _cf.get(('_bp', key)) != b_ptr:
Bw = Bs.view(torch.uint8).reshape(N // 16, (K // 2) * 16)
bs_shape = Bss.shape; Bsc = Bss.view(torch.uint8).reshape(bs_shape[0] // 32, bs_shape[1] * 32)
c_list = list(c); c_list[3] = Bw; c_list[4] = Bsc; c = tuple(c_list); _cf[key] = c
_cf[('_bp', key)] = b_ptr
target = y_pp if NK > 1 else out
sk_off = y_pp.stride(0) if NK > 1 else 0
cm_s = (y_pp.stride(1) if NK > 1 else out.stride(0))
cn_s = (y_pp.stride(2) if NK > 1 else out.stride(1))
_ps_nx2[(grid_size,)](A, Bw, target, Bsc, M, N, K_kernel,
A.stride(0), A.stride(1), Bw.stride(0), Bw.stride(1),
sk_off, cm_s, cn_s, Bsc.stride(0), Bsc.stride(1),
BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN, BLOCK_SIZE_K=BSK,
GROUP_SIZE_M=GSM, NUM_KSPLIT=NK, SPLITK_BLOCK_SIZE=SBS,
num_warps=nw, num_stages=ns, waves_per_eu=wpe,
matrix_instr_nonkdim=mind, PREQUANT=True, cache_modifier=cm)
if NK > 1:
_gluon_reduce_kernel[rgr](y_pp, out, M, N, y_pp.stride(0), y_pp.stride(1), y_pp.stride(2), out.stride(0), out.stride(1), 16, 64, AK, MK)
return out
scrolls · 176 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 553074.
#!POPCORN leaderboard amd-mxfp4-mm#!POPCORN gpu MI355X-- """- v150: v120 + lean quantization op.- Replace _mxfp4_quant_op with a custom Triton JIT function that:- 1. Uses branchless FP4 E2M1 quantization (no 3-way if/else)- 2. Uses bitcast for exp2 (from v138)- 3. Minimizes total VALU instructions- 4. Uses integer-only rounding for max (avoids log2/floor/exp2 chain)- """+ """v911: NK=14 for split-K shape — maximize CU utilization.+ NK=7 → 119 blocks (46% CU). NK=14 → 238 blocks (93% CU).+ Each block does 1 K-iter instead of 2. More parallelism for memory-bound shape.+ Trade: reduce sums 14 partials instead of 7."""+ import os+ os.environ['DISABLE_LLVM_OPT'] = 'disable-lsr,disable-machine-licm,disable-machine-sink'+ os.environ['TRITON_HIP_USE_BLOCK_PINGPONG'] = '1'+ import torch, triton, triton.language as tl+ from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid+ from aiter.ops.triton.gluon.gemm_afp4wfp4 import _gemm_afp4wfp4_reduce_kernel as _gluon_reduce_kernel+ import aiter.ops.triton.gemm_afp4wfp4 as _gm+ _gs=_gm.get_splitkfrom task import input_t, output_t- import torch- import triton- import triton.language as tl- import aiter- from aiter import dtypes- from aiter.ops.gemm_op_a4w4 import get_GEMM_config- from aiter.ops.gemm_op_common import get_padded_m-- import aiter.ops.triton.gemm_afp4wfp4 as _gemm_mod- _reduce_kernel = _gemm_mod._gemm_afp4wfp4_reduce_kernel- _get_splitk_fn = _gemm_mod.get_splitk-- _fp4x2 = dtypes.fp4x2- _fp8_e8m0 = dtypes.fp8_e8m0- _bf16 = dtypes.bf16--+ # ===== NX=2 XCD remap helper =====@triton.jit- def _lean_mxfp4_quant_op(- x, # [BLOCK_M, BLOCK_K] float32- BLOCK_K: tl.constexpr,- BLOCK_M: tl.constexpr,- QUANT_BLOCK: tl.constexpr,- ):- """Lean MXFP4 quantization — fewer VALU instructions, branchless E2M1.- Produces bit-identical output to aiter's _mxfp4_quant_op.- """- NUM_BLOCKS: tl.constexpr = BLOCK_K // QUANT_BLOCK+ def _rx(p, n, NX: tl.constexpr):+ cs = tl.cdiv(n, NX)+ return (p % NX) * cs + p // NX- x_3d = x.reshape(BLOCK_M, NUM_BLOCKS, QUANT_BLOCK)-- # --- Scale computation (integer-only, no log2/floor) ---- amax = tl.max(tl.abs(x_3d), axis=-1, keep_dims=True)- # Round amax up to next power-of-2 (clear mantissa, round up exponent)- amax_i32 = amax.to(tl.int32, bitcast=True)- amax_rounded = ((amax_i32 + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000)- amax_f32 = amax_rounded.to(tl.float32, bitcast=True)-- # Extract exponent directly: E8M0 scale = exponent_of(amax_rounded) - 2 - 127 + 127- # = exponent_field - 2 = (amax_rounded >> 23) - 2- # But we need to handle amax=0 => scale should be 0 (biased), i.e. -127 unbiased- amax_exp = (amax_rounded >> 23).to(tl.int32)- # scale_unbiased = amax_exp - 127 - 2 = amax_exp - 129- scale_unbiased = amax_exp - 129- # Clamp to [-127, 127] — tl.clamp only supports float, use tl.where- scale_unbiased = tl.where(scale_unbiased < -127, -127, scale_unbiased)- scale_unbiased = tl.where(scale_unbiased > 127, 127, scale_unbiased)- # Handle amax=0: amax_exp=0, scale_unbiased=-129 clamped to -127. OK.-- # E8M0 biased scale (uint8)- bs_e8m0 = (scale_unbiased + 127).to(tl.uint8)-- # --- Inverse scale via bitcast (fast exp2) ---- # quant_scale = exp2(-scale_unbiased) = bitcast(((-scale_unbiased) + 127) << 23)- inv_exp = (127 - scale_unbiased)- quant_scale = (inv_exp << 23).to(tl.float32, bitcast=True)-- # --- Quantize: scale input ---- qx = x_3d * quant_scale-- # --- Convert to FP4 E2M1 (branchless) ---- # Extract sign and abs- qx_i32 = qx.to(tl.int32, bitcast=True)- sign_bit = ((qx_i32 >> 31) & 0x8).to(tl.uint8) # sign at bit 3 for FP4- qx_abs = (qx_i32 & 0x7FFFFFFF).to(tl.float32, bitcast=True)-- # Saturate: clamp abs to [0, 6.0] — values >= 6.0 become 0x7 (max E2M1 = 1.5 * 2^2 = 6.0)- # After clamping, all values are in representable range, no saturation branch needed- qx_clamped = tl.minimum(qx_abs, 6.0)- qx_clamped_i32 = qx_clamped.to(tl.int32, bitcast=True)-- # Denormal path: values < 1.0 need special handling- # E2M1 denormals: 0.0 (0b000), 0.5 (0b001)- # Normal E2M1: 1.0 (0b010), 1.5 (0b011), 2.0 (0b100), 3.0 (0b101), 4.0 (0b110), 6.0 (0b111)- #- # For denormals (< 1.0): add magic number to round, extract low bits- denorm_exp: tl.constexpr = (127 - 1) + (23 - 1) + 1- denorm_magic: tl.constexpr = denorm_exp << 23- denorm_magic_f: tl.constexpr = tl.cast(denorm_magic, tl.float32, bitcast=True)- denormal_result = (qx_clamped + denorm_magic_f).to(tl.int32, bitcast=True) - denorm_magic- denormal_result = denormal_result.to(tl.uint8)-- # Normal path (>= 1.0): round to nearest E2M1- # IEEE float32 mantissa has 23 bits, E2M1 mantissa has 1 bit- # So we need to round at bit 22 (keep only 1 mantissa bit)- # Bias adjustment: subtract (127-1) from exponent to get E2M1 exponent- qx_clamped_abs_i32 = qx_clamped_i32- mant_odd = (qx_clamped_abs_i32 >> 22) & 1- val_to_add: tl.constexpr = ((1 - 127) << 23) + (1 << 21) - 1- normal_result = (qx_clamped_abs_i32 + val_to_add + mant_odd) >> 22- normal_result = normal_result.to(tl.uint8)-- # Select: denormal if < 1.0, normal otherwise- is_normal = qx_abs >= 1.0- e2m1 = tl.where(is_normal, normal_result, denormal_result)-- # Apply sign- e2m1 = e2m1 | sign_bit-- # Pack 2 FP4 values per byte- e2m1 = tl.reshape(e2m1, [BLOCK_M, NUM_BLOCKS, QUANT_BLOCK // 2, 2])- evens, odds = tl.split(e2m1)- x_fp4 = evens | (odds << 4)- x_fp4 = x_fp4.reshape(BLOCK_M, BLOCK_K // 2)-- return x_fp4, bs_e8m0.reshape(BLOCK_M, NUM_BLOCKS)--+ # ===== FORKED preshuffle kernel with NX=2 XCD remap =====+ @triton.heuristics({+ "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0) and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0) and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),+ "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"]) * triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),+ })@triton.jit- def _remap_xcd(pid, num_pids, NUM_XCDS: tl.constexpr):- chunk_size = tl.cdiv(num_pids, NUM_XCDS)- xcd = pid % NUM_XCDS- pid_in_xcd = pid // NUM_XCDS- return xcd * chunk_size + pid_in_xcd+ def _ps_nx2(a_ptr, b_ptr, c_ptr, b_scales_ptr, M, N, K,+ stride_am, stride_ak, stride_bn, stride_bk, stride_ck, stride_cm, stride_cn, stride_bsn, stride_bsk,+ BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,+ GROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,+ EVEN_K: tl.constexpr, num_warps: tl.constexpr, num_stages: tl.constexpr,+ waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,+ GRID_MN: tl.constexpr, PREQUANT: tl.constexpr, cache_modifier: tl.constexpr):+ tl.assume(stride_am > 0); tl.assume(stride_ak > 0); tl.assume(stride_bk > 0)+ tl.assume(stride_bn > 0); tl.assume(stride_cm > 0); tl.assume(stride_cn > 0)+ tl.assume(stride_bsk > 0); tl.assume(stride_bsn > 0)- @triton.jit- def _pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M: tl.constexpr):- num_pid_in_group = GROUP_SIZE_M * num_pid_n- group_id = pid // num_pid_in_group- first_pid_m = group_id * GROUP_SIZE_M- group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)- pid_m = first_pid_m + (pid % num_pid_in_group) % group_size_m- pid_n = (pid % num_pid_in_group) // group_size_m- return pid_m, pid_n--- @triton.jit- def _fused_quant_gemm_kernel(- a_ptr, b_ptr, c_ptr, b_scales_ptr,- M, N, K_real,- stride_am, stride_ak,- stride_bk, stride_bn,- stride_ck, stride_cm, stride_cn,- stride_bsn, stride_bsk,- BLOCK_SIZE_M: tl.constexpr,- BLOCK_SIZE_N: tl.constexpr,- BLOCK_SIZE_K: tl.constexpr,- GROUP_SIZE_M: tl.constexpr,- NUM_KSPLIT: tl.constexpr,- SPLITK_BLOCK_SIZE: tl.constexpr,- QUANT_BLOCK: tl.constexpr,- ):- SCALE_GROUP_SIZE: tl.constexpr = 32- K_packed = K_real // 2- GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)- total_pids = GRID_MN * NUM_KSPLIT- total_pids_padded = ((total_pids + 7) // 8) * 8-pid_unified = tl.program_id(axis=0)- pid_unified = _remap_xcd(pid_unified, total_pids_padded, NUM_XCDS=8)- if pid_unified < total_pids:+ if pid_unified < GRID_MN * NUM_KSPLIT:pid_k = pid_unified % NUM_KSPLITpid = pid_unified // NUM_KSPLITnum_pid_m = tl.cdiv(M, BLOCK_SIZE_M)num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)-if NUM_KSPLIT == 1:- pid_m, pid_n = _pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)+ pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)else:- pid_m = pid // num_pid_n- pid_n = pid % num_pid_n+ pid_m = pid // num_pid_n; pid_n = pid % num_pid_n+ tl.assume(pid_m >= 0); tl.assume(pid_n >= 0); tl.assume(pid_k >= 0)- tl.assume(pid_m >= 0)- tl.assume(pid_n >= 0)-- if (pid_k * SPLITK_BLOCK_SIZE) < K_real:- num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE, BLOCK_SIZE_K)-+ SCALE_GROUP_SIZE: tl.constexpr = 32+ if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:+ num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)+ offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)+ offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M- offs_k = pid_k * SPLITK_BLOCK_SIZE + tl.arange(0, BLOCK_SIZE_K)- a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak-- offs_k_packed = pid_k * (SPLITK_BLOCK_SIZE // 2) + tl.arange(0, BLOCK_SIZE_K // 2)- offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N- b_ptrs = b_ptr + offs_k_packed[:, None] * stride_bk + offs_bn[None, :] * stride_bn-- offs_bsn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)) % N- offs_ks_scale = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(- 0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32- )- b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks_scale[None, :] * stride_bsk-+ a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak)+ offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)+ offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr+ offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N+ b_ptrs = b_ptr + (offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk)+ offs_bsn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))) % N+ offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)+ b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bskaccumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)-- for k in tl.range(0, num_k_iter):- a_bf16 = tl.load(a_ptrs).to(tl.float32)- a_fp4, a_scales = _lean_mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, QUANT_BLOCK)-- b_fp4 = tl.load(b_ptrs)-- b_scales = (- tl.load(b_scale_ptrs)+ for k in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):+ b_scales = (tl.load(b_scale_ptrs, cache_modifier=cache_modifier).reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)- .permute(0, 5, 3, 1, 4, 2, 6)- .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)- )-- accumulator = tl.dot_scaled(a_fp4, a_scales, "e2m1", b_fp4, b_scales, "e2m1", accumulator)-+ .permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE))+ if EVEN_K:+ a_bf16 = tl.load(a_ptrs)+ b = tl.load(b_ptrs, cache_modifier=cache_modifier)+ b = (b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)+ .permute(0, 1, 4, 2, 3, 5).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2).trans(1, 0))+ if PREQUANT:+ a, a_scales = _hw_cvt_quant(a_bf16.to(tl.float32), BLOCK_SIZE_K, BLOCK_SIZE_M, 32)+ accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")a_ptrs += BLOCK_SIZE_K * stride_ak- b_ptrs += (BLOCK_SIZE_K // 2) * stride_bk+ b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bkb_scale_ptrs += BLOCK_SIZE_K * stride_bsk-c = accumulator.to(c_ptr.type.element_ty)offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)⋯ 1 unchanged linesc_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)tl.store(c_ptrs, c, mask=c_mask)-+ # ===== v710 kernel for M>32 =====@triton.jit- def _fused_quant_shuffle_kernel(- x_ptr, x_fp4_ptr, bs_ptr,- stride_x_m, stride_x_n,- stride_x_fp4_m, stride_x_fp4_n,- M, N, scale_n_valid,- SCALE_N: tl.constexpr,- BLOCK_SIZE_M: tl.constexpr,- BLOCK_SIZE_N: tl.constexpr,- NUM_ITER: tl.constexpr,- NUM_STAGES: tl.constexpr,- MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,- ):- pid_m = tl.program_id(0)- start_n = tl.program_id(1) * NUM_ITER- NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE-- for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):- x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)- x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)- x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n- x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]- x = tl.load(x_ptr + x_offs, mask=x_mask, other=0.0).to(tl.float32)-- out_tensor, bs_e8m0 = _lean_mxfp4_quant_op(- x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE- )-- out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)- out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)- out_offs = out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n- out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]- tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)-- bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)- bs_offs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)- m_idx = bs_offs_m[:, None]- n_idx = bs_offs_n[None, :]- i0 = m_idx // 32- i1 = (m_idx // 16) % 2- i2 = m_idx % 16- i3 = n_idx // 8- i4 = (n_idx // 4) % 2- i5 = n_idx % 4- shuffled_offset = (i0 * (SCALE_N * 32) + i3 * 256 + i5 * 64 + i2 * 4 + i4 * 2 + i1)- bs_valid = (bs_offs_m < M)[:, None] & (bs_offs_n < scale_n_valid)[None, :]- bs_e8m0 = tl.where(bs_valid, bs_e8m0, 127)- bs_store_mask = (m_idx < (M + 255) // 256 * 256) & (n_idx < SCALE_N)- tl.store(bs_ptr + shuffled_offset, bs_e8m0, mask=bs_store_mask)--- _cache_asm = {}- _cache_fused = {}- _gemm_asm = None- _warmup_done = False--- def custom_kernel(data: input_t) -> output_t:- global _gemm_asm, _warmup_done-- A, B, B_q, B_shuffle, B_scale_sh = data- M, K = A.shape- N = B_shuffle.shape[0]-- use_fused = (M <= 64)-- # Warmup: use ASM path to init aiter module- if not _warmup_done:- scale_n_valid = (K + 31) // 32- SCALE_M = ((M + 255) // 256) * 256- SCALE_N = ((scale_n_valid + 7) // 8) * 8- BSM = triton.next_power_of_2(M) if M <= 32 else 16- grid = (triton.cdiv(M, BSM), triton.cdiv(K, 32))-- x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=A.device)- bs_sh = torch.full((SCALE_M, SCALE_N), 127, dtype=torch.uint8, device=A.device)-- _fused_quant_shuffle_kernel[grid](- A, x_fp4, bs_sh,- A.stride(0), A.stride(1),- x_fp4.stride(0), x_fp4.stride(1),- M, K, scale_n_valid,- SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=32,- NUM_ITER=1, NUM_STAGES=1, MXFP4_QUANT_BLOCK_SIZE=32,- num_warps=1, waves_per_eu=0, num_stages=1,- )-- result = aiter.gemm_a4w4(- x_fp4.view(_fp4x2), B_shuffle,- bs_sh.view(_fp8_e8m0), B_scale_sh,- dtype=_bf16, bpreshuffle=True,- )- _warmup_done = True- try:- _gemm_asm = torch.ops.aiter.gemm_a4w4_asm- except Exception:- try:- import aiter.jit.core as _jc- _gemm_asm = getattr(_jc, 'gemm_a4w4_asm', None)- except Exception:- pass- return result-- if use_fused:- # --- Fused quant+GEMM: single kernel launch ---- key = (M, K, N)- c = _cache_fused.get(key)- if c is None:- K_packed = K // 2- scale_n = (K + 31) // 32- SCALE_N_B = ((scale_n + 7) // 8) * 8-- BLOCK_SIZE_M = 16- BLOCK_SIZE_N = 64 if M <= 16 else 128- BLOCK_SIZE_K = 512-- base_blocks = triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N)- target_ksplit = max(1, 256 // max(1, base_blocks))-- if target_ksplit > 1:- SPLITK_BLOCK_SIZE, BLOCK_SIZE_K_adj, NUM_KSPLIT = _get_splitk_fn(- K_packed, BLOCK_SIZE_K, target_ksplit- )- if BLOCK_SIZE_K_adj < 512:- BLOCK_SIZE_K_adj = 512- SPLITK_BLOCK_SIZE = 2 * K_packed- NUM_KSPLIT = 1- else:- BLOCK_SIZE_K = BLOCK_SIZE_K_adj- else:- NUM_KSPLIT = 1- SPLITK_BLOCK_SIZE = 2 * K_packed-- if NUM_KSPLIT > 1:- y_pp = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=A.device)- else:- y_pp = None- SPLITK_BLOCK_SIZE = 2 * K_packed-- y = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)-- total_blocks_raw = NUM_KSPLIT * triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N)- total_blocks = ((total_blocks_raw + 7) // 8) * 8-- bs_stride_n = 32 * SCALE_N_B- bs_stride_k = 1-- c = (K_packed, SCALE_N_B, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K,- NUM_KSPLIT, SPLITK_BLOCK_SIZE,- y, y_pp, total_blocks, bs_stride_n, bs_stride_k)- _cache_fused[key] = c-- (K_packed, SCALE_N_B, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K,- NUM_KSPLIT, SPLITK_BLOCK_SIZE,- y, y_pp, total_blocks, bs_stride_n, bs_stride_k) = c-- B_q_u8 = B_q.view(torch.uint8) if B_q.dtype != torch.uint8 else B_q- B_q_T = B_q_u8.T- B_scale_u8 = B_scale_sh.view(torch.uint8)-- out_tensor = y if NUM_KSPLIT == 1 else y_pp-- _fused_quant_gemm_kernel[(total_blocks,)](- A, B_q_T, out_tensor, B_scale_u8,- M, N, K,- A.stride(0), A.stride(1),- B_q_T.stride(0), B_q_T.stride(1),- 0 if NUM_KSPLIT == 1 else y_pp.stride(0),- y.stride(0) if NUM_KSPLIT == 1 else y_pp.stride(1),- y.stride(1) if NUM_KSPLIT == 1 else y_pp.stride(2),- bs_stride_n, bs_stride_k,- BLOCK_SIZE_M=BLOCK_SIZE_M,- BLOCK_SIZE_N=BLOCK_SIZE_N,- BLOCK_SIZE_K=BLOCK_SIZE_K,- GROUP_SIZE_M=8,- NUM_KSPLIT=NUM_KSPLIT,- SPLITK_BLOCK_SIZE=SPLITK_BLOCK_SIZE,- QUANT_BLOCK=32,- num_warps=8,- num_stages=2,- waves_per_eu=0,- )-- if NUM_KSPLIT > 1:- ACTUAL_KSPLIT = triton.cdiv(K_packed, (SPLITK_BLOCK_SIZE // 2))- grid_reduce = (triton.cdiv(M, 16), triton.cdiv(N, 64))- _reduce_kernel[grid_reduce](- y_pp, y, M, N,- y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),- y.stride(0), y.stride(1),- 16, 64, ACTUAL_KSPLIT,- triton.next_power_of_2(NUM_KSPLIT),- )-- return y-+ def _hw_cvt_quant(x, BLOCK_K: tl.constexpr, BLOCK_M: tl.constexpr, QUANT_BLOCK: tl.constexpr):+ NUM_BLOCKS: tl.constexpr = BLOCK_K // QUANT_BLOCK+ HALF_QB: tl.constexpr = QUANT_BLOCK // 2+ x_3d = x.reshape(BLOCK_M, NUM_BLOCKS, QUANT_BLOCK)+ amax = tl.max(tl.abs(x_3d), axis=-1, keep_dims=True)+ amax_i32 = amax.to(tl.int32, bitcast=True)+ amax_rounded = ((amax_i32 + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000)+ amax_exp = (amax_rounded >> 23).to(tl.int32)+ su = amax_exp - 129;su = tl.where(su < -126, -126, su);su = tl.where(su > 127, 127, su)+ bs_e8m0 = (su + 127).to(tl.uint8)+ # Fused CVT scale: pass 2^(su) so hardware does srcExp -= su → FP4(value * 2^(-su))+ cvt_exp = su + 127+ cvt_scale = (cvt_exp << 23).to(tl.float32, bitcast=True)+ cvt_scale_bc = tl.broadcast_to(cvt_scale, (BLOCK_M, NUM_BLOCKS, HALF_QB))+ cvt_scale_flat = cvt_scale_bc.reshape(BLOCK_M, BLOCK_K // 2)+ # Raw x pairs — NO pre-multiply! Hardware CVT handles scaling.+ x_flat = x_3d.reshape(BLOCK_M, BLOCK_K)+ x_pairs = x_flat.reshape(BLOCK_M, BLOCK_K // 2, 2)+ evens, odds = tl.split(x_pairs)+ fp4_packed = tl.inline_asm_elementwise("v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3","=v,v,v,v",args=[evens, odds, cvt_scale_flat],dtype=tl.int32,is_pure=True,pack=1)+ x_fp4 = fp4_packed.to(tl.uint8)+ return x_fp4.reshape(BLOCK_M, BLOCK_K // 2), bs_e8m0.reshape(BLOCK_M, NUM_BLOCKS)+ # ===== Configs =====+ def _get_ps_config(M, N, K):+ if K > 4096:+ # v911: NK=14 → 238 blocks (93% CU) with 1 K-iter each+ return {"BLOCK_SIZE_M":16,"BLOCK_SIZE_N":128,"BLOCK_SIZE_K":512,"GROUP_SIZE_M":1,"num_warps":4,"num_stages":2,"waves_per_eu":2,"matrix_instr_nonkdim":16,"cache_modifier":".cg","NUM_KSPLIT":14}+ if M <= 4:+ return {"BLOCK_SIZE_M":4,"BLOCK_SIZE_N":128,"BLOCK_SIZE_K":256,"GROUP_SIZE_M":1,"num_warps":4,"num_stages":2,"waves_per_eu":0,"matrix_instr_nonkdim":16,"cache_modifier":".cg","NUM_KSPLIT":1}+ elif M <= 8:+ return {"BLOCK_SIZE_M":8,"BLOCK_SIZE_N":128,"BLOCK_SIZE_K":256,"GROUP_SIZE_M":1,"num_warps":4,"num_stages":2,"waves_per_eu":0,"matrix_instr_nonkdim":16,"cache_modifier":".cg","NUM_KSPLIT":1}+ elif K <= 1024:+ return {"BLOCK_SIZE_M":8,"BLOCK_SIZE_N":128,"BLOCK_SIZE_K":256,"GROUP_SIZE_M":1,"num_warps":4,"num_stages":2,"waves_per_eu":2,"matrix_instr_nonkdim":16,"cache_modifier":None,"NUM_KSPLIT":1}+ elif M <= 32:+ return {"BLOCK_SIZE_M":32,"BLOCK_SIZE_N":64,"BLOCK_SIZE_K":512,"GROUP_SIZE_M":1,"num_warps":8,"num_stages":1,"waves_per_eu":2,"matrix_instr_nonkdim":16,"cache_modifier":None,"NUM_KSPLIT":1}+ elif M <= 64:+ return {"BLOCK_SIZE_M":8,"BLOCK_SIZE_N":128,"BLOCK_SIZE_K":512,"GROUP_SIZE_M":1,"num_warps":4,"num_stages":2,"waves_per_eu":2,"matrix_instr_nonkdim":16,"cache_modifier":".cg","NUM_KSPLIT":1}else:- # --- ASM GEMM path ---- key = (M, K, N)- c = _cache_asm.get(key)- if c is None:- scale_n_valid = (K + 31) // 32- SCALE_M = ((M + 255) // 256) * 256- SCALE_N = ((scale_n_valid + 7) // 8) * 8- padded_m = get_padded_m(M, N, K, 0)+ # M>64: BSK=256 (v895 confirmed -0.4µs on M=256), BSM=16, nw=4/wpe=2+ return {"BLOCK_SIZE_M":16,"BLOCK_SIZE_N":128,"BLOCK_SIZE_K":256,"GROUP_SIZE_M":1,"num_warps":4,"num_stages":2,"waves_per_eu":2,"matrix_instr_nonkdim":16,"cache_modifier":".cg","NUM_KSPLIT":1}- BSM = triton.next_power_of_2(M) if M <= 32 else 16- NW = 1- BSN = 32- NUM_ITER_Q = 2- grid = (triton.cdiv(M, BSM), triton.cdiv(K, BSN * NUM_ITER_Q))-- ck_config = get_GEMM_config(M, N, K)- kernel_name = ""- split_k = 0- if ck_config is not None:- split_k = ck_config.get("splitK", 0) or 0- kernel_name = ck_config["kernelName"]-- x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=A.device)- bs_sh = torch.full((SCALE_M, SCALE_N), 127, dtype=torch.uint8, device=A.device)- out = torch.empty((padded_m, N), dtype=torch.bfloat16, device=A.device)-- x_fp4_view = x_fp4.view(_fp4x2)- bs_sh_view = bs_sh.view(_fp8_e8m0)- out_view = out[:M] if M < padded_m else out-- c = (scale_n_valid, SCALE_N, BSM, BSN, NUM_ITER_Q, grid,- x_fp4, bs_sh, out, x_fp4_view, bs_sh_view, out_view,- kernel_name, split_k,- A.stride(0), A.stride(1), x_fp4.stride(0), x_fp4.stride(1))- _cache_asm[key] = c-- (scale_n_valid, SCALE_N, BSM, BSN, NUM_ITER_Q, grid,- x_fp4, bs_sh, out, x_fp4_view, bs_sh_view, out_view,- kernel_name, split_k,- stride_a0, stride_a1, stride_fp4_0, stride_fp4_1) = c-- _fused_quant_shuffle_kernel[grid](- A, x_fp4, bs_sh,- stride_a0, stride_a1,- stride_fp4_0, stride_fp4_1,- M, K, scale_n_valid,- SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN,- NUM_ITER=NUM_ITER_Q, NUM_STAGES=NUM_ITER_Q, MXFP4_QUANT_BLOCK_SIZE=32,- num_warps=1, waves_per_eu=0, num_stages=NUM_ITER_Q,- )-- if _gemm_asm is not None:- _gemm_asm(x_fp4_view, B_shuffle, bs_sh_view, B_scale_sh,- out, kernel_name, None, 1.0, 0.0, True, split_k)- return out_view-- return aiter.gemm_a4w4(- x_fp4_view, B_shuffle, bs_sh_view, B_scale_sh,- dtype=_bf16, bpreshuffle=True,- )+ _cf={}+ def custom_kernel(data:input_t)->output_t:+ A,B,Bq,Bs,Bss=data;M,K=A.shape;N=Bs.shape[0]+ key=(M,K,N);c=_cf.get(key)+ if c is None:+ config = _get_ps_config(M, N, K)+ K_kernel = K // 2; BSK = config["BLOCK_SIZE_K"]; BSN = max(config["BLOCK_SIZE_N"], 32); BSM = config["BLOCK_SIZE_M"]+ NK = config["NUM_KSPLIT"]+ if NK > 1: SBS, BSK, NK = _gs(K_kernel, BSK, NK)+ else: SBS = 2 * K_kernel+ grid_mn = triton.cdiv(M, BSM) * triton.cdiv(N, BSN)+ grid_size = grid_mn * NK+ out = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)+ y_pp = torch.empty((NK, M, N), dtype=torch.float32, device=A.device) if NK > 1 else None+ AK = triton.cdiv(K_kernel, (SBS // 2)) if NK > 1 else 1+ c = ('ps', out, y_pp, None, None, grid_size, K_kernel, BSM, BSN, BSK,+ config["GROUP_SIZE_M"], NK, SBS, config["num_warps"], config["num_stages"],+ config["waves_per_eu"], config["matrix_instr_nonkdim"], config["cache_modifier"],+ AK, triton.next_power_of_2(NK) if NK > 1 else 1,+ (triton.cdiv(M,16), triton.cdiv(N,64)) if NK > 1 else None)+ _cf[key] = c+ _, out, y_pp, Bw, Bsc, grid_size, K_kernel, BSM, BSN, BSK, GSM, NK, SBS, nw, ns, wpe, mind, cm, AK, MK, rgr = c+ b_ptr = Bs.data_ptr()+ if Bw is None or _cf.get(('_bp', key)) != b_ptr:+ Bw = Bs.view(torch.uint8).reshape(N // 16, (K // 2) * 16)+ bs_shape = Bss.shape; Bsc = Bss.view(torch.uint8).reshape(bs_shape[0] // 32, bs_shape[1] * 32)+ c_list = list(c); c_list[3] = Bw; c_list[4] = Bsc; c = tuple(c_list); _cf[key] = c+ _cf[('_bp', key)] = b_ptr+ target = y_pp if NK > 1 else out+ sk_off = y_pp.stride(0) if NK > 1 else 0+ cm_s = (y_pp.stride(1) if NK > 1 else out.stride(0))+ cn_s = (y_pp.stride(2) if NK > 1 else out.stride(1))+ _ps_nx2[(grid_size,)](A, Bw, target, Bsc, M, N, K_kernel,+ A.stride(0), A.stride(1), Bw.stride(0), Bw.stride(1),+ sk_off, cm_s, cn_s, Bsc.stride(0), Bsc.stride(1),+ BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN, BLOCK_SIZE_K=BSK,+ GROUP_SIZE_M=GSM, NUM_KSPLIT=NK, SPLITK_BLOCK_SIZE=SBS,+ num_warps=nw, num_stages=ns, waves_per_eu=wpe,+ matrix_instr_nonkdim=mind, PREQUANT=True, cache_modifier=cm)+ if NK > 1:+ _gluon_reduce_kernel[rgr](y_pp, out, M, N, y_pp.stride(0), y_pp.stride(1), y_pp.stride(2), out.stride(0), out.stride(1), 16, 64, AK, MK)+ return out
scrolls · 625 diff lines total
Best evidence level for this revision: reported
JSON