submission 721756
Sikuan Wang · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 84 lines, June 9 Researcher Reciprocity License v1.0.
submission-12.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-721756?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:be37673435b1efc827230414901b4cc2e0e048a07f3c4519434ad4ce005ab638
license declaredunknown
license concludedunknown
authorsSikuan Wang
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Optimized MXFP4 GEMM: stride-aware Triton e8m0 shuffle.Kernel source
submission-12.py84 lines
"""
Optimized MXFP4 GEMM: stride-aware Triton e8m0 shuffle.
The scale tensor from dynamic_mxfp4_quant is column-major (stride=(1, 8)).
The reference e8m0_shuffle does: alloc pad → copy → permute → .contiguous()
= 2 allocations + 2 kernel launches.
This replaces it with a single stride-aware Triton kernel that reads from
the column-major input and writes the shuffled+padded output in one pass:
= 1 allocation + 1 kernel launch.
"""
from __future__ import annotations
import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from task import input_t, output_t
_FP4 = dtypes.fp4x2
_E8M0 = dtypes.fp8_e8m0
_BF16 = dtypes.bf16
@triton.jit
def _e8m0_shuffle_k(
in_ptr, out_ptr,
M, K_SCALE, STRIDE_0, STRIDE_1, SN,
TOTAL,
BLOCK: tl.constexpr,
):
pid = tl.program_id(0)
offs = pid * BLOCK + tl.arange(0, BLOCK)
mask = offs < TOTAL
s0 = 32 * SN
a = offs // s0
r = offs % s0
d = r // 256
r = r % 256
f = r // 64
r = r % 64
c = r // 4
r = r % 4
e = r // 2
b = r % 2
in_row = a * 32 + b * 16 + c
in_col = d * 8 + e * 4 + f
in_bounds = mask & (in_row < M) & (in_col < K_SCALE)
in_idx = in_row * STRIDE_0 + in_col * STRIDE_1
val = tl.load(in_ptr + in_idx, mask=in_bounds, other=0)
tl.store(out_ptr + offs, val, mask=mask)
def _shuffle(scale: torch.Tensor) -> torch.Tensor:
m, ks = scale.shape
s0, s1 = scale.stride()
sm = (m + 255) // 256 * 256
sn = (ks + 7) // 8 * 8
out = torch.empty(sm, sn, dtype=scale.dtype, device=scale.device)
total = sm * sn
_e8m0_shuffle_k[(triton.cdiv(total, 1024),)](
scale, out, m, ks, s0, s1, sn, total, BLOCK=1024,
)
return out
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
a = data[0]
if not a.is_contiguous():
a = a.contiguous()
aq, asc = dynamic_mxfp4_quant(a)
return aiter.gemm_a4w4(
aq.view(_FP4), data[3],
_shuffle(asc).view(_E8M0), data[4],
dtype=_BF16, bpreshuffle=True,
)
scrolls · 84 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