Skip to content
KernelIndex
Search⌘K

submission 676749

chineseman · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 195 lines, June 9 Researcher Reciprocity License v1.0.

v196_fused_quant_shuffle.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-676749?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
14.6µs
#523 of 1143
2026-03-31

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:275b4a44dbfe7ab11d16aec114a7b15e6fceb89cd3066a5711506b69c2427ad9
license declaredunknown
license concludedunknown
authorschineseman
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4mant_odd = (normal_u32 >> 22) & 1 # bit 22 = fp4 mantissa bit
num-warps = 4num_warps=4, num_stages=1,
stages = 1num_warps=4, num_stages=1,
tile-m = 32BLOCK_M = 32
tile-n = 32BLOCK_N = 32

Kernel source

v196_fused_quant_shuffle.py195 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""v196: Custom fused quant+shuffle Triton kernel — eliminates intermediate
blockscale buffer and second kernel launch."""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
from aiter import dtypes
import aiter

_bf16 = dtypes.bf16
_fp4x2 = dtypes.fp4x2
_fp8_e8m0 = dtypes.fp8_e8m0

MXFP4_QUANT_BLOCK_SIZE = 32


@triton.jit
def _fused_quant_shuffle_kernel(
    x_ptr, x_fp4_ptr, shuffled_scale_ptr,
    stride_x_m: tl.int64, stride_x_n: tl.int64,
    stride_fp4_m: tl.int64, stride_fp4_n: tl.int64,
    M, N,
    N_scale,
    N_pad_scale,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)

    m_base = pid_m * BLOCK_SIZE_M
    n_base = pid_n * BLOCK_SIZE_N

    offs_m = m_base + tl.arange(0, BLOCK_SIZE_M)
    offs_n = n_base + tl.arange(0, BLOCK_SIZE_N)

    mask_m = offs_m < M
    mask_n = offs_n < N
    mask = mask_m[:, None] & mask_n[None, :]

    # Load input tile [BLOCK_SIZE_M, BLOCK_SIZE_N]
    x_ptrs = x_ptr + offs_m[:, None] * stride_x_m + offs_n[None, :] * stride_x_n
    x = tl.load(x_ptrs, mask=mask, other=0.0).to(tl.float32)

    # --- Quantization (matching aiter _mxfp4_quant_op exactly) ---
    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
    x = tl.reshape(x, (BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE))

    # Blockscale: max abs per group of 32, rounded up to power of 2
    amax = tl.max(tl.abs(x), axis=2, keep_dims=True)
    amax_i = amax.to(tl.int32, bitcast=True)
    amax_i = (amax_i + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    amax = amax_i.to(tl.float32, bitcast=True)

    # Unbiased exponent: floor(log2(amax)) - 2
    scale_unb = tl.log2(amax)
    scale_unb = scale_unb.to(tl.int32) - 2
    scale_unb = tl.maximum(scale_unb, -127)
    scale_unb = tl.minimum(scale_unb, 127)

    # e8m0 byte = unbiased + 127
    bs_e8m0 = (scale_unb + 127).to(tl.uint8)
    bs_e8m0_2d = tl.reshape(bs_e8m0, (BLOCK_SIZE_M, NUM_QUANT_BLOCKS))

    # Quantize: multiply by 2^(-scale_unbiased)
    quant_scale = tl.exp2((-scale_unb).to(tl.float32))
    qx = x * quant_scale

    # --- FP32 -> FP4 e2m1 conversion ---
    qx_u32 = qx.to(tl.uint32, bitcast=True)
    sign = qx_u32 & 0x80000000
    qx_u32 = qx_u32 ^ sign  # absolute value
    qx_f32 = qx_u32.to(tl.float32, bitcast=True)

    saturate_mask = qx_f32 >= 6.0
    denormal_mask = (~saturate_mask) & (qx_f32 < 1.0)
    normal_mask = ~(saturate_mask | denormal_mask)

    # Denormal path: value in [0, 1.0)
    # Magic number: (127 - 1 + 23 - 1 + 1) << 23 = 149 << 23
    DENORM_MAGIC = tl.constexpr(149 << 23)
    denorm_magic_u32 = tl.full([1], 149 << 23, dtype=tl.uint32)
    denorm_magic_f32 = denorm_magic_u32.to(tl.float32, bitcast=True)
    denorm_x = (qx_f32 + denorm_magic_f32).to(tl.uint32, bitcast=True)
    denorm_x = denorm_x - denorm_magic_u32
    denorm_x = denorm_x.to(tl.uint8)

    # Normal path: value in [1.0, 6.0), round-to-nearest-even
    normal_u32 = qx_u32
    mant_odd = (normal_u32 >> 22) & 1  # bit 22 = fp4 mantissa bit
    # val_to_add = ((1 - 127) << 23) + (1 << 21) - 1 = -126*8388608 + 2097151 = -1054867457
    # As uint32: 0xC0FFFFFF - let's compute directly
    VAL_ADD = tl.full([1], ((1 - 127) << 23) + (1 << 21) - 1, dtype=tl.int32)
    VAL_ADD_U = VAL_ADD.to(tl.uint32, bitcast=True)
    normal_u32 = normal_u32 + VAL_ADD_U + mant_odd
    normal_u32 = normal_u32 >> 22  # shift mantissa into low bits
    normal_x = normal_u32.to(tl.uint8)

    # Merge all paths
    e2m1 = tl.full(qx.shape, 7, dtype=tl.uint8)  # saturate default
    e2m1 = tl.where(normal_mask, normal_x, e2m1)
    e2m1 = tl.where(denormal_mask, denorm_x, e2m1)

    # Apply sign (bit 3 of fp4)
    sign_fp4 = (sign >> 28).to(tl.uint8)
    e2m1 = e2m1 | sign_fp4

    # Pack consecutive pairs: even in low nibble, odd in high nibble
    # Reshape to [..., 16, 2], split, pack
    e2m1 = tl.reshape(e2m1, (BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2))
    evens, odds = tl.split(e2m1)
    packed = evens | (odds << 4)
    packed = tl.reshape(packed, (BLOCK_SIZE_M, BLOCK_SIZE_N // 2))

    # Store packed fp4 output
    fp4_offs_n = n_base // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
    fp4_mask = mask_m[:, None] & (fp4_offs_n[None, :] < (N // 2))
    fp4_ptrs = x_fp4_ptr + offs_m[:, None] * stride_fp4_m + fp4_offs_n[None, :] * stride_fp4_n
    tl.store(fp4_ptrs, packed, mask=fp4_mask)

    # Store shuffled scales (inline shuffle — no intermediate buffer)
    scale_n_base = n_base // MXFP4_QUANT_BLOCK_SIZE
    scale_offs_n = scale_n_base + tl.arange(0, NUM_QUANT_BLOCKS)

    sm = offs_m[:, None]  # [BLOCK_SIZE_M, 1]
    sn = scale_offs_n[None, :]  # [1, NUM_QUANT_BLOCKS]

    i0 = sm // 32
    i1 = (sm % 32) // 16
    i2 = sm % 16
    i3 = sn // 8
    i4 = (sn % 8) // 4
    i5 = sn % 4

    out_idx = i0 * (N_pad_scale // 8 * 256) + i3 * 256 + i5 * 64 + i2 * 4 + i4 * 2 + i1

    scale_mask = mask_m[:, None] & (scale_offs_n[None, :] < N_scale)
    tl.store(shuffled_scale_ptr + out_idx, bs_e8m0_2d, mask=scale_mask)


BLOCK_M = 32
BLOCK_N = 32

_cache = {}


def _quant_fused(x):
    M, N = x.shape
    key = (M, N)

    if key not in _cache:
        x_fp4_buf = torch.empty(M, N // 2, dtype=torch.uint8, device=x.device)
        N_scale = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE

        M_pad = (M + 255) // 256 * 256
        N_pad_scale = (N_scale + 7) // 8 * 8
        shuffled_buf = torch.zeros(M_pad * N_pad_scale, dtype=torch.uint8, device=x.device)

        grid = (
            triton.cdiv(M, BLOCK_M),
            triton.cdiv(N, BLOCK_N),
        )

        _cache[key] = (
            x_fp4_buf, shuffled_buf,
            N_scale, M_pad, N_pad_scale, grid,
        )

    (x_fp4_buf, shuffled_buf,
     N_scale, M_pad, N_pad_scale, grid) = _cache[key]

    _fused_quant_shuffle_kernel[grid](
        x, x_fp4_buf, shuffled_buf,
        x.stride(0), x.stride(1),
        x_fp4_buf.stride(0), x_fp4_buf.stride(1),
        M, N, N_scale, N_pad_scale,
        BLOCK_SIZE_M=BLOCK_M, BLOCK_SIZE_N=BLOCK_N,
        MXFP4_QUANT_BLOCK_SIZE=MXFP4_QUANT_BLOCK_SIZE,
        num_warps=4, num_stages=1,
    )

    return x_fp4_buf.view(_fp4x2), shuffled_buf.view(M_pad, N_pad_scale).view(_fp8_e8m0)


def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    A_q, A_scale_sh = _quant_fused(A)
    return aiter.gemm_a4w4(
        A_q, B_shuffle, A_scale_sh, B_scale_sh,
        dtype=_bf16, bpreshuffle=True,
    )
scrolls · 195 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