Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
24.0µs
#866 of 1143
2026-03-25

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.

fp4FP4 quant + FP4 GEMM highly optimized implementation for AMD MI355X.
persistent-kernel6. Persistent kernel design for batch processing
split-k5. Split-K for large K dimensions

Kernel 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