submission 754350
NinoHeather · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 704 lines, June 9 Researcher Reciprocity License v1.0.
my_submission_46.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-754350?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:a876e76d3d093db5f81e387e0fff7b74994a67cde2c388a58163d785d5cb1957
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",Kernel source
my_submission_46.py704 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
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
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"
)
_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"
@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=16, ksplit=2)
yield ShapePlan(
512,
model_dim,
256,
257,
block_m=32,
ksplit=0,
kernel1=_K1_256x32,
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_FLY_64_REDUCE, block_m=32)
out = fmoe.MOEMetadata(s1, s2, 32, 1, False, out.has_bias, False)
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()
@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
_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
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),
)
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 = fused_dynamic_mxfp4_quant_moe_sort(
hidden_states,
sorted_ids=ws.sorted_ids,
num_valid_ids=ws.num_valid_ids,
token_num=ws.token,
topk=1,
block_size=ws.block_m,
)
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 = fused_dynamic_mxfp4_quant_moe_sort(
ws.a2_buf.view(-1, ws.inter_dim),
sorted_ids=ws.sorted_ids,
num_valid_ids=ws.num_valid_ids,
token_num=ws.token,
topk=ws.topk,
block_size=ws.block_m,
)
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 · 704 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 683939.
⋯ 2 unchanged linesfrom __future__ import annotationsimport functools+ import inspectimport mathimport osfrom dataclasses import dataclass⋯ 7 unchanged linesimport aiterimport aiter.fused_moe as fmoeimport aiter.ops.flydsl.moe_kernels as flydsl_kernels- from aiter import ActivationType, QuantType+ from aiter import ActivationType, QuantType, dtypesfrom 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_sorttorch.set_grad_enabled(False)⋯ 12 unchanged linesdef _register_missing_flydsl_kernels() -> None:- # Some builds omit this key, but config rows may reference it.flydsl_kernels._KERNEL_PARAMS.setdefault("flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic",{⋯ 42 unchanged lines"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"@dataclass(frozen=True)⋯ 46 unchanged linesreturn 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 = 7168yield ShapePlan(16, model_dim, 256, 257, block_m=16, ksplit=2)⋯ 47 unchanged linesif col in merged.columns:merged[col] = merged[col].fillna(0).astype(int)- # Keep cktile for tiny utilization.util = (merged["token"] * merged["topk"]) // merged["expert"].clip(lower=1)merged.loc[(merged["expert"] > 0) & (merged["topk"] > 0) & (util <= 10), "ksplit"] = 4- # Remove known weak rows; force manual rows below.kill = (merged["expert"] == 257) & (merged["inter_dim"] == 256) & (merged["token"].isin([16, 128]))merged = merged[~kill]⋯ 40 unchanged linesdef _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(⋯ 51 unchanged linesdoweight_stage1,)row = fmoe.cfg_2stages.get(sig) if isinstance(fmoe.cfg_2stages, dict) else None- if not row or row.get("use_non_temporal_load") is None:- return meta- flag = bool(row["use_non_temporal_load"])+ 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):+ 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- 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(meta.stage1)- s2 = maybe_toggle(meta.stage2)- if s1 is meta.stage1 and s2 is meta.stage2:- return meta- return fmoe.MOEMetadata(s1, s2, meta.block_m, meta.ksplit, meta.run_1stage, meta.has_bias, flag)+ 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_FLY_64_REDUCE, block_m=32)+ out = fmoe.MOEMetadata(s1, s2, 32, 1, False, out.has_bias, False)++ return out+fmoe.get_2stage_cfgs = wrapped # type: ignore[assignment]⋯ 7 unchanged linesif hasattr(fmoe.get_2stage_cfgs, "cache_clear"):fmoe.get_2stage_cfgs.cache_clear()_patch_runtime_nt_from_cfg()- # keep cheap deterministic helper cache only (not an input-result cache)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()+ @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+++ _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+ 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),+ )+++ 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 = fused_dynamic_mxfp4_quant_moe_sort(+ hidden_states,+ sorted_ids=ws.sorted_ids,+ num_valid_ids=ws.num_valid_ids,+ token_num=ws.token,+ topk=1,+ block_size=ws.block_m,+ )+ 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 = fused_dynamic_mxfp4_quant_moe_sort(+ ws.a2_buf.view(-1, ws.inter_dim),+ sorted_ids=ws.sorted_ids,+ num_valid_ids=ws.num_valid_ids,+ token_num=ws.token,+ topk=ws.topk,+ block_size=ws.block_m,+ )+ 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,⋯ 11 unchanged lines) = datahidden_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,
scrolls · 456 diff lines total
Best evidence level for this revision: reported
JSON