Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
14.5µs
#514 of 1143
2026-03-24

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.

fp4MXFP4 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