submission 533545
josusanmartin · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 287 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-533545?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:48d7a8501210e3d4b0bdb38d6e8b55e5eeafc82ccd92a823db9432e7e538cde3
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15
Kernel source
submission.py287 lines
# Write your code here# Write your code here# Write your code here#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
from __future__ import annotations
import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op as _mxfp4_quant_op_even
from task import input_t, output_t
import aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 as _kernel_module
_kernel_module._mxfp4_quant_op = _mxfp4_quant_op_even
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
_BF16 = dtypes.bf16
_FP4X2 = dtypes.fp4x2
_FP8_E8M0 = dtypes.fp8_e8m0
_KERNEL_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_PUBLIC_SMALL = {
(4, 2880, 512): {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
(16, 2112, 7168): {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 8,
},
(32, 4096, 512): {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
(32, 2880, 512): {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
},
}
_PUBLIC_LARGE = {
(64, 7168, 2048): 2,
(256, 3072, 1536): 1,
}
_HIDDEN_SHAPES = {
(8, 2112, 7168),
(16, 3072, 1536),
(64, 3072, 1536),
(256, 2880, 512),
}
_QUANT_BLOCK = 32
_QUANT_TILE = 128
_OUT_PAD_BF16 = 32
_BUFS = {}
@triton.jit
def _dynamic_mxfp4_quant_kernel_even_asm_layout(
x_ptr,
x_fp4_ptr,
bs_ptr,
stride_x_m,
stride_x_n,
stride_x_fp4_m,
stride_x_fp4_n,
stride_bs_m,
stride_bs_n,
M: tl.constexpr,
N: tl.constexpr,
scaleN: tl.constexpr,
scaleM_pad: tl.constexpr,
scaleN_pad: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
SHUFFLE: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
stride_x_m = tl.cast(stride_x_m, tl.int64)
stride_x_n = tl.cast(stride_x_n, tl.int64)
stride_x_fp4_m = tl.cast(stride_x_fp4_m, tl.int64)
stride_x_fp4_n = tl.cast(stride_x_fp4_n, tl.int64)
x_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
x_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE)
x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
x = tl.load(x_ptr + x_offs, mask=x_mask).to(tl.float32)
out_tensor, bs_e8m0 = _mxfp4_quant_op_even(
x,
MXFP4_QUANT_BLOCK_SIZE,
BLOCK_SIZE,
MXFP4_QUANT_BLOCK_SIZE,
)
out_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
out_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE // 2 + tl.arange(
0, MXFP4_QUANT_BLOCK_SIZE // 2
)
out_offs = (
out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
)
out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)
bs_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
bs_offs_n = pid_n
if SHUFFLE:
bs_offs_0 = bs_offs_m[:, None] // 32
bs_offs_1 = bs_offs_m[:, None] % 32
bs_offs_2 = bs_offs_1 % 16
bs_offs_1 = bs_offs_1 // 16
bs_offs_3 = bs_offs_n[None, :] // 8
bs_offs_4 = bs_offs_n[None, :] % 8
bs_offs_5 = bs_offs_4 % 4
bs_offs_4 = bs_offs_4 // 4
bs_offs = (
bs_offs_1
+ bs_offs_4 * 2
+ bs_offs_2 * 4
+ bs_offs_5 * 64
+ bs_offs_3 * 256
+ bs_offs_0 * 32 * scaleN
)
bs_mask1 = (bs_offs_m < M)[:, None] & (bs_offs_n < scaleN)[None, :]
bs_mask2 = (bs_offs_m < scaleM_pad)[:, None] & (bs_offs_n < scaleN_pad)[None, :]
bs_e8m0 = tl.where(bs_mask1, bs_e8m0, 127)
tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask2)
else:
bs_offs = bs_offs_m[:, None] * stride_bs_m + bs_offs_n[None, :] * stride_bs_n
bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < N)[None, :]
tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask)
def _e8m0_shuffle_safe(scale: torch.Tensor) -> torch.Tensor:
m, n = scale.shape
scale_padded = torch.empty(
((m + 255) // 256) * 256,
((n + 7) // 8) * 8,
dtype=scale.dtype,
device=scale.device,
)
scale_padded.fill_(0x7F)
scale_padded[:m, :n] = scale
sm, sn = scale_padded.shape
return (
scale_padded.view(sm // 32, 2, 16, sn // 8, 2, 4)
.permute(0, 3, 5, 2, 4, 1)
.contiguous()
.view(sm, sn)
)
def _safe_wrapper(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor):
a_q_raw, a_scale = dynamic_mxfp4_quant(a.contiguous())
a_scale_sh = _e8m0_shuffle_safe(a_scale)
return aiter.gemm_a4w4(
a_q_raw.view(_FP4X2),
b_shuffle,
a_scale_sh.view(_FP8_E8M0),
b_scale_sh,
dtype=_BF16,
bpreshuffle=True,
)
def _get_large_bufs(m: int, k: int, n: int, device):
x_fp4 = torch.empty((m, k >> 1), dtype=torch.uint8, device=device)
scale_n = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK
scale_n_pad = ((scale_n + 7) >> 3) << 3
scale_m_pad = ((m + 255) >> 8) << 8
scale = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)
padded_m = ((m + 31) >> 5) << 5
out = torch.empty_strided(
(padded_m, n),
(n + _OUT_PAD_BF16, 1),
dtype=_BF16,
device=device,
)
return x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
a, b, _b_q, b_shuffle, b_scale_sh = data
m, k = a.shape
n = b.shape[0]
key = (int(m), int(n), int(k))
if key in _HIDDEN_SHAPES or (key not in _PUBLIC_SMALL and key not in _PUBLIC_LARGE):
return _safe_wrapper(a, b_shuffle, b_scale_sh)
if key in _PUBLIC_LARGE:
if key not in _BUFS:
_BUFS[key] = ("large", _get_large_bufs(m, k, n, a.device))
_, (x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out) = _BUFS[key]
grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, scale_n_pad)
_dynamic_mxfp4_quant_kernel_even_asm_layout[grid](
a,
x_fp4,
scale,
a.stride(0),
a.stride(1),
x_fp4.stride(0),
x_fp4.stride(1),
scale.stride(0),
scale.stride(1),
M=m,
N=k,
scaleN=scale_n,
scaleM_pad=scale_m_pad,
scaleN_pad=scale_n_pad,
BLOCK_SIZE=_QUANT_TILE,
MXFP4_QUANT_BLOCK_SIZE=_QUANT_BLOCK,
SHUFFLE=True,
)
gemm_a4w4_asm(
x_fp4.view(_FP4X2),
b_shuffle,
scale.view(_FP8_E8M0),
b_scale_sh,
out,
_KERNEL_32X128,
bpreshuffle=True,
log2_k_split=_PUBLIC_LARGE[key],
)
return out[:m]
if key not in _BUFS:
_BUFS[key] = ("small", torch.empty((m, n), dtype=_BF16, device=a.device))
_, out = _BUFS[key]
w = b_shuffle.view(torch.uint8).reshape(n // 16, k // 2 * 16)
sm, sn = b_scale_sh.shape
w_scales = b_scale_sh.view(torch.uint8).reshape(sm // 32, sn * 32)
return gemm_a16wfp4_preshuffle(
a,
w,
w_scales,
prequant=True,
y=out,
config=_PUBLIC_SMALL.get(key),
)
scrolls · 287 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 532764.
- #!POPCORN leaderboard amd-mxfp4-mm+ # Write your code here# Write your code here# Write your code here#!POPCORN leaderboard amd-mxfp4-mm#!POPCORN gpu MI355Xfrom __future__ import annotations⋯ 43 unchanged lines"waves_per_eu": 1,"matrix_instr_nonkdim": 16,"cache_modifier": ".cg",- "NUM_KSPLIT": 7,+ "NUM_KSPLIT": 8,},(32, 4096, 512): {"BLOCK_SIZE_M": 8,⋯ 13 unchanged lines"BLOCK_SIZE_K": 512,"GROUP_SIZE_M": 1,"num_warps": 4,- "num_stages": 3,+ "num_stages": 2,"waves_per_eu": 1,"matrix_instr_nonkdim": 16,- "cache_modifier": ".cg",+ "cache_modifier": None,"NUM_KSPLIT": 1,},}⋯ 12 unchanged lines_QUANT_BLOCK = 32_QUANT_TILE = 128+ _OUT_PAD_BF16 = 32_BUFS = {}⋯ 117 unchanged linesscale_m_pad = ((m + 255) >> 8) << 8scale = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)padded_m = ((m + 31) >> 5) << 5- out = torch.empty((padded_m, n), dtype=_BF16, device=device)+ out = torch.empty_strided(+ (padded_m, n),+ (n + _OUT_PAD_BF16, 1),+ dtype=_BF16,+ device=device,+ )return x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out
scrolls · 49 diff lines total
Best evidence level for this revision: reported
JSON