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
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.
fp4
mant_odd = (normal_u32 >> 22) & 1 # bit 22 = fp4 mantissa bitnum-warps = 4
num_warps=4, num_stages=1,stages = 1
num_warps=4, num_stages=1,tile-m = 32
BLOCK_M = 32tile-n = 32
BLOCK_N = 32Kernel 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