submission 715655
xiehuanyi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 89 lines, June 9 Researcher Reciprocity License v1.0.
submission_v18.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-715655?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:6342cf42287e73bf012cfd933f20e900fa150aeaaac721e6a829cbc52b9c4837
license declaredunknown
license concludedunknown
authorsxiehuanyi
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 GEMM v18: Optimal hybrid - Triton for small K+M, ASM for rest.Kernel source
submission_v18.py89 lines
"""
MXFP4 GEMM v18: Optimal hybrid - Triton for small K+M, ASM for rest.
Based on empirical benchmarks:
- K <= 2048 and M <= 32: Triton gemm_afp4wfp4 is 15-21% faster
- K > 2048 or M > 32: ASM gemm_a4w4 with log2_k_split is better
"""
from task import input_t, output_t
import torch
import sys
_first_call = True
def custom_kernel(data: input_t) -> output_t:
global _first_call
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
B = B.contiguous()
m, k = A.shape
n = B.shape[0]
A_fp4_raw, A_scale_raw = dynamic_mxfp4_quant(A)
# Trigger JIT on first call
if _first_call:
_first_call = False
A_scale_sh = e8m0_shuffle(A_scale_raw).view(dtypes.fp8_e8m0)
_ = aiter.gemm_a4w4(
A_fp4_raw.view(dtypes.fp4x2), B_shuffle,
A_scale_sh, B_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True,
)
# Decision: Triton for small K+M, ASM for rest
use_triton = (m <= 32 and k <= 2048)
if use_triton:
_, B_scale_raw = dynamic_mxfp4_quant(B)
try:
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import gemm_afp4wfp4
n_blocks = (n + 63) // 64
num_ksplit = max(1, min(304 // max(n_blocks, 1), k // 128))
config = {
'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 128,
'GROUP_SIZE_M': 8, 'NUM_KSPLIT': num_ksplit,
'num_warps': 4, 'num_stages': 2, 'waves_per_eu': 0,
'matrix_instr_nonkdim': 16, 'cache_modifier': '.cg',
}
out = gemm_afp4wfp4(
A_fp4_raw, B_q.view(A_fp4_raw.dtype),
A_scale_raw, B_scale_raw,
dtype=torch.bfloat16, config=config,
)
return out[:m, :n]
except Exception as e:
print(f"[OPT] Triton failed: {e}", file=sys.stderr)
# ASM path with k_split for small M, default for large M
A_fp4 = A_fp4_raw.view(dtypes.fp4x2)
A_scale_sh = e8m0_shuffle(A_scale_raw).view(dtypes.fp8_e8m0)
if m < 64:
kname = '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E'
n_blocks = (n + 127) // 128
best_ks = 0
for ks in range(1, 8):
if (k >> ks) < 64:
break
best_ks = ks
if n_blocks * (1 << ks) >= 304:
break
out = torch.empty((m, n), dtype=torch.bfloat16, device='cuda')
torch.ops.aiter.gemm_a4w4_asm(
A_fp4, B_shuffle, A_scale_sh, B_scale_sh,
out, kname, bpreshuffle=True, log2_k_split=best_ks,
)
return out
return aiter.gemm_a4w4(
A_fp4, B_shuffle, A_scale_sh, B_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True,
)
scrolls · 89 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