Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 MoEsuite of 7 cases
AMD Instinct MI355X
185.3µs
#616 of 782
2026-03-20

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 = 16BLOCK_SIZE_M = 16
tile-n = 4BLOCK_SIZE_N = 4

Kernel 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