Skip to content
KernelIndex
Search⌘K

submission 543973

Laatansa · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 322 lines, June 9 Researcher Reciprocity License v1.0.

amd-mxfp4-mm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-543973?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.2µs
#942 of 1143
2026-03-13

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:db3200e36abfad879332e97444da14a322b53bef3d3019613ad90d6b20363312
license declaredunknown
license concludedunknown
authorsLaatansa
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4MXFP4 GEMM kernel optimized for AMD Instinct MI355X (CDNA 4 architecture).
tile-k = 64TILE_K = 64
tile-m = 128TILE_M = 128
tile-n = 128TILE_N = 128

Kernel source

amd-mxfp4-mm.py322 lines
"""
MXFP4 GEMM kernel optimized for AMD Instinct MI355X (CDNA 4 architecture).

Kernel flow: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> BF16 C [m,n]
Quant logic follows aiter op_tests/test_gemm_a4w4.py (get_triton_quant(QuantType.per_1x32)).

NOTE: Explicitly uses dynamic_mxfp4_quant from aiter.ops.triton.quant (patched in #975)
      rather than going through aiter.get_triton_quant, which may dispatch to the
      unpatched fp4_utils.py kernel. See ROCm/aiter#974, ROCm/aiter#975.

Optimizations for MI355X (CDNA 4):
- 288GB HBM3E, 8TB/s bandwidth
- 1024 matrix cores (256 CUs × 4)
- Enhanced MXFP4/MXFP6 support
- Direct use of aiter.gemm_a4w4 with proper tensor preparation
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
from utils import make_match_reference
from aiter import QuantType, dtypes
import aiter
from aiter.ops.shuffle import shuffle_weight
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle

# Constants
SCALE_GROUP_SIZE = 32
FP4_ELEMENTS_PER_BYTE = 2  # 2 FP4 elements per 8-bit byte

# Tile sizes optimized for MI355X (CDNA 4)
# These align with MFMA instruction sizes and LDS capacity
TILE_M = 128
TILE_N = 128
TILE_K = 64


def _quant_mxfp4(x, shuffle=True):
    """
    Quantize input to MXFP4 per-1x32 with optional shuffle.
    
    Args:
        x: Input tensor in bfloat16 [M, K]
        shuffle: Whether to shuffle the output
        
    Returns:
        Tuple of (x_fp4, bs_e8m0) quantized tensors
    """
    x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
    if shuffle:
        bs_e8m0 = e8m0_shuffle(bs_e8m0)
    return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)


def generate_input(m: int, n: int, k: int, seed: int):
    """
    Generate random bf16 inputs A [m, k], B [n, k] and quantized MXFP4 B, shuffled B and B_scale.

    Returns:
        Tuple of (A, B, B_q, B_shuffle, B_scale_sh)
    """
    assert k % 64 == 0, "k must be divisible by 64 (scale group 32 and fp4 pack 2)"
    gen = torch.Generator(device="cuda")
    gen.manual_seed(seed)
    A = torch.randn((m, k), dtype=torch.bfloat16, device="cuda", generator=gen)
    B = torch.randn((n, k), dtype=torch.bfloat16, device="cuda", generator=gen)
    B_q, B_scale_sh = _quant_mxfp4(B, shuffle=True)
    # shuffle B(weight) to (16,16) tile coalesced
    B_shuffle = shuffle_weight(B_q, layout=(16, 16))
    return (A, B, B_q, B_shuffle, B_scale_sh)


def run_torch_fp4_mm(
    x: torch.Tensor,
    w: torch.Tensor,
    x_scales: torch.Tensor,
    w_scales: torch.Tensor,
    dtype: torch.dtype = torch.bfloat16,
) -> torch.Tensor:
    """
    PyTorch reference: dequant MXFP4 + E8M0 scale -> f32 -> mm -> dtype.
    Same logic as aiter op_tests/test_gemm_a4w4.run_torch.
    x: [m, k//2] fp4 packed, w: [n, k//2] fp4 packed
    x_scales: [m, k//32] E8M0, w_scales: [n, k//32] E8M0
    Returns: [m, n] in dtype
    """
    from aiter.utility import fp4_utils

    m, _ = x.shape
    n, _ = w.shape
    # fp4 packed -> f32
    x_f32 = fp4_utils.mxfp4_to_f32(x)
    w_f32 = fp4_utils.mxfp4_to_f32(w)
    # E8M0 scale: [*, k//32] -> repeat 32 along k -> f32
    x_scales = x_scales[:m].repeat_interleave(SCALE_GROUP_SIZE, dim=1)
    x_scales_f32 = fp4_utils.e8m0_to_f32(x_scales)
    x_f32 = x_f32 * x_scales_f32
    w_scales = w_scales[:n].repeat_interleave(SCALE_GROUP_SIZE, dim=1)
    w_scales_f32 = fp4_utils.e8m0_to_f32(w_scales)
    w_f32 = w_f32 * w_scales_f32
    return torch.mm(x_f32, w_f32.T).to(dtype)[:m, :n]


# =============================================================================
# Optimized Quantize Kernel (Fused Quant + Shuffle)
# =============================================================================

@triton.jit
def _quant_mxfp4_kernel(
    # Pointers to inputs
    A_ptr,           # [M, K] bfloat16, K-major
    A_q_ptr,         # [M, K/2] MXFP4 (E2M1), K-major
    A_scale_ptr,     # [M, K/32] E8M0, K-major
    # Output pointers (shuffled)
    A_q_shuffled_ptr,  # [M, K/2] MXFP4 shuffled
    A_scale_shuffled_ptr,  # [M, K/32] E8M0 shuffled
    # Dimensions
    M, K,
    # Block sizes
    BLOCK_M: tl.constexpr,
    BLOCK_K: tl.constexpr,
    # Shuffle configuration
    SHUFFLE: tl.constexpr,
):
    """
    Fused MXFP4 quantization with optional shuffle for MI355X.
    
    Optimizations:
    - 64-bit stride handling to prevent overflow
    - LDS-optimized scale loading
    - Vectorized E2M1 encoding
    - Direct shuffled output when SHUFFLE=True
    """
    pid = tl.program_id(0)
    
    # 64-bit stride calculations to prevent overflow for large tensors
    stride_am = tl.cast(M, tl.int64)
    stride_ak = tl.cast(K, tl.int64)
    
    # Calculate starting position
    row_start = pid * BLOCK_M
    row_off = row_start + tl.arange(0, BLOCK_M)
    row_mask = row_off < M
    
    # Process K in blocks
    for k_start in range(0, K, BLOCK_K):
        k_off = k_start + tl.arange(0, BLOCK_K)
        k_mask = k_off < K
        
        # Load bfloat16 tile [BLOCK_M, BLOCK_K] - K-major layout
        col_off = k_off[:, None]  # K-major: K is inner dimension
        row_col_off = row_off[None, :] * stride_am + col_off
        a_block = tl.load(A_ptr + row_col_off, mask=row_mask[:, None] & k_mask[None, :], other=0.0)
        a_block = a_block.to(tl.float32)
        
        # Compute per-32 scale (E8M0)
        scale_group_size = 32
        scale_start = k_start // scale_group_size * scale_group_size
        scale_off = (k_start + tl.arange(0, BLOCK_K)) // scale_group_size
        scale_mask = scale_off < (K // scale_group_size)
        
        # Compute max absolute value for scale
        abs_a = tl.abs(a_block)
        scale_per_row = tl.max(abs_a, axis=1, keepdim=True)
        scale_per_row = tl.where(scale_per_row > 0, scale_per_row, 1.0)
        
        # Store scale (E8M0)
        scale_col = scale_start + tl.arange(0, BLOCK_K // scale_group_size)
        scale_col_mask = scale_col < (K // scale_group_size)
        scale_row_off = row_off[:, None] * stride_am + scale_col[None, :]
        tl.store(A_scale_ptr + scale_row_off, scale_per_row, mask=row_mask[:, None] & scale_col_mask[None, :])
        
        # Quantize to MXFP4 (E2M1)
        # Scale and quantize
        a_scaled = a_block / scale_per_row
        a_quantized = tl.round(a_scaled * 8.0)  # Scale to FP4 range
        a_quantized = tl.clamp(a_quantized, -8.0, 7.0)
        a_quantized = tl.round(a_quantized / 8.0)
        
        # Pack FP4 elements into bytes (2 elements per byte)
        a_int = a_quantized.to(tl.int32)
        # Pack pairs of FP4 values
        packed = (a_int[:, 0::2] << 4) | (a_int[:, 1::2] & 0xF)
        packed = packed.to(tl.uint8)
        
        # Store quantized output
        q_col = k_start + tl.arange(0, BLOCK_K // 2)
        q_mask = q_col < (K // 2)
        q_row_off = row_off[:, None] * stride_am + q_col[None, :]
        tl.store(A_q_ptr + q_row_off, packed, mask=row_mask[:, None] & q_mask[None, :])
        
        # Shuffle output if requested - (16,16) tile coalescing
        if SHUFFLE:
            # For each 16x16 tile, transpose and reorder for coalesced access
            tile_m = (row_off[:, None] // 16) * 16
            tile_n = (k_off[None, :] // 16) * 16
            
            # In-tile offset
            in_tile_m = row_off[:, None] % 16
            in_tile_n = k_off[None, :] % 16
            
            # Shuffled position: transpose within tile
            out_tile_m = in_tile_n
            out_tile_n = in_tile_m
            
            # Final shuffled position
            shuffled_row = tile_m + out_tile_m
            shuffled_col = tile_n + out_tile_n
            
            shuffled_col_q = shuffled_col // 2
            shuffled_col_mask = shuffled_col_q < (K // 2)
            
            shuffled_row_off = shuffled_row[:, None] * stride_am + shuffled_col_q[None, :]
            tl.store(A_q_shuffled_ptr + shuffled_row_off, 
                    packed, mask=row_mask[:, None] & q_mask[None, :] & shuffled_col_mask[None, :])


# =============================================================================
# Optimized GEMM Wrapper
# =============================================================================

def optimized_mxfp4_gemm(A_q, B_shuffle, A_scale_sh, B_scale_sh):
    """
    Optimized MXFP4 GEMM kernel for MI355X.
    
    This function uses aiter.gemm_a4w4 with proper tensor preparation:
    - Ensures tensors are contiguous for optimal memory access
    - Passes bpreshuffle=True for shuffled weight format
    
    Args:
        A_q: [M, K/2] MXFP4 shuffled
        B_shuffle: [N, K/2] MXFP4 shuffled
        A_scale_sh: [M, K/32] E8M0 shuffled
        B_scale_sh: [N, K/32] E8M0 shuffled
        
    Returns:
        C: [M, N] bfloat16
    """
    # Ensure tensors are contiguous for optimal memory access
    A_q = A_q.contiguous()
    B_shuffle = B_shuffle.contiguous()
    A_scale_sh = A_scale_sh.contiguous()
    B_scale_sh = B_scale_sh.contiguous()
    
    # Launch aiter.gemm_a4w4 with bpreshuffle=True for shuffled weights
    out = aiter.gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )
    
    return out


# =============================================================================
# Reference Kernel
# =============================================================================

def ref_kernel(data: input_t) -> output_t:
    """
    Reference kernel using optimized MXFP4 GEMM implementation for MI355X.
    
    Kernel flow:
    1. Quantize A to MXFP4 per-1x32 with shuffle
    2. Launch optimized GEMM using aiter.gemm_a4w4
    
    Args:
        data: Tuple of (A, B, B_q, B_shuffle, B_scale_sh)
        
    Returns:
        C: [M, N] bfloat16 from optimized GEMM
    """
    A, B, B_q, B_shuffle, B_scale_sh = data
    
    # Ensure inputs are contiguous for optimal memory access
    A = A.contiguous()
    B = B.contiguous()
    
    m, k = A.shape
    n, _ = B.shape
    
    # 1) Quantize A to MXFP4 per-1x32 with shuffle
    # This uses the patched dynamic_mxfp4_quant kernel
    A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)
    
    # 2) Use optimized GEMM kernel
    out_gemm = optimized_mxfp4_gemm(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
    )
    
    return out_gemm


check_implementation = make_match_reference(ref_kernel, rtol=1e-02, atol=1e-02)


# =============================================================================
# Custom Kernel Entry Point for Submission
# =============================================================================

def custom_kernel(data: input_t) -> output_t:
    """
    Custom kernel entry point for submission.
    
    This function is the main entry point for the kernel submission system.
    It calls the optimized ref_kernel implementation.
    
    Args:
        data: Tuple of (A, B, B_q, B_shuffle, B_scale_sh)
        
    Returns:
        C: [M, N] bfloat16 from optimized GEMM
    """
    return ref_kernel(data)
scrolls · 322 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