Skip to content
KernelIndex
Search⌘K

submission 754899

Itay Etelis · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e998523329a859fa0c4e2f23dcd0e148ad65d10b54b913497f7c41193e244d90
license declaredunknown
license concludedunknown
authorsItay Etelis
imported2026-08-15

Techniques

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

split-kactivation=_S, split_k=plan["ksplit"])

Kernel source

sub_v79_inline_ck2stage.py159 lines
"""
V79: Full inline CK 2-stage with pre-allocated a2 buffer.

V76 uses fused_moe_2stages() for CK path which allocates a2 internally.
This version inlines the CK 2-stage path, pre-allocating a2 and calling
stage1/stage2 directly via resolved metadata.

Saves: a2 allocation (~2-3us) + fused_moe_2stages dispatch (~1-2us)
"""

import torch
from task import input_t, output_t
import aiter
from aiter import ActivationType, QuantType, dtypes
import aiter.fused_moe as _fmoe
from aiter.fused_moe import (
    fused_moe, get_inter_dim, get_2stage_cfgs, get_padded_M,
    fused_dynamic_mxfp4_quant_moe_sort,
    cktile_moe_stage1, cktile_moe_stage2,
)

_S = ActivationType.Silu
_Q = QuantType.per_1x32
_BF16 = torch.bfloat16
_CK1_32 = "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_CK1_32_4CU = "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_CK1_64 = "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_CK2_d512 = "moe_ck2stages_gemm2_256x32x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
_FLY_32r = "flydsl_moe2_afp4_wfp4_bf16_t32x256x256_reduce"
_FLY_64r = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"

_first_call = True
_plans = {}
_sort_bufs = {}
_a2_bufs = {}

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)

def _inject():
    if _fmoe.cfg_2stages is None: return
    ek2 = {"run_1stage": False, "ksplit": 2}
    e = {"run_1stage": False, "ksplit": 0}
    _fmoe.cfg_2stages.update({
        _mk(16, 512, 33): {**ek2, "block_m": 32, "kernelName1": "", "kernelName2": ""},
        _mk(128, 512, 33): {**ek2, "block_m": 32, "kernelName1": "", "kernelName2": ""},
        _mk(512, 512, 33): {**e, "block_m": 32, "kernelName1": _CK1_32, "kernelName2": _FLY_32r},
        _mk(512, 2048, 33): {**e, "block_m": 64, "kernelName1": _CK1_64, "kernelName2": _FLY_64r},
        _mk(16, 256, 257): {**ek2, "block_m": 16, "kernelName1": "", "kernelName2": ""},
        _mk(128, 256, 257): {**ek2, "block_m": 16, "kernelName1": "", "kernelName2": ""},
        _mk(512, 256, 257): {**e, "block_m": 32, "kernelName1": _CK1_32_4CU, "kernelName2": _FLY_32r},
    })
    _fmoe.get_2stage_cfgs.cache_clear()

def _get_sort_bufs(M, topk, E, model_dim, block_m, device):
    key = (M, topk, E, block_m)
    if key not in _sort_bufs:
        mp = int(M * topk + E * block_m - topk); mb = (mp + block_m - 1) // block_m
        _sort_bufs[key] = (
            torch.empty(mp, dtype=dtypes.i32, device=device),
            torch.empty(mp, dtype=dtypes.fp32, device=device),
            torch.empty(mb, dtype=dtypes.i32, device=device),
            torch.empty(2, dtype=dtypes.i32, device=device),
            torch.empty((M, model_dim), dtype=_BF16, device=device))
    return _sort_bufs[key]

def _get_a2_buf(M, topk, inter_dim, device):
    key = (M, topk, inter_dim)
    if key not in _a2_bufs:
        _a2_bufs[key] = torch.empty((M, topk, inter_dim), dtype=_BF16, device=device)
    return _a2_bufs[key]

def _get_plan(M, topk, E, model_dim, inter_dim, hp, ip):
    key = (M, E, inter_dim)
    if key in _plans: return _plans[key]
    pm = get_padded_M(M)
    meta = get_2stage_cfgs(pm, model_dim, inter_dim, E, topk, dtypes.bf16,
        dtypes.fp4x2, dtypes.fp4x2, _Q, True, _S, False, hp, ip, True)
    _plans[key] = {"block_m": meta.block_m, "ksplit": meta.ksplit,
        "use_cktile": meta.ksplit > 1,
        "stage1": meta.stage1, "stage2": meta.stage2,
        "hp": hp, "ip": ip}
    return _plans[key]

def _run(hidden_states, w1, w2, w1_scale, w2_scale, topk_weights, topk_ids, config):
    M = hidden_states.shape[0]; topk = topk_ids.shape[1]
    E, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)
    hp = config["d_hidden_pad"] - config["d_hidden"]; ip = config["d_expert_pad"] - config["d_expert"]
    device = hidden_states.device
    plan = _get_plan(M, topk, E, model_dim, inter_dim, hp, ip)
    block_m = plan["block_m"]

    si, sw, se, nv, moe_buf = _get_sort_bufs(M, topk, E, model_dim, block_m, device)
    aiter.moe_sorting_fwd(topk_ids, topk_weights, si, sw, se, nv, moe_buf,
        E, int(block_m), None, None, 0)

    if plan["use_cktile"]:
        n1 = ip // 64 * 64 * 2; k1 = hp // 128 * 128
        n2 = hp // 64 * 64; k2 = ip // 128 * 128
        a2 = cktile_moe_stage1(hidden_states, w1, w2, si, se, nv,
            None, topk, block_m, None, w1_scale.view(dtypes.fp8_e8m0),
            sorted_weights=None, n_pad_zeros=n1, k_pad_zeros=k1,
            activation=_S, split_k=plan["ksplit"])
        cktile_moe_stage2(a2, w1, w2, si, se, nv, moe_buf, topk,
            w2_scale.view(dtypes.fp8_e8m0), None, block_m,
            activation=_S, sorted_weights=sw, n_pad_zeros=n2, k_pad_zeros=k2)
    else:
        # INLINE CK 2-stage: quant → stage1 → re-quant → stage2
        # Step 1: Quant hidden_states
        a1, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
            hidden_states, sorted_ids=si, num_valid_ids=nv,
            token_num=M, topk=1, block_size=block_m)

        # Step 2: Stage1 with pre-allocated a2
        a2_buf = _get_a2_buf(M, topk, inter_dim, device)
        a2 = plan["stage1"](
            a1, w1, w2, si, se, nv, a2_buf, topk,
            block_m=block_m, a1_scale=a1_scale,
            w1_scale=w1_scale.view(dtypes.fp8_e8m0),
            sorted_weights=None)

        # Step 3: Re-quant intermediate
        a2_flat = a2.view(-1, inter_dim)
        a2_q, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
            a2_flat, sorted_ids=si, num_valid_ids=nv,
            token_num=M, topk=topk, block_size=block_m)
        a2_q = a2_q.view(M, topk, -1)

        # Step 4: Stage2
        plan["stage2"](
            a2_q, w1, w2, si, se, nv, moe_buf, topk,
            w2_scale=w2_scale.view(dtypes.fp8_e8m0),
            a2_scale=a2_scale,
            block_m=block_m,
            sorted_weights=sw)

    return moe_buf

def custom_kernel(data: input_t) -> output_t:
    global _first_call
    (hidden_states, _, _, _, _, w1, w2, w1_scale, w2_scale,
     topk_weights, topk_ids, config) = data
    hp = config["d_hidden_pad"] - config["d_hidden"]; ip = config["d_expert_pad"] - config["d_expert"]
    if _first_call:
        _first_call = False
        r = fused_moe(hidden_states, w1, w2, topk_weights, topk_ids,
            activation=_S, quant_type=_Q, w1_scale=w1_scale, w2_scale=w2_scale,
            hidden_pad=hp, intermediate_pad=ip)
        _inject()
        return r
    try:
        return _run(hidden_states, w1, w2, w1_scale, w2_scale, topk_weights, topk_ids, config)
    except Exception:
        return fused_moe(hidden_states, w1, w2, topk_weights, topk_ids,
            activation=_S, quant_type=_Q, w1_scale=w1_scale, w2_scale=w2_scale,
            hidden_pad=hp, intermediate_pad=ip)
scrolls · 159 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