Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
10.5µs
#271 of 1143
2026-03-25

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.

fp4fp4 = tl.reshape(ev | (od << 4), [BLOCK_M, BLOCK_K // 2])
split-k_OLD_SCALE = """ offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)) + tl.arange(

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