submission 622982
dannywillowliu-uchi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 80 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-622982?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:003a196fea323dd4a96f0ad2b633252cdd60f2356c3707e8f63596a9bc71ef67
license declaredunknown
license concludedunknown
authorsdannywillowliu-uchi
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 GEMM v6 - Preshuffle path: single kernel launch with zero-copy format conversion.Kernel source
submission.py80 lines
"""
MXFP4 GEMM v6 - Preshuffle path: single kernel launch with zero-copy format conversion.
gemm_a16wfp4_preshuffle does fused bf16->MXFP4 quant + GEMM, accepts pre-shuffled weights
and scales directly. No separate unshuffle or quant kernel needed.
Fallback: V5 hybrid for environments without preshuffle support.
"""
import torch
from task import input_t, output_t
_preshuffle_gemm = None
_fused_gemm = None
_triton_gemm = None
try:
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle as _pg
_preshuffle_gemm = _pg
except ImportError:
pass
try:
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4 as _fg
_fused_gemm = _fg
except ImportError:
pass
try:
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import gemm_afp4wfp4 as _tg
_triton_gemm = _tg
except ImportError:
pass
from aiter.ops.triton.quant import dynamic_mxfp4_quant
_cache = {}
def _unshuffle_e8m0(scale_sh):
sm, sn = scale_sh.shape
return (scale_sh.reshape(sm // 32, sn // 8, 4, 16, 2, 2)
.permute(0, 5, 3, 1, 4, 2).contiguous().reshape(sm, sn))
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape
n = B_q.shape[0]
key = (m, n, k)
out = _cache.get(key)
if out is None:
out = torch.empty(m, n, dtype=torch.bfloat16, device=A.device)
_cache[key] = out
if _preshuffle_gemm is not None:
# Zero-copy conversion to preshuffle format:
# weights: (N, K//2) -> (N//16, K//2*16) via reshape (same bytes)
# scales: e8m0_shuffle format (padded_N, K//32) -> (N//32, K) via reshape+narrow
k_fp4 = k // 2
w_pre = B_shuffle.view(torch.uint8).reshape(n // 16, k_fp4 * 16)
ws_pre = B_scale_sh.view(torch.uint8).reshape(-1, k)[:n // 32]
_preshuffle_gemm(A, w_pre, ws_pre, dtype=torch.bfloat16, y=out)
return out
# Fallback: V5 hybrid path
bq_u8 = B_q.view(torch.uint8)
bscale_raw = _unshuffle_e8m0(B_scale_sh.view(torch.uint8))
if _fused_gemm is not None and m * k <= 50000:
_fused_gemm(A, bq_u8, bscale_raw, dtype=torch.bfloat16, y=out)
elif _triton_gemm is not None:
A_q, A_scale = dynamic_mxfp4_quant(A)
_triton_gemm(
A_q.view(torch.uint8), bq_u8,
A_scale.view(torch.uint8), bscale_raw,
dtype=torch.bfloat16, y=out,
)
elif _fused_gemm is not None:
_fused_gemm(A, bq_u8, bscale_raw, dtype=torch.bfloat16, y=out)
return out
scrolls · 80 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