Skip to content
KernelIndex
Search⌘K

submission 749778

coderwhisper · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-749778?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
149.5µs
#175 of 782
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:26ba3775ec1b519f0d8a587038fda5e4db5152e1d5e1eed8a59a981e3c08b12e
license declaredunknown
license concludedunknown
authorscoderwhisper
imported2026-08-15

Techniques

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

fp4_meta_cache[key] = ('mxfp4', block_m, metadata, tile_m, use_flydsl_s2,

Kernel source

submission.py278 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X

import os
import sys
import pathlib
import functools

# ── Patch FlyDSL bug BEFORE importing aiter ──
_FLYDSL_PATCH_TARGET = pathlib.Path(
    "/home/runner/aiter/aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py"
)
_PATCHED = False
if _FLYDSL_PATCH_TARGET.exists():
    _src = _FLYDSL_PATCH_TARGET.read_text()
    if "b_tile_in_gate," in _src:
        _lines = _src.split('\n')
        _fixed_lines = []
        _skip_next = 0
        for _i, _line in enumerate(_lines):
            if _skip_next > 0:
                _skip_next -= 1
                continue
            if 'b_tile_in_gate,' in _line:
                _indent = _line[:len(_line) - len(_line.lstrip())]
                _fixed_lines.append(f"{_indent}b_tile_in,")
                _skip_next = 1
            elif 'b_scale_gate=None,' in _line:
                _indent = _line[:len(_line) - len(_line.lstrip())]
                _fixed_lines.append(f"{_indent}b_scale=None,")
                _skip_next = 1
            else:
                _fixed_lines.append(_line)
        _fixed = '\n'.join(_fixed_lines)
        if _fixed != _src:
            try:
                _FLYDSL_PATCH_TARGET.write_text(_fixed)
                import shutil
                for _pyc in _FLYDSL_PATCH_TARGET.parent.rglob("__pycache__"):
                    shutil.rmtree(_pyc, ignore_errors=True)
                _PATCHED = True
            except PermissionError:
                pass
    else:
        _PATCHED = True

import torch
import aiter
from task import input_t, output_t
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import (
    get_2stage_cfgs,
    get_inter_dim,
    get_padded_M,
    get_ksplit,
    use_nt,
)
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort
from aiter.ops.flydsl import is_flydsl_available

_HAS_FLYDSL = is_flydsl_available() and _PATCHED
if _HAS_FLYDSL:
    from aiter.ops.flydsl import flydsl_moe_stage2
    print(f"[submission] FlyDSL available for stage2", file=sys.stderr)

_meta_cache = {}


def _has_tuned_stage2(metadata):
    """Check if metadata.stage2 has a tuned (non-empty) kernel name."""
    if isinstance(metadata.stage2, functools.partial):
        kn = metadata.stage2.keywords.get('kernelName', '')
        return bool(kn)
    return False


def custom_kernel(data: input_t) -> output_t:
    (
        hidden_states,
        _gate_up_weight,
        _down_weight,
        _gate_up_weight_scale,
        _down_weight_scale,
        w1,
        w2,
        w1_scale,
        w2_scale,
        topk_weights,
        topk_ids,
        config,
    ) = data

    hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
    intermediate_pad = config["d_expert_pad"] - config["d_expert"]

    M = hidden_states.shape[0]
    E = w1.shape[0]
    topk = topk_ids.shape[1]
    _, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)
    is_shuffled = getattr(w1, "is_shuffled", False)
    dtype = hidden_states.dtype

    key = (M, E, model_dim, inter_dim)

    if key not in _meta_cache:
        device = hidden_states.device

        if M <= 128:
            # ksplit=2 bf16 cktile path
            os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1"
            os.environ["AITER_KSPLIT"] = "2"
            get_ksplit.cache_clear()
            use_nt.cache_clear()

            metadata = get_2stage_cfgs(
                get_padded_M(M), model_dim, inter_dim, E, topk, dtype,
                dtypes.fp4x2, dtypes.fp4x2, QuantType.per_1x32,
                True, ActivationType.Silu, False,
                hidden_pad, intermediate_pad, is_shuffled,
            )
            block_m = int(metadata.block_m)
            ksplit = int(metadata.ksplit)

            n_pad_zeros_1 = intermediate_pad // 64 * 64 * 2
            k_pad_zeros_1 = hidden_pad // 128 * 128
            n_pad_zeros_2 = hidden_pad // 64 * 64
            k_pad_zeros_2 = intermediate_pad // 128 * 128
            _, n1, k1 = w1.shape
            _, k2, n2 = w2.shape
            D = n2 if k2 == k1 else n2 * 2

            max_num_tokens_padded = int(M * topk + E * block_m - topk)
            max_num_m_blocks = (max_num_tokens_padded + block_m - 1) // block_m
            s_ids = torch.empty(max_num_tokens_padded, dtype=torch.int32, device=device)
            s_weights = torch.empty(max_num_tokens_padded, dtype=torch.float32, device=device)
            s_expert_ids = torch.empty(max_num_m_blocks, dtype=torch.int32, device=device)
            n_valid = torch.empty(2, dtype=torch.int32, device=device)
            moe_buf = torch.empty((M, model_dim), dtype=dtype, device=device)
            tmp_out = torch.empty((M, topk, w1.shape[1]), dtype=dtype, device=device)
            a2 = torch.empty((M, topk, D), dtype=dtype, device=device)
            _meta_cache[key] = ('direct', block_m, ksplit,
                s_ids, s_weights, s_expert_ids, n_valid, moe_buf,
                n_pad_zeros_1, k_pad_zeros_1, n_pad_zeros_2, k_pad_zeros_2,
                tmp_out, a2)
        else:
            # M>128: MXFP4 with CK stage1
            os.environ["AITER_BYPASS_TUNE_CONFIG"] = "0"
            os.environ["AITER_KSPLIT"] = "0"
            get_ksplit.cache_clear()
            use_nt.cache_clear()

            metadata = get_2stage_cfgs(
                get_padded_M(M), model_dim, inter_dim, E, topk, dtype,
                dtypes.fp4x2, dtypes.fp4x2, QuantType.per_1x32,
                True, ActivationType.Silu, False,
                hidden_pad, intermediate_pad, is_shuffled,
            )
            block_m = int(metadata.block_m)

            # Decision: use tuned CK stage2 if available, else FlyDSL stage2
            has_tuned_s2 = _has_tuned_stage2(metadata)
            use_flydsl_s2 = _HAS_FLYDSL and not has_tuned_s2

            # For FlyDSL stage2: cap block_m at 64 (max supported tile_m)
            if use_flydsl_s2 and block_m > 64:
                block_m = 64
            # Verify FlyDSL compatibility after capping
            use_flydsl_s2 = use_flydsl_s2 and block_m <= 64
            tile_m = block_m if use_flydsl_s2 else 0

            max_num_tokens_padded = int(M * topk + E * block_m - topk)
            max_num_m_blocks = (max_num_tokens_padded + block_m - 1) // block_m
            s_ids = torch.empty(max_num_tokens_padded, dtype=torch.int32, device=device)
            s_weights = torch.empty(max_num_tokens_padded, dtype=torch.float32, device=device)
            s_expert_ids = torch.empty(max_num_m_blocks, dtype=torch.int32, device=device)
            n_valid = torch.empty(2, dtype=torch.int32, device=device)
            moe_buf = torch.empty((M, model_dim), dtype=dtype, device=device)
            a2 = torch.empty((M, topk, inter_dim), dtype=dtype, device=device)

            _meta_cache[key] = ('mxfp4', block_m, metadata, tile_m, use_flydsl_s2,
                s_ids, s_weights, s_expert_ids, n_valid, moe_buf, a2)

    cached = _meta_cache[key]

    if cached[0] == 'direct':
        (_, block_m, ksplit,
         s_ids, s_weights, s_expert_ids, n_valid, moe_buf,
         n_pad_zeros_1, k_pad_zeros_1, n_pad_zeros_2, k_pad_zeros_2,
         tmp_out, a2) = cached

        w1s = w1_scale.view(dtypes.fp8_e8m0)
        w2s = w2_scale.view(dtypes.fp8_e8m0)

        aiter.moe_sorting_fwd(
            topk_ids, topk_weights,
            s_ids, s_weights, s_expert_ids, n_valid, moe_buf,
            E, block_m, None, None, 0,
        )

        tmp_out.zero_()
        aiter.moe_cktile2stages_gemm1(
            hidden_states, w1, tmp_out,
            s_ids, s_expert_ids, n_valid,
            topk, n_pad_zeros_1, k_pad_zeros_1,
            None, None, w1s, None,
            ActivationType.Silu, block_m, ksplit,
        )
        aiter.silu_and_mul(a2, tmp_out)

        aiter.moe_cktile2stages_gemm2(
            a2, w2, moe_buf,
            s_ids, s_expert_ids, n_valid,
            topk, n_pad_zeros_2, k_pad_zeros_2,
            s_weights, None, w2s, None,
            ActivationType.Silu, block_m,
        )
        return moe_buf

    else:
        (_, block_m, metadata, tile_m, use_flydsl_s2,
         s_ids, s_weights, s_expert_ids, n_valid, moe_buf, a2) = cached

        w1s = w1_scale.view(dtypes.fp8_e8m0)
        w2s = w2_scale.view(dtypes.fp8_e8m0)

        aiter.moe_sorting_fwd(
            topk_ids, topk_weights,
            s_ids, s_weights, s_expert_ids, n_valid, moe_buf,
            E, block_m, None, None, 0,
        )

        # Quantize activations to MXFP4
        a1, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
            hidden_states, sorted_ids=s_ids, num_valid_ids=n_valid,
            token_num=M, topk=1, block_size=block_m,
        )

        # Stage 1: CK GEMM (uses tuned kernel if available)
        a2 = metadata.stage1(
            a1, w1, w2,
            s_ids, s_expert_ids, n_valid,
            a2, topk,
            block_m=block_m,
            a1_scale=a1_scale,
            w1_scale=w1s,
            sorted_weights=None,
        )

        # Quantize intermediate to MXFP4
        a2_flat = a2.view(-1, a2.shape[-1])
        a2_q, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
            a2_flat, sorted_ids=s_ids, num_valid_ids=n_valid,
            token_num=M, topk=topk, block_size=block_m,
        )
        a2_q = a2_q.view(M, topk, -1)

        # Stage 2: tuned CK or FlyDSL
        if use_flydsl_s2:
            flydsl_moe_stage2(
                a2_q, w2, s_ids, s_expert_ids, n_valid,
                out=moe_buf, topk=topk,
                tile_m=tile_m, tile_n=128, tile_k=256,
                a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
                mode="reduce",
                w2_scale=w2s, a2_scale=a2_scale, sorted_weights=s_weights,
            )
        else:
            metadata.stage2(
                a2_q, w1, w2,
                s_ids, s_expert_ids, n_valid,
                moe_buf, topk,
                w2_scale=w2s,
                a2_scale=a2_scale,
                block_m=block_m,
                sorted_weights=s_weights,
            )
        return moe_buf
scrolls · 278 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