submission 634224
olezhka_007 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 177 lines, June 9 Researcher Reciprocity License v1.0.
submission_patched_a16wfp4.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-634224?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:83009a4a2d10d15d9f9ac9e445da517311bfab66c7ab9ee4430ad1947e8be86d
license declaredunknown
license concludedunknown
authorsolezhka_007
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
submission_patched_a16wfp4.py177 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
Monkey-patch _gemm_a16wfp4_kernel to read shuffled B scales inline.
Eliminates B quant kernel — single Triton kernel for K≤512 shapes.
"""
import os
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
# === MONKEY-PATCH: Replace scale pointer computation with shuffled version ===
_KF = "/home/runner/aiter/aiter/ops/triton/_triton_kernels/gemm/basic/gemm_a16wfp4.py"
try:
with open(_KF, 'r') as f:
_src = f.read()
# Original scale pointer setup (non-preshuffle kernel)
_OLD_SCALE = """ offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)) + tl.arange(
0, BLOCK_SIZE_K // SCALE_GROUP_SIZE
)
# B scales are N x K even though B operand is K x N.
b_scale_ptrs = (
b_scales_ptr + offs_bn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
)"""
# Replacement: compute shuffled flat offsets directly
_NEW_SCALE = """ offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)) + tl.arange(
0, BLOCK_SIZE_K // SCALE_GROUP_SIZE
)
# Shuffled scale: compute flat byte offset using e8m0_shuffle mapping
_SN = 2 * K // SCALE_GROUP_SIZE # total scale cols (K is K_half)
_n = offs_bn # shape (BLOCK_SIZE_N,)
_kg = offs_ks # shape (BLOCK_SIZE_K // SCALE_GROUP_SIZE,)
_d0 = _n // 32
_d1 = (_n % 32) // 16
_d2 = _n % 16
_d3 = _kg // 8
_d4 = (_kg % 8) // 4
_d5 = _kg % 4
b_scale_ptrs = b_scales_ptr + (
_d0[:, None] * (_SN * 32) + _d3[None, :] * 256 +
_d5[None, :] * 64 + _d2[:, None] * 4 + _d4[None, :] * 2 + _d1[:, None]
)"""
# Original scale pointer advance in loop
_OLD_ADV = " b_scale_ptrs += BLOCK_SIZE_K // SCALE_GROUP_SIZE * stride_bsk"
# Replacement: recompute shuffled offsets for new K position
_NEW_ADV = """ offs_ks += BLOCK_SIZE_K // SCALE_GROUP_SIZE
_d3 = offs_ks // 8
_d4 = (offs_ks % 8) // 4
_d5 = offs_ks % 4
b_scale_ptrs = b_scales_ptr + (
_d0[:, None] * (_SN * 32) + _d3[None, :] * 256 +
_d5[None, :] * 64 + _d2[:, None] * 4 + _d4[None, :] * 2 + _d1[:, None]
)"""
patched = False
has_old = _OLD_SCALE in _src
has_new = '_SN = 2 * K // SCALE_GROUP_SIZE' in _src # check if already patched
if has_old:
_src = _src.replace(_OLD_SCALE, _NEW_SCALE, 1)
if _OLD_ADV in _src:
_src = _src.replace(_OLD_ADV, _NEW_ADV, 1)
with open(_KF, 'w') as f:
f.write(_src)
patched = True
elif has_new:
patched = True # already patched from previous run
except Exception:
pass
# === END PATCH ===
# Clear Triton cache to force recompile with patched kernel
import shutil, glob
for d in glob.glob("/home/runner/.triton/cache/*") + glob.glob("/tmp/triton_*"):
try: shutil.rmtree(d)
except: pass
import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
from task import input_t, output_t
QUANT_GROUP = 32
ASM_KERNEL = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
K_THRESHOLD = 512
# ---- Triton quant for ASM path (K>512) ----
@triton.jit
def _mxfp4_quant_shuffled_kernel(
X, Out_fp4, Out_scale, M, K, M_pad,
stride_xm, stride_xk, stride_fm, stride_fk,
stride_sm: tl.constexpr, SN: tl.constexpr,
BLOCK_M: tl.constexpr, BLOCK_K: tl.constexpr,
):
QUANT_GROUP: tl.constexpr = 32
NUM_GROUPS: tl.constexpr = BLOCK_K // QUANT_GROUP
pid_m = tl.program_id(0); pid_k = tl.program_id(1)
rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
rk = pid_k * BLOCK_K + tl.arange(0, BLOCK_K)
mask = (rm < M)[:, None] & (rk < K)[None, :]
x = tl.load(X + rm[:, None] * stride_xm + rk[None, :] * stride_xk, mask=mask, other=0.0).to(tl.float32)
x = tl.reshape(x, [BLOCK_M, NUM_GROUPS, QUANT_GROUP])
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
ai = amax.to(tl.int32, bitcast=True)
ai = (ai + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = ai.to(tl.float32, bitcast=True)
su = tl.clamp(tl.log2(amax).floor() - 2, min=-127, max=127)
e8 = su.to(tl.uint8) + 127
qs = tl.exp2(-su)
qx = (x * qs).to(tl.uint32, bitcast=True)
s = qx & 0x80000000; qx = qx ^ s
qf = qx.to(tl.float32, bitcast=True)
DMI: tl.constexpr = 1249902592
DMF: tl.constexpr = tl.cast(1249902592, tl.float32, bitcast=True)
sat = qf >= 6.0; den = (~sat) & (qf < 1.0); nor = ~(sat | den)
dx = (qf + DMF).to(tl.int32, bitcast=True) - DMI
nx = qx.to(tl.int32, bitcast=True); mo = (nx >> 22) & 1; nx = ((nx + (-1054867457) + mo) >> 22).to(tl.uint8)
e = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
e = tl.where(nor, nx, e); e = tl.where(den, dx.to(tl.uint8), e)
e = e | (s >> 28).to(tl.uint8)
e = tl.reshape(e, [BLOCK_M, NUM_GROUPS, QUANT_GROUP // 2, 2])
ev, od = tl.split(e)
fp4 = tl.reshape(ev | (od << 4), [BLOCK_M, BLOCK_K // 2])
rh = pid_k * (BLOCK_K // 2) + tl.arange(0, BLOCK_K // 2)
tl.store(Out_fp4 + rm[:, None] * stride_fm + rh[None, :] * stride_fk, fp4, mask=(rm < M)[:, None] & (rh < K // 2)[None, :])
sc = tl.reshape(e8, [BLOCK_M, NUM_GROUPS])
rs = pid_k * NUM_GROUPS + tl.arange(0, NUM_GROUPS)
d0=rm//32; d1=(rm%32)//16; d2=rm%16; d3=rs//8; d4=(rs%8)//4; d5=rs%4
sf = d0[:,None]*(SN*32)+d3[None,:]*256+d5[None,:]*64+d2[:,None]*4+d4[None,:]*2+d1[:,None]
tl.store(Out_scale + (sf // SN) * stride_sm + (sf % SN), sc, mask=(rm < M_pad)[:, None] & (rs < SN)[None, :])
_BUFS = {}
def _get_bufs(m, k, n, device):
key = (m, k, n)
if key not in _BUFS:
mp = ((m+255)//256)*256; sn = k//QUANT_GROUP
fp4 = torch.empty((m, k//2), dtype=torch.uint8, device=device)
sc = torch.zeros((mp, sn), dtype=torch.uint8, device=device)
out = torch.empty((m, n), dtype=torch.bfloat16, device=device)
BM = min(32, mp); BK = min(256, k)
while k % BK != 0: BK //= 2
grid = (triton.cdiv(m, BM), triton.cdiv(k, BK))
_BUFS[key] = (fp4, sc, out, grid, BM, BK, mp, sn)
return _BUFS[key]
_ROUTE = {}
def custom_kernel(data: input_t) -> output_t:
A, _B, _B_q, B_shuffle, B_scale_sh = data
if not A.is_contiguous():
A = A.contiguous()
m, k = int(A.shape[0]), int(A.shape[1])
n = int(B_shuffle.shape[0])
key = (m, k, n)
if key not in _ROUTE:
_ROUTE[key] = "a16" if k <= K_THRESHOLD else "asm"
if _ROUTE[key] == "a16":
# PATCHED: pass shuffled B_scale_sh directly — kernel reads shuffled inline
b_q_u8 = _B_q.view(torch.uint8)
return gemm_a16wfp4(A, b_q_u8, B_scale_sh.view(torch.uint8), dtype=torch.bfloat16)
else:
fp4, sc, out, grid, BM, BK, mp, sn = _get_bufs(m, k, n, A.device)
_mxfp4_quant_shuffled_kernel[grid](A, fp4, sc, m, k, mp, A.stride(0), A.stride(1), fp4.stride(0), fp4.stride(1), stride_sm=sc.stride(0), SN=sn, BLOCK_M=BM, BLOCK_K=BK)
return aiter.gemm_a4w4_asm(fp4.view(dtypes.fp4x2), B_shuffle, sc.view(dtypes.fp8_e8m0), B_scale_sh, out, ASM_KERNEL, bpreshuffle=True)
scrolls · 177 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON