submission 558039
Eurafat45 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 882 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-558039?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:6dbd86f1a153879c249ae4fe4676c6c76a35c5fb35da910c41c026b55af6fdde
license declaredunknown
license concludedunknown
authorsEurafat45
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Submission template for DeepSeek-R1 MXFP4 MoE kernel.split-k
splitk=0,tile-m = 16
BLOCK_SIZE_M=16,tile-n = 4
BLOCK_SIZE_N=4,Kernel source
submission.py882 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
import importlib
import functools
import os
import time
# Must be set before importing aiter.fused_moe (read at import time).
os.environ.setdefault("AITER_USE_OPUS_MOE_SORTING", "1")
import torch
from task import input_t, output_t
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe
_fused_moe_mod = importlib.import_module("aiter.fused_moe")
_orig_get_2stage_cfgs = _fused_moe_mod.get_2stage_cfgs
try:
_fused_mxfp4_quant_moe_sort_fn = (
importlib.import_module("aiter.ops.triton.quant.fused_mxfp4_quant")
.fused_dynamic_mxfp4_quant_moe_sort
)
except Exception:
_fused_mxfp4_quant_moe_sort_fn = getattr(
_fused_moe_mod, "fused_dynamic_mxfp4_quant_moe_sort", None
)
_sort_workspace_cache = {}
_fastpath_workspace_cache = {}
_scale_sort_workspace_cache = {}
_route_bucket_cache = {}
_runtime_cfg_cache = {}
_runtime_state_cache = {}
_fastpath_disabled = False
_fastpath_error_printed = False
_fused_quant_sort_disabled = _fused_mxfp4_quant_moe_sort_fn is None
_fused_quant_sort_error_printed = False
_FUSED_QUANT_MOE_SORT_TOKEN_SWITCH = 1024
_RUNTIME_EXPLORE_SAMPLES = 0
_RUNTIME_RECHECK_INTERVAL = 0
_K_STAGE1_S32 = (
"moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_"
"MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_K_STAGE1_W32 = (
"moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_"
"MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_K_STAGE1_W64 = (
"moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_"
"MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_K_STAGE1_W128 = (
"moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_"
"MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_K_STAGE2_S32 = (
"moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_"
"Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
)
_K_STAGE2_W128 = (
"moe_ck2stages_gemm2_256x128x128x128_1x4_MulABScaleExpertWeightShuffled_v3_"
"Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
)
_K_STAGE2_FLYDSL_64 = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
_K_STAGE2_FLYDSL_128 = "flydsl_moe2_afp4_wfp4_bf16_t128x256x256_reduce"
_FORCED_E257_CFG = {
(128, 256): (
32,
_K_STAGE1_S32,
_K_STAGE2_S32,
),
}
_FORCED_E33_CFG = {
(16, 512): (
32,
_K_STAGE1_W32,
_K_STAGE2_S32,
),
(128, 512): (
32,
_K_STAGE1_W32,
_K_STAGE2_S32,
),
(512, 512): (
32,
_K_STAGE1_W32,
_K_STAGE2_S32,
),
(512, 2048): (
128,
_K_STAGE1_W128,
_K_STAGE2_W128,
),
}
def _use_non_temporal_load(token: int, topk: int, expert: int) -> bool:
use_nt_env = int(os.environ.get("AITER_USE_NT", "-1"))
if use_nt_env != -1:
return bool(use_nt_env)
if int(expert) == 33:
return True
return (int(token) * int(topk) // int(expert)) < 64
def _workspace_moe_sorting(
topk_ids,
topk_weights,
num_experts,
model_dim,
moebuf_dtype,
block_size=32,
expert_mask=None,
num_local_tokens=None,
dispatch_policy=0,
):
device = topk_ids.device
m, topk = topk_ids.shape
block_size = int(block_size)
num_experts = int(num_experts)
model_dim = int(model_dim)
max_num_tokens_padded = int(topk_ids.numel() + num_experts * block_size - topk)
max_num_m_blocks = int((max_num_tokens_padded + block_size - 1) // block_size)
key = (
int(device.index if device.index is not None else -1),
str(device.type),
int(m),
int(topk),
num_experts,
model_dim,
str(moebuf_dtype),
block_size,
)
ws = _sort_workspace_cache.get(key)
if ws is None:
ws = (
torch.empty(max_num_tokens_padded, dtype=torch.int32, device=device),
torch.empty(max_num_tokens_padded, dtype=torch.float32, device=device),
torch.empty(max_num_m_blocks, dtype=torch.int32, device=device),
torch.empty(2, dtype=torch.int32, device=device),
torch.empty((m, model_dim), dtype=moebuf_dtype, device=device),
)
_sort_workspace_cache[key] = ws
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = ws
use_opus = os.environ.get("AITER_USE_OPUS_MOE_SORTING", "0") == "1"
fwd_fn = (
_fused_moe_mod.aiter.moe_sorting_opus_fwd
if use_opus and hasattr(_fused_moe_mod.aiter, "moe_sorting_opus_fwd")
else _fused_moe_mod.aiter.moe_sorting_fwd
)
fwd_fn(
topk_ids,
topk_weights,
sorted_ids,
sorted_weights,
sorted_expert_ids,
num_valid_ids,
moe_buf,
num_experts,
block_size,
expert_mask,
num_local_tokens,
dispatch_policy,
)
return sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf
def _cdiv(x: int, y: int) -> int:
return (int(x) + int(y) - 1) // int(y)
def _alloc_fp8_e8m0(shape, device):
return torch.empty(shape, dtype=torch.uint8, device=device).view(
_fused_moe_mod.dtypes.fp8_e8m0
)
def _get_fastpath_workspace(token_num, topk, model_dim, inter_dim, dtype, device):
key = (
int(device.index if device.index is not None else -1),
str(device.type),
int(token_num),
int(topk),
int(model_dim),
int(inter_dim),
str(dtype),
)
ws = _fastpath_workspace_cache.get(key)
if ws is None:
ws = {
"a1_q": torch.empty(
(token_num, model_dim // 2),
dtype=_fused_moe_mod.dtypes.fp4x2,
device=device,
),
"a1_scale": _alloc_fp8_e8m0((token_num, model_dim // 32), device),
"a2_inter": torch.empty(
(token_num, topk, inter_dim),
dtype=dtype,
device=device,
),
"a2_q": torch.empty(
(token_num * topk, inter_dim // 2),
dtype=_fused_moe_mod.dtypes.fp4x2,
device=device,
),
"a2_scale": _alloc_fp8_e8m0((token_num * topk, inter_dim // 32), device),
}
_fastpath_workspace_cache[key] = ws
return ws
def _moe_mxfp4_sort_into(blockscale_e8m0, sorted_ids, num_valid_ids, token_num, tag):
topk = 1
src = blockscale_e8m0
if len(src.shape) == 3:
topk = int(src.shape[1])
src = src.view(-1, src.shape[-1])
m_i, n_i = src.shape
m_o, n_o = int(sorted_ids.shape[0]), int(n_i)
shape_u32 = (_cdiv(m_o, 32), _cdiv(n_o, 8), 4, 16)
key = (
int(src.device.index if src.device.index is not None else -1),
str(src.device.type),
int(m_o),
int(n_o),
str(tag),
)
out_u32 = _scale_sort_workspace_cache.get(key)
if out_u32 is None:
out_u32 = torch.empty(shape_u32, dtype=torch.uint32, device=src.device)
_scale_sort_workspace_cache[key] = out_u32
grid = (_cdiv(m_o, 32), _cdiv(n_i, 8))
_fused_moe_mod.fp4_utils._moe_mxfp4_sort_kernel[grid](
src.view(torch.uint8),
sorted_ids,
num_valid_ids,
out_u32,
*src.stride(),
*out_u32.stride(),
token_num=token_num,
M_i=m_i,
N_i=n_i,
BLOCK_SIZE_M=16,
BLOCK_SIZE_N=4,
TOPK=topk,
)
return out_u32.view(_fused_moe_mod.dtypes.fp8_e8m0).view(-1, n_o)
def _try_fused_mxfp4_quant_moe_sort(
x, sorted_ids, num_valid_ids, token_num, topk, block_m
):
global _fused_quant_sort_disabled, _fused_quant_sort_error_printed
if _fused_quant_sort_disabled or int(token_num) > _FUSED_QUANT_MOE_SORT_TOKEN_SWITCH:
return None, None
try:
return _fused_mxfp4_quant_moe_sort_fn(
x,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=int(token_num),
topk=int(topk),
block_size=int(block_m),
)
except Exception as e:
_fused_quant_sort_disabled = True
if not _fused_quant_sort_error_printed:
print(f"[submission fused quant+sort disabled] {type(e).__name__}: {e}")
_fused_quant_sort_error_printed = True
return None, None
def _route_density_bucket(topk_ids: torch.Tensor, num_experts: int) -> str:
key = (
int(topk_ids.data_ptr()),
int(topk_ids._version),
int(topk_ids.shape[0]),
int(topk_ids.shape[1]),
int(num_experts),
)
cached = _route_bucket_cache.get(key)
if cached is not None:
return cached
flat = topk_ids.reshape(-1).to(torch.int64)
hist = torch.bincount(flat, minlength=int(num_experts))
max_load = int(hist.max().item())
active = int((hist > 0).sum().item())
mean_load = float(flat.numel()) / float(max(num_experts, 1))
max_ratio = max_load / (mean_load + 1e-6)
active_ratio = float(active) / float(max(num_experts, 1))
bucket = "skewed" if (max_ratio >= 3.0 or active_ratio <= 0.45) else "balanced"
if len(_route_bucket_cache) > 128:
_route_bucket_cache.clear()
_route_bucket_cache[key] = bucket
return bucket
def _stage2_callable(kernel_name, activation, q_type, use_non_temporal_load):
if (
isinstance(kernel_name, str)
and kernel_name.startswith("flydsl_")
and hasattr(_fused_moe_mod, "is_flydsl_available")
and _fused_moe_mod.is_flydsl_available()
and hasattr(_fused_moe_mod, "_flydsl_stage2_wrapper")
):
return functools.partial(
_fused_moe_mod._flydsl_stage2_wrapper,
kernelName=kernel_name,
)
return functools.partial(
_fused_moe_mod.aiter.ck_moe_stage2_fwd,
kernelName=kernel_name,
activation=activation,
quant_type=q_type,
use_non_temporal_load=use_non_temporal_load,
)
def _build_metadata_from_cfg(cfg, token, expert, topk, activation, q_type, dtype):
block_m, kernel1, kernel2 = cfg
use_nt = _use_non_temporal_load(int(token), int(topk), int(expert))
return _fused_moe_mod.MOEMetadata(
functools.partial(
_fused_moe_mod.ck_moe_stage1,
kernelName=kernel1,
activation=activation,
quant_type=q_type,
dtype=dtype,
splitk=0,
use_non_temporal_load=use_nt,
),
_stage2_callable(kernel2, activation, q_type, use_nt),
int(block_m),
0,
False,
)
def _runtime_candidate_cfgs(token, inter_dim, expert, route_bucket):
token = int(token)
inter_dim = int(inter_dim)
expert = int(expert)
if expert == 257 and inter_dim == 256:
if token >= 512:
return [(32, _K_STAGE1_W32, _K_STAGE2_S32)]
return [(32, _K_STAGE1_S32, _K_STAGE2_S32)]
if expert == 33 and inter_dim == 512:
if token >= 512:
if (
hasattr(_fused_moe_mod, "is_flydsl_available")
and _fused_moe_mod.is_flydsl_available()
and hasattr(_fused_moe_mod, "_flydsl_stage2_wrapper")
):
return [(64, _K_STAGE1_W64, _K_STAGE2_FLYDSL_64)]
return [(32, _K_STAGE1_W32, _K_STAGE2_S32)]
return [(32, _K_STAGE1_S32, _K_STAGE2_S32)]
if expert == 33 and inter_dim == 2048:
return [(128, _K_STAGE1_W128, _K_STAGE2_W128)]
return []
def _init_runtime_state(key, cands):
if len(_runtime_state_cache) > 128:
_runtime_state_cache.clear()
state = {
"cands": tuple(cands),
"calls": 0,
"stats": {cfg: [0, 0.0] for cfg in cands},
}
_runtime_state_cache[key] = state
return state
def _pick_runtime_cfg(key, cands):
state = _runtime_state_cache.get(key)
if state is None or state.get("cands") != tuple(cands):
state = _init_runtime_state(key, cands)
state["calls"] += 1
stats = state["stats"]
for cfg in cands:
if stats[cfg][0] < _RUNTIME_EXPLORE_SAMPLES:
return cfg, True
if (
len(cands) > 1
and _RUNTIME_RECHECK_INTERVAL > 0
and (state["calls"] % _RUNTIME_RECHECK_INTERVAL == 0)
):
cfg = min(cands, key=lambda c: stats[c][0])
return cfg, True
cfg = min(cands, key=lambda c: stats[c][1] / max(stats[c][0], 1))
return cfg, False
def _update_runtime_stats(key, cfg, elapsed_us):
if key is None or cfg is None:
return
state = _runtime_state_cache.get(key)
if state is None:
return
entry = state["stats"].get(cfg)
if entry is None:
state["stats"][cfg] = [1, float(elapsed_us)]
return
entry[0] += 1
entry[1] += float(elapsed_us)
def _run_fastpath_once(
hidden_states,
w1,
w2,
topk_weights,
topk_ids,
w1_scale,
w2_scale,
metadata,
):
token_num, topk = topk_ids.shape
expert, model_dim, inter_dim = _fused_moe_mod.get_inter_dim(w1.shape, w2.shape)
expert = int(expert)
model_dim = int(model_dim)
inter_dim = int(inter_dim)
block_m = int(metadata.block_m)
dtype = hidden_states.dtype
device = hidden_states.device
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = (
_workspace_moe_sorting(
topk_ids=topk_ids,
topk_weights=topk_weights,
num_experts=expert,
model_dim=model_dim,
moebuf_dtype=dtype,
block_size=block_m,
expert_mask=None,
num_local_tokens=None,
dispatch_policy=0,
)
)
ws = _get_fastpath_workspace(
token_num=token_num,
topk=topk,
model_dim=model_dim,
inter_dim=inter_dim,
dtype=dtype,
device=device,
)
a1_q, a1_scale_sorted = _try_fused_mxfp4_quant_moe_sort(
hidden_states,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
topk=1,
block_m=block_m,
)
if a1_q is None:
_fused_moe_mod.aiter.dynamic_per_group_scaled_quant_fp4(
ws["a1_q"],
hidden_states,
ws["a1_scale"],
32,
False,
None,
1,
)
a1_q = ws["a1_q"]
a1_scale_sorted = _moe_mxfp4_sort_into(
ws["a1_scale"],
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
tag=("a1", token_num, model_dim, block_m),
)
a2_inter = metadata.stage1(
a1_q,
w1,
w2,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
ws["a2_inter"],
topk,
block_m=block_m,
a1_scale=a1_scale_sorted,
w1_scale=(
w1_scale.view(_fused_moe_mod.dtypes.fp8_e8m0)
if w1.dtype == _fused_moe_mod.dtypes.fp4x2
else w1_scale
),
sorted_weights=None,
)
a2_q, a2_scale_sorted = _try_fused_mxfp4_quant_moe_sort(
a2_inter.view(-1, inter_dim),
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
topk=topk,
block_m=block_m,
)
if a2_q is None:
_fused_moe_mod.aiter.dynamic_per_group_scaled_quant_fp4(
ws["a2_q"],
a2_inter.view(-1, inter_dim),
ws["a2_scale"],
32,
False,
None,
int(topk),
)
a2_q = ws["a2_q"]
a2_scale_sorted = _moe_mxfp4_sort_into(
ws["a2_scale"].view(token_num, topk, -1),
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
tag=("a2", token_num, topk, inter_dim, block_m),
)
metadata.stage2(
a2_q.view(token_num, topk, -1),
w1,
w2,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
moe_buf,
topk,
w2_scale=(
w2_scale.view(_fused_moe_mod.dtypes.fp8_e8m0)
if w2.dtype == _fused_moe_mod.dtypes.fp4x2
else w2_scale
),
a2_scale=a2_scale_sorted,
block_m=block_m,
sorted_weights=sorted_weights,
)
return moe_buf
def _select_runtime_metadata(
token,
model_dim,
inter_dim,
expert,
topk,
dtype,
activation,
q_type,
route_bucket,
):
key = (int(token), int(inter_dim), int(expert), str(route_bucket))
cands = _runtime_candidate_cfgs(token, inter_dim, expert, route_bucket)
if not cands:
metadata = _patched_get_2stage_cfgs(
token=token,
model_dim=model_dim,
inter_dim=inter_dim,
expert=expert,
topk=topk,
dtype=dtype,
q_dtype_a=_fused_moe_mod.dtypes.fp4x2,
q_dtype_w=_fused_moe_mod.dtypes.fp4x2,
q_type=q_type,
use_g1u1=True,
activation=activation,
doweight_stage1=False,
hidden_pad=0,
intermediate_pad=0,
is_shuffled=True,
)
return metadata, None, None, False
if len(cands) == 1:
_runtime_cfg_cache[key] = cands[0]
metadata = _build_metadata_from_cfg(
cands[0], token, expert, topk, activation, q_type, dtype
)
return metadata, key, cands[0], False
cfg, need_timing = _pick_runtime_cfg(key, cands)
_runtime_cfg_cache[key] = cfg
metadata = _build_metadata_from_cfg(
cfg, token, expert, topk, activation, q_type, dtype
)
return metadata, key, cfg, need_timing
def _fallback_fused_moe(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
topk_weights,
topk_ids,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
hidden_pad,
intermediate_pad,
):
return fused_moe(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
topk_weights,
topk_ids,
expert_mask=None,
activation=ActivationType.Silu,
quant_type=QuantType.per_1x32,
doweight_stage1=False,
w1_scale=gate_up_weight_scale_shuffled,
w2_scale=down_weight_scale_shuffled,
a1_scale=None,
a2_scale=None,
hidden_pad=hidden_pad,
intermediate_pad=intermediate_pad,
)
def _patched_get_2stage_cfgs(
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,
):
if (
int(model_dim) == 7168
and int(expert) == 257
and int(topk) == 9
and activation == ActivationType.Silu
and q_type == QuantType.per_1x32
and dtype == torch.bfloat16
and use_g1u1
and not doweight_stage1
and is_shuffled
):
cfg = _FORCED_E257_CFG.get((int(token), int(inter_dim)))
if cfg is not None:
block_m, kernel1, kernel2 = cfg
use_nt = _use_non_temporal_load(int(token), int(topk), int(expert))
return _fused_moe_mod.MOEMetadata(
functools.partial(
_fused_moe_mod.ck_moe_stage1,
kernelName=kernel1,
activation=activation,
quant_type=q_type,
dtype=dtype,
splitk=0,
use_non_temporal_load=use_nt,
),
_stage2_callable(kernel2, activation, q_type, use_nt),
int(block_m),
0,
False,
)
cfg = _FORCED_E33_CFG.get((int(token), int(inter_dim)))
if (
cfg is not None
and int(model_dim) == 7168
and int(expert) == 33
and int(topk) == 9
and activation == ActivationType.Silu
and q_type == QuantType.per_1x32
and dtype == torch.bfloat16
and use_g1u1
and not doweight_stage1
and is_shuffled
):
block_m, kernel1, kernel2 = cfg
use_nt = _use_non_temporal_load(int(token), int(topk), int(expert))
return _fused_moe_mod.MOEMetadata(
functools.partial(
_fused_moe_mod.ck_moe_stage1,
kernelName=kernel1,
activation=activation,
quant_type=q_type,
dtype=dtype,
splitk=0,
use_non_temporal_load=use_nt,
),
_stage2_callable(kernel2, activation, q_type, use_nt),
int(block_m),
0,
False,
)
return _orig_get_2stage_cfgs(
token=token,
model_dim=model_dim,
inter_dim=inter_dim,
expert=expert,
topk=topk,
dtype=dtype,
q_dtype_a=q_dtype_a,
q_dtype_w=q_dtype_w,
q_type=q_type,
use_g1u1=use_g1u1,
activation=activation,
doweight_stage1=doweight_stage1,
hidden_pad=hidden_pad,
intermediate_pad=intermediate_pad,
is_shuffled=is_shuffled,
)
# Keep only kernel-path override; do not cache cross-call intermediate/final results.
_fused_moe_mod.moe_sorting = _workspace_moe_sorting
_fused_moe_mod.get_2stage_cfgs = _patched_get_2stage_cfgs
def custom_kernel(data: input_t) -> output_t:
"""
Submission template for DeepSeek-R1 MXFP4 MoE kernel.
Input data tuple:
hidden_states: [M, d_hidden] bf16
gate_up_weight: [E, 2*d_expert_pad, d_hidden_pad//2] fp4x2 (raw)
down_weight: [E, d_hidden_pad, d_expert_pad//2] fp4x2 (raw)
gate_up_weight_scale: [E, 2*d_expert_pad, scale_K] e8m0 (raw)
down_weight_scale: [E, d_hidden_pad, scale_K] e8m0 (raw)
gate_up_weight_shuffled: [E, 2*d_expert_pad, d_hidden_pad//2] fp4x2 (shuffled)
down_weight_shuffled: [E, d_hidden_pad, d_expert_pad//2] fp4x2 (shuffled)
gate_up_weight_scale_shuffled:[padded, flat] e8m0 (shuffled)
down_weight_scale_shuffled: [padded, flat] e8m0 (shuffled)
topk_weights: [M, total_top_k] float32
topk_ids: [M, total_top_k] int32
config: dict
Returns:
output: [M, d_hidden] bf16
"""
(
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
hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
intermediate_pad = config["d_expert_pad"] - config["d_expert"]
global _fastpath_disabled
if _fastpath_disabled:
return _fallback_fused_moe(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
topk_weights,
topk_ids,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
hidden_pad,
intermediate_pad,
)
try:
token_num = int(hidden_states.shape[0])
token = int(_fused_moe_mod.get_padded_M(token_num))
topk = int(topk_ids.shape[1])
expert, model_dim, inter_dim = _fused_moe_mod.get_inter_dim(
gate_up_weight_shuffled.shape, down_weight_shuffled.shape
)
expert = int(expert)
model_dim = int(model_dim)
inter_dim = int(inter_dim)
use_fastpath = (expert == 257 and inter_dim == 256) or (
expert == 33 and inter_dim == 512
)
if not use_fastpath:
return _fallback_fused_moe(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
topk_weights,
topk_ids,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
hidden_pad,
intermediate_pad,
)
route_bucket = "balanced"
run_args = (
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
topk_weights,
topk_ids,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
)
metadata, runtime_key, runtime_cfg, need_timing = _select_runtime_metadata(
token=token,
model_dim=model_dim,
inter_dim=inter_dim,
expert=expert,
topk=topk,
dtype=hidden_states.dtype,
activation=ActivationType.Silu,
q_type=QuantType.per_1x32,
route_bucket=route_bucket,
)
if need_timing:
if hidden_states.is_cuda:
torch.cuda.synchronize(hidden_states.device)
t0 = time.perf_counter()
out = _run_fastpath_once(*run_args, metadata=metadata)
if hidden_states.is_cuda:
torch.cuda.synchronize(hidden_states.device)
_update_runtime_stats(
runtime_key, runtime_cfg, (time.perf_counter() - t0) * 1e6
)
return out
return _run_fastpath_once(*run_args, metadata=metadata)
except Exception as e:
global _fastpath_error_printed
if not _fastpath_error_printed:
print(f"[submission fastpath disabled] {type(e).__name__}: {e}")
_fastpath_error_printed = True
_fastpath_disabled = True
return _fallback_fused_moe(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
topk_weights,
topk_ids,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
hidden_pad,
intermediate_pad,
)
scrolls · 882 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