Skip to content
KernelIndex
Search⌘K

submission 754384

flower2123 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 224 lines, June 9 Researcher Reciprocity License v1.0.

submission_v1_flower_moe.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-754384?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
125.4µs
#69 of 782
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3570996e2cc98ca626ca8eaa376f7e84afd6777e045624a844cebe09d9261513
license declaredunknown
license concludedunknown
authorsflower2123
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4{"stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16", "tile_m": 32, "tile_n": 128, "tile_k": 128, "mode": "atomic", "MPerBlock": 32}),
tile-n = 32BLK, BIG, BM, BN = 32, 128, 32, 8

Kernel source

submission_v1_flower_moe.py224 lines
# Author: flower

import os
import functools
import torch
import triton
from typing import Dict, Tuple, Optional
from task import input_t, output_t

import aiter
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import (
    get_2stage_cfgs, get_padded_M, get_inter_dim,
    ck_moe_stage1, cktile_moe_stage1, cktile_moe_stage2,
    _flydsl_stage2_wrapper,
)
import aiter.fused_moe as _fmoe
import aiter.ops.flydsl.moe_kernels as _fkern
from aiter.ops.triton._triton_kernels.quant.fused_mxfp4_quant import (
    _fused_dynamic_mxfp4_quant_moe_sort_kernel,
)
from aiter.utility import fp4_utils

for _nm, _cfg in [
    ("flydsl_moe2_afp4_wfp4_bf16_t32x128x128_atomic",
     {"stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16", "tile_m": 32, "tile_n": 128, "tile_k": 128, "mode": "atomic", "MPerBlock": 32}),
    ("flydsl_moe2_afp4_wfp4_bf16_t32x256x128_atomic",
     {"stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16", "tile_m": 32, "tile_n": 256, "tile_k": 128, "mode": "atomic", "MPerBlock": 32}),
    ("flydsl_moe2_afp4_wfp4_bf16_t16x256x128_atomic",
     {"stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16", "tile_m": 16, "tile_n": 256, "tile_k": 128, "mode": "atomic", "MPerBlock": 16}),
    ("flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic",
     {"stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16", "tile_m": 16, "tile_n": 128, "tile_k": 128, "mode": "atomic", "MPerBlock": 16}),
]:
    _fkern._KERNEL_PARAMS[_nm] = _cfg

_CK_S1_128 = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_CK_S1_32 = "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_FLY_16_128 = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"
_FLY_16_256 = "flydsl_moe2_afp4_wfp4_bf16_t16x256x128_atomic"


def _mk(t, i, e):
    return (256, t, 7168, i, e, 9, "ActivationType.Silu", "torch.bfloat16",
            "torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",
            "QuantType.per_1x32", True, False)


_SHAPE_MAP = {
    _mk(16, 512, 33):   dict(block_m=32, ksplit=2, kernelName1="", kernelName2="", run_1stage=False),
    _mk(128, 512, 33):  dict(block_m=64, ksplit=0, kernelName1=_CK_S1_128, kernelName2=_FLY_16_128, run_1stage=False),
    _mk(512, 512, 33):  dict(block_m=64, ksplit=0, kernelName1=_CK_S1_128, kernelName2=_FLY_16_128, run_1stage=False),
    _mk(512, 2048, 33): dict(block_m=64, ksplit=0, kernelName1=_CK_S1_128, kernelName2=_FLY_16_256, run_1stage=False),
    _mk(16, 256, 257):  dict(block_m=16, ksplit=2, kernelName1="", kernelName2="", run_1stage=False),
    _mk(128, 256, 257): dict(block_m=16, ksplit=2, kernelName1="", kernelName2="", run_1stage=False),
    _mk(512, 256, 257): dict(block_m=32, ksplit=0, kernelName1=_CK_S1_32, kernelName2=_FLY_16_128, run_1stage=False, use_non_temporal_load=True),
}

_mem = {}


def _sorting_tensors(n_tok, n_exp, k, dim, bm, dev):
    tag = ("s", n_tok, n_exp, k, dim, bm)
    if tag in _mem:
        return _mem[tag]
    cap = int(n_tok * k + n_exp * bm - k)
    nb = int((cap + bm - 1) // bm)
    d = {
        "si": torch.empty(cap, dtype=dtypes.i32, device=dev),
        "sw": torch.empty(cap, dtype=dtypes.fp32, device=dev),
        "se": torch.empty(nb, dtype=dtypes.i32, device=dev),
        "nv": torch.empty(2, dtype=dtypes.i32, device=dev),
        "ob": torch.empty((n_tok, dim), dtype=torch.bfloat16, device=dev),
    }
    _mem[tag] = d
    return d


def _inter_buf(n_tok, k, inter, dev):
    tag = ("i", n_tok, k, inter)
    if tag in _mem:
        return _mem[tag]
    t = torch.empty((n_tok, k, inter), dtype=torch.bfloat16, device=dev)
    _mem[tag] = t
    return t


def _qt_alloc(rows, cols, sid_n, k, dev):
    tag = ("q", rows, cols, sid_n, k)
    if tag in _mem:
        return _mem[tag]
    BLK, BM, BN, BM2, BN2 = 32, 32, 8, 16, 4
    sn = triton.cdiv(cols, BLK)
    r = {
        "f": torch.empty((rows, cols // 2), dtype=torch.uint8, device=dev),
        "s": torch.empty(
            (triton.cdiv(sid_n, BM), triton.cdiv(sn, BN), BN2, BM2, 4),
            dtype=torch.uint8, device=dev),
    }
    _mem[tag] = r
    return r


def _run_quant(x, sid, nval, ntok, k, bm, dev):
    rows, cols = x.shape
    BLK, BIG, BM, BN = 32, 128, 32, 8
    sn = triton.cdiv(cols, BLK)
    sid_n = sid.shape[0]
    qb = _qt_alloc(rows, cols, sid_n, k, dev)
    n_pid = triton.cdiv(rows, BIG) * sn + triton.cdiv(sid_n, BM) * triton.cdiv(sn, BN)
    _fused_dynamic_mxfp4_quant_moe_sort_kernel[(n_pid,)](
        x, qb["f"], sid, nval, qb["s"],
        rows, cols, sn,
        *x.stride(), *qb["f"].stride(), *qb["s"].stride(),
        token_num=ntok, M_i=rows, N_i=sn,
        MXFP4_QUANT_BLOCK_SIZE=BLK, BLOCK_SIZE_Mx=BIG,
        BLOCK_SIZE_M=BM // 2, BLOCK_SIZE_N=BN // 2, TOPK=k,
    )
    return qb["f"].view(dtypes.fp4x2), qb["s"].view(dtypes.fp8_e8m0).view(-1, sn)


_loaded = False


def _setup():
    global _loaded
    if _loaded:
        return
    _loaded = True

    if _fmoe.cfg_2stages is None:
        import pandas as pd
        from aiter.jit.core import AITER_CONFIGS
        fp = AITER_CONFIGS.AITER_CONFIG_FMOE_FILE
        if os.path.exists(fp):
            cols = ["cu_num", "token", "model_dim", "inter_dim", "expert", "topk",
                    "act_type", "dtype", "q_dtype_a", "q_dtype_w", "q_type",
                    "use_g1u1", "doweight_stage1"]
            df = pd.read_csv(fp)
            if "_tag" in df.columns:
                df = df[df["_tag"].fillna("") == ""]
            _fmoe.cfg_2stages = df.set_index(cols).to_dict("index")
        else:
            _fmoe.cfg_2stages = {}

    _fmoe.cfg_2stages.update(_SHAPE_MAP)

    orig = _fmoe.get_2stage_cfgs

    @functools.lru_cache(maxsize=2048)
    def _hook(token, model_dim, inter_dim, expert, topk,
              dtype, q_dtype_a, q_dtype_w, q_type, use_g1u1,
              activation, doweight_stage1, hidden_pad, intermediate_pad, is_shuffled=True):
        md = orig(token, model_dim, inter_dim, expert, topk,
                  dtype, q_dtype_a, q_dtype_w, q_type, use_g1u1,
                  activation, doweight_stage1, hidden_pad, intermediate_pad, is_shuffled)
        from aiter.jit.utils.chip_info import get_cu_num
        lk = (get_cu_num(), token, model_dim, inter_dim, expert, topk,
              str(activation), str(dtype), str(q_dtype_a), str(q_dtype_w),
              str(q_type), use_g1u1, doweight_stage1)
        entry = _fmoe.cfg_2stages.get(lk)
        if entry and entry.get("use_non_temporal_load") is not None:
            nt = entry["use_non_temporal_load"]
            s1 = md.stage1
            if hasattr(s1, 'func') and s1.func is not None and 'use_non_temporal_load' in (s1.keywords or {}):
                kw = dict(s1.keywords); kw['use_non_temporal_load'] = nt
                md = _fmoe.MOEMetadata(
                    functools.partial(s1.func, **{k: v for k, v in kw.items()}),
                    md.stage2, md.block_m, md.ksplit, md.run_1stage, md.has_bias, nt)
                s2 = md.stage2
                if s2 and hasattr(s2, 'keywords') and 'use_non_temporal_load' in (s2.keywords or {}):
                    kw2 = dict(s2.keywords); kw2['use_non_temporal_load'] = nt
                    md = _fmoe.MOEMetadata(
                        md.stage1,
                        functools.partial(s2.func, **{k: v for k, v in kw2.items()}),
                        md.block_m, md.ksplit, md.run_1stage, md.has_bias, nt)
        return md

    _fmoe.get_2stage_cfgs = _hook


def custom_kernel(data: input_t) -> output_t:
    (h, gu_w, d_w, gu_sc, d_sc, gu_ws, d_ws, gu_scs, d_scs, tw, ti, cfg) = data
    _setup()

    hp = cfg["d_hidden_pad"] - cfg["d_hidden"]
    ip = cfg["d_expert_pad"] - cfg["d_expert"]
    n = h.shape[0]
    k = ti.shape[1]
    dev = ti.device
    ne, mdim, idim = get_inter_dim(gu_ws.shape, d_ws.shape)

    md = get_2stage_cfgs(
        get_padded_M(n), mdim, idim, ne, k,
        torch.bfloat16, dtypes.fp4x2, dtypes.fp4x2,
        QuantType.per_1x32, True, ActivationType.Silu,
        False, hp, ip, True)

    bm = int(md.block_m)
    sb = _sorting_tensors(n, ne, k, mdim, bm, dev)
    aiter.moe_sorting_fwd(ti, tw, sb["si"], sb["sw"], sb["se"], sb["nv"], sb["ob"], ne, bm, None, None, 0)

    w1s = gu_scs.view(dtypes.fp8_e8m0)
    w2s = d_scs.view(dtypes.fp8_e8m0)

    if md.ksplit > 1:
        a = h.to(torch.bfloat16)
        inter = md.stage1(a, gu_ws, d_ws, sb["si"], sb["se"], sb["nv"],
                          _inter_buf(n, k, idim, dev), k,
                          block_m=bm, a1_scale=None, w1_scale=w1s, sorted_weights=None)
        md.stage2(inter, gu_ws, d_ws, sb["si"], sb["se"], sb["nv"],
                  sb["ob"], k, w2_scale=w2s, a2_scale=None, block_m=bm, sorted_weights=sb["sw"])
    else:
        qa, qas = _run_quant(h, sb["si"], sb["nv"], n, 1, bm, dev)
        inter = _inter_buf(n, k, idim, dev)
        inter = md.stage1(qa, gu_ws, d_ws, sb["si"], sb["se"], sb["nv"],
                          inter, k, block_m=bm, a1_scale=qas, w1_scale=w1s, sorted_weights=None)
        flat = inter.view(-1, idim)
        qi, qis = _run_quant(flat, sb["si"], sb["nv"], n, k, bm, dev)
        qi = qi.view(n, k, -1)
        md.stage2(qi, gu_ws, d_ws, sb["si"], sb["se"], sb["nv"],
                  sb["ob"], k, w2_scale=w2s, a2_scale=qis, block_m=bm, sorted_weights=sb["sw"])

    return sb["ob"]
scrolls · 224 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