submission 530427
josusanmartin · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 161 lines, June 9 Researcher Reciprocity License v1.0.
mxfp4_v87_best_hybrid.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-530427?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:2c1d9519b5fe6e7795383818d22e49109080137c5b044a7811bb24e2179a2466
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
- M=64,K=2048: ASM splitK=2 → 13.7µs (v62 finding)tile-m = 8
- M=4,K=512: Fused GEMM with BLOCK_SIZE_M=8 → 7.38µs (v80 finding)Kernel source
mxfp4_v87_best_hybrid.py161 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
Version 87: Best hybrid combining ALL discoveries:
- M=4,K=512: Fused GEMM with BLOCK_SIZE_M=8 → 7.38µs (v80 finding)
- M=16,K=7168: Fused GEMM with KSPLIT=14 → 16.8µs (v79 finding)
- M=32,K=512: Fused GEMM with DEFAULT config → 9.47/9.50µs (v77 finding)
- M=64,K=2048: ASM splitK=2 → 13.7µs (v62 finding)
- M=256,K=1536: ASM splitK=1 → 12.5µs (v62 finding)
Expected geomean: ~11.1µs (vs 12.7µs baseline)
"""
import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.utility.fp4_utils import _dynamic_mxfp4_quant_kernel_asm_layout
from task import input_t, output_t
@triton.jit
def _mxfp4_quant_op_asm_exact(
x,
BLOCK_SIZE_N,
BLOCK_SIZE_M,
MXFP4_QUANT_BLOCK_SIZE,
):
E8_BIAS: tl.constexpr = 127
E2_BIAS: tl.constexpr = 1
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
scale_e8m0_unbiased = tl.log2(amax).floor() - 2
scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
quant_scale = tl.exp2(-scale_e8m0_unbiased)
qx = x * quant_scale
qx = qx.to(tl.uint32, bitcast=True)
s = qx & 0x80000000
e = (qx >> 23) & 0xFF
m = qx & 0x7FFFFF
adjusted_exponents = tl.core.sub(E8_BIAS, e + 1, sanitize_overflow=False)
m = tl.where(e < E8_BIAS, (0x400000 | (m >> 1)) >> adjusted_exponents, m)
e = tl.maximum(e, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)
e2m1_tmp = tl.minimum((((e << 2) | (m >> 21)) + 1) >> 1, 0x7)
e2m1_value = ((s >> 28) | e2m1_tmp).to(tl.uint8)
e2m1_value = tl.reshape(
e2m1_value, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2]
)
evens, odds = tl.split(e2m1_value)
x_fp4 = evens | (odds << 4)
x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
import aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 as _kernel_module
_kernel_module._mxfp4_quant_op = _mxfp4_quant_op_asm_exact
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
_bf16 = dtypes.bf16
_fp4x2 = dtypes.fp4x2
_fp8_e8m0 = dtypes.fp8_e8m0
_kernel_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
# ASM splitK for large-M shapes
_ASM_SPLITK = {
(64, 7168, 2048): 2, # Proven best: 4 K-splits
(256, 3072, 1536): 1, # Proven best: 2 K-splits
}
# Fused configs - only override where needed
_FUSED_CONFIGS = {
# M=4: BLOCK_SIZE_M=8 reduces padding waste (4→8 vs 4→32)
(4, 2880, 512): {
"BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 1,
"waves_per_eu": 2, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": 1,
},
# M=16: Split-K=14 for better CU utilization
(16, 2112, 7168): {
"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 1,
"waves_per_eu": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": 14,
},
# M=32: Use default config (BLOCK_SIZE_N=64 gives more N-tiles → better CU util)
# No config override → uses default
}
_QUANT_BLOCK = 32
_QUANT_TILE = 128
_bufs = {}
def _get_asm_bufs(m, k, n, device):
x_fp4 = torch.empty((m, k >> 1), dtype=torch.uint8, device=device)
sN = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK
sN_pad = ((sN + 7) >> 3) << 3
sM_pad = ((m + 255) >> 8) << 8
scale = torch.empty((sM_pad, sN_pad), dtype=torch.uint8, device=device)
padded_m = ((m + 31) >> 5) << 5
out = torch.empty((padded_m, n), dtype=_bf16, device=device)
return x_fp4, scale, sN, sN_pad, sM_pad, out, padded_m
@torch.inference_mode()
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]
key = (m, n, k)
if key in _ASM_SPLITK:
# ASM 2-kernel path for M>=64 shapes
if key not in _bufs:
_bufs[key] = ('asm', _get_asm_bufs(m, k, n, A.device))
_, (x_fp4, scale, sN, sN_pad, sM_pad, out, padded_m) = _bufs[key]
grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, sN_pad)
_dynamic_mxfp4_quant_kernel_asm_layout[grid](
A, x_fp4, scale,
A.stride(0), A.stride(1),
x_fp4.stride(0), x_fp4.stride(1),
scale.stride(0), scale.stride(1),
M=m, N=k, scaleN=sN,
scaleM_pad=sM_pad, scaleN_pad=sN_pad,
BLOCK_SIZE=_QUANT_TILE,
MXFP4_QUANT_BLOCK_SIZE=_QUANT_BLOCK,
SCALING_MODE=0, SHUFFLE=True,
)
splitK = _ASM_SPLITK[key]
gemm_a4w4_asm(
x_fp4.view(_fp4x2), B_shuffle, scale.view(_fp8_e8m0), B_scale_sh,
out, _kernel_32x128,
bpreshuffle=True, log2_k_split=splitK,
)
return out[:m]
else:
# Fused preshuffle GEMM for M<=32 shapes
if key not in _bufs:
_bufs[key] = ('fused', torch.empty((m, n), dtype=torch.bfloat16, device=A.device))
_, out = _bufs[key]
w = B_shuffle.view(torch.uint8).reshape(n // 16, k // 2 * 16)
sm, sn = B_scale_sh.shape
w_scales = B_scale_sh.view(torch.uint8).reshape(sm // 32, sn * 32)
config = _FUSED_CONFIGS.get(key)
return gemm_a16wfp4_preshuffle(A, w, w_scales, prequant=True, y=out, config=config)
scrolls · 161 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 529480.
⋯ 1 unchanged lines#!POPCORN gpu MI355X"""- Autogenerated MXFP4 candidate: cached_quant_h8_2112_7168_h8_32x128_s3- Strategy: cached_quant- Generated by autosubmit_amd_mxfp4_mm.py.- """- from __future__ import annotations+ Version 87: Best hybrid combining ALL discoveries:+ - M=4,K=512: Fused GEMM with BLOCK_SIZE_M=8 → 7.38µs (v80 finding)+ - M=16,K=7168: Fused GEMM with KSPLIT=14 → 16.8µs (v79 finding)+ - M=32,K=512: Fused GEMM with DEFAULT config → 9.47/9.50µs (v77 finding)+ - M=64,K=2048: ASM splitK=2 → 13.7µs (v62 finding)+ - M=256,K=1536: ASM splitK=1 → 12.5µs (v62 finding)- import csv- import os- from pathlib import Path-+ Expected geomean: ~11.1µs (vs 12.7µs baseline)+ """import torch+ import triton+ import triton.language as tl+ import aiter+ from aiter import dtypes+ from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm+ from aiter.utility.fp4_utils import _dynamic_mxfp4_quant_kernel_asm_layout+ from task import input_t, output_t- _AITER_BASE = Path("/home/runner/aiter")- _CUSTOM_CONFIG_PATH = "/tmp/submission_20260311_auto_cached_quant_h8_2112_7168_h8_32x128_s3_6a1a37d4.csv"- _CUSTOM_ROWS = (- {'cu_num': '256', 'M': '4', 'N': '2880', 'K': '512', 'kernelName': '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E', 'splitK': '2'},- {'cu_num': '256', 'M': '16', 'N': '2112', 'K': '7168', 'kernelName': '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E', 'splitK': '3'},- {'cu_num': '256', 'M': '32', 'N': '2880', 'K': '512', 'kernelName': '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E', 'splitK': '2'},- {'cu_num': '256', 'M': '32', 'N': '4096', 'K': '512', 'kernelName': '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E', 'splitK': '2'},- {'cu_num': '256', 'M': '8', 'N': '2112', 'K': '7168', 'kernelName': '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E', 'splitK': '3'},- )- _CUSTOM_KEYS = {(row["M"], row["N"], row["K"]) for row in _CUSTOM_ROWS}-- def _looks_like_config(path: Path) -> bool:- try:- with path.open(newline="") as f:- reader = csv.DictReader(f)- return reader.fieldnames is not None and "kernelName" in reader.fieldnames and "splitK" in reader.fieldnames- except Exception:- return False--- def _find_default_config() -> Path | None:- candidates = (- _AITER_BASE / "aiter" / "configs" / "a4w4_blockscale_tuned_gemm.csv",- _AITER_BASE / "hsa" / "configs" / "a4w4_tuned_gemm.csv",- _AITER_BASE / "configs" / "a4w4_tuned_gemm.csv",+ @triton.jit+ def _mxfp4_quant_op_asm_exact(+ x,+ BLOCK_SIZE_N,+ BLOCK_SIZE_M,+ MXFP4_QUANT_BLOCK_SIZE,+ ):+ E8_BIAS: tl.constexpr = 127+ E2_BIAS: tl.constexpr = 1+ NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE+ x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)+ amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)+ amax = amax.to(tl.int32, bitcast=True)+ amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000+ amax = amax.to(tl.float32, bitcast=True)+ scale_e8m0_unbiased = tl.log2(amax).floor() - 2+ scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)+ bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127+ quant_scale = tl.exp2(-scale_e8m0_unbiased)+ qx = x * quant_scale+ qx = qx.to(tl.uint32, bitcast=True)+ s = qx & 0x80000000+ e = (qx >> 23) & 0xFF+ m = qx & 0x7FFFFF+ adjusted_exponents = tl.core.sub(E8_BIAS, e + 1, sanitize_overflow=False)+ m = tl.where(e < E8_BIAS, (0x400000 | (m >> 1)) >> adjusted_exponents, m)+ e = tl.maximum(e, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)+ e2m1_tmp = tl.minimum((((e << 2) | (m >> 21)) + 1) >> 1, 0x7)+ e2m1_value = ((s >> 28) | e2m1_tmp).to(tl.uint8)+ e2m1_value = tl.reshape(+ e2m1_value, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2])- for path in candidates:- if path.exists() and _looks_like_config(path):- return path- for path in sorted(_AITER_BASE.rglob("*a4w4*.csv")):- if _looks_like_config(path):- return path- return None+ evens, odds = tl.split(e2m1_value)+ x_fp4 = evens | (odds << 4)+ x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)+ return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)- def _write_merged_config():- base_path = _find_default_config()- fieldnames = ["cu_num", "M", "N", "K", "kernelName", "splitK"]- rows = []+ import aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 as _kernel_module+ _kernel_module._mxfp4_quant_op = _mxfp4_quant_op_asm_exact- if base_path is not None:- with base_path.open(newline="") as f:- reader = csv.DictReader(f)- if reader.fieldnames:- fieldnames = list(reader.fieldnames)- for row in reader:- if (row.get("M", ""), row.get("N", ""), row.get("K", "")) in _CUSTOM_KEYS:- continue- rows.append(row)+ from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle- rows.extend(_CUSTOM_ROWS)- with open(_CUSTOM_CONFIG_PATH, "w", newline="") as f:- writer = csv.DictWriter(f, fieldnames=fieldnames)- writer.writeheader()- writer.writerows(rows)+ _bf16 = dtypes.bf16+ _fp4x2 = dtypes.fp4x2+ _fp8_e8m0 = dtypes.fp8_e8m0+ _kernel_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"- _write_merged_config()- os.environ["AITER_CONFIG_GEMM_A4W4"] = _CUSTOM_CONFIG_PATH+ # ASM splitK for large-M shapes+ _ASM_SPLITK = {+ (64, 7168, 2048): 2, # Proven best: 4 K-splits+ (256, 3072, 1536): 1, # Proven best: 2 K-splits+ }- import aiter- from aiter import QuantType, dtypes- from aiter.utility import fp4_utils+ # Fused configs - only override where needed+ _FUSED_CONFIGS = {+ # M=4: BLOCK_SIZE_M=8 reduces padding waste (4→8 vs 4→32)+ (4, 2880, 512): {+ "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 1,+ "waves_per_eu": 2, "matrix_instr_nonkdim": 16,+ "cache_modifier": ".cg", "NUM_KSPLIT": 1,+ },+ # M=16: Split-K=14 for better CU utilization+ (16, 2112, 7168): {+ "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 1,+ "waves_per_eu": 1, "matrix_instr_nonkdim": 16,+ "cache_modifier": ".cg", "NUM_KSPLIT": 14,+ },+ # M=32: Use default config (BLOCK_SIZE_N=64 gives more N-tiles → better CU util)+ # No config override → uses default+ }- from task import input_t, output_t+ _QUANT_BLOCK = 32+ _QUANT_TILE = 128- _DTYPES = dtypes- _FP4_UTILS = fp4_utils- _TRITON_QUANT = aiter.get_triton_quant(QuantType.per_1x32)- _GEMM = aiter.gemm_a4w4- _A_QUANT_BUFFERS = {}- _OUT_CACHE = {}+ _bufs = {}- _QUANT_BLOCK_SIZE_BY_MK = {- (4, 512): 32,- (32, 512): 32,- }+ def _get_asm_bufs(m, k, n, device):+ x_fp4 = torch.empty((m, k >> 1), dtype=torch.uint8, device=device)+ sN = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK+ sN_pad = ((sN + 7) >> 3) << 3+ sM_pad = ((m + 255) >> 8) << 8+ scale = torch.empty((sM_pad, sN_pad), dtype=torch.uint8, device=device)+ padded_m = ((m + 31) >> 5) << 5+ out = torch.empty((padded_m, n), dtype=_bf16, device=device)+ return x_fp4, scale, sN, sN_pad, sM_pad, out, padded_m- def _ceil_div(x: int, y: int) -> int:- return (x + y - 1) // y+ @torch.inference_mode()+ 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]+ key = (m, n, k)- def _get_quant_buffers(m: int, k: int, device):- scale_n = _ceil_div(k, 32)- scale_n_pad = _ceil_div(scale_n, 8) * 8- scale_m_pad = _ceil_div(m, 256) * 256- key = (device.type, device.index, m, k, scale_m_pad, scale_n_pad)- buffers = _A_QUANT_BUFFERS.get(key)- if buffers is None:- a_q = torch.empty((m, k // 2), dtype=torch.uint8, device=device)- a_scale = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)- buffers = (a_q, a_scale, scale_n, scale_m_pad, scale_n_pad)- _A_QUANT_BUFFERS[key] = buffers- return buffers+ if key in _ASM_SPLITK:+ # ASM 2-kernel path for M>=64 shapes+ if key not in _bufs:+ _bufs[key] = ('asm', _get_asm_bufs(m, k, n, A.device))+ _, (x_fp4, scale, sN, sN_pad, sM_pad, out, padded_m) = _bufs[key]+ grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, sN_pad)+ _dynamic_mxfp4_quant_kernel_asm_layout[grid](+ A, x_fp4, scale,+ A.stride(0), A.stride(1),+ x_fp4.stride(0), x_fp4.stride(1),+ scale.stride(0), scale.stride(1),+ M=m, N=k, scaleN=sN,+ scaleM_pad=sM_pad, scaleN_pad=sN_pad,+ BLOCK_SIZE=_QUANT_TILE,+ MXFP4_QUANT_BLOCK_SIZE=_QUANT_BLOCK,+ SCALING_MODE=0, SHUFFLE=True,+ )- def _quantize_a(a):- m = int(a.shape[0])- k = int(a.shape[1])- try:- a_q, a_scale, scale_n, scale_m_pad, scale_n_pad = _get_quant_buffers(m, k, a.device)- block_size = _QUANT_BLOCK_SIZE_BY_MK.get((m, k), 128)- grid = (_ceil_div(m, block_size), scale_n_pad)- _FP4_UTILS._dynamic_mxfp4_quant_kernel_asm_layout[grid](- a,- a_q,- a_scale,- *a.stride(),- *a_q.stride(),- *a_scale.stride(),- M=m,- N=k,- scaleN=scale_n,- scaleM_pad=scale_m_pad,- scaleN_pad=scale_n_pad,- BLOCK_SIZE=block_size,- MXFP4_QUANT_BLOCK_SIZE=32,- SCALING_MODE=0,- SHUFFLE=True,+ splitK = _ASM_SPLITK[key]+ gemm_a4w4_asm(+ x_fp4.view(_fp4x2), B_shuffle, scale.view(_fp8_e8m0), B_scale_sh,+ out, _kernel_32x128,+ bpreshuffle=True, log2_k_split=splitK,)- return a_q.view(_DTYPES.fp4x2), a_scale.view(_DTYPES.fp8_e8m0)- except Exception:- return _TRITON_QUANT(a, shuffle=True)+ return out[:m]+ else:+ # Fused preshuffle GEMM for M<=32 shapes+ if key not in _bufs:+ _bufs[key] = ('fused', torch.empty((m, n), dtype=torch.bfloat16, device=A.device))+ _, out = _bufs[key]- def custom_kernel(data: input_t) -> output_t:- a, _b, _b_q, b_shuffle, b_scale_sh = data- a_q, a_scale_sh = _quantize_a(a)- return _GEMM(- a_q,- b_shuffle,- a_scale_sh,- b_scale_sh,- dtype=dtypes.bf16,- bpreshuffle=True,- )+ w = B_shuffle.view(torch.uint8).reshape(n // 16, k // 2 * 16)+ sm, sn = B_scale_sh.shape+ w_scales = B_scale_sh.view(torch.uint8).reshape(sm // 32, sn * 32)++ config = _FUSED_CONFIGS.get(key)+ return gemm_a16wfp4_preshuffle(A, w, w_scales, prequant=True, y=out, config=config)
scrolls · 286 diff lines total
Best evidence level for this revision: reported
JSON