submission 623450
Arkadip Maitra · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 166 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-623450?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:3b246e659295bdfb4dd6a9a856683af55afc79d042e991f35dea55704a434af5
license declaredunknown
license concludedunknown
authorsArkadip Maitra
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
submission.py166 lines
"""
Optimized MXFP4 GEMM: bf16 A, MXFP4 B -> fused quant+GEMM -> bf16 C.
Primary optimization: Use the Triton gemm_a16wfp4 preshuffle kernel that fuses
A quantization (bf16 -> mxfp4) into the GEMM kernel itself, eliminating:
- dynamic_mxfp4_quant kernel launch for A
- e8m0_shuffle kernel launch for A scales
- HBM round-trip for A_q and A_scale intermediate buffers
Secondary optimization: Output tensor caching — reuse the same output buffer
for repeated calls with the same (m, n) shape, avoiding torch.empty overhead.
Fallback: Reference approach (separate quant + gemm_a4w4).
"""
from task import input_t, output_t
_initialized = False
_fused_fn = None
_output_cache = {}
_fallback_imports = None
def _try_init_fused():
global _initialized, _fused_fn
if _initialized:
return
_initialized = True
# Try 1: aiter top-level API
try:
import aiter
from aiter import dtypes
for name in ('gemm_a16wfp4', 'gemm_a16w4'):
fn = getattr(aiter, name, None)
if fn is not None:
def _wrap_toplevel(A, B_sh, B_sc, m, n, k, _fn=fn):
return _fn(A, B_sh, B_sc, dtype=dtypes.bf16, bpreshuffle=True)
_fused_fn = _wrap_toplevel
return
except Exception:
pass
# Try 2: Triton gemm module wrapper
for import_path in [
'aiter.ops.triton.gemm.gemm_a16wfp4',
'aiter.ops.triton.gemm',
]:
try:
import importlib
mod = importlib.import_module(import_path)
for attr in ('gemm_a16wfp4_preshuffle', 'gemm_a16wfp4'):
fn = getattr(mod, attr, None)
if fn is not None:
from aiter import dtypes
def _wrap_module(A, B_sh, B_sc, m, n, k, _fn=fn):
return _fn(A, B_sh, B_sc, dtype=dtypes.bf16, bpreshuffle=True)
_fused_fn = _wrap_module
return
except (ImportError, AttributeError):
pass
# Try 3: Direct Triton kernel launch
try:
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
_gemm_a16wfp4_preshuffle_kernel,
_get_config,
get_splitk,
)
import triton
import torch
def _direct_kernel(A, B_shuffle, B_scale_sh, m, n, k):
config = _get_config(m, n, k, shuffle=True)
if config is not None:
BSM = int(config.get('BLOCK_SIZE_M', 32))
BSN = int(config.get('BLOCK_SIZE_N', 64))
BSK = int(config.get('BLOCK_SIZE_K', 256))
GROUP_SIZE_M = int(config.get('GROUP_SIZE_M', 4))
NUM_KSPLIT = int(config.get('NUM_KSPLIT', 1))
nw = int(config.get('num_warps', 4))
ns = int(config.get('num_stages', 2))
wpe = int(config.get('waves_per_eu', 2))
mink = int(config.get('matrix_instr_nonkdim', 16))
cm = str(config.get('cache_modifier', '.cg'))
else:
BSM = 16 if m <= 16 else 32
BSN = 64
BSK = 256
GROUP_SIZE_M = 4
NUM_KSPLIT = max(1, 4 if k >= 4096 else 2 if k >= 1024 else 1)
nw, ns, wpe, mink = 4, 2, 2, 16
cm = '.cg'
K_half = k // 2
SPLITK_BLOCK_SIZE, BSK, NUM_KSPLIT = get_splitk(K_half, BSK, NUM_KSPLIT)
grid = (triton.cdiv(m, BSM) * triton.cdiv(n, BSN) * NUM_KSPLIT,)
if NUM_KSPLIT > 1:
c = torch.zeros((NUM_KSPLIT, m, n), dtype=torch.bfloat16, device='cuda')
stride_ck = m * n
stride_cm = n
else:
c = torch.empty((m, n), dtype=torch.bfloat16, device='cuda')
stride_ck = 0
stride_cm = n
_gemm_a16wfp4_preshuffle_kernel[grid](
A, B_shuffle, c, B_scale_sh,
m, n, K_half,
A.stride(0), A.stride(1),
B_shuffle.stride(0), B_shuffle.stride(1),
stride_ck, stride_cm, 1,
B_scale_sh.stride(0), B_scale_sh.stride(1),
BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN, BLOCK_SIZE_K=BSK,
GROUP_SIZE_M=GROUP_SIZE_M,
NUM_KSPLIT=NUM_KSPLIT,
SPLITK_BLOCK_SIZE=SPLITK_BLOCK_SIZE,
PREQUANT=True,
num_warps=nw, num_stages=ns,
waves_per_eu=wpe,
matrix_instr_nonkdim=mink,
cache_modifier=cm,
)
if NUM_KSPLIT > 1:
return c.sum(dim=0)
return c
_fused_fn = _direct_kernel
return
except (ImportError, AttributeError):
pass
def _get_fallback_imports():
global _fallback_imports
if _fallback_imports is None:
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
_fallback_imports = (aiter, dtypes, dynamic_mxfp4_quant, e8m0_shuffle)
return _fallback_imports
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[0]
_try_init_fused()
if _fused_fn is not None:
try:
return _fused_fn(A, B_shuffle, B_scale_sh, m, n, k)
except Exception:
pass
aiter, dtypes, dynamic_mxfp4_quant, e8m0_shuffle = _get_fallback_imports()
A_q, A_scale = dynamic_mxfp4_quant(A)
A_q = A_q.view(dtypes.fp4x2)
A_scale = e8m0_shuffle(A_scale).view(dtypes.fp8_e8m0)
return aiter.gemm_a4w4(
A_q, B_shuffle, A_scale, B_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True,
)
scrolls · 166 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