submission 737563
chenxingqiang · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 190 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-737563?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:99efdae2c170eabe5b610f766ab95145ad5a9f08bb974e64b981a41c2ac01c0b
license declaredunknown
license concludedunknown
authorschenxingqiang
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
FP4 GEMM: hybrid with pre-allocated quant buffers.num-warps = 1
NUM_WARPS = 1stages = 1
NUM_STAGES = 1tile-k = 256
M <= 32: a16wfp4_preshuffle prequant with BLOCK_K=256, ns=2, .cg for K<=1024.tile-m = 64
BLOCK_SIZE_M = 64tile-n = 32
BLOCK_SIZE_N = 32Kernel source
submission.py190 lines
"""
FP4 GEMM: hybrid with pre-allocated quant buffers.
M <= 32: a16wfp4_preshuffle prequant with BLOCK_K=256, ns=2, .cg for K<=1024.
M >= 64: gemm_afp4wfp4 with pre-allocated A quant buffers to avoid torch.empty.
"""
import torch
import triton
from aiter.ops.triton._triton_kernels.quant.quant import _dynamic_mxfp4_quant_kernel
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import gemm_afp4wfp4_
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle_, gemm_a16wfp4_
from aiter.ops.triton.utils.common_utils import serialize_dict
from task import input_t, output_t
_QUANT_BUF: dict = {}
_OUT_BUF: dict = {}
_PRESHUFFLE_CACHE: dict = {}
_UNSHUFFLE_CACHE: dict = {}
_VIEW_CACHE_LIMIT = 4
_SMALL_K_CONFIG = {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
}
_SMALL_K_CONFIG_SER = serialize_dict(_SMALL_K_CONFIG)
_LARGE_K_CONFIG = {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 4,
}
_LARGE_K_CONFIG_SER = serialize_dict(_LARGE_K_CONFIG)
def _cache_put(cache: dict, key, value, limit: int):
if len(cache) >= limit and key not in cache:
cache.clear()
cache[key] = value
def _e8m0_unshuffle(scale_sh: torch.Tensor, orig_n: int, k_groups: int) -> torch.Tensor:
s = scale_sh.view(torch.uint8)
sm, sn = s.shape
s = (
s.view(sm // 32, sn // 8, 4, 16, 2, 2)
.permute(0, 5, 3, 1, 4, 2)
.contiguous()
.view(sm, sn)
)
return s[:orig_n, :k_groups].contiguous()
def _fast_mxfp4_quant(x: torch.Tensor):
M, N = x.shape
buf_key = (M, N)
cached = _QUANT_BUF.get(buf_key)
if cached is not None:
x_fp4, blockscale = cached
else:
x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=x.device)
k_groups = (N + 31) // 32
blockscale = torch.empty((k_groups, M), dtype=torch.uint8, device=x.device).T
if len(_QUANT_BUF) > 4:
_QUANT_BUF.clear()
_QUANT_BUF[buf_key] = (x_fp4, blockscale)
if M <= 32:
BLOCK_SIZE_M = triton.next_power_of_2(M)
BLOCK_SIZE_N = 32
NUM_WARPS = 1
NUM_STAGES = 1
NUM_ITER = 1
else:
NUM_ITER = 4
BLOCK_SIZE_M = 64
BLOCK_SIZE_N = 64
NUM_WARPS = 4
NUM_STAGES = 2
if N <= 16384:
BLOCK_SIZE_M = 32
BLOCK_SIZE_N = 128
if N <= 1024:
NUM_ITER = 1
NUM_STAGES = 1
NUM_WARPS = 4
BLOCK_SIZE_N = min(256, triton.next_power_of_2(N))
BLOCK_SIZE_N = max(32, BLOCK_SIZE_N)
BLOCK_SIZE_M = min(8, triton.next_power_of_2(M))
grid = (
triton.cdiv(M, BLOCK_SIZE_M),
triton.cdiv(N, BLOCK_SIZE_N * NUM_ITER),
)
_dynamic_mxfp4_quant_kernel[grid](
x, x_fp4, blockscale,
*x.stride(), *x_fp4.stride(), *blockscale.stride(),
M=M, N=N,
MXFP4_QUANT_BLOCK_SIZE=32,
SCALING_MODE=0,
NUM_ITER=NUM_ITER,
BLOCK_SIZE_M=BLOCK_SIZE_M,
BLOCK_SIZE_N=BLOCK_SIZE_N,
NUM_STAGES=NUM_STAGES,
num_warps=NUM_WARPS,
waves_per_eu=0,
num_stages=1,
)
return x_fp4, blockscale
def custom_kernel(data: input_t) -> output_t:
A, _, B_q, B_shuffle, B_scale_sh = data
if not A.is_contiguous():
A = A.contiguous()
m, k = A.shape
n = B_q.shape[0]
out_key = (m, n)
y = _OUT_BUF.get(out_key)
if y is None:
y = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
if len(_OUT_BUF) > 8:
_OUT_BUF.clear()
_OUT_BUF[out_key] = y
if m <= 32:
preshuffle_key = (id(B_shuffle), id(B_scale_sh), n, k)
cached = _PRESHUFFLE_CACHE.get(preshuffle_key)
if cached is None:
B_sh_u8 = B_shuffle.view(torch.uint8)
B_sh_ps = B_sh_u8.reshape(n // 16, (k // 2) * 16)
bs_u8 = B_scale_sh.view(torch.uint8)
sm, sn = bs_u8.shape
B_scale_ps = bs_u8.view(sm // 32, sn * 32)
_cache_put(_PRESHUFFLE_CACHE, preshuffle_key, (B_sh_ps, B_scale_ps), _VIEW_CACHE_LIMIT)
else:
B_sh_ps, B_scale_ps = cached
if k <= 1024:
cfg = _SMALL_K_CONFIG_SER
elif k >= 4096:
cfg = _LARGE_K_CONFIG_SER
else:
cfg = None
return gemm_a16wfp4_preshuffle_(
A, B_sh_ps, B_scale_ps,
prequant=True, dtype=torch.bfloat16,
y=y,
config=cfg,
)
else:
unshuffle_key = (id(B_scale_sh), n, k // 32)
B_scale_raw = _UNSHUFFLE_CACHE.get(unshuffle_key)
if B_scale_raw is None:
B_scale_raw = _e8m0_unshuffle(B_scale_sh, n, k // 32)
_cache_put(_UNSHUFFLE_CACHE, unshuffle_key, B_scale_raw, _VIEW_CACHE_LIMIT)
B_q_u8 = B_q.view(torch.uint8)
if m <= 128:
return gemm_a16wfp4_(
A, B_q_u8, B_scale_raw,
dtype=torch.bfloat16,
y=y,
config=None,
)
else:
A_fp4, A_scale = _fast_mxfp4_quant(A)
return gemm_afp4wfp4_(
A_fp4, B_q_u8, A_scale, B_scale_raw,
dtype=torch.bfloat16,
y=y,
config=None,
)
scrolls · 190 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