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
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.
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 linesimport pandas as pdimport torch+ try:+ import triton+ except Exception:+ triton = Nonefrom task import input_t, output_t⋯ 4 unchanged linesfrom aiter.fused_moe import fused_moefrom aiter.ops.moe_sorting import moe_sorting_fwdfrom 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 = Nonetorch.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 linesdef _manual_plans() -> Iterable[ShapePlan]:model_dim = 7168yield 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 linesand _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),+ )++@dataclassclass _DirectBS512Expert33Workspace:meta: Any⋯ 12 unchanged linesnum_valid_ids: torch.Tensormoe_out: torch.Tensora2_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 linespadded_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 linesnum_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 linesNone,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 linesw1_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