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
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.
fp4
MXFP4 GEMM kernel optimized for AMD Instinct MI355X (CDNA 4 architecture).tile-k = 64
TILE_K = 64tile-m = 128
TILE_M = 128tile-n = 128
TILE_N = 128Kernel 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