Skip to content
KernelIndex
Search⌘K

submission 525898

anuragj0803 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

aj_moe_sub_v2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-525898?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
177.9µs
#397 of 782
2026-03-10

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a7159f438cbbdea82440d62a25993db789fb1ae54deb93f05a72831f63ed3bb1
license declaredunknown
license concludedunknown
authorsanuragj0803
imported2026-08-26

Techniques

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

fp4raise RuntimeError("AITER is required for this MXFP4 fused MoE kernel.")

Kernel source

aj_moe_sub_v2.py406 lines
import os

# Set early, before torch / aiter import
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
os.environ.setdefault("TORCH_BLAS_PREFER_HIPBLASLT", "1")

import torch
import torch.nn as nn
from typing import Dict, Tuple, Optional
from task import input_t, output_t

try:
    import aiter
    from aiter import ActivationType, QuantType
    from aiter.fused_moe import fused_moe
    _HAS_AITER = True
except Exception:
    _HAS_AITER = False


# -----------------------------
# Global caches
# -----------------------------
_CFG_CACHE = {}
_MOE_CACHE = {}
_NEEDS_CROP_CACHE = {}


def _shape_key(config: Dict):
    return (
        config["bs"],
        config["d_hidden"],
        config["d_expert"],
        config["d_hidden_pad"],
        config["d_expert_pad"],
        config["n_routed_experts"],
        config["n_shared_experts"],
        config["n_experts_per_token"],
        config["total_top_k"],
    )


def _problem_family(config: Dict) -> str:
    """
    Bucket shapes into families for specialized tuning.
    """
    bs = config["bs"]
    d_hidden = config["d_hidden"]
    d_expert = config["d_expert"]
    n_routed = config["n_routed_experts"]
    n_shared = config["n_shared_experts"]
    total_top_k = config["total_top_k"]

    # Leaderboard / benchmark families
    if d_hidden == 7168 and n_shared == 1 and total_top_k == 9:
        if n_routed == 256 and d_expert == 256:
            if bs <= 32:
                return "257e_k9_smallm"
            elif bs <= 128:
                return "257e_k9_midm"
            else:
                return "257e_k9_largem"

        if n_routed == 32 and d_expert == 512:
            if bs <= 32:
                return "33e_512_k9_smallm"
            elif bs <= 128:
                return "33e_512_k9_midm"
            else:
                return "33e_512_k9_largem"

        if n_routed == 32 and d_expert == 2048:
            return "33e_2048_k9_stage2heavy"

    # Correctness / hidden eval families
    if d_hidden == 4096:
        if d_expert >= 1536:
            return "h4096_stage2heavy"
        return "h4096_generic"

    return "generic"


def _get_kernel_cfg(config: Dict) -> Dict:
    key = _shape_key(config)
    cfg = _CFG_CACHE.get(key)
    if cfg is not None:
        return cfg

    hidden_pad = max(0, config["d_hidden_pad"] - config["d_hidden"])
    intermediate_pad = max(0, config["d_expert_pad"] - config["d_expert"])
    family = _problem_family(config)

    # Keep correctness-preserving defaults first.
    cfg = {
        "family": family,
        "hidden_pad": hidden_pad,
        "intermediate_pad": intermediate_pad,
        "doweight_stage1": False,   # known-correct baseline
        "force_nt_like": False,     # placeholder only
        "blockmwide_like": False,   # placeholder only
        "stage2_heavy_like": False, # placeholder only
        "presorted_like": False,    # placeholder only
        "d_hidden": config["d_hidden"],
    }

    # Shape-specialized policy hooks
    if family in ("257e_k9_smallm", "257e_k9_midm", "33e_512_k9_smallm", "33e_512_k9_midm"):
        cfg["blockmwide_like"] = True

    if family in ("33e_2048_k9_stage2heavy", "h4096_stage2heavy"):
        cfg["stage2_heavy_like"] = True

    # These are placeholders for experimentation.
    # They do not change correctness unless you wire them to a proven fast path.
    if family in ("257e_k9_smallm", "33e_512_k9_smallm"):
        cfg["force_nt_like"] = True

    _CFG_CACHE[key] = cfg
    return cfg


class Expert(nn.Module):
    def __init__(self, config: Dict, d_expert: Optional[int] = None):
        super().__init__()
        self.config = config
        self.d_hidden = config["d_hidden"]
        self.d_expert = config["d_expert"] if d_expert is None else d_expert

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        raise RuntimeError("Expert.forward is not used in this fused kernel path.")


class MoEGate(nn.Module):
    def __init__(self, config: Dict):
        super().__init__()
        self.top_k = config["n_experts_per_token"]
        self.num_experts = config["n_routed_experts"]
        self.d_hidden = config["d_hidden"]

    def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
        raise RuntimeError("MoEGate.forward is not used; harness provides topk.")


def _maybe_cast_inputs(hidden_states, topk_weights, topk_ids):
    if hidden_states.dtype != torch.bfloat16:
        hidden_states = hidden_states.to(torch.bfloat16)
    if topk_weights.dtype != torch.float32:
        topk_weights = topk_weights.to(torch.float32)
    if topk_ids.dtype != torch.int32:
        topk_ids = topk_ids.to(torch.int32)
    return hidden_states, topk_weights, topk_ids


def _fused_moe_baseline(
    hidden_states,
    gate_up_weight_shuffled,
    down_weight_shuffled,
    gate_up_weight_scale_shuffled,
    down_weight_scale_shuffled,
    topk_weights,
    topk_ids,
    cfg,
):
    """
    Current known-correct path.
    """
    return fused_moe(
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        topk_weights,
        topk_ids,
        w1_scale=gate_up_weight_scale_shuffled,
        w2_scale=down_weight_scale_shuffled,
        quant_type=QuantType.per_1x32,
        activation=ActivationType.Silu,
        doweight_stage1=cfg["doweight_stage1"],
        intermediate_pad=cfg["intermediate_pad"],
        hidden_pad=cfg["hidden_pad"],
        bias1=None,
        bias2=None,
    )


def _fused_moe_blockmwide_candidate(
    hidden_states,
    gate_up_weight_shuffled,
    down_weight_shuffled,
    gate_up_weight_scale_shuffled,
    down_weight_scale_shuffled,
    topk_weights,
    topk_ids,
    cfg,
):
    """
    Placeholder for a future block-M-wide optimized path.

    Right now this falls back to the known-correct baseline.
    Replace this only when you have verified a faster variant.
    """
    return _fused_moe_baseline(
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        gate_up_weight_scale_shuffled,
        down_weight_scale_shuffled,
        topk_weights,
        topk_ids,
        cfg,
    )


def _fused_moe_stage2_candidate(
    hidden_states,
    gate_up_weight_shuffled,
    down_weight_shuffled,
    gate_up_weight_scale_shuffled,
    down_weight_scale_shuffled,
    topk_weights,
    topk_ids,
    cfg,
):
    """
    Placeholder for stage2-focused tuning.

    Candidates later:
    - toggled doweight_stage1 for only stage2-heavy shapes
    - forced alt AITER tuned instance if exposed
    - separated shared expert handling if correctness can be preserved
    """
    return _fused_moe_baseline(
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        gate_up_weight_scale_shuffled,
        down_weight_scale_shuffled,
        topk_weights,
        topk_ids,
        cfg,
    )


def _fused_moe_force_nt_candidate(
    hidden_states,
    gate_up_weight_shuffled,
    down_weight_shuffled,
    gate_up_weight_scale_shuffled,
    down_weight_scale_shuffled,
    topk_weights,
    topk_ids,
    cfg,
):
    """
    Placeholder for a forced layout/orientation path.

    AITER release notes mention standardized GEMM weight shape and default TN layout,
    which makes a force-NT/TN style experiment plausible. Keep baseline for now.
    """
    return _fused_moe_baseline(
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        gate_up_weight_scale_shuffled,
        down_weight_scale_shuffled,
        topk_weights,
        topk_ids,
        cfg,
    )


class MoE(nn.Module):
    def __init__(self, config: Dict):
        super().__init__()
        self.config = config
        self._cfg = _get_kernel_cfg(config)
        self._shape_key = _shape_key(config)

    @torch.no_grad()
    def forward(
        self,
        hidden_states: torch.Tensor,
        gate_up_weight: torch.Tensor,
        down_weight: torch.Tensor,
        gate_up_weight_scale: torch.Tensor,
        down_weight_scale: torch.Tensor,
        gate_up_weight_shuffled: torch.Tensor,
        down_weight_shuffled: torch.Tensor,
        gate_up_weight_scale_shuffled: torch.Tensor,
        down_weight_scale_shuffled: torch.Tensor,
        topk_weights: torch.Tensor,
        topk_ids: torch.Tensor,
    ) -> torch.Tensor:
        if not _HAS_AITER:
            raise RuntimeError("AITER is required for this MXFP4 fused MoE kernel.")

        hidden_states, topk_weights, topk_ids = _maybe_cast_inputs(
            hidden_states, topk_weights, topk_ids
        )

        family = self._cfg["family"]

        # Family-specific dispatch
        if self._cfg["stage2_heavy_like"]:
            output = _fused_moe_stage2_candidate(
                hidden_states,
                gate_up_weight_shuffled,
                down_weight_shuffled,
                gate_up_weight_scale_shuffled,
                down_weight_scale_shuffled,
                topk_weights,
                topk_ids,
                self._cfg,
            )
        elif self._cfg["blockmwide_like"]:
            output = _fused_moe_blockmwide_candidate(
                hidden_states,
                gate_up_weight_shuffled,
                down_weight_shuffled,
                gate_up_weight_scale_shuffled,
                down_weight_scale_shuffled,
                topk_weights,
                topk_ids,
                self._cfg,
            )
        elif self._cfg["force_nt_like"]:
            output = _fused_moe_force_nt_candidate(
                hidden_states,
                gate_up_weight_shuffled,
                down_weight_shuffled,
                gate_up_weight_scale_shuffled,
                down_weight_scale_shuffled,
                topk_weights,
                topk_ids,
                self._cfg,
            )
        else:
            output = _fused_moe_baseline(
                hidden_states,
                gate_up_weight_shuffled,
                down_weight_shuffled,
                gate_up_weight_scale_shuffled,
                down_weight_scale_shuffled,
                topk_weights,
                topk_ids,
                self._cfg,
            )

        need_crop = _NEEDS_CROP_CACHE.get(self._shape_key)
        if need_crop is None:
            need_crop = output.shape[-1] != self._cfg["d_hidden"]
            _NEEDS_CROP_CACHE[self._shape_key] = need_crop

        if need_crop:
            output = output[..., : self._cfg["d_hidden"]]

        if output.dtype != torch.bfloat16:
            output = output.to(torch.bfloat16)

        return output


def _get_moe(config: Dict) -> MoE:
    key = _shape_key(config)
    moe = _MOE_CACHE.get(key)
    if moe is None:
        moe = MoE(config).eval()
        _MOE_CACHE[key] = moe
    return moe


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

    moe = _get_moe(config)

    with torch.inference_mode():
        output = moe(
            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,
        )

    return output
scrolls · 406 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