submission 692216
ZainHaider20 · 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-692216?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:d1e4f35889ccc6a6a4e2302ff4d045c64c616783aad55c9b46e9f7b27d1d1e34
license declaredunknown
license concludedunknown
authorsZainHaider20
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
_CUSTOM_CSV = """cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatioKernel source
submission.py190 lines
"""
Use injected-config wrapper GEMM only for the m=4 and m=16 public benchmark
shapes, while keeping the stable direct ASM path for the two m=32 shapes and
the large shapes.
"""
from __future__ import annotations
import os
from pathlib import Path
from task import input_t, output_t
_CUSTOM_CSV = """cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio
256,4,2880,512,29,0,0.0,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E,0.0,0.0,0.0
256,16,2112,7168,29,1,0.0,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E,0.0,0.0,0.0
"""
_CONFIG_PATH = Path("/tmp/aiter_mxfp4_mm_cfg_m4_m16_only.csv")
if not _CONFIG_PATH.exists():
_CONFIG_PATH.write_text(_CUSTOM_CSV, encoding="utf-8")
os.environ["AITER_CONFIG_GEMM_A4W4"] = str(_CONFIG_PATH)
import aiter
import torch
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
_WRAPPER_SHAPES = {
(4, 2880, 512),
(8, 2112, 7168),
(16, 2112, 7168),
}
_ASM_KERNEL_MAP: dict[tuple[int, int, int], tuple[str, int]] = {
(32, 4096, 512): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
0,
),
(32, 2880, 512): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
0,
),
(64, 7168, 2048): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
0,
),
(256, 3072, 1536): (
"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
0,
),
}
_EXACT_CACHED_QUANT_SHAPES = {
(64, 7168, 2048),
(256, 3072, 1536),
}
_EXACT_QUANT_CONFIGS = {
(64, 2048): (4, 32, 128, 4, 2),
(64, 1536): (4, 32, 128, 4, 2),
(256, 1536): (4, 32, 128, 4, 2),
}
_OUT_CACHE: dict[tuple[int, int, torch.device], torch.Tensor] = {}
_QUANT_CACHE: dict[tuple[int, int, torch.device], tuple[torch.Tensor, torch.Tensor]] = {}
_TRITON_BITS = None
def _get_triton_bits():
global _TRITON_BITS
if _TRITON_BITS is not None:
return _TRITON_BITS
import triton
from aiter.ops.triton._triton_kernels.quant.quant import _dynamic_mxfp4_quant_kernel
_TRITON_BITS = (triton, _dynamic_mxfp4_quant_kernel)
return _TRITON_BITS
@torch.inference_mode()
def _get_quant_buffers(m: int, k: int, device: torch.device):
key = (m, k, device)
cached = _QUANT_CACHE.get(key)
if cached is not None:
return cached
x_fp4 = torch.empty((m, k // 2), dtype=torch.uint8, device=device)
blockscale_e8m0 = torch.empty((m, k // 32), dtype=torch.uint8, device=device)
_QUANT_CACHE[key] = (x_fp4, blockscale_e8m0)
return x_fp4, blockscale_e8m0
@torch.inference_mode()
def _quant_mxfp4(x: torch.Tensor):
x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
return x_fp4.view(dtypes.fp4x2), e8m0_shuffle(bs_e8m0).view(dtypes.fp8_e8m0)
_compiled_quant_mxfp4 = torch.compile(_quant_mxfp4)
@torch.inference_mode()
def _quant_mxfp4_cached_exact(x: torch.Tensor):
triton, kernel = _get_triton_bits()
m, k = x.shape
x_fp4, blockscale_e8m0 = _get_quant_buffers(m, k, x.device)
num_iter, block_size_m, block_size_n, num_warps, num_stages = _EXACT_QUANT_CONFIGS[(m, k)]
grid = (
triton.cdiv(m, block_size_m),
triton.cdiv(k, block_size_n * num_iter),
)
kernel[grid](
x,
x_fp4,
blockscale_e8m0,
*x.stride(),
*x_fp4.stride(),
*blockscale_e8m0.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,
num_stages=num_stages,
num_warps=num_warps,
waves_per_eu=0,
)
return x_fp4.view(dtypes.fp4x2), e8m0_shuffle(blockscale_e8m0).view(dtypes.fp8_e8m0)
@torch.inference_mode()
def _get_out(m: int, n: int, device: torch.device) -> torch.Tensor:
key = (m, n, device)
out = _OUT_CACHE.get(key)
if out is None:
out = torch.empty(((m + 31) // 32 * 32, n), dtype=dtypes.bf16, device=device)
_OUT_CACHE[key] = out
return out
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
A, _, _, B_shuffle, B_scale_sh = data
if not A.is_contiguous():
A = A.contiguous()
m, k = A.shape
n = B_shuffle.shape[0]
shape = (m, n, k)
if shape in _WRAPPER_SHAPES:
A_q, A_scale_sh = _compiled_quant_mxfp4(A)
return aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
if shape in _EXACT_CACHED_QUANT_SHAPES:
A_q, A_scale_sh = _quant_mxfp4_cached_exact(A)
else:
A_q, A_scale_sh = _compiled_quant_mxfp4(A)
if shape not in _ASM_KERNEL_MAP:
return aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
kernel_name, split_k = _ASM_KERNEL_MAP[shape]
out = _get_out(m, n, A.device)
aiter.gemm_a4w4_asm(
A_q.view(m, k // 2),
B_shuffle,
A_scale_sh,
B_scale_sh,
out,
kernel_name,
bpreshuffle=True,
log2_k_split=split_k,
)
return out[:m]
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