submission 634613
brandonin · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 61 lines, June 9 Researcher Reciprocity License v1.0.
submission_v36.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-634613?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:e4a4ac8ad872e3e774bff833a0aeb5f8c2eeb5f7abf279e676b4be41b8e15d96
license declaredunknown
license concludedunknown
authorsbrandonin
imported2026-08-26
Kernel source
submission_v36.py61 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
# v36: Fused BF16→FP4 quant + GEMM via Triton gemm_a16wfp4_preshuffle.
#
# Key insight: eliminate the separate A-quantization step entirely.
# gemm_a16wfp4_preshuffle takes BF16 A directly and quantizes on-the-fly.
#
# Scale format conversion (O(1)):
# B_scale_sh (CK e8m0_shuffle format): (N_pad, sn_pad) in uint8
# Triton shuffle_scales format: (N_pad//32, sn_pad*32) in uint8
# Both apply the SAME permutation to scale bytes — only the 2D view differs.
# So: B_scale_sh.view(uint8).reshape(N_pad//32, sn_pad*32) gives Triton format.
#
# Weight format conversion (O(1)):
# B_shuffle: shuffle_weight applied, shape (N, K//2)
# gemm_a16wfp4_preshuffle expects: (N//16, K//2*16)
# So: B_shuffle.view(uint8).reshape(N//16, K//2*16)
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ['CXX'] = 'clang++'
import torch
import sys
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
from task import input_t, output_t
_shape_cache = {}
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
m, k = A.shape
n = B_shuffle.shape[0]
shape_key = (m, n, k)
# Reshape B_shuffle: (N, K//2) → (N//16, K//2*16) for preshuffle format
w_triton = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
# Convert B_scale_sh from CK view (N_pad, sn_pad) to Triton view (N_pad//32, sn_pad*32).
# Same bytes in memory — just a different 2D interpretation.
sm, sn = B_scale_sh.shape
w_scales_triton = B_scale_sh.view(torch.uint8).reshape(sm // 32, sn * 32)
# Cache output buffer per shape
if shape_key not in _shape_cache:
_shape_cache[shape_key] = torch.empty((m, n), dtype=torch.bfloat16, device='cuda')
print(f"[v36] New shape ({m},{n},{k}): w_triton={w_triton.shape} "
f"w_scales={w_scales_triton.shape}", file=sys.stderr)
out = _shape_cache[shape_key]
gemm_a16wfp4_preshuffle(
A, w_triton, w_scales_triton,
prequant=True,
dtype=torch.bfloat16,
y=out,
)
return out
scrolls · 61 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