submission 595640
wsxhjnb1 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 355 lines, June 9 Researcher Reciprocity License v1.0.
submission_mxfp4_mm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-595640?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:7c8264e6ac0fd7fe0117df1061318ef6977ab7df81617fa2eaa4d128ce4015f9
license declaredunknown
license concludedunknown
authorswsxhjnb1
imported2026-08-26
Kernel source
submission_mxfp4_mm.py355 lines
from __future__ import annotations
import os
from typing import Tuple
import torch
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
from task import input_t, output_t
# -----------------------------------------------------------------------------
# MXFP4 GEMM skeleton for AMD MI355X qualification round.
#
# Safe default:
# patched dynamic_mxfp4_quant + aiter.gemm_a4w4(bpreshuffle=True)
#
# Where to optimize next:
# 1) _dispatch_bucket
# 2) _run_custom_bucket_kernel
# -----------------------------------------------------------------------------
SCALE_GROUP_SIZE = 32
_WARMED_BUCKETS: set[Tuple[str, int, int, int]] = set()
def _env_flag(name: str, default: bool = False) -> bool:
value = os.getenv(name)
if value is None:
return default
return value.lower() in {"1", "true", "yes", "on"}
def _env_str(name: str, default: str) -> str:
value = os.getenv(name)
return default if value is None else value
def _env_int(name: str, default: int) -> int:
value = os.getenv(name)
if value is None or value == "":
return default
try:
return int(value)
except ValueError:
return default
MM_ENABLE_WARMUP = _env_flag("MM_ENABLE_WARMUP", True)
MM_WARMUP_SYNC = _env_flag("MM_WARMUP_SYNC", False)
MM_WARMUP_SCOPE = _env_str("MM_WARMUP_SCOPE", "bucket_k")
MM_CONTIG_MODE = _env_str("MM_CONTIG_MODE", "never")
MM_BSHUFFLE_CONTIG_MODE = _env_str("MM_BSHUFFLE_CONTIG_MODE", "never")
MM_BSCALE_CONTIG_MODE = _env_str("MM_BSCALE_CONTIG_MODE", "never")
MM_IMPL = _env_str("MM_IMPL", "aiter_ck")
MM_FORCE_TORCH_DEBUG = _env_flag("MM_FORCE_TORCH_DEBUG", False)
MM_ASM_KERNEL_SMALL_K_ALIGNED_N = _env_str("MM_ASM_KERNEL_SMALL_K_ALIGNED_N", "")
MM_ASM_KERNEL_SMALL_K_UNALIGNED_N = _env_str("MM_ASM_KERNEL_SMALL_K_UNALIGNED_N", "")
MM_ASM_KERNEL_VERY_LARGE_K = _env_str("MM_ASM_KERNEL_VERY_LARGE_K", "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E")
MM_ASM_KERNEL_THIN_M_WIDE_N = _env_str("MM_ASM_KERNEL_THIN_M_WIDE_N", "")
MM_ASM_KERNEL_LARGE_M = _env_str("MM_ASM_KERNEL_LARGE_M", "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E")
MM_ASM_KERNEL_DEFAULT = _env_str("MM_ASM_KERNEL_DEFAULT", "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E")
MM_ASM_LOG2_SPLIT_SMALL_K_ALIGNED_N = _env_int("MM_ASM_LOG2_SPLIT_SMALL_K_ALIGNED_N", 0)
MM_ASM_LOG2_SPLIT_SMALL_K_UNALIGNED_N = _env_int("MM_ASM_LOG2_SPLIT_SMALL_K_UNALIGNED_N", 0)
MM_ASM_LOG2_SPLIT_VERY_LARGE_K = _env_int("MM_ASM_LOG2_SPLIT_VERY_LARGE_K", 1)
MM_ASM_LOG2_SPLIT_THIN_M_WIDE_N = _env_int("MM_ASM_LOG2_SPLIT_THIN_M_WIDE_N", 0)
MM_ASM_LOG2_SPLIT_LARGE_M = _env_int("MM_ASM_LOG2_SPLIT_LARGE_M", 0)
MM_ASM_LOG2_SPLIT_DEFAULT = _env_int("MM_ASM_LOG2_SPLIT_DEFAULT", 0)
MM_FAST_CK_PATH = (
MM_IMPL == "aiter_ck"
and not (
MM_ASM_KERNEL_SMALL_K_ALIGNED_N
or MM_ASM_KERNEL_SMALL_K_UNALIGNED_N
or MM_ASM_KERNEL_VERY_LARGE_K
or MM_ASM_KERNEL_THIN_M_WIDE_N
or MM_ASM_KERNEL_LARGE_M
or MM_ASM_KERNEL_DEFAULT
)
and MM_ASM_LOG2_SPLIT_SMALL_K_ALIGNED_N == 0
and MM_ASM_LOG2_SPLIT_SMALL_K_UNALIGNED_N == 0
and MM_ASM_LOG2_SPLIT_VERY_LARGE_K == 0
and MM_ASM_LOG2_SPLIT_THIN_M_WIDE_N == 0
and MM_ASM_LOG2_SPLIT_LARGE_M == 0
and MM_ASM_LOG2_SPLIT_DEFAULT == 0
)
def _maybe_contiguous(x: torch.Tensor, mode: str) -> torch.Tensor:
if mode == "always":
return x.contiguous()
if mode == "never":
return x
return x if x.is_contiguous() else x.contiguous()
def _shape_bucket(m: int, n: int, k: int) -> str:
# Tuned for the public benchmark regimes.
if k <= 512:
if n % 512 == 0:
return "small_k_aligned_n"
return "small_k_unaligned_n"
if k >= 4096:
return "very_large_k"
if m <= 32 and n >= 4096:
return "thin_m_wide_n"
if m >= 128:
return "large_m"
return "default"
def _quant_mxfp4_shuffled(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
# Important: this uses the patched kernel from aiter.ops.triton.quant.
x_fp4, scale_e8m0 = dynamic_mxfp4_quant(x)
scale_e8m0 = e8m0_shuffle(scale_e8m0)
return x_fp4.view(dtypes.fp4x2), scale_e8m0.view(dtypes.fp8_e8m0)
def _run_aiter_ck(
a_q: torch.Tensor,
b_shuffle: torch.Tensor,
a_scale_sh: torch.Tensor,
b_scale_sh: torch.Tensor,
) -> torch.Tensor:
return aiter.gemm_a4w4(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
def _run_aiter_asm(
a_q: torch.Tensor,
b_shuffle: torch.Tensor,
a_scale_sh: torch.Tensor,
b_scale_sh: torch.Tensor,
kernel_name: str = "",
log2_k_split: int | None = None,
) -> torch.Tensor:
m = a_q.shape[0]
n = b_shuffle.shape[0]
out = torch.empty(((m + 31) // 32) * 32, n, dtype=dtypes.bf16, device=a_q.device)
split = log2_k_split if log2_k_split is not None and log2_k_split > 0 else None
aiter.gemm_a4w4_asm(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
out,
kernel_name,
None,
bpreshuffle=True,
log2_k_split=split,
)
return out[:m]
def _warmup_key(bucket: str, m: int, n: int, k: int) -> tuple:
if MM_WARMUP_SCOPE == "bucket":
return (bucket,)
if MM_WARMUP_SCOPE == "bucket_k":
return (bucket, k)
return (bucket, m, n, k)
def _default_asm_kernel(bucket: str) -> str:
if bucket == "small_k_aligned_n":
return MM_ASM_KERNEL_SMALL_K_ALIGNED_N
if bucket == "small_k_unaligned_n":
return MM_ASM_KERNEL_SMALL_K_UNALIGNED_N
if bucket == "very_large_k":
return MM_ASM_KERNEL_VERY_LARGE_K
if bucket == "thin_m_wide_n":
return MM_ASM_KERNEL_THIN_M_WIDE_N
if bucket == "large_m":
return MM_ASM_KERNEL_LARGE_M
return MM_ASM_KERNEL_DEFAULT
def _default_asm_split(bucket: str) -> int:
if bucket == "small_k_aligned_n":
return MM_ASM_LOG2_SPLIT_SMALL_K_ALIGNED_N
if bucket == "small_k_unaligned_n":
return MM_ASM_LOG2_SPLIT_SMALL_K_UNALIGNED_N
if bucket == "very_large_k":
return MM_ASM_LOG2_SPLIT_VERY_LARGE_K
if bucket == "thin_m_wide_n":
return MM_ASM_LOG2_SPLIT_THIN_M_WIDE_N
if bucket == "large_m":
return MM_ASM_LOG2_SPLIT_LARGE_M
return MM_ASM_LOG2_SPLIT_DEFAULT
def _resolve_asm_spec(bucket: str) -> tuple[str, int]:
return _default_asm_kernel(bucket), _default_asm_split(bucket)
def _run_custom_bucket_kernel(
bucket: str,
a: torch.Tensor,
a_q: torch.Tensor,
a_scale_sh: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
) -> torch.Tensor:
# TODO(user): replace this function with your bucket-specific Triton / Helion /
# asm implementation. Keep the signature stable so the dispatcher and warmup
# logic do not need to change.
_ = (bucket, a)
kernel_name, log2_k_split = _resolve_asm_spec(bucket)
return _run_aiter_asm(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
kernel_name=kernel_name,
log2_k_split=log2_k_split,
)
def _dispatch_bucket(
a: torch.Tensor,
a_q: torch.Tensor,
a_scale_sh: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
) -> torch.Tensor:
if MM_FAST_CK_PATH:
return _run_aiter_ck(a_q, b_shuffle, a_scale_sh, b_scale_sh)
m, k = a.shape
n = b_shuffle.shape[0]
bucket = _shape_bucket(m, n, k)
impl = MM_IMPL
kernel_name, log2_k_split = _resolve_asm_spec(bucket)
if impl == "custom":
return _run_custom_bucket_kernel(bucket, a, a_q, a_scale_sh, b_shuffle, b_scale_sh)
if kernel_name:
return _run_aiter_asm(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
kernel_name=kernel_name,
log2_k_split=log2_k_split,
)
if impl == "aiter_asm":
return _run_aiter_asm(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
kernel_name="",
log2_k_split=log2_k_split,
)
return _run_aiter_ck(a_q, b_shuffle, a_scale_sh, b_scale_sh)
def _run_selected_impl(
a: torch.Tensor,
a_q: torch.Tensor,
a_scale_sh: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
) -> torch.Tensor:
if MM_IMPL == "aiter_ck":
m, k = a.shape
n = b_shuffle.shape[0]
bucket = _shape_bucket(m, n, k)
kernel_name, log2_k_split = _resolve_asm_spec(bucket)
if kernel_name or log2_k_split > 0:
return _run_aiter_asm(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
kernel_name=kernel_name,
log2_k_split=log2_k_split,
)
return _run_aiter_ck(a_q, b_shuffle, a_scale_sh, b_scale_sh)
return _dispatch_bucket(a, a_q, a_scale_sh, b_shuffle, b_scale_sh)
def _maybe_warmup(
a: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
) -> None:
if not MM_ENABLE_WARMUP:
return
m, k = a.shape
n = b_shuffle.shape[0]
bucket = _warmup_key(_shape_bucket(m, n, k), m, n, k)
if bucket in _WARMED_BUCKETS:
return
a_q, a_scale_sh = _quant_mxfp4_shuffled(a)
if MM_FAST_CK_PATH:
_ = _run_aiter_ck(a_q, b_shuffle, a_scale_sh, b_scale_sh)
else:
_ = _run_selected_impl(a, a_q, a_scale_sh, b_shuffle, b_scale_sh)
if MM_WARMUP_SYNC:
torch.cuda.synchronize()
_WARMED_BUCKETS.add(bucket)
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
"""
Expected input:
(A, B, B_q, B_shuffle, B_scale_sh)
Timed-path rules for this task:
- do NOT reshuffle B inside custom_kernel
- do NOT rebuild B_scale_sh inside custom_kernel
- keep Python-side branching minimal
"""
a, b, b_q, b_shuffle, b_scale_sh = data
# Only A is needed on the timed path. Keep B/B_q available for debug checks.
bshuffle_contig_mode = MM_BSHUFFLE_CONTIG_MODE if MM_BSHUFFLE_CONTIG_MODE != "if_needed" else MM_CONTIG_MODE
bscale_contig_mode = MM_BSCALE_CONTIG_MODE if MM_BSCALE_CONTIG_MODE != "if_needed" else MM_CONTIG_MODE
a = _maybe_contiguous(a, MM_CONTIG_MODE)
b_shuffle = _maybe_contiguous(b_shuffle, bshuffle_contig_mode)
b_scale_sh = _maybe_contiguous(b_scale_sh, bscale_contig_mode)
_maybe_warmup(a, b_shuffle, b_scale_sh)
a_q, a_scale_sh = _quant_mxfp4_shuffled(a)
if MM_FORCE_TORCH_DEBUG:
# Kept as a non-failing knob so you can switch on extra instrumentation
# in local experiments without changing the submission surface.
_ = (b, b_q)
if MM_FAST_CK_PATH:
return _run_aiter_ck(a_q, b_shuffle, a_scale_sh, b_scale_sh)
return _run_selected_impl(a, a_q, a_scale_sh, b_shuffle, b_scale_sh)
scrolls · 355 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