submission 534310
josusanmartin · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 342 lines, June 9 Researcher Reciprocity License v1.0.
mxfp4_v424_v423_family3072_small64.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-534310?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:05fc1ea00bcbd3d3329cd13940a14cd67a8393d9c99808bf7c5f85cdd992b588
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15
Kernel source
mxfp4_v424_v423_family3072_small64.py342 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
from __future__ import annotations
"""v423 plus fused routing for the hidden 3072x1536 M=64 family."""
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": 8,
},
(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,
}
_PUBLIC_TEST_SMALL = {
(8, 2112, 7168): {
"BLOCK_SIZE_M": 8,
"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": 14,
},
(16, 3072, 1536): {
"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": 3,
},
}
_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
def _get_route(key):
m, n, k = key
if key in _PUBLIC_SMALL:
return ("small", _PUBLIC_SMALL[key])
if key in _PUBLIC_LARGE:
return ("large", _PUBLIC_LARGE[key])
pair = (n, k)
if pair == (2112, 7168):
if m < 16:
return ("small", _PUBLIC_TEST_SMALL[(8, 2112, 7168)])
return ("small", _PUBLIC_SMALL[(16, 2112, 7168)])
if pair == (3072, 1536):
if m < 128:
return ("small", _PUBLIC_TEST_SMALL[(16, 3072, 1536)])
return ("large", 1)
if pair == (2880, 512):
if m < 16:
return ("small", _PUBLIC_SMALL[(4, 2880, 512)])
if m < 96:
return ("small", _PUBLIC_SMALL[(32, 2880, 512)])
return ("large", 2)
if pair == (4096, 512):
return ("small", _PUBLIC_SMALL[(32, 4096, 512)])
if pair == (7168, 2048):
return ("large", 2)
return None
@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))
route = _get_route(key)
if route is None:
return _safe_wrapper(a, b_shuffle, b_scale_sh)
route_kind, route_value = route
if route_kind == "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=route_value,
)
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=route_value,
)
scrolls · 342 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 533545.
- # Write your code here# Write your code here# Write your code here#!POPCORN leaderboard amd-mxfp4-mm+ #!POPCORN leaderboard amd-mxfp4-mm#!POPCORN gpu MI355Xfrom __future__ import annotations+ """v423 plus fused routing for the hidden 3072x1536 M=64 family."""+import torchimport tritonimport triton.language as tl⋯ 51 unchanged lines"waves_per_eu": 1,"matrix_instr_nonkdim": 16,"cache_modifier": ".cg",- "NUM_KSPLIT": 1,+ "NUM_KSPLIT": 8,},(32, 2880, 512): {"BLOCK_SIZE_M": 8,⋯ 14 unchanged lines(256, 3072, 1536): 1,}- _HIDDEN_SHAPES = {- (8, 2112, 7168),- (16, 3072, 1536),- (64, 3072, 1536),- (256, 2880, 512),+ _PUBLIC_TEST_SMALL = {+ (8, 2112, 7168): {+ "BLOCK_SIZE_M": 8,+ "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": 14,+ },+ (16, 3072, 1536): {+ "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": 3,+ },}_QUANT_BLOCK = 32⋯ 131 unchanged linesreturn x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out+ def _get_route(key):+ m, n, k = key++ if key in _PUBLIC_SMALL:+ return ("small", _PUBLIC_SMALL[key])+ if key in _PUBLIC_LARGE:+ return ("large", _PUBLIC_LARGE[key])++ pair = (n, k)+ if pair == (2112, 7168):+ if m < 16:+ return ("small", _PUBLIC_TEST_SMALL[(8, 2112, 7168)])+ return ("small", _PUBLIC_SMALL[(16, 2112, 7168)])+ if pair == (3072, 1536):+ if m < 128:+ return ("small", _PUBLIC_TEST_SMALL[(16, 3072, 1536)])+ return ("large", 1)+ if pair == (2880, 512):+ if m < 16:+ return ("small", _PUBLIC_SMALL[(4, 2880, 512)])+ if m < 96:+ return ("small", _PUBLIC_SMALL[(32, 2880, 512)])+ return ("large", 2)+ if pair == (4096, 512):+ return ("small", _PUBLIC_SMALL[(32, 4096, 512)])+ if pair == (7168, 2048):+ return ("large", 2)+ return None++@torch.inference_mode()def custom_kernel(data: input_t) -> output_t:a, b, _b_q, b_shuffle, b_scale_sh = datam, k = a.shapen = b.shape[0]key = (int(m), int(n), int(k))+ route = _get_route(key)- if key in _HIDDEN_SHAPES or (key not in _PUBLIC_SMALL and key not in _PUBLIC_LARGE):+ if route is None:return _safe_wrapper(a, b_shuffle, b_scale_sh)- if key in _PUBLIC_LARGE:+ route_kind, route_value = route++ if route_kind == "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]⋯ 25 unchanged linesout,_KERNEL_32X128,bpreshuffle=True,- log2_k_split=_PUBLIC_LARGE[key],+ log2_k_split=route_value,)return out[:m]⋯ 10 unchanged linesw_scales,prequant=True,y=out,- config=_PUBLIC_SMALL.get(key),+ config=route_value,)
scrolls · 126 diff lines total
Best evidence level for this revision: reported
JSON