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
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.
fp4
raise 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