submission 720595
yszheda · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 127 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-720595?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:f190666b85611495fb5d95e278c33c388f8a7607e21cbdb7be9b2d6a96a5eb6f
license declaredunknown
license concludedunknown
authorsyszheda
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
a_dtype="fp4",Kernel source
submission.py127 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
import sys
import os
try:
import flydsl
_flydsl_root = os.path.dirname(os.path.dirname(flydsl.__file__))
if _flydsl_root not in sys.path:
sys.path.insert(0, _flydsl_root)
except Exception:
pass
def _patch_triton_fp4():
try:
import triton
import triton.language as tl
if hasattr(tl, "core") and hasattr(tl.core, "dtype"):
if not hasattr(tl.core.dtype, "SUPPORTED_TENSOR_DTYPES"):
tl.core.dtype.SUPPORTED_TENSOR_DTYPES = set()
tl.core.dtype.SUPPORTED_TENSOR_DTYPES.add("float4_e2m1fn_x2")
try:
import triton._utils as _tu
if hasattr(_tu, "type_canonicalisation_dict"):
_tu.type_canonicalisation_dict["float4_e2m1fn_x2"] = (
"*kfloat4_e2m1fn_x2"
)
except Exception:
pass
except Exception:
pass
_patch_triton_fp4()
import torch
from task import input_t, output_t
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
_flydsl_available = False
try:
from kernels.preshuffle_gemm import compile_preshuffle_gemm_w4
_flydsl_available = True
except Exception:
pass
_TILE_CONFIGS = [
(8, 32, 128, 256, False, None),
(16, 32, 128, 256, False, None),
(32, 32, 128, 256, False, None),
(64, 64, 256, 256, False, None),
(128, 64, 256, 256, False, None),
(256, 64, 256, 256, False, None),
(512, 64, 256, 256, False, None),
(float("inf"), 64, 256, 256, False, None),
]
def _get_tile_config(M, N, K):
for m_max, tm, tn, tk, async_copy, wpe in _TILE_CONFIGS:
if M <= m_max:
return tm, tn, tk, async_copy, wpe
return 64, 256, 256, False, None
_kernel_cache = {}
def _get_cached_kernel(M, N, K):
cache_key = (M, N, K)
if cache_key in _kernel_cache:
return _kernel_cache[cache_key]
tile_m, tile_n, tile_k, use_async_copy, waves_per_eu = _get_tile_config(M, N, K)
launch_fn = compile_preshuffle_gemm_w4(
M=M,
N=N,
K=K,
tile_m=tile_m,
tile_n=tile_n,
tile_k=tile_k,
a_dtype="fp4",
b_dtype="fp4",
out_dtype="bf16",
lds_stage=2,
use_cshuffle_epilog=False,
waves_per_eu=waves_per_eu,
use_async_copy=use_async_copy,
dsrd_preload=2,
dvmem_preload=2,
)
_kernel_cache[cache_key] = launch_fn
return launch_fn
def _quant_mxfp4(x, shuffle=True):
x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
if shuffle:
bs_e8m0 = e8m0_shuffle(bs_e8m0)
return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape
n, _ = B.shape
A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)
return torch.ops.aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
scrolls · 127 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