submission 715459
.jonnss · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 746 lines, June 9 Researcher Reciprocity License v1.0.
Submission_v333.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-715459?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:f49a282ca38d6df3d53f57b8c3a331f3e9736934670acf8490f2d7c6f4d81b8a
license declaredunknown
license concludedunknown
authors.jonnss
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
Submission_v333.py746 lines
import dataclasses
import functools
import gc
import inspect
import importlib
import os
import torch
from task import input_t, output_t
os.environ["VLLM_MOE_CHUNK_SIZE"] = "512"
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe
gc.disable()
_PATCH_DONE = False
_DIRECT_FUSED_MOE = None
_FUSED_MOE_MODULE = None
_CURRENT_SORT_BLOCK_SIZE = 32
_CURRENT_SORT_EXPERT = -1
_ADAPTIVE_QUANT_THRESHOLD = 1_000_000
_CFG_CACHE = {}
_SORT_BLOCK_CACHE = {}
_FLYDSL_STAGE2_KEY = "flydsl_moe2_afp4_wfp4_bf16_t16x256x128_atomic"
_FLYDSL_BASE_KEY = "flydsl_moe2_afp4_wfp4_bf16_t32x256x256_atomic"
_E33_STAGE2_SHAPES = {
(128, 7168, 512, 33, 9): 64,
(512, 7168, 512, 33, 9): 128,
(512, 7168, 2048, 33, 9): 64,
}
_DIRECT_PATH_SHAPES = frozenset(
{
(512, 7168, 256, 257, 9),
(128, 7168, 512, 33, 9),
(512, 7168, 512, 33, 9),
(512, 7168, 2048, 33, 9),
}
)
def _replace_meta(meta, *, block_m=None, ksplit=None, stage1=None, stage2=None):
try:
if dataclasses.is_dataclass(meta):
updates = {}
if block_m is not None:
updates["block_m"] = block_m
if ksplit is not None:
updates["ksplit"] = ksplit
if stage1 is not None:
updates["stage1"] = stage1
if stage2 is not None:
updates["stage2"] = stage2
if updates:
return dataclasses.replace(meta, **updates)
except Exception:
pass
try:
if block_m is not None:
meta.block_m = block_m
if ksplit is not None:
meta.ksplit = ksplit
if stage1 is not None:
meta.stage1 = stage1
if stage2 is not None:
meta.stage2 = stage2
return meta
except Exception:
pass
if all(hasattr(meta, name) for name in ("stage1", "stage2", "block_m", "ksplit", "run_1stage")):
has_bias = getattr(meta, "has_bias", False)
use_non_temporal_load = getattr(meta, "use_non_temporal_load", True)
try:
return type(meta)(
meta.stage1 if stage1 is None else stage1,
meta.stage2 if stage2 is None else stage2,
meta.block_m if block_m is None else block_m,
meta.ksplit if ksplit is None else ksplit,
meta.run_1stage,
has_bias,
use_non_temporal_load,
)
except Exception:
pass
return meta
def _split_kw_name(fn) -> str:
if isinstance(fn, functools.partial):
keywords = dict(fn.keywords or {})
if "splitk" in keywords:
return "splitk"
if "split_k" in keywords:
return "split_k"
target = fn.func if isinstance(fn, functools.partial) else fn
name = getattr(target, "__name__", "")
if name in {"cktile_moe_gemm1", "cktile_moe_stage1"}:
return "split_k"
return "splitk"
def _retune_partial(fn, *, block_m=None, splitk=None):
if not isinstance(fn, functools.partial):
return fn
keywords = dict(fn.keywords or {})
if block_m is not None:
keywords["block_m"] = block_m
if splitk is not None:
keywords.pop("splitk", None)
keywords.pop("split_k", None)
keywords[_split_kw_name(fn)] = splitk
return functools.partial(fn.func, *(fn.args or ()), **keywords)
def _wrap_flydsl_stage2(stage2_wrapper, original_stage2, kernel_name: str):
keywords = {}
if isinstance(original_stage2, functools.partial):
keywords.update(original_stage2.keywords or {})
keywords["kernelName"] = kernel_name
return functools.partial(stage2_wrapper, **keywords)
def _register_flydsl_t16(moe_kernels) -> None:
kernel_params = getattr(moe_kernels, "_KERNEL_PARAMS", None)
if not isinstance(kernel_params, dict):
return
if _FLYDSL_STAGE2_KEY in kernel_params:
return
base = dict(kernel_params.get(_FLYDSL_BASE_KEY, {}))
if not base:
return
base.update(
tile_m=16,
tile_n=256,
tile_k=128,
mode="atomic",
MPerBlock=16,
)
kernel_params[_FLYDSL_STAGE2_KEY] = base
def _precompile_flydsl_defs(moe_kernels) -> None:
get_compiled_stage2 = getattr(moe_kernels, "_get_compiled_stage2", None)
if get_compiled_stage2 is None:
return
for inter_dim in (512, 2048):
try:
get_compiled_stage2(
model_dim=7168,
inter_dim=inter_dim,
experts=33,
topk=9,
tile_m=16,
tile_n=256,
tile_k=128,
doweight=True,
a_dtype="fp4",
b_dtype="fp4",
out_dtype="bf16",
)
except Exception:
pass
def _choose_block_size_m(hidden_states: torch.Tensor, gate_up_weight_shuffled: torch.Tensor) -> int | None:
m = int(hidden_states.shape[0])
if m <= 16:
return 16
return None
def _should_use_direct_path(
hidden_states: torch.Tensor,
gate_up_weight_shuffled: torch.Tensor,
topk_ids: torch.Tensor,
config,
) -> bool:
shape = (
int(hidden_states.shape[0]),
int(config["d_hidden"]),
int(config["d_expert"]),
int(gate_up_weight_shuffled.shape[0]),
int(topk_ids.shape[1]),
)
return shape in _DIRECT_PATH_SHAPES
def _use_bypass(token: int, expert: int) -> bool:
return token <= 16 or (token <= 128 and expert > 64)
def _desired_ksplit(token: int, expert: int) -> int:
return 2 if _use_bypass(token, expert) else 0
def _cacheable(value):
if isinstance(value, (str, int, float, bool, type(None))):
return value
if isinstance(value, tuple):
return tuple(_cacheable(item) for item in value)
return repr(value)
def _make_cfg_cache_key(sig: inspect.Signature | None, use_bypass: bool, call_args, call_kwargs):
if sig is not None:
try:
bound = sig.bind_partial(*call_args, **call_kwargs)
return (
use_bypass,
tuple((name, _cacheable(value)) for name, value in bound.arguments.items()),
)
except Exception:
pass
return (
use_bypass,
tuple(_cacheable(value) for value in call_args),
tuple(sorted((key, _cacheable(value)) for key, value in call_kwargs.items())),
)
def _call_with_bypass_env(use_bypass: bool, fn, *args, **kwargs):
previous = os.environ.get("AITER_BYPASS_TUNE_CONFIG")
os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1" if use_bypass else "0"
try:
return fn(*args, **kwargs)
finally:
if previous is None:
os.environ.pop("AITER_BYPASS_TUNE_CONFIG", None)
else:
os.environ["AITER_BYPASS_TUNE_CONFIG"] = previous
def _install_patch() -> None:
global _PATCH_DONE, _DIRECT_FUSED_MOE, _FUSED_MOE_MODULE
if _PATCH_DONE:
return
_PATCH_DONE = True
try:
aiter_mod = importlib.import_module("aiter")
fused_moe_module = importlib.import_module("aiter.fused_moe")
_FUSED_MOE_MODULE = fused_moe_module
fp4_utils = importlib.import_module("aiter.utility.fp4_utils")
moe_kernels = importlib.import_module("aiter.ops.flydsl.moe_kernels")
_register_flydsl_t16(moe_kernels)
_precompile_flydsl_defs(moe_kernels)
original_cfgs = getattr(fused_moe_module, "get_2stage_cfgs", None)
raw_cfgs = getattr(original_cfgs, "__wrapped__", original_cfgs)
original_cfgs_sig = inspect.signature(original_cfgs) if original_cfgs is not None else None
raw_cfgs_sig = inspect.signature(raw_cfgs) if raw_cfgs is not None else None
original_get_ksplit = getattr(fused_moe_module, "get_ksplit", None)
original_get_ksplit_sig = inspect.signature(original_get_ksplit) if original_get_ksplit is not None else None
original_quant_sort = getattr(fused_moe_module, "fused_dynamic_mxfp4_quant_moe_sort", None)
quant_hip = getattr(aiter_mod, "per_1x32_f4_quant_hip", None)
moe_mxfp4_sort = getattr(fp4_utils, "moe_mxfp4_sort", None)
flydsl_stage2_wrapper = getattr(fused_moe_module, "_flydsl_stage2_wrapper", None)
fused_moe_internal = getattr(fused_moe_module, "fused_moe_", None)
fused_moe_2stages = getattr(fused_moe_module, "fused_moe_2stages", None)
direct_fused_moe = None
if fused_moe_internal is not None:
try:
direct_fused_moe = inspect.unwrap(fused_moe_internal)
except Exception:
direct_fused_moe = getattr(fused_moe_internal, "__wrapped__", None)
if direct_fused_moe is None:
direct_fused_moe = fused_moe_internal
_DIRECT_FUSED_MOE = direct_fused_moe
if original_cfgs is None or flydsl_stage2_wrapper is None:
return
def _extract_int(
sig: inspect.Signature | None,
call_args,
call_kwargs,
name: str,
fallback_index: int,
) -> int:
if sig is not None:
try:
bound_args = sig.bind_partial(*call_args, **call_kwargs).arguments
if name in bound_args:
return int(bound_args[name])
except Exception:
pass
value = call_kwargs.get(name, call_args[fallback_index] if len(call_args) > fallback_index else -1)
try:
return int(value)
except Exception:
return -1
def patched_get_2stage_cfgs(*args, **kwargs):
global _CURRENT_SORT_BLOCK_SIZE, _CURRENT_SORT_EXPERT
token = _extract_int(original_cfgs_sig, args, kwargs, "token", 0)
model_dim = _extract_int(original_cfgs_sig, args, kwargs, "model_dim", 1)
inter_dim = _extract_int(original_cfgs_sig, args, kwargs, "inter_dim", 2)
expert = _extract_int(original_cfgs_sig, args, kwargs, "expert", 3)
topk = _extract_int(original_cfgs_sig, args, kwargs, "topk", 4)
shape = (token, model_dim, inter_dim, expert, topk)
use_bypass = _use_bypass(token, expert)
cache_key = _make_cfg_cache_key(raw_cfgs_sig, use_bypass, args, kwargs)
meta = _CFG_CACHE.get(cache_key)
if meta is None:
raw_meta = _call_with_bypass_env(use_bypass, raw_cfgs, *args, **kwargs)
ksplit = _desired_ksplit(token, expert)
block_m = getattr(raw_meta, "block_m", None)
stage1 = _retune_partial(getattr(raw_meta, "stage1", None), splitk=ksplit)
stage2 = getattr(raw_meta, "stage2", None)
override_block_m = _E33_STAGE2_SHAPES.get(shape)
if override_block_m is not None:
block_m = override_block_m
stage1 = _retune_partial(stage1, block_m=override_block_m, splitk=0)
stage2 = _retune_partial(stage2, block_m=override_block_m)
stage2 = _wrap_flydsl_stage2(flydsl_stage2_wrapper, stage2, _FLYDSL_STAGE2_KEY)
ksplit = 0
if shape == (512, 7168, 256, 257, 9):
ksplit = 0
stage1 = _retune_partial(stage1, splitk=0)
meta = _replace_meta(raw_meta, block_m=block_m, ksplit=ksplit, stage1=stage1, stage2=stage2)
_CFG_CACHE[cache_key] = meta
_CURRENT_SORT_BLOCK_SIZE = int(getattr(meta, "block_m", None) or 32)
_CURRENT_SORT_EXPERT = expert
return meta
def patched_get_ksplit(*args, **kwargs):
token = _extract_int(original_get_ksplit_sig, args, kwargs, "token", 0)
expert = _extract_int(original_get_ksplit_sig, args, kwargs, "expert", 3)
if token >= 0 and expert >= 0:
return _desired_ksplit(token, expert)
if original_get_ksplit is None:
return 0
return original_get_ksplit(*args, **kwargs)
fused_moe_module.get_2stage_cfgs = patched_get_2stage_cfgs
if original_get_ksplit is not None:
fused_moe_module.get_ksplit = patched_get_ksplit
for maybe_globals in (
getattr(fused_moe, "__globals__", None),
getattr(fused_moe_internal, "__globals__", None),
getattr(fused_moe_2stages, "__globals__", None),
getattr(direct_fused_moe, "__globals__", None),
):
if isinstance(maybe_globals, dict) and maybe_globals.get("get_2stage_cfgs") is original_cfgs:
maybe_globals["get_2stage_cfgs"] = patched_get_2stage_cfgs
if isinstance(maybe_globals, dict) and maybe_globals.get("get_ksplit") is original_get_ksplit:
maybe_globals["get_ksplit"] = patched_get_ksplit
if original_quant_sort is None or quant_hip is None or moe_mxfp4_sort is None:
return
def split_quant_sort(
x: torch.Tensor,
sorted_ids: torch.Tensor,
num_valid_ids: torch.Tensor,
token_num: int,
topk: int,
block_size: int = 32,
scaling_mode: str = "even",
):
if scaling_mode != "even":
return original_quant_sort(
x,
sorted_ids,
num_valid_ids,
token_num,
topk,
block_size=block_size,
scaling_mode=scaling_mode,
)
quant_out = quant_hip(x, scale=None, shuffle=False)
x_fp4, blockscale_e8m0 = quant_out[:2]
requested_block_size = int(_CURRENT_SORT_BLOCK_SIZE or block_size)
cache_key = (requested_block_size, int(_CURRENT_SORT_EXPERT))
candidates = []
for candidate in (
_SORT_BLOCK_CACHE.get(cache_key),
requested_block_size,
int(block_size),
max(requested_block_size, 128),
64,
128,
256,
512,
):
if candidate and candidate not in candidates:
candidates.append(candidate)
last_error = None
for candidate in candidates:
try:
blockscale_sorted = moe_mxfp4_sort(
blockscale_e8m0,
sorted_ids,
num_valid_ids,
token_num,
candidate,
)
_SORT_BLOCK_CACHE[cache_key] = candidate
return x_fp4, blockscale_sorted
except AssertionError as exc:
last_error = exc
if last_error is not None:
raise last_error
return x_fp4, blockscale_e8m0
def adaptive_quant_sort(*args, **kwargs):
x = kwargs.get("x", args[0] if len(args) > 0 else None)
scaling_mode = kwargs.get("scaling_mode", args[6] if len(args) > 6 else "even")
if x is None or scaling_mode != "even":
return original_quant_sort(*args, **kwargs)
if int(x.numel()) > _ADAPTIVE_QUANT_THRESHOLD:
return split_quant_sort(*args, **kwargs)
return original_quant_sort(*args, **kwargs)
fused_moe_module.fused_dynamic_mxfp4_quant_moe_sort = adaptive_quant_sort
for maybe_globals in (
getattr(fused_moe_internal, "__globals__", None),
getattr(fused_moe_2stages, "__globals__", None),
getattr(direct_fused_moe, "__globals__", None),
):
if (
isinstance(maybe_globals, dict)
and maybe_globals.get("fused_dynamic_mxfp4_quant_moe_sort") is original_quant_sort
):
maybe_globals["fused_dynamic_mxfp4_quant_moe_sort"] = adaptive_quant_sort
except Exception:
pass
def _call_moe(
hidden_states: torch.Tensor,
gate_up_weight_shuffled: torch.Tensor,
down_weight_shuffled: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
gate_up_weight_scale_shuffled: torch.Tensor,
down_weight_scale_shuffled: torch.Tensor,
hidden_pad: int,
intermediate_pad: int,
block_size_m: int | None,
):
direct_fused_moe = _DIRECT_FUSED_MOE
if direct_fused_moe is None:
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,
block_size_M=block_size_m,
hidden_pad=hidden_pad,
intermediate_pad=intermediate_pad,
)
block_size_arg = -1 if not block_size_m else int(block_size_m)
try:
return direct_fused_moe(
hidden_states=hidden_states,
w1=gate_up_weight_shuffled,
w2=down_weight_shuffled,
topk_weight=topk_weights,
topk_ids=topk_ids,
expert_mask=None,
activation=ActivationType.Silu.value,
quant_type=QuantType.per_1x32.value,
doweight_stage1=False,
w1_scale=gate_up_weight_scale_shuffled,
w2_scale=down_weight_scale_shuffled,
a1_scale=None,
a2_scale=None,
block_size_M=block_size_arg,
num_local_tokens=None,
moe_sorting_dispatch_policy=0,
dtype=None,
hidden_pad=hidden_pad,
intermediate_pad=intermediate_pad,
bias1=None,
bias2=None,
)
except Exception:
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,
block_size_M=block_size_m,
hidden_pad=hidden_pad,
intermediate_pad=intermediate_pad,
)
def _call_moe_inlined(
hidden_states: torch.Tensor,
gate_up_weight_shuffled: torch.Tensor,
down_weight_shuffled: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
gate_up_weight_scale_shuffled: torch.Tensor,
down_weight_scale_shuffled: torch.Tensor,
hidden_pad: int,
intermediate_pad: int,
block_size_m: int | None,
):
fused_moe_module = _FUSED_MOE_MODULE
if fused_moe_module is None:
return _call_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,
block_size_m,
)
try:
get_inter_dim = getattr(fused_moe_module, "get_inter_dim")
dtypes = getattr(fused_moe_module, "dtypes")
quant_remap = getattr(fused_moe_module, "quant_remap")
get_gfx = getattr(fused_moe_module, "get_gfx")
get_padded_M = getattr(fused_moe_module, "get_padded_M")
get_2stage_cfgs = getattr(fused_moe_module, "get_2stage_cfgs")
moe_sorting = getattr(fused_moe_module, "moe_sorting")
fused_moe_2stages = getattr(fused_moe_module, "fused_moe_2stages")
activation = ActivationType.Silu
quant_type = QuantType.per_1x32
M, topk = topk_ids.shape
E, model_dim, inter_dim = get_inter_dim(
gate_up_weight_shuffled.shape,
down_weight_shuffled.shape,
)
assert gate_up_weight_shuffled.shape[1] in [inter_dim, inter_dim * 2]
is_g1u1 = inter_dim != gate_up_weight_shuffled.shape[1]
is_shuffled = getattr(gate_up_weight_shuffled, "is_shuffled", False)
dtype = hidden_states.dtype
global_e = E
q_dtype_w = gate_up_weight_shuffled.dtype
q_dtype_a = gate_up_weight_shuffled.dtype if gate_up_weight_shuffled.dtype != torch.uint32 else dtypes.fp8
quant_type = quant_remap.get(quant_type, quant_type)
if quant_type == QuantType.per_1x32:
if activation == ActivationType.Swiglu:
if get_gfx() != "gfx950" or M < 512:
q_dtype_a = dtypes.bf16
else:
q_dtype_a = dtypes.fp8
else:
q_dtype_a = dtypes.fp4x2
metadata = get_2stage_cfgs(
get_padded_M(M),
model_dim,
inter_dim,
E,
topk,
dtype,
q_dtype_a,
q_dtype_w,
quant_type,
is_g1u1,
activation,
False,
hidden_pad,
intermediate_pad,
is_shuffled,
)
block_size_eff = getattr(metadata, "block_m", None) if block_size_m is None else block_size_m
if block_size_eff is not None:
block_size_eff = int(block_size_eff)
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = moe_sorting(
topk_ids,
topk_weights,
global_e,
model_dim,
dtype,
block_size_eff,
None,
None,
0,
)
if getattr(metadata, "run_1stage", False):
return metadata.stage1(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
topk,
sorted_ids,
sorted_weights,
sorted_expert_ids,
num_valid_ids,
moe_buf,
is_g1u1,
block_size_eff,
q_dtype_a=q_dtype_a,
q_dtype_w=q_dtype_w,
w1_scale=gate_up_weight_scale_shuffled,
w2_scale=down_weight_scale_shuffled,
a1_scale=None,
a2_scale=None,
num_local_tokens=None,
M=M,
device=topk_ids.device,
doweight_stage1=False,
)
return fused_moe_2stages(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
topk,
sorted_ids,
sorted_weights,
sorted_expert_ids,
num_valid_ids,
moe_buf,
is_g1u1,
block_size_eff,
activation=activation,
quant_type=quant_type,
doweight_stage1=False,
q_dtype_a=q_dtype_a,
q_dtype_w=q_dtype_w,
w1_scale=gate_up_weight_scale_shuffled,
w2_scale=down_weight_scale_shuffled,
a1_scale=None,
a2_scale=None,
num_local_tokens=None,
hidden_pad=hidden_pad,
intermediate_pad=intermediate_pad,
bias1=None,
bias2=None,
)
except Exception:
return _call_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,
block_size_m,
)
@torch.inference_mode()
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
_install_patch()
hidden_pad = int(config["d_hidden_pad"]) - int(config["d_hidden"])
intermediate_pad = int(config["d_expert_pad"]) - int(config["d_expert"])
block_size_m = _choose_block_size_m(hidden_states, gate_up_weight_shuffled)
if not _should_use_direct_path(hidden_states, gate_up_weight_shuffled, topk_ids, config):
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,
block_size_M=block_size_m,
hidden_pad=hidden_pad,
intermediate_pad=intermediate_pad,
)
return _call_moe_inlined(
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,
block_size_m,
)
scrolls · 746 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 711471.
⋯ 19 unchanged lines_PATCH_DONE = False+ _DIRECT_FUSED_MOE = None+ _FUSED_MOE_MODULE = None_CURRENT_SORT_BLOCK_SIZE = 32_CURRENT_SORT_EXPERT = -1_ADAPTIVE_QUANT_THRESHOLD = 1_000_000⋯ 6 unchanged lines(512, 7168, 512, 33, 9): 128,(512, 7168, 2048, 33, 9): 64,}+ _DIRECT_PATH_SHAPES = frozenset(+ {+ (512, 7168, 256, 257, 9),+ (128, 7168, 512, 33, 9),+ (512, 7168, 512, 33, 9),+ (512, 7168, 2048, 33, 9),+ }+ )def _replace_meta(meta, *, block_m=None, ksplit=None, stage1=None, stage2=None):⋯ 132 unchanged linesreturn None+ def _should_use_direct_path(+ hidden_states: torch.Tensor,+ gate_up_weight_shuffled: torch.Tensor,+ topk_ids: torch.Tensor,+ config,+ ) -> bool:+ shape = (+ int(hidden_states.shape[0]),+ int(config["d_hidden"]),+ int(config["d_expert"]),+ int(gate_up_weight_shuffled.shape[0]),+ int(topk_ids.shape[1]),+ )+ return shape in _DIRECT_PATH_SHAPES++def _use_bypass(token: int, expert: int) -> bool:return token <= 16 or (token <= 128 and expert > 64)⋯ 40 unchanged linesdef _install_patch() -> None:- global _PATCH_DONE+ global _PATCH_DONE, _DIRECT_FUSED_MOE, _FUSED_MOE_MODULEif _PATCH_DONE:return_PATCH_DONE = True⋯ 1 unchanged linestry:aiter_mod = importlib.import_module("aiter")fused_moe_module = importlib.import_module("aiter.fused_moe")+ _FUSED_MOE_MODULE = fused_moe_modulefp4_utils = importlib.import_module("aiter.utility.fp4_utils")moe_kernels = importlib.import_module("aiter.ops.flydsl.moe_kernels")⋯ 12 unchanged linesflydsl_stage2_wrapper = getattr(fused_moe_module, "_flydsl_stage2_wrapper", None)fused_moe_internal = getattr(fused_moe_module, "fused_moe_", None)fused_moe_2stages = getattr(fused_moe_module, "fused_moe_2stages", None)+ direct_fused_moe = None+ if fused_moe_internal is not None:+ try:+ direct_fused_moe = inspect.unwrap(fused_moe_internal)+ except Exception:+ direct_fused_moe = getattr(fused_moe_internal, "__wrapped__", None)+ if direct_fused_moe is None:+ direct_fused_moe = fused_moe_internal+ _DIRECT_FUSED_MOE = direct_fused_moeif original_cfgs is None or flydsl_stage2_wrapper is None:return⋯ 71 unchanged linesgetattr(fused_moe, "__globals__", None),getattr(fused_moe_internal, "__globals__", None),getattr(fused_moe_2stages, "__globals__", None),+ getattr(direct_fused_moe, "__globals__", None),):if isinstance(maybe_globals, dict) and maybe_globals.get("get_2stage_cfgs") is original_cfgs:maybe_globals["get_2stage_cfgs"] = patched_get_2stage_cfgs⋯ 73 unchanged linesfor maybe_globals in (getattr(fused_moe_internal, "__globals__", None),getattr(fused_moe_2stages, "__globals__", None),+ getattr(direct_fused_moe, "__globals__", None),):if (isinstance(maybe_globals, dict)⋯ 4 unchanged linespass+ def _call_moe(+ hidden_states: torch.Tensor,+ gate_up_weight_shuffled: torch.Tensor,+ down_weight_shuffled: torch.Tensor,+ topk_weights: torch.Tensor,+ topk_ids: torch.Tensor,+ gate_up_weight_scale_shuffled: torch.Tensor,+ down_weight_scale_shuffled: torch.Tensor,+ hidden_pad: int,+ intermediate_pad: int,+ block_size_m: int | None,+ ):+ direct_fused_moe = _DIRECT_FUSED_MOE+ if direct_fused_moe is None:+ 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,+ block_size_M=block_size_m,+ hidden_pad=hidden_pad,+ intermediate_pad=intermediate_pad,+ )++ block_size_arg = -1 if not block_size_m else int(block_size_m)+ try:+ return direct_fused_moe(+ hidden_states=hidden_states,+ w1=gate_up_weight_shuffled,+ w2=down_weight_shuffled,+ topk_weight=topk_weights,+ topk_ids=topk_ids,+ expert_mask=None,+ activation=ActivationType.Silu.value,+ quant_type=QuantType.per_1x32.value,+ doweight_stage1=False,+ w1_scale=gate_up_weight_scale_shuffled,+ w2_scale=down_weight_scale_shuffled,+ a1_scale=None,+ a2_scale=None,+ block_size_M=block_size_arg,+ num_local_tokens=None,+ moe_sorting_dispatch_policy=0,+ dtype=None,+ hidden_pad=hidden_pad,+ intermediate_pad=intermediate_pad,+ bias1=None,+ bias2=None,+ )+ except Exception:+ 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,+ block_size_M=block_size_m,+ hidden_pad=hidden_pad,+ intermediate_pad=intermediate_pad,+ )+++ def _call_moe_inlined(+ hidden_states: torch.Tensor,+ gate_up_weight_shuffled: torch.Tensor,+ down_weight_shuffled: torch.Tensor,+ topk_weights: torch.Tensor,+ topk_ids: torch.Tensor,+ gate_up_weight_scale_shuffled: torch.Tensor,+ down_weight_scale_shuffled: torch.Tensor,+ hidden_pad: int,+ intermediate_pad: int,+ block_size_m: int | None,+ ):+ fused_moe_module = _FUSED_MOE_MODULE+ if fused_moe_module is None:+ return _call_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,+ block_size_m,+ )++ try:+ get_inter_dim = getattr(fused_moe_module, "get_inter_dim")+ dtypes = getattr(fused_moe_module, "dtypes")+ quant_remap = getattr(fused_moe_module, "quant_remap")+ get_gfx = getattr(fused_moe_module, "get_gfx")+ get_padded_M = getattr(fused_moe_module, "get_padded_M")+ get_2stage_cfgs = getattr(fused_moe_module, "get_2stage_cfgs")+ moe_sorting = getattr(fused_moe_module, "moe_sorting")+ fused_moe_2stages = getattr(fused_moe_module, "fused_moe_2stages")++ activation = ActivationType.Silu+ quant_type = QuantType.per_1x32+ M, topk = topk_ids.shape+ E, model_dim, inter_dim = get_inter_dim(+ gate_up_weight_shuffled.shape,+ down_weight_shuffled.shape,+ )++ assert gate_up_weight_shuffled.shape[1] in [inter_dim, inter_dim * 2]+ is_g1u1 = inter_dim != gate_up_weight_shuffled.shape[1]+ is_shuffled = getattr(gate_up_weight_shuffled, "is_shuffled", False)+ dtype = hidden_states.dtype+ global_e = E+ q_dtype_w = gate_up_weight_shuffled.dtype+ q_dtype_a = gate_up_weight_shuffled.dtype if gate_up_weight_shuffled.dtype != torch.uint32 else dtypes.fp8+ quant_type = quant_remap.get(quant_type, quant_type)+ if quant_type == QuantType.per_1x32:+ if activation == ActivationType.Swiglu:+ if get_gfx() != "gfx950" or M < 512:+ q_dtype_a = dtypes.bf16+ else:+ q_dtype_a = dtypes.fp8+ else:+ q_dtype_a = dtypes.fp4x2++ metadata = get_2stage_cfgs(+ get_padded_M(M),+ model_dim,+ inter_dim,+ E,+ topk,+ dtype,+ q_dtype_a,+ q_dtype_w,+ quant_type,+ is_g1u1,+ activation,+ False,+ hidden_pad,+ intermediate_pad,+ is_shuffled,+ )++ block_size_eff = getattr(metadata, "block_m", None) if block_size_m is None else block_size_m+ if block_size_eff is not None:+ block_size_eff = int(block_size_eff)++ sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = moe_sorting(+ topk_ids,+ topk_weights,+ global_e,+ model_dim,+ dtype,+ block_size_eff,+ None,+ None,+ 0,+ )++ if getattr(metadata, "run_1stage", False):+ return metadata.stage1(+ hidden_states,+ gate_up_weight_shuffled,+ down_weight_shuffled,+ topk,+ sorted_ids,+ sorted_weights,+ sorted_expert_ids,+ num_valid_ids,+ moe_buf,+ is_g1u1,+ block_size_eff,+ q_dtype_a=q_dtype_a,+ q_dtype_w=q_dtype_w,+ w1_scale=gate_up_weight_scale_shuffled,+ w2_scale=down_weight_scale_shuffled,+ a1_scale=None,+ a2_scale=None,+ num_local_tokens=None,+ M=M,+ device=topk_ids.device,+ doweight_stage1=False,+ )++ return fused_moe_2stages(+ hidden_states,+ gate_up_weight_shuffled,+ down_weight_shuffled,+ topk,+ sorted_ids,+ sorted_weights,+ sorted_expert_ids,+ num_valid_ids,+ moe_buf,+ is_g1u1,+ block_size_eff,+ activation=activation,+ quant_type=quant_type,+ doweight_stage1=False,+ q_dtype_a=q_dtype_a,+ q_dtype_w=q_dtype_w,+ w1_scale=gate_up_weight_scale_shuffled,+ w2_scale=down_weight_scale_shuffled,+ a1_scale=None,+ a2_scale=None,+ num_local_tokens=None,+ hidden_pad=hidden_pad,+ intermediate_pad=intermediate_pad,+ bias1=None,+ bias2=None,+ )+ except Exception:+ return _call_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,+ block_size_m,+ )++@torch.inference_mode()def custom_kernel(data: input_t) -> output_t:(⋯ 17 unchanged linesintermediate_pad = int(config["d_expert_pad"]) - int(config["d_expert"])block_size_m = _choose_block_size_m(hidden_states, gate_up_weight_shuffled)- return fused_moe(+ if not _should_use_direct_path(hidden_states, gate_up_weight_shuffled, topk_ids, config):+ 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,+ block_size_M=block_size_m,+ hidden_pad=hidden_pad,+ intermediate_pad=intermediate_pad,+ )++ return _call_moe_inlined(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,- block_size_M=block_size_m,- hidden_pad=hidden_pad,- intermediate_pad=intermediate_pad,+ gate_up_weight_scale_shuffled,+ down_weight_scale_shuffled,+ hidden_pad,+ intermediate_pad,+ block_size_m,)
scrolls · 393 diff lines total
Best evidence level for this revision: reported
JSON