submission 596978
sangmin7b · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 293 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-596978?include=source"interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, fp32, fp8_e8m0, int32, 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:ba2a9f9a1f4fa2385c256a08e26b2736ca76f15b14511ea16925027cebfe9163
license declaredunknown
license concludedunknown
authorssangmin7b
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
tile-m = 16
BLOCK_SIZE_M = 16tile-n = 4
BLOCK_SIZE_N = 4Kernel source
submission.py293 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
import torch
import triton
import triton.language as tl
from typing import Dict
from task import input_t, output_t
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import fused_moe
# ---------------------------------------------------------------------------
# MXFP4 (E2M1) quantization helper
# ---------------------------------------------------------------------------
@triton.jit
def _mxfp4_quant_op(x, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_M: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr):
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
amax = tl.max(tl.abs(x), axis=2, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
scale_e8m0_unbiased = tl.floor(tl.log2(amax)) - 2
scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
bs_e8m0 = (scale_e8m0_unbiased.to(tl.uint8) + 127).reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
quant_scale = tl.exp2(-scale_e8m0_unbiased)
qx = (x * quant_scale).reshape(BLOCK_SIZE_M, BLOCK_SIZE_N).to(tl.uint32, bitcast=True)
s = qx & 0x80000000
e = (qx >> 23) & 0xFF
m = qx & 0x7FFFFF
E8_BIAS: tl.constexpr = 127
E2_BIAS: tl.constexpr = 1
adjusted_exp = tl.core.sub(E8_BIAS, e + 1, sanitize_overflow=False)
m = tl.where(e < E8_BIAS, (0x400000 | (m >> 1)) >> adjusted_exp, m)
e = tl.maximum(e, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)
e2m1_tmp = tl.minimum((((e << 2) | (m >> 21)) + 1) >> 1, 0x7)
e2m1 = ((s >> 28) | e2m1_tmp).to(tl.uint8)
e2m1 = e2m1.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2, 2)
evens, odds = tl.split(e2m1)
out_fp4 = evens | (odds << 4)
return out_fp4, bs_e8m0
# ---------------------------------------------------------------------------
# Optimized fused quant+sort kernel
#
# Opt 1: Phase 1 saves E8M0 scales to a temp buffer [M, scaleN].
# Phase 2 gathers from that buffer instead of reloading the full bf16
# activation tensor and recomputing scales.
#
# Phase 2 data read (bs=128, N=7168, total_top_k=9):
# Before: 1,008 programs × 16 KB (bf16 x slices) = 16.1 MB
# After: 1,008 programs × 8 B (uint8 scales) = 0.26 MB (~62× less)
# ---------------------------------------------------------------------------
@triton.jit
def _fused_mxfp4_quant_moe_sort_kernel(
x_ptr,
x_fp4_ptr,
unsorted_scale_ptr,
sorted_ids_ptr,
num_valid_ids_ptr,
blockscale_e8m0_sorted_ptr,
Mx, Nx, scaleNx,
stride_x_m, stride_x_n,
stride_x_fp4_m, stride_x_fp4_n,
stride_o3, stride_o2, stride_o1, stride_o0, stride_o4,
token_num, M_i, N_i,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
BLOCK_SIZE_Mx: tl.constexpr,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
TOPK: tl.constexpr,
):
pid = tl.program_id(0)
num_pid_x = tl.cdiv(Mx, BLOCK_SIZE_Mx) * scaleNx
# ---- Phase 1: quantize all tokens → x_fp4 + unsorted_scale ----
if pid < num_pid_x:
pid_m = pid // scaleNx
pid_n = pid % scaleNx
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_Mx + tl.arange(0, BLOCK_SIZE_Mx)
x_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE)
x = tl.load(
x_ptr + x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n,
mask=(x_offs_m < Mx)[:, None] & (x_offs_n < Nx)[None, :],
other=0.0,
).to(tl.float32)
out_fp4, bs_e8m0 = _mxfp4_quant_op(
x, MXFP4_QUANT_BLOCK_SIZE, BLOCK_SIZE_Mx, MXFP4_QUANT_BLOCK_SIZE
)
# bs_e8m0: [BLOCK_SIZE_Mx, 1]
# Store packed fp4
out_offs_n = pid_n * (MXFP4_QUANT_BLOCK_SIZE // 2) + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE // 2)
tl.store(
x_fp4_ptr + x_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n,
out_fp4,
mask=(x_offs_m < Mx)[:, None] & (out_offs_n < (Nx // 2))[None, :],
)
# Store E8M0 scale to temp buffer — one per (token, col-block)
tl.store(
unsorted_scale_ptr + x_offs_m * scaleNx + pid_n,
bs_e8m0[:, 0],
mask=x_offs_m < Mx,
)
return
# ---- Phase 2: gather saved scales → CK-tile shuffle → store ----
pid -= num_pid_x
BLOCK_SIZE_M_EFF: tl.constexpr = BLOCK_SIZE_M * 2
BLOCK_SIZE_N_EFF: tl.constexpr = BLOCK_SIZE_N * 2
num_pid_n = tl.cdiv(N_i, BLOCK_SIZE_N_EFF)
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
num_valid_ids = tl.load(num_valid_ids_ptr)
if pid_m * BLOCK_SIZE_M_EFF >= num_valid_ids:
return
stride_o0 = tl.cast(stride_o0, tl.int64)
stride_o1 = tl.cast(stride_o1, tl.int64)
stride_o2 = tl.cast(stride_o2, tl.int64)
stride_o3 = tl.cast(stride_o3, tl.int64)
stride_o4 = tl.cast(stride_o4, tl.int64)
sorted_ids_offs = pid_m * BLOCK_SIZE_M_EFF + tl.arange(0, BLOCK_SIZE_M_EFF)
packed_ids = tl.load(
sorted_ids_ptr + sorted_ids_offs,
mask=sorted_ids_offs < num_valid_ids,
other=token_num,
)
topk_ids = packed_ids >> 24
token_ids = packed_ids & 0xFFFFFF
if TOPK == 1:
x_row_ids = token_ids
else:
x_row_ids = token_ids * TOPK + topk_ids
# Gather scales from temp buffer — replaces reloading x + recomputing scales
scale_col_offs = pid_n * BLOCK_SIZE_N_EFF + tl.arange(0, BLOCK_SIZE_N_EFF)
scales = tl.load(
unsorted_scale_ptr + x_row_ids[:, None] * scaleNx + scale_col_offs[None, :],
mask=(token_ids < token_num)[:, None] & (scale_col_offs < N_i)[None, :],
other=127,
)
# scales: [32, 8] uint8
# CK-tile shuffle: [32, 8] → [16, 4, 4]
bs_e8m0 = (
scales
.reshape(2, BLOCK_SIZE_M, 2, BLOCK_SIZE_N)
.permute(1, 3, 2, 0)
.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N, 4)
)
offs_0 = tl.arange(0, BLOCK_SIZE_M)
offs_1 = tl.arange(0, BLOCK_SIZE_N)
offs_4 = tl.arange(0, 4)
offs = (
offs_0[:, None, None] * stride_o0
+ offs_1[None, :, None] * stride_o1
+ pid_n * stride_o2
+ pid_m * stride_o3
+ offs_4[None, None, :] * stride_o4
)
tl.store(blockscale_e8m0_sorted_ptr + offs, bs_e8m0)
def fused_dynamic_mxfp4_quant_moe_sort(
x: torch.Tensor,
sorted_ids: torch.Tensor,
num_valid_ids: torch.Tensor,
token_num: int,
topk: int,
block_size: int = 32,
scaling_mode: str = "even",
):
M, N = x.shape
MXFP4_QUANT_BLOCK_SIZE = 32
BLOCK_SIZE_Mx = 128
BLOCK_SIZE_M = 16
BLOCK_SIZE_N = 4
scaleN = triton.cdiv(N, MXFP4_QUANT_BLOCK_SIZE)
M_o = sorted_ids.shape[0]
N_i = scaleN
x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=x.device)
unsorted_scales = torch.empty((M, scaleN), dtype=torch.uint8, device=x.device)
blockscale_e8m0_sorted = torch.empty(
(triton.cdiv(M_o, BLOCK_SIZE_M), triton.cdiv(N_i, BLOCK_SIZE_N),
BLOCK_SIZE_N, BLOCK_SIZE_M, 4),
dtype=torch.uint8, device=x.device,
)
num_pid_phase1 = triton.cdiv(M, BLOCK_SIZE_Mx) * scaleN
num_pid_phase2 = triton.cdiv(M_o, BLOCK_SIZE_M) * triton.cdiv(N_i, BLOCK_SIZE_N)
_fused_mxfp4_quant_moe_sort_kernel[(num_pid_phase1 + num_pid_phase2,)](
x,
x_fp4,
unsorted_scales,
sorted_ids,
num_valid_ids,
blockscale_e8m0_sorted,
M, N, scaleN,
*x.stride(),
*x_fp4.stride(),
*blockscale_e8m0_sorted.stride(),
token_num=token_num,
M_i=M_o,
N_i=N_i,
MXFP4_QUANT_BLOCK_SIZE=MXFP4_QUANT_BLOCK_SIZE,
BLOCK_SIZE_Mx=BLOCK_SIZE_Mx,
BLOCK_SIZE_M=BLOCK_SIZE_M,
BLOCK_SIZE_N=BLOCK_SIZE_N,
TOPK=topk,
)
return (
x_fp4.view(dtypes.fp4x2),
blockscale_e8m0_sorted.view(dtypes.fp8_e8m0).view(-1, N_i),
)
# Monkey-patch into aiter so fused_moe_2stages picks up our kernel
import aiter.ops.triton.quant.fused_mxfp4_quant as _aiter_quant_mod
_aiter_quant_mod.fused_dynamic_mxfp4_quant_moe_sort = fused_dynamic_mxfp4_quant_moe_sort
# ---------------------------------------------------------------------------
# Submission entry point
# ---------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
(
hidden_states,
gate_up_weight,
down_weight,
gate_up_weight_scale,
down_weight_scale,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
config,
) = data
hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
intermediate_pad = config["d_expert_pad"] - config["d_expert"]
return fused_moe(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
topk_weights,
topk_ids,
expert_mask=None,
activation=ActivationType.Silu,
quant_type=QuantType.per_1x32,
doweight_stage1=False,
w1_scale=gate_up_weight_scale_shuffled,
w2_scale=down_weight_scale_shuffled,
a1_scale=None,
a2_scale=None,
hidden_pad=hidden_pad,
intermediate_pad=intermediate_pad,
)
scrolls · 293 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