submission 585029
zhaohb · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 223 lines, June 9 Researcher Reciprocity License v1.0.
submission_opt.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-585029?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:99e4949dbfbc42aab141442194727f8ea9828cb5011b93992101b013dd077334
license declaredunknown
license concludedunknown
authorszhaohb
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Optimized FP4 quant + FP4 GEMM path for MI355X.num-warps = 1
num_warps = 1split-k
split_k = 0stages = 1
num_stages=1,Kernel source
submission_opt.py223 lines
"""
Optimized FP4 quant + FP4 GEMM path for MI355X.
Optimizations:
1. Move imports/constants to module scope to trim Python overhead.
2. Reuse quant/output buffers across repeated calls with the same shapes.
3. Bypass `aiter.gemm_a4w4()`'s per-call output allocation when internals are available.
4. Fall back to the public aiter path if any internal fast path is unavailable.
"""
from task import input_t, output_t
import aiter
import torch
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant as _public_dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle as _public_e8m0_shuffle
try:
import triton
from aiter.ops.triton._triton_kernels.quant.quant import _dynamic_mxfp4_quant_kernel
_HAS_FAST_QUANT = True
except Exception:
triton = None
_dynamic_mxfp4_quant_kernel = None
_HAS_FAST_QUANT = False
try:
from aiter.ops.gemm_op_a4w4 import (
gemm_a4w4_asm,
gemm_a4w4_blockscale,
get_GEMM_config,
)
_HAS_FAST_GEMM = True
except Exception:
gemm_a4w4_asm = None
gemm_a4w4_blockscale = None
get_GEMM_config = None
_HAS_FAST_GEMM = False
BF16 = dtypes.bf16
FP4X2 = dtypes.fp4x2
FP8_E8M0 = dtypes.fp8_e8m0
_PUBLIC_GEMM_A4W4 = aiter.gemm_a4w4
_FAST_PATH_ENABLED = _HAS_FAST_QUANT and _HAS_FAST_GEMM
_QUANT_CACHE: dict[tuple[object, int, int], tuple[torch.Tensor, torch.Tensor, torch.Tensor]] = {}
_SCALE_SHUFFLE_CACHE: dict[
tuple[object, int, int],
tuple[torch.Tensor, torch.Tensor, torch.Tensor],
] = {}
_OUT_CACHE: dict[tuple[object, int, int], torch.Tensor] = {}
def _device_key(device: torch.device) -> tuple[str, int | None]:
return (device.type, device.index)
def _quant_mxfp4_cached(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
m, k = x.shape
cache_key = (_device_key(x.device), m, k)
cached = _QUANT_CACHE.get(cache_key)
if cached is None:
x_fp4_u8 = torch.empty((m, k // 2), dtype=torch.uint8, device=x.device)
scale_u8 = torch.empty((m, (k + 31) // 32), dtype=torch.uint8, device=x.device)
x_fp4 = x_fp4_u8.view(FP4X2)
_QUANT_CACHE[cache_key] = (x_fp4_u8, scale_u8, x_fp4)
else:
x_fp4_u8, scale_u8, x_fp4 = cached
if m <= 32:
num_iter = 1
block_size_m = triton.next_power_of_2(m)
block_size_n = 32
num_warps = 1
num_stages_cfg = 1
else:
num_iter = 4
block_size_m = 64
block_size_n = 64
num_warps = 4
num_stages_cfg = 2
if k <= 16384:
block_size_m = 32
block_size_n = 128
if k <= 1024:
num_iter = 1
num_stages_cfg = 1
num_warps = 4
block_size_n = max(32, min(256, triton.next_power_of_2(k)))
block_size_m = min(8, triton.next_power_of_2(m))
grid = (
triton.cdiv(m, block_size_m),
triton.cdiv(k, block_size_n * num_iter),
)
_dynamic_mxfp4_quant_kernel[grid](
x,
x_fp4_u8,
scale_u8,
*x.stride(),
*x_fp4_u8.stride(),
*scale_u8.stride(),
M=m,
N=k,
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_cfg,
num_warps=num_warps,
waves_per_eu=0,
num_stages=1,
)
return x_fp4, scale_u8
def _shuffle_e8m0_cached(scale_u8: torch.Tensor) -> torch.Tensor:
m, n = scale_u8.shape
padded_m = (m + 255) // 256 * 256
padded_n = (n + 7) // 8 * 8
cache_key = (_device_key(scale_u8.device), padded_m, padded_n)
cached = _SCALE_SHUFFLE_CACHE.get(cache_key)
if cached is None:
scale_pad = torch.empty((padded_m, padded_n), dtype=torch.uint8, device=scale_u8.device)
scale_sh = torch.empty((padded_m, padded_n), dtype=torch.uint8, device=scale_u8.device)
scale_sh_view = scale_sh.view(padded_m // 32, padded_n // 8, 4, 16, 2, 2)
_SCALE_SHUFFLE_CACHE[cache_key] = (scale_pad, scale_sh, scale_sh_view)
else:
scale_pad, scale_sh, scale_sh_view = cached
scale_pad[:m, :n] = scale_u8
scale_sh_view.copy_(
scale_pad.view(padded_m // 32, 2, 16, padded_n // 8, 2, 4).permute(0, 3, 5, 2, 4, 1)
)
return scale_sh.view(FP8_E8M0)
def _get_out_buffer(device: torch.device, m: int, n: int) -> torch.Tensor:
padded_m = (m + 31) // 32 * 32
cache_key = (_device_key(device), padded_m, n)
out = _OUT_CACHE.get(cache_key)
if out is None:
out = torch.empty((padded_m, n), dtype=BF16, device=device)
_OUT_CACHE[cache_key] = out
return out
def _gemm_cached(
a_q: torch.Tensor,
b_shuffle: torch.Tensor,
a_scale_sh: torch.Tensor,
b_scale_sh: torch.Tensor,
) -> torch.Tensor:
m = a_q.shape[0]
n = b_shuffle.shape[0]
k = a_q.shape[1] * 2
out = _get_out_buffer(a_q.device, m, n)
ck_config = get_GEMM_config(m, n, k)
split_k = 0
kernel_name = ""
if ck_config is not None:
split_k = ck_config.get("splitK", None)
kernel_name = ck_config["kernelName"]
if ck_config is not None and "_ZN" not in kernel_name:
split_k = 0 if split_k is None else split_k
gemm_a4w4_blockscale(a_q, b_shuffle, a_scale_sh, b_scale_sh, out, splitK=split_k)
else:
gemm_a4w4_asm(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
out,
kernel_name,
None,
1.0,
0.0,
True,
log2_k_split=split_k,
)
return out[:m]
def _public_path(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor) -> torch.Tensor:
x_fp4, scale_e8m0 = _public_dynamic_mxfp4_quant(a)
a_q = x_fp4.view(FP4X2)
a_scale_sh = _public_e8m0_shuffle(scale_e8m0).view(FP8_E8M0)
return _PUBLIC_GEMM_A4W4(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
dtype=BF16,
bpreshuffle=True,
)
def custom_kernel(data: input_t) -> output_t:
global _FAST_PATH_ENABLED
a, _, _, b_shuffle, b_scale_sh = data
if not a.is_contiguous():
a = a.contiguous()
if _FAST_PATH_ENABLED:
try:
a_q, a_scale = _quant_mxfp4_cached(a)
a_scale_sh = _shuffle_e8m0_cached(a_scale)
return _gemm_cached(a_q, b_shuffle, a_scale_sh, b_scale_sh)
except Exception:
_FAST_PATH_ENABLED = False
return _public_path(a, b_shuffle, b_scale_sh)
scrolls · 223 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