Skip to content
KernelIndex
Search⌘K

submission 754981

NinoHeather · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:0eb14a5ea11e9e1d3e5998c2b9cae09f806b2042c3d0a6983d7391a856945785
license declaredunknown
license concludedunknown
authorsNinoHeather
imported2026-08-15

Techniques

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

fp4"a_dtype": "fp4",
tile-m = 32_QUANT_BLOCK_SIZE_M = 32
tile-n = 8_QUANT_BLOCK_SIZE_N = 8

Kernel source

my_submission_46.py830 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
from __future__ import annotations

import functools
import inspect
import math
import os
from dataclasses import dataclass
from typing import Any, Dict, Iterable, Tuple

import pandas as pd
import torch
try:
    import triton
except Exception:
    triton = None

from task import input_t, output_t

import aiter
import aiter.fused_moe as fmoe
import aiter.ops.flydsl.moe_kernels as flydsl_kernels
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import fused_moe
from aiter.ops.moe_sorting import moe_sorting_fwd
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort
try:
    from aiter.ops.triton._triton_kernels.quant.fused_mxfp4_quant import (
        _fused_dynamic_mxfp4_quant_moe_sort_kernel,
    )
except Exception:
    _fused_dynamic_mxfp4_quant_moe_sort_kernel = None


torch.set_grad_enabled(False)


def _apply_runtime_env() -> None:
    defaults = {
        "PYTORCH_ROCM_ARCH": "gfx950",
        "AITER_USE_NT": "1",
        "AITER_USE_FLYDSL_MOE": "1",
        "AITER_USE_FLYDSL_MOE_STAGE2": "1",
        "AITER_USE_OPUS_MOE_SORTING": "1",
    }
    for k, v in defaults.items():
        os.environ.setdefault(k, v)


def _register_missing_flydsl_kernels() -> None:
    flydsl_kernels._KERNEL_PARAMS.setdefault(
        "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,
        },
    )


_SIG_FIELDS = (
    "cu_num",
    "token",
    "model_dim",
    "inter_dim",
    "expert",
    "topk",
    "act_type",
    "dtype",
    "q_dtype_a",
    "q_dtype_w",
    "q_type",
    "use_g1u1",
    "doweight_stage1",
)
_SIG_META = (
    "ActivationType.Silu",
    "torch.bfloat16",
    "torch.float4_e2m1fn_x2",
    "torch.float4_e2m1fn_x2",
    "QuantType.per_1x32",
)

_K1_256x128 = (
    "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_"
    "Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_K2_F16_128_A = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"
_K2_FLY_64_REDUCE = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
_MXFP4_QUANT_BLOCK_SIZE = 32
_QUANT_BLOCK_SIZE_MX = 128
_QUANT_BLOCK_SIZE_M = 32
_QUANT_BLOCK_SIZE_N = 8
_QUANT_BLOCK_SIZE_M_U32 = 16
_QUANT_BLOCK_SIZE_N_U32 = 4


@dataclass(frozen=True)
class ShapePlan:
    token: int
    model_dim: int
    inter_dim: int
    expert: int
    topk: int = 9
    block_m: int = 32
    ksplit: int = 0
    kernel1: str = ""
    kernel2: str = ""
    run_1stage: bool = False
    use_nt: bool | None = None


def _signature(plan: ShapePlan) -> Tuple[Any, ...]:
    return (
        256,
        plan.token,
        plan.model_dim,
        plan.inter_dim,
        plan.expert,
        plan.topk,
        *_SIG_META,
        1,
        0,
    )


def _row_data(plan: ShapePlan) -> Dict[str, Any]:
    row = {
        "block_m": plan.block_m,
        "ksplit": plan.ksplit,
        "kernelName1": plan.kernel1,
        "kernelName2": plan.kernel2,
        "run_1stage": plan.run_1stage,
        "us1": 0.0,
        "err1": "0.0%",
        "us2": 0.0,
        "err2": "0.0%",
        "us": -1.0,
        "tflops": 0.0,
        "bw": 0.0,
        "_tag": "",
    }
    if plan.use_nt is not None:
        row["use_non_temporal_load"] = bool(plan.use_nt)
    return row


def _retarget_stage_kernel(stage: Any, *, kernel_name: str, block_m: int) -> Any:
    if not isinstance(stage, functools.partial):
        return stage
    try:
        params = set(inspect.signature(stage.func).parameters)
    except Exception:
        params = set()
    kw = dict(stage.keywords or {})
    if "kernelName" in kw or "kernelName" in params or not params:
        kw["kernelName"] = kernel_name
    if "block_m" in kw or "block_m" in params or not params:
        kw["block_m"] = block_m
    return functools.partial(stage.func, *stage.args, **kw)


def _supports_kernel_retarget(stage: Any) -> bool:
    if not isinstance(stage, functools.partial):
        return False
    try:
        params = set(inspect.signature(stage.func).parameters)
    except Exception:
        params = set()
    kw = dict(stage.keywords or {})
    return ("kernelName" in params) or ("kernelName" in kw) or (not params)


def _manual_plans() -> Iterable[ShapePlan]:
    model_dim = 7168
    yield ShapePlan(16, model_dim, 256, 257, block_m=16, ksplit=2)
    yield ShapePlan(
        128,
        model_dim,
        256,
        257,
        block_m=32,
        ksplit=0,
        kernel1=_K1_256x128,
        kernel2=_K2_F16_128_A,
        use_nt=True,
    )
    yield ShapePlan(
        512,
        model_dim,
        256,
        257,
        block_m=64,
        ksplit=0,
        kernel1=_K1_256x128,
        kernel2=_K2_F16_128_A,
        use_nt=True,
    )
    yield ShapePlan(16, model_dim, 512, 33, block_m=32, ksplit=2)
    yield ShapePlan(128, model_dim, 512, 33, block_m=64, ksplit=0, kernel1=_K1_256x128, kernel2=_K2_F16_128_A)
    yield ShapePlan(512, model_dim, 512, 33, block_m=64, ksplit=0, kernel1=_K1_256x128, kernel2=_K2_F16_128_A)
    yield ShapePlan(512, model_dim, 2048, 33, block_m=64, ksplit=0, kernel1=_K1_256x128, kernel2=_K2_F16_128_A)


def _read_csv_configs() -> Dict[Tuple[Any, ...], Dict[str, Any]]:
    cfg_root = os.path.join(os.path.dirname(aiter.__file__), "configs")
    model_cfg_root = os.path.join(cfg_root, "model_configs")
    csv_files = [os.path.join(cfg_root, "tuned_fmoe.csv")]
    if os.path.isdir(model_cfg_root):
        for name in sorted(os.listdir(model_cfg_root)):
            if name.endswith(".csv") and "fmoe" in name:
                csv_files.append(os.path.join(model_cfg_root, name))

    frames = [pd.read_csv(path) for path in csv_files if os.path.exists(path)]
    if not frames:
        existing = getattr(fmoe, "cfg_2stages", None)
        return dict(existing) if isinstance(existing, dict) else {}

    merged = pd.concat(frames, ignore_index=True)
    merged = merged.drop_duplicates(subset=list(_SIG_FIELDS), keep="last")
    for col in (
        "ksplit",
        "block_m",
        "cu_num",
        "token",
        "model_dim",
        "inter_dim",
        "expert",
        "topk",
        "use_g1u1",
        "doweight_stage1",
        "run_1stage",
    ):
        if col in merged.columns:
            merged[col] = merged[col].fillna(0).astype(int)

    util = (merged["token"] * merged["topk"]) // merged["expert"].clip(lower=1)
    merged.loc[(merged["expert"] > 0) & (merged["topk"] > 0) & (util <= 10), "ksplit"] = 4

    kill = (merged["expert"] == 257) & (merged["inter_dim"] == 256) & (merged["token"].isin([16, 128]))
    merged = merged[~kill]

    cfg = merged.set_index(list(_SIG_FIELDS)).to_dict("index")
    for row in cfg.values():
        for k, v in list(row.items()):
            if isinstance(v, float) and math.isnan(v):
                row[k] = ""
            elif isinstance(v, float) and v == int(v):
                row[k] = int(v)
    return cfg


def _apply_plans_into_cfg(cfg: Dict[Tuple[Any, ...], Dict[str, Any]]) -> None:
    for plan in _manual_plans():
        cfg[_signature(plan)] = _row_data(plan)


def _patch_scheduler_heuristics() -> None:
    fmoe.use_nt = lambda token, topk, e: True  # type: ignore[assignment]

    def ksplit_rule(token: int, topk: int, expert: int, inter_dim: int, model_dim: int) -> int:
        if expert == 33 and token >= 512:
            return 0
        dense = (token * topk) // max(expert, 1)
        if dense <= 10:
            return 4
        if dense <= 64:
            return 2
        return 0

    @functools.lru_cache(maxsize=1024)
    def block_m_rule(token: int, topk: int, expert: int, inter_dim: int) -> int:
        dense = (token * topk) // max(expert, 1)
        if dense >= 100:
            return 128
        if dense >= 20:
            return 64
        return 32

    fmoe.get_ksplit = ksplit_rule  # type: ignore[assignment]
    fmoe.get_block_size_M = block_m_rule  # type: ignore[assignment]


def _patch_runtime_nt_from_cfg() -> None:
    base_get_cfg = fmoe.get_2stage_cfgs
    enable_expert257_bs128_tuned = True

    @functools.lru_cache(maxsize=2048)
    def wrapped(
        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,
    ):
        meta = base_get_cfg(
            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,
        )
        try:
            from aiter.jit.utils.chip_info import get_cu_num
        except Exception:
            return meta

        sig = (
            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,
        )
        row = fmoe.cfg_2stages.get(sig) if isinstance(fmoe.cfg_2stages, dict) else None
        out = meta
        if row and row.get("use_non_temporal_load") is not None:
            flag = bool(row["use_non_temporal_load"])

            def maybe_toggle(stage):
                if not isinstance(stage, functools.partial):
                    return stage
                kw = dict(stage.keywords or {})
                if "use_non_temporal_load" in kw:
                    kw["use_non_temporal_load"] = flag
                    return functools.partial(stage.func, *stage.args, **kw)
                if "non_temporal_load" in kw:
                    kw["non_temporal_load"] = flag
                    return functools.partial(stage.func, *stage.args, **kw)
                return stage

            s1 = maybe_toggle(out.stage1)
            s2 = maybe_toggle(out.stage2)
            out = fmoe.MOEMetadata(s1, s2, out.block_m, out.ksplit, out.run_1stage, out.has_bias, flag)

        if (
            enable_expert257_bs128_tuned
            and token == 128
            and model_dim == 7168
            and inter_dim == 256
            and expert == 257
            and topk == 9
            and _supports_kernel_retarget(out.stage1)
            and _supports_kernel_retarget(out.stage2)
        ):
            s1 = _retarget_stage_kernel(out.stage1, kernel_name=_K1_256x128, block_m=32)
            s2 = _retarget_stage_kernel(out.stage2, kernel_name=_K2_F16_128_A, block_m=32)
            out = fmoe.MOEMetadata(s1, s2, 32, 0, False, out.has_bias, True)

        return out

    fmoe.get_2stage_cfgs = wrapped  # type: ignore[assignment]


def _bootstrap() -> None:
    _apply_runtime_env()
    _register_missing_flydsl_kernels()
    _patch_scheduler_heuristics()
    cfg = _read_csv_configs()
    _apply_plans_into_cfg(cfg)
    fmoe.cfg_2stages = cfg
    if hasattr(fmoe.get_2stage_cfgs, "cache_clear"):
        fmoe.get_2stage_cfgs.cache_clear()
    _patch_runtime_nt_from_cfg()
    original_get_padded_m = getattr(fmoe, "get_padded_M", None)
    if original_get_padded_m is not None:
        fmoe.get_padded_M = functools.lru_cache(maxsize=256)(original_get_padded_m)  # type: ignore[assignment]


def _install_expert33_dp2_sorting() -> None:
    orig_impl = getattr(fmoe, "_moe_sorting_impl", None)
    sort_fn = getattr(aiter, "moe_sorting_opus_fwd", None)
    if orig_impl is None or sort_fn is None:
        return

    workspace_cache: Dict[Tuple[Any, ...], Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]] = {}

    def _wrapped(
        topk_ids,
        topk_weights,
        num_experts,
        model_dim,
        moebuf_dtype,
        block_size,
        expert_mask=None,
        num_local_tokens=None,
        dispatch_policy=0,
        use_opus=False,
    ):
        if int(num_experts) != 33:
            return orig_impl(
                topk_ids,
                topk_weights,
                num_experts,
                model_dim,
                moebuf_dtype,
                block_size,
                expert_mask,
                num_local_tokens,
                dispatch_policy,
                use_opus,
            )

        m_tokens, topk = topk_ids.shape
        max_padded = int(topk_ids.numel() + num_experts * block_size - topk)
        max_blocks = int((max_padded + block_size - 1) // block_size)
        device = topk_ids.device
        key = (
            device.index if device.type == "cuda" else -1,
            max_padded,
            max_blocks,
            int(m_tokens),
            int(model_dim),
            str(moebuf_dtype),
            int(block_size),
            int(num_experts),
        )
        cached = workspace_cache.get(key)
        if cached is None:
            cached = (
                torch.empty(max_padded, dtype=torch.int32, device=device),
                torch.empty(max_padded, dtype=torch.float32, device=device),
                torch.empty(max_blocks, dtype=torch.int32, device=device),
                torch.empty(2, dtype=torch.int32, device=device),
                torch.empty((m_tokens, model_dim), dtype=moebuf_dtype, device=device),
            )
            workspace_cache[key] = cached

        sid, sw, sei, nvi, mb = cached
        sort_fn(
            topk_ids,
            topk_weights,
            sid,
            sw,
            sei,
            nvi,
            mb,
            num_experts,
            block_size,
            expert_mask,
            num_local_tokens,
            2,
        )
        return sid, sw, sei, nvi, mb

    fmoe._moe_sorting_impl = _wrapped


_bootstrap()
_install_expert33_dp2_sorting()


def _cdiv(x: int, y: int) -> int:
    return (x + y - 1) // y


def _alloc_quant_workspace(
    rows: int,
    cols: int,
    sorted_rows: int,
    device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor]:
    scale_n = _cdiv(cols, _MXFP4_QUANT_BLOCK_SIZE)
    fp4 = torch.empty((rows, cols // 2), dtype=torch.uint8, device=device)
    bs = torch.empty(
        (
            _cdiv(sorted_rows, _QUANT_BLOCK_SIZE_M),
            _cdiv(scale_n, _QUANT_BLOCK_SIZE_N),
            _QUANT_BLOCK_SIZE_N_U32,
            _QUANT_BLOCK_SIZE_M_U32,
            4,
        ),
        dtype=torch.uint8,
        device=device,
    )
    return fp4, bs


def _quant_with_preallocated_buffers(
    x: torch.Tensor,
    sorted_ids: torch.Tensor,
    num_valid_ids: torch.Tensor,
    token_num: int,
    topk_factor: int,
    block_m: int,
    fp4_out: torch.Tensor,
    bs_out: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
    if _fused_dynamic_mxfp4_quant_moe_sort_kernel is None or triton is None:
        return fused_dynamic_mxfp4_quant_moe_sort(
            x,
            sorted_ids=sorted_ids,
            num_valid_ids=num_valid_ids,
            token_num=token_num,
            topk=topk_factor,
            block_size=block_m,
        )

    rows, cols = x.shape
    scale_n = triton.cdiv(cols, _MXFP4_QUANT_BLOCK_SIZE)
    sorted_rows = sorted_ids.shape[0]
    num_pid = triton.cdiv(rows, _QUANT_BLOCK_SIZE_MX) * scale_n + triton.cdiv(
        sorted_rows, _QUANT_BLOCK_SIZE_M
    ) * triton.cdiv(scale_n, _QUANT_BLOCK_SIZE_N)

    _fused_dynamic_mxfp4_quant_moe_sort_kernel[(num_pid,)](
        x,
        fp4_out,
        sorted_ids,
        num_valid_ids,
        bs_out,
        rows,
        cols,
        scale_n,
        *x.stride(),
        *fp4_out.stride(),
        *bs_out.stride(),
        token_num=token_num,
        M_i=rows,
        N_i=scale_n,
        MXFP4_QUANT_BLOCK_SIZE=_MXFP4_QUANT_BLOCK_SIZE,
        BLOCK_SIZE_Mx=_QUANT_BLOCK_SIZE_MX,
        BLOCK_SIZE_M=_QUANT_BLOCK_SIZE_M // 2,
        BLOCK_SIZE_N=_QUANT_BLOCK_SIZE_N // 2,
        TOPK=topk_factor,
    )
    return (
        fp4_out.view(dtypes.fp4x2),
        bs_out.view(dtypes.fp8_e8m0).view(-1, scale_n),
    )


@dataclass
class _DirectBS512Expert33Workspace:
    meta: Any
    block_m: int
    token: int
    topk: int
    inter_dim: int
    total_experts: int
    gate_up_shuffled: torch.Tensor
    down_shuffled: torch.Tensor
    w1_scale: torch.Tensor
    w2_scale: torch.Tensor
    sorted_ids: torch.Tensor
    sorted_weights: torch.Tensor
    sorted_expert_ids: torch.Tensor
    num_valid_ids: torch.Tensor
    moe_out: torch.Tensor
    a2_buf: torch.Tensor
    hidden_quant_fp4: torch.Tensor
    hidden_quant_scale: torch.Tensor
    a2_quant_fp4: torch.Tensor
    a2_quant_scale: torch.Tensor


_DIRECT_BS512_EXPERT33_REGISTRY: Dict[Tuple[Any, ...], _DirectBS512Expert33Workspace | None] = {}


def _is_bs512_expert33_target(cfg: Dict[str, Any], token: int, expert: int, inter_dim: int) -> bool:
    return (
        token == 512
        and expert == 33
        and inter_dim in (512, 2048)
        and int(cfg.get("n_routed_experts", 0)) == 32
        and int(cfg.get("n_shared_experts", 0)) == 1
        and int(cfg.get("d_hidden", 0)) == 7168
    )


def _direct_registry_key(data: input_t) -> Tuple[Any, ...]:
    hs, w1, w2, s1, s2, cfg = data[0], data[5], data[6], data[7], data[8], data[11]
    dev = hs.device
    return (
        dev.type,
        dev.index if dev.type == "cuda" else -1,
        int(cfg["d_expert"]),
        int(w1.data_ptr()),
        int(w2.data_ptr()),
        int(s1.data_ptr()),
        int(s2.data_ptr()),
    )


def _as_fp8_e8m0_scale(weight: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
    return scale.view(dtypes.fp8_e8m0) if weight.dtype == dtypes.fp4x2 else scale


def _make_direct_workspace(data: input_t) -> _DirectBS512Expert33Workspace | None:
    hs, w1, w2, s1, s2, tw, cfg = data[0], data[5], data[6], data[7], data[8], data[9], data[11]
    token = int(hs.shape[0])
    total_experts = int(w1.shape[0])
    topk = int(tw.shape[1])
    hidden_pad = int(cfg["d_hidden_pad"] - cfg["d_hidden"])
    inter_pad = int(cfg["d_expert_pad"] - cfg["d_expert"])
    _, model_dim, inter_dim = fmoe.get_inter_dim(w1.shape, w2.shape)

    meta = fmoe.get_2stage_cfgs(
        fmoe.get_padded_M(token),
        model_dim,
        inter_dim,
        total_experts,
        topk,
        hs.dtype,
        dtypes.fp4x2,
        w1.dtype,
        QuantType.per_1x32,
        inter_dim != w1.shape[1],
        ActivationType.Silu,
        False,
        hidden_pad,
        inter_pad,
        getattr(w1, "is_shuffled", False),
    )
    if meta.run_1stage:
        return None

    block_m = int(meta.block_m)
    padded_tokens = int(token * topk + total_experts * block_m - topk)
    padded_blocks = int((padded_tokens + block_m - 1) // block_m)
    device = hs.device
    hidden_quant_fp4, hidden_quant_scale = _alloc_quant_workspace(
        token,
        int(model_dim),
        padded_tokens,
        device,
    )
    a2_quant_fp4, a2_quant_scale = _alloc_quant_workspace(
        token * topk,
        int(inter_dim),
        padded_tokens,
        device,
    )
    return _DirectBS512Expert33Workspace(
        meta=meta,
        block_m=block_m,
        token=token,
        topk=topk,
        inter_dim=int(inter_dim),
        total_experts=total_experts,
        gate_up_shuffled=w1,
        down_shuffled=w2,
        w1_scale=_as_fp8_e8m0_scale(w1, s1),
        w2_scale=_as_fp8_e8m0_scale(w2, s2),
        sorted_ids=torch.empty(padded_tokens, dtype=torch.int32, device=device),
        sorted_weights=torch.empty(padded_tokens, dtype=torch.float32, device=device),
        sorted_expert_ids=torch.empty(padded_blocks, dtype=torch.int32, device=device),
        num_valid_ids=torch.empty(2, dtype=torch.int32, device=device),
        moe_out=torch.empty((token, int(model_dim)), dtype=hs.dtype, device=device),
        a2_buf=torch.empty((token, topk, int(inter_dim)), dtype=hs.dtype, device=device),
        hidden_quant_fp4=hidden_quant_fp4,
        hidden_quant_scale=hidden_quant_scale,
        a2_quant_fp4=a2_quant_fp4,
        a2_quant_scale=a2_quant_scale,
    )


def _execute_direct_workspace(
    ws: _DirectBS512Expert33Workspace,
    hidden_states: torch.Tensor,
    topk_weights: torch.Tensor,
    topk_ids: torch.Tensor,
) -> output_t:
    moe_sorting_fwd(
        topk_ids,
        topk_weights,
        ws.sorted_ids,
        ws.sorted_weights,
        ws.sorted_expert_ids,
        ws.num_valid_ids,
        ws.moe_out,
        ws.total_experts,
        ws.block_m,
        None,
        None,
        0,
    )
    q_hidden, q_hidden_scale = _quant_with_preallocated_buffers(
        hidden_states,
        ws.sorted_ids,
        ws.num_valid_ids,
        token_num=ws.token,
        topk_factor=1,
        block_m=ws.block_m,
        fp4_out=ws.hidden_quant_fp4,
        bs_out=ws.hidden_quant_scale,
    )
    ws.meta.stage1(
        q_hidden,
        ws.gate_up_shuffled,
        ws.down_shuffled,
        ws.sorted_ids,
        ws.sorted_expert_ids,
        ws.num_valid_ids,
        ws.a2_buf,
        ws.topk,
        block_m=ws.block_m,
        a1_scale=q_hidden_scale,
        w1_scale=ws.w1_scale,
        sorted_weights=None,
    )
    q_a2, q_a2_scale = _quant_with_preallocated_buffers(
        ws.a2_buf.view(-1, ws.inter_dim),
        ws.sorted_ids,
        ws.num_valid_ids,
        token_num=ws.token,
        topk_factor=ws.topk,
        block_m=ws.block_m,
        fp4_out=ws.a2_quant_fp4,
        bs_out=ws.a2_quant_scale,
    )
    ws.meta.stage2(
        q_a2.view(ws.token, ws.topk, -1),
        ws.gate_up_shuffled,
        ws.down_shuffled,
        ws.sorted_ids,
        ws.sorted_expert_ids,
        ws.num_valid_ids,
        ws.moe_out,
        ws.topk,
        w2_scale=ws.w2_scale,
        a2_scale=q_a2_scale,
        block_m=ws.block_m,
        sorted_weights=ws.sorted_weights,
    )
    return ws.moe_out


def _try_direct_bs512_expert33(data: input_t) -> output_t | None:
    cache_key = _direct_registry_key(data)
    try:
        cached_workspace = _DIRECT_BS512_EXPERT33_REGISTRY[cache_key]
    except KeyError:
        cached_workspace = _make_direct_workspace(data)
        _DIRECT_BS512_EXPERT33_REGISTRY[cache_key] = cached_workspace

    if cached_workspace is None:
        return None

    hidden_states = data[0]
    routed_weights = data[9]
    routed_ids = data[10]

    try:
        return _execute_direct_workspace(cached_workspace, hidden_states, routed_weights, routed_ids)
    except Exception:
        _DIRECT_BS512_EXPERT33_REGISTRY[cache_key] = None
        return None


def custom_kernel(data: input_t) -> output_t:
    (
        hidden_states,
        _raw_gate_up,
        _raw_down,
        _raw_gate_scale,
        _raw_down_scale,
        gate_up_shuffled,
        down_shuffled,
        gate_up_scale_shuffled,
        down_scale_shuffled,
        topk_weights,
        topk_ids,
        config,
    ) = data
    hidden_pad = int(config["d_hidden_pad"] - config["d_hidden"])
    inter_pad = int(config["d_expert_pad"] - config["d_expert"])

    token = int(hidden_states.shape[0])
    inter_dim = int(config["d_expert"])
    expert = int(config["n_routed_experts"] + config["n_shared_experts"])
    if _is_bs512_expert33_target(config, token=token, expert=expert, inter_dim=inter_dim):
        direct_out = _try_direct_bs512_expert33(data)
        if direct_out is not None:
            return direct_out

    return fused_moe(
        hidden_states,
        gate_up_shuffled,
        down_shuffled,
        topk_weights,
        topk_ids,
        expert_mask=None,
        activation=ActivationType.Silu,
        quant_type=QuantType.per_1x32,
        doweight_stage1=False,
        w1_scale=gate_up_scale_shuffled,
        w2_scale=down_scale_shuffled,
        a1_scale=None,
        a2_scale=None,
        hidden_pad=hidden_pad,
        intermediate_pad=inter_pad,
    )

scrolls · 830 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 754350.

⋯ 10 unchanged lines
import pandas as pd
import torch
+ try:
+ import triton
+ except Exception:
+ triton = None
from task import input_t, output_t
⋯ 4 unchanged lines
from aiter.fused_moe import fused_moe
from aiter.ops.moe_sorting import moe_sorting_fwd
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort
+ try:
+ from aiter.ops.triton._triton_kernels.quant.fused_mxfp4_quant import (
+ _fused_dynamic_mxfp4_quant_moe_sort_kernel,
+ )
+ except Exception:
+ _fused_dynamic_mxfp4_quant_moe_sort_kernel = None
torch.set_grad_enabled(False)
⋯ 55 unchanged lines
"moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_"
"Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
- _K1_256x32 = (
- "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_"
- "Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
- )
_K2_F16_128_A = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"
_K2_FLY_64_REDUCE = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
+ _MXFP4_QUANT_BLOCK_SIZE = 32
+ _QUANT_BLOCK_SIZE_MX = 128
+ _QUANT_BLOCK_SIZE_M = 32
+ _QUANT_BLOCK_SIZE_N = 8
+ _QUANT_BLOCK_SIZE_M_U32 = 16
+ _QUANT_BLOCK_SIZE_N_U32 = 4
@dataclass(frozen=True)
⋯ 75 unchanged lines
def _manual_plans() -> Iterable[ShapePlan]:
model_dim = 7168
yield ShapePlan(16, model_dim, 256, 257, block_m=16, ksplit=2)
- yield ShapePlan(128, model_dim, 256, 257, block_m=16, ksplit=2)
yield ShapePlan(
- 512,
+ 128,
model_dim,
256,
257,
block_m=32,
ksplit=0,
- kernel1=_K1_256x32,
+ kernel1=_K1_256x128,
kernel2=_K2_F16_128_A,
use_nt=True,
)
+ yield ShapePlan(
+ 512,
+ model_dim,
+ 256,
+ 257,
+ block_m=64,
+ ksplit=0,
+ kernel1=_K1_256x128,
+ kernel2=_K2_F16_128_A,
+ use_nt=True,
+ )
yield ShapePlan(16, model_dim, 512, 33, block_m=32, ksplit=2)
yield ShapePlan(128, model_dim, 512, 33, block_m=64, ksplit=0, kernel1=_K1_256x128, kernel2=_K2_F16_128_A)
yield ShapePlan(512, model_dim, 512, 33, block_m=64, ksplit=0, kernel1=_K1_256x128, kernel2=_K2_F16_128_A)
⋯ 170 unchanged lines
and _supports_kernel_retarget(out.stage2)
):
s1 = _retarget_stage_kernel(out.stage1, kernel_name=_K1_256x128, block_m=32)
- s2 = _retarget_stage_kernel(out.stage2, kernel_name=_K2_FLY_64_REDUCE, block_m=32)
- out = fmoe.MOEMetadata(s1, s2, 32, 1, False, out.has_bias, False)
+ s2 = _retarget_stage_kernel(out.stage2, kernel_name=_K2_F16_128_A, block_m=32)
+ out = fmoe.MOEMetadata(s1, s2, 32, 0, False, out.has_bias, True)
return out
⋯ 98 unchanged lines
_install_expert33_dp2_sorting()
+ def _cdiv(x: int, y: int) -> int:
+ return (x + y - 1) // y
+
+
+ def _alloc_quant_workspace(
+ rows: int,
+ cols: int,
+ sorted_rows: int,
+ device: torch.device,
+ ) -> tuple[torch.Tensor, torch.Tensor]:
+ scale_n = _cdiv(cols, _MXFP4_QUANT_BLOCK_SIZE)
+ fp4 = torch.empty((rows, cols // 2), dtype=torch.uint8, device=device)
+ bs = torch.empty(
+ (
+ _cdiv(sorted_rows, _QUANT_BLOCK_SIZE_M),
+ _cdiv(scale_n, _QUANT_BLOCK_SIZE_N),
+ _QUANT_BLOCK_SIZE_N_U32,
+ _QUANT_BLOCK_SIZE_M_U32,
+ 4,
+ ),
+ dtype=torch.uint8,
+ device=device,
+ )
+ return fp4, bs
+
+
+ def _quant_with_preallocated_buffers(
+ x: torch.Tensor,
+ sorted_ids: torch.Tensor,
+ num_valid_ids: torch.Tensor,
+ token_num: int,
+ topk_factor: int,
+ block_m: int,
+ fp4_out: torch.Tensor,
+ bs_out: torch.Tensor,
+ ) -> tuple[torch.Tensor, torch.Tensor]:
+ if _fused_dynamic_mxfp4_quant_moe_sort_kernel is None or triton is None:
+ return fused_dynamic_mxfp4_quant_moe_sort(
+ x,
+ sorted_ids=sorted_ids,
+ num_valid_ids=num_valid_ids,
+ token_num=token_num,
+ topk=topk_factor,
+ block_size=block_m,
+ )
+
+ rows, cols = x.shape
+ scale_n = triton.cdiv(cols, _MXFP4_QUANT_BLOCK_SIZE)
+ sorted_rows = sorted_ids.shape[0]
+ num_pid = triton.cdiv(rows, _QUANT_BLOCK_SIZE_MX) * scale_n + triton.cdiv(
+ sorted_rows, _QUANT_BLOCK_SIZE_M
+ ) * triton.cdiv(scale_n, _QUANT_BLOCK_SIZE_N)
+
+ _fused_dynamic_mxfp4_quant_moe_sort_kernel[(num_pid,)](
+ x,
+ fp4_out,
+ sorted_ids,
+ num_valid_ids,
+ bs_out,
+ rows,
+ cols,
+ scale_n,
+ *x.stride(),
+ *fp4_out.stride(),
+ *bs_out.stride(),
+ token_num=token_num,
+ M_i=rows,
+ N_i=scale_n,
+ MXFP4_QUANT_BLOCK_SIZE=_MXFP4_QUANT_BLOCK_SIZE,
+ BLOCK_SIZE_Mx=_QUANT_BLOCK_SIZE_MX,
+ BLOCK_SIZE_M=_QUANT_BLOCK_SIZE_M // 2,
+ BLOCK_SIZE_N=_QUANT_BLOCK_SIZE_N // 2,
+ TOPK=topk_factor,
+ )
+ return (
+ fp4_out.view(dtypes.fp4x2),
+ bs_out.view(dtypes.fp8_e8m0).view(-1, scale_n),
+ )
+
+
@dataclass
class _DirectBS512Expert33Workspace:
meta: Any
⋯ 12 unchanged lines
num_valid_ids: torch.Tensor
moe_out: torch.Tensor
a2_buf: torch.Tensor
+ hidden_quant_fp4: torch.Tensor
+ hidden_quant_scale: torch.Tensor
+ a2_quant_fp4: torch.Tensor
+ a2_quant_scale: torch.Tensor
_DIRECT_BS512_EXPERT33_REGISTRY: Dict[Tuple[Any, ...], _DirectBS512Expert33Workspace | None] = {}
⋯ 61 unchanged lines
padded_tokens = int(token * topk + total_experts * block_m - topk)
padded_blocks = int((padded_tokens + block_m - 1) // block_m)
device = hs.device
+ hidden_quant_fp4, hidden_quant_scale = _alloc_quant_workspace(
+ token,
+ int(model_dim),
+ padded_tokens,
+ device,
+ )
+ a2_quant_fp4, a2_quant_scale = _alloc_quant_workspace(
+ token * topk,
+ int(inter_dim),
+ padded_tokens,
+ device,
+ )
return _DirectBS512Expert33Workspace(
meta=meta,
block_m=block_m,
⋯ 11 unchanged lines
num_valid_ids=torch.empty(2, dtype=torch.int32, device=device),
moe_out=torch.empty((token, int(model_dim)), dtype=hs.dtype, device=device),
a2_buf=torch.empty((token, topk, int(inter_dim)), dtype=hs.dtype, device=device),
+ hidden_quant_fp4=hidden_quant_fp4,
+ hidden_quant_scale=hidden_quant_scale,
+ a2_quant_fp4=a2_quant_fp4,
+ a2_quant_scale=a2_quant_scale,
)
⋯ 17 unchanged lines
None,
0,
)
- q_hidden, q_hidden_scale = fused_dynamic_mxfp4_quant_moe_sort(
+ q_hidden, q_hidden_scale = _quant_with_preallocated_buffers(
hidden_states,
- sorted_ids=ws.sorted_ids,
- num_valid_ids=ws.num_valid_ids,
+ ws.sorted_ids,
+ ws.num_valid_ids,
token_num=ws.token,
- topk=1,
- block_size=ws.block_m,
+ topk_factor=1,
+ block_m=ws.block_m,
+ fp4_out=ws.hidden_quant_fp4,
+ bs_out=ws.hidden_quant_scale,
)
ws.meta.stage1(
q_hidden,
⋯ 9 unchanged lines
w1_scale=ws.w1_scale,
sorted_weights=None,
)
- q_a2, q_a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
+ q_a2, q_a2_scale = _quant_with_preallocated_buffers(
ws.a2_buf.view(-1, ws.inter_dim),
- sorted_ids=ws.sorted_ids,
- num_valid_ids=ws.num_valid_ids,
+ ws.sorted_ids,
+ ws.num_valid_ids,
token_num=ws.token,
- topk=ws.topk,
- block_size=ws.block_m,
+ topk_factor=ws.topk,
+ block_m=ws.block_m,
+ fp4_out=ws.a2_quant_fp4,
+ bs_out=ws.a2_quant_scale,
)
ws.meta.stage2(
q_a2.view(ws.token, ws.topk, -1),
scrolls · 256 diff lines total

Best evidence level for this revision: reported

JSON