submission 574967
HorizonLiang · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 357 lines, June 9 Researcher Reciprocity License v1.0.
amd_202602_mxfp4_mm_large_self_gemm_64only.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-574967?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:c668a9fe8ed0354066f9965a6328cc55eee600d7eb8fd2f42155b91f72ab54e3
license declaredunknown
license concludedunknown
authorsHorizonLiang
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
amd_202602_mxfp4_mm_large_self_gemm_64only.py357 lines
"""
Hybrid MXFP4 GEMM submission for MI355X.
Small-M keeps the current strongest path. Large benchmark shapes reuse the fast
quantization path, then switch GEMM to Triton's AFP4xWFP4 preshuffled kernel so
we can isolate the custom GEMM ceiling without also changing quantization.
This variant only swaps the 64x7168x2048 benchmark shape.
"""
from __future__ import annotations
import triton
import triton.language as tl
import torch
from task import input_t, output_t
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
@triton.heuristics(
{
"EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0
and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0,
}
)
@triton.jit
def _dynamic_mxfp4_quant_kernel_shuffled(
x_ptr,
x_fp4_ptr,
scale_sh_ptr,
stride_x_m_in,
stride_x_n_in,
stride_x_fp4_m_in,
stride_x_fp4_n_in,
M,
N,
SCALE_N_VALID: tl.constexpr,
SCALE_M_PAD: tl.constexpr,
SCALE_N_PAD: tl.constexpr,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
NUM_ITER: tl.constexpr,
NUM_STAGES: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
EVEN_M_N: tl.constexpr,
):
pid_m = tl.program_id(0)
start_n = tl.program_id(1) * NUM_ITER
stride_x_m = tl.cast(stride_x_m_in, tl.int64)
stride_x_n = tl.cast(stride_x_n_in, tl.int64)
stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
num_quant_blocks: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
if EVEN_M_N:
x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
else:
x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(
tl.float32
)
x_fp4, scale_e8m0 = _mxfp4_quant_op(
x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
)
out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
out_offs = (
out_offs_m[:, None] * stride_x_fp4_m
+ out_offs_n[None, :] * stride_x_fp4_n
)
if EVEN_M_N:
tl.store(x_fp4_ptr + out_offs, x_fp4)
else:
out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
tl.store(x_fp4_ptr + out_offs, x_fp4, mask=out_mask)
scale_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
scale_offs_n = pid_n * num_quant_blocks + tl.arange(0, num_quant_blocks)
scale_offs_0 = scale_offs_m[:, None] // 32
scale_offs_1 = scale_offs_m[:, None] % 32
scale_offs_2 = scale_offs_1 % 16
scale_offs_1 = scale_offs_1 // 16
scale_offs_3 = scale_offs_n[None, :] // 8
scale_offs_4 = scale_offs_n[None, :] % 8
scale_offs_5 = scale_offs_4 % 4
scale_offs_4 = scale_offs_4 // 4
scale_offs = (
scale_offs_1
+ scale_offs_4 * 2
+ scale_offs_2 * 4
+ scale_offs_5 * 64
+ scale_offs_3 * 256
+ scale_offs_0 * 32 * SCALE_N_VALID
)
scale_mask_valid = (scale_offs_m < M)[:, None] & (
scale_offs_n < ((N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE)
)[None, :]
scale_mask_pad = (scale_offs_m < SCALE_M_PAD)[:, None] & (
scale_offs_n < SCALE_N_PAD
)[None, :]
scale_e8m0 = tl.where(scale_mask_valid, scale_e8m0, 127)
tl.store(scale_sh_ptr + scale_offs, scale_e8m0, mask=scale_mask_pad)
_QUANT_CACHE: dict[tuple[tuple[str, int | None], int, int], tuple[torch.Tensor, torch.Tensor]] = {}
_OUTPUT_CACHE: dict[tuple[tuple[str, int | None], int, int], torch.Tensor] = {}
_KERNEL_OVERRIDES: dict[tuple[int, int, int], str] = {
(4, 2880, 512): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",
(16, 2112, 7168): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
(32, 4096, 512): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",
(32, 2880, 512): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",
(64, 7168, 2048): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
(256, 3072, 1536): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
}
_SMALL_M_CONFIG_OVERRIDES: dict[tuple[int, int, int], dict[str, int | str | None]] = {
(4, 2880, 512): {
"BLOCK_SIZE_M": 4,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 2,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
},
(16, 2112, 7168): {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 7,
},
}
_LARGE_SELF_GEMM_CONFIGS: dict[tuple[int, int, int], dict[str, int | str | None]] = {
(64, 7168, 2048): {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 16,
"BLOCK_SIZE_K": 1024,
"GROUP_SIZE_M": 1,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
},
}
def _device_key(device: torch.device) -> tuple[str, int | None]:
return device.type, device.index
def _get_quant_outputs(
m: int,
n: int,
device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor]:
scale_n_valid = triton.cdiv(n, 32)
scale_n_pad = triton.cdiv(scale_n_valid, 8) * 8
scale_m_pad = triton.cdiv(m, 256) * 256
key = (_device_key(device), m, n)
cached = _QUANT_CACHE.get(key)
if cached is None:
cached = (
torch.empty((m, n // 2), dtype=torch.uint8, device=device),
torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device),
)
_QUANT_CACHE[key] = cached
return cached
def _get_output(m: int, n: int, device: torch.device) -> torch.Tensor:
key = (_device_key(device), m, n)
cached = _OUTPUT_CACHE.get(key)
if cached is None:
cached = torch.empty((m, n), dtype=torch.bfloat16, device=device)
_OUTPUT_CACHE[key] = cached
return cached
def _quant_mxfp4_shuffled(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
m, n = x.shape
x_fp4, scale_sh = _get_quant_outputs(m, n, x.device)
if m <= 32:
num_iter = 1
block_size_m = triton.next_power_of_2(m)
block_size_n = 32
num_stages = 1
else:
num_iter = 4
block_size_m = 64
block_size_n = 64
num_stages = 2
if n <= 16384:
block_size_m = 32
block_size_n = 128
if n <= 1024:
num_iter = 1
num_stages = 1
block_size_n = max(32, min(256, triton.next_power_of_2(n)))
block_size_m = min(8, triton.next_power_of_2(m))
grid = (
triton.cdiv(m, block_size_m),
triton.cdiv(n, block_size_n * num_iter),
)
_dynamic_mxfp4_quant_kernel_shuffled[grid](
x,
x_fp4,
scale_sh,
*x.stride(),
*x_fp4.stride(),
M=m,
N=n,
SCALE_N_VALID=triton.cdiv(n, 32),
SCALE_M_PAD=triton.cdiv(m, 256) * 256,
SCALE_N_PAD=triton.cdiv(triton.cdiv(n, 32), 8) * 8,
BLOCK_SIZE_M=block_size_m,
BLOCK_SIZE_N=block_size_n,
NUM_ITER=num_iter,
NUM_STAGES=num_stages,
MXFP4_QUANT_BLOCK_SIZE=32,
num_warps=4,
num_stages=1,
)
return x_fp4, scale_sh
def _view_preshuffled_weight(
weight_shuffle: torch.Tensor,
weight_scale_shuffle: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
n, k_half = weight_shuffle.shape
weight_u8 = (
weight_shuffle.view(torch.uint8)
if weight_shuffle.dtype is not torch.uint8
else weight_shuffle
)
scale_u8 = (
weight_scale_shuffle.view(torch.uint8)
if weight_scale_shuffle.dtype is not torch.uint8
else weight_scale_shuffle
)
scale_m, scale_n = scale_u8.shape
weight_triton = weight_u8.view(n // 16, k_half * 16)
scale_triton = scale_u8.view(scale_m // 32, scale_n * 32)[: n // 32]
return weight_triton, scale_triton
def _view_preshuffled_activation_scales(
scale_shuffle: torch.Tensor,
m: int,
) -> torch.Tensor:
scale_u8 = (
scale_shuffle.view(torch.uint8)
if scale_shuffle.dtype is not torch.uint8
else scale_shuffle
)
scale_m, scale_n = scale_u8.shape
return scale_u8.view(scale_m // 32, scale_n * 32)[: m // 32]
def custom_kernel(data: input_t) -> output_t:
import aiter
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import gemm_afp4wfp4_preshuffle
a, b, b_q, b_shuffle, b_scale_sh = data
del b, b_q
a = a.contiguous()
m, k = a.shape
n = b_shuffle.shape[0]
shape = (m, n, k)
if m <= 32:
w_triton, w_scale_triton = _view_preshuffled_weight(b_shuffle, b_scale_sh)
out = _get_output(m, n, a.device)
config = _SMALL_M_CONFIG_OVERRIDES.get(shape)
return gemm_a16wfp4_preshuffle(
a,
w_triton,
w_scale_triton,
prequant=True,
dtype=torch.bfloat16,
y=out,
config=config,
)
a_q, a_scale_sh = _quant_mxfp4_shuffled(a)
large_config = _LARGE_SELF_GEMM_CONFIGS.get(shape)
if large_config is not None:
w_triton, w_scale_triton = _view_preshuffled_weight(b_shuffle, b_scale_sh)
a_scale_triton = _view_preshuffled_activation_scales(a_scale_sh, m)
out = _get_output(m, n, a.device)
return gemm_afp4wfp4_preshuffle(
a_q,
w_triton,
a_scale_triton,
w_scale_triton,
dtype=torch.bfloat16,
y=out,
config=large_config,
)
kernel_name = _KERNEL_OVERRIDES.get(shape)
if kernel_name is None:
return aiter.gemm_a4w4(
a_q.view(dtypes.fp4x2),
b_shuffle,
a_scale_sh.view(dtypes.fp8_e8m0),
b_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
out = _get_output(m, n, a.device)
gemm_a4w4_asm(
a_q.view(dtypes.fp4x2),
b_shuffle,
a_scale_sh.view(dtypes.fp8_e8m0),
b_scale_sh,
out,
kernel_name,
None,
1.0,
0.0,
True,
0,
)
return out
scrolls · 357 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