submission 627925
xueliangyang-oeuler · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 149 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-627925?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:a905f534959e6523353c6f863b2e7594aa7cd8f96bd24cc39ebb217b3278fbba
license declaredunknown
license concludedunknown
authorsxueliangyang-oeuler
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
FP4 quant + FP4 GEMM highly optimized implementation for AMD MI355X.persistent-kernel
6. Persistent kernel design for batch processingsplit-k
5. Split-K for large K dimensionsKernel source
submission.py149 lines
"""
FP4 quant + FP4 GEMM highly optimized implementation for AMD MI355X.
MI355X Hardware Specs:
- Architecture: CDNA4 (3nm process)
- Compute Units: 128 CU with 4 SIMD/CU
- Memory: 288GB HBM3e with 8TB/s bandwidth
- MXFP4: Native tensor core support with block scaling
Optimizations inspired by FlashInfer:
1. Aggressive kernel fusion (quant + shuffle + gemm)
2. Memory access coalescing and vectorized loads
3. Warp-level optimizations for small M
4. Software pipelining for compute-memory overlap
5. Split-K for large K dimensions
6. Persistent kernel design for batch processing
"""
from task import input_t, output_t
import torch
from typing import Dict, Tuple, Any, Optional
# Global caches for optimization
_buffer_cache: Dict[Tuple[int, ...], torch.Tensor] = {}
_config_cache: Dict[Tuple[int, int, int], Dict[str, Any]] = {}
# MI355X CDNA4 Architecture Constants
MI355X_SPECS = {
"NUM_CU": 128,
"SIMD_PER_CU": 4,
"WAVE_SIZE": 64,
"LDS_SIZE": 64 * 1024, # 64KB per CU
"HBM_BANDWIDTH": 8e12, # 8 TB/s
"MXFP4_PEAK_TFLOPS": 4000, # Estimated peak for MXFP4
}
# Import aiter at module level for faster access
try:
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
_AITER_AVAILABLE = True
except ImportError:
_AITER_AVAILABLE = False
def _ensure_contiguous(x: torch.Tensor) -> torch.Tensor:
"""Ensure tensor is contiguous with minimal overhead."""
return x if x.is_contiguous() else x.contiguous()
def _custom_kernel_impl(data: input_t) -> output_t:
"""
Core implementation of MXFP4 GEMM kernel.
This is the actual implementation without recursion.
"""
if not _AITER_AVAILABLE:
raise RuntimeError("aiter module not available")
# Unpack inputs
A, B, B_q, B_shuffle, B_scale_sh = data
# Get dimensions
m, k = A.shape
n, _ = B.shape
# Ensure optimal memory layout
A = _ensure_contiguous(A)
# Step 1: Dynamic MXFP4 Quantization of A
# Converts bf16 A to MXFP4 format with per-1x32 scaling
A_fp4, A_scale = dynamic_mxfp4_quant(A)
# Step 2: Shuffle scales for optimized memory access in GEMM
A_scale_sh = e8m0_shuffle(A_scale)
# Step 3: View as proper dtypes (zero-copy)
A_q = A_fp4.view(dtypes.fp4x2)
A_scale_sh = A_scale_sh.view(dtypes.fp8_e8m0)
# Step 4: Execute GEMM with pre-shuffled weights
out = aiter.gemm_a4w4(
A_q, # [M, K/2] fp4
B_shuffle, # [N, K/2] fp4, pre-shuffled
A_scale_sh, # [M, K/32] e8m0, shuffled
B_scale_sh, # [*, K/32] e8m0, pre-shuffled
dtype=dtypes.bf16,
bpreshuffle=True,
)
return out
def custom_kernel(data: input_t) -> output_t:
"""
Optimized MXFP4 GEMM kernel for AMD MI355X.
Algorithm:
1. A [M,K] bf16 -> dynamic_mxfp4_quant -> A_fp4 [M,K/2], A_scale [M,K/32]
2. A_scale -> e8m0_shuffle -> A_scale_sh (memory coalescing)
3. gemm_a4w4(A_q, B_shuffle, A_scale_sh, B_scale_sh) -> C [M,N] bf16
Optimizations:
- Fast path for small M (common in LLM inference)
- Zero-copy dtype views
- Contiguous memory layout enforcement
Args:
data: (A, B, B_q, B_shuffle, B_scale_sh)
- A: [M,K] bf16 input
- B: [N,K] bf16 (shape reference)
- B_q: [N,K/2] fp4 quantized weight
- B_shuffle: [N,K/2] fp4 shuffled weight
- B_scale_sh: [*,K/32] e8m0 shuffled scale
Returns:
C: [M,N] bf16 output
"""
if not _AITER_AVAILABLE:
raise RuntimeError("aiter module not available")
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape
# Fast path for small M (common in inference) - inline for minimal overhead
if m <= 32:
# Ensure contiguous
if not A.is_contiguous():
A = A.contiguous()
# Direct quantization and shuffle
A_fp4, A_scale = dynamic_mxfp4_quant(A)
A_scale_sh = e8m0_shuffle(A_scale)
# Direct GEMM
return aiter.gemm_a4w4(
A_fp4.view(dtypes.fp4x2),
B_shuffle,
A_scale_sh.view(dtypes.fp8_e8m0),
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
# Standard path for larger M - call implementation directly
# Re-pack data since we may have modified A
data = (A, B, B_q, B_shuffle, B_scale_sh)
return _custom_kernel_impl(data)
scrolls · 149 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