submission 564099
oofbaroomf · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 620 lines, June 9 Researcher Reciprocity License v1.0.
submission_543848_257bs512k1_noopus.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-564099?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:b289f11bccbc6ac288f3daded1449241360ab8b75658a2dde52855203f1b1b28
license declaredunknown
license concludedunknown
authorsoofbaroomf
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
split_k=ksplit,Kernel source
submission_543848_257bs512k1_noopus.py620 lines
import functools
import importlib.util
import inspect
from pathlib import Path
import torch
from task import input_t, output_t
def _set_fmoe_tune_path() -> None:
spec = importlib.util.find_spec("aiter")
if spec is None or spec.origin is None:
return
aiter_root = Path(spec.origin).resolve().parent
tune_file = aiter_root / "configs" / "model_configs" / "dsv3_fp4_tuned_fmoe.csv"
if tune_file.exists():
import os
os.environ["AITER_CONFIG_FMOE"] = str(tune_file)
_set_fmoe_tune_path()
import aiter # noqa: E402
import aiter.fused_moe as _fused_moe_mod # noqa: E402
from aiter import ActivationType, QuantType # noqa: E402
_QSORT_THRESHOLD = 256
_FORCE_SEP_QSORT_257_TOKENS = (16, 128)
_KERNEL_33_512_STAGE1 = (
"moe_ck2stages_gemm1_256x64x128x128_1x4_"
"MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_"
"FP4X2_FP4X2_B16"
)
_KERNEL_33_512_STAGE2 = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
_KERNEL_33_2048_STAGE1 = (
"moe_ck2stages_gemm1_256x128x128x128_1x4_"
"MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_"
"FP4X2_FP4X2_B16"
)
_KERNEL_33_2048_STAGE2 = (
"moe_ck2stages_gemm2_256x128x128x128_1x4_"
"MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_"
"FP4X2_FP4X2_B16"
)
_A2_CACHE: dict[tuple[str, tuple[int, ...], str], torch.Tensor] = {}
_SORT_CACHE: dict[
tuple[str, int, int, int, int, str, int],
tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor],
] = {}
def _pad_rows(x: torch.Tensor, target_rows: int) -> torch.Tensor:
if x.shape[0] >= target_rows:
return x
pad_rows = target_rows - x.shape[0]
pad = x[-1:].expand(pad_rows, *x.shape[1:]).clone()
return torch.cat((x, pad), dim=0)
def _repartial_with(stage, **updates):
fn = getattr(stage, "func", None)
if fn is None:
return stage
params = inspect.signature(fn).parameters
kwargs = dict(stage.keywords or {})
for key, value in updates.items():
if key in params:
kwargs[key] = value
return functools.partial(fn, *(stage.args or ()), **kwargs)
def _make_k2_metadata(*args):
(
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,
) = args
ksplit = 2
return _fused_moe_mod.MOEMetadata(
functools.partial(
_fused_moe_mod.cktile_moe_stage1,
n_pad_zeros=intermediate_pad // 64 * 64 * (2 if use_g1u1 else 1),
k_pad_zeros=hidden_pad // 128 * 128,
activation=activation,
split_k=ksplit,
),
functools.partial(
_fused_moe_mod.cktile_moe_stage2,
n_pad_zeros=hidden_pad // 64 * 64,
k_pad_zeros=intermediate_pad // 128 * 128,
activation=activation,
),
16 if token < 2048 else 32 if token < 16384 else 64,
ksplit,
False,
False,
True,
)
def _make_exact_ck_metadata(*args, block_m: int, kernel1: str, kernel2: str):
(
_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,
) = args
if kernel2.startswith("flydsl_"):
stage2 = functools.partial(
_fused_moe_mod._flydsl_stage2_wrapper,
kernelName=kernel2,
)
else:
stage2 = functools.partial(
aiter.ck_moe_stage2_fwd,
kernelName=kernel2,
activation=activation,
quant_type=q_type,
use_non_temporal_load=True,
)
return _fused_moe_mod.MOEMetadata(
functools.partial(
_fused_moe_mod.ck_moe_stage1,
kernelName=kernel1,
activation=activation,
quant_type=q_type,
splitk=0,
use_non_temporal_load=True,
dtype=dtype,
),
stage2,
block_m,
0,
False,
False,
True,
)
def _make_k1_large_metadata(md):
if md.run_1stage:
return md
stage1 = _repartial_with(md.stage1, splitk=1, split_k=1, use_non_temporal_load=True)
stage2 = _repartial_with(md.stage2, use_non_temporal_load=True)
return _fused_moe_mod.MOEMetadata(
stage1=stage1,
stage2=stage2,
block_m=md.block_m,
ksplit=1,
run_1stage=md.run_1stage,
has_bias=md.has_bias,
use_non_temporal_load=True,
)
def _patch_metadata() -> None:
if getattr(
_fused_moe_mod,
"_oof_force_nt_sep256_257k2_33d512flydsl_33d2048bm64",
False,
):
return
orig_get_2stage_cfgs = _fused_moe_mod.get_2stage_cfgs
@functools.lru_cache(maxsize=2048)
def wrapped_get_2stage_cfgs(*args):
(
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,
) = args
if (
expert == 257
and inter_dim == 256
and token <= 128
and q_type == QuantType.per_1x32
and q_dtype_w == aiter.dtypes.fp4x2
and is_shuffled
):
return _make_k2_metadata(*args)
if (
expert == 33
and inter_dim == 512
and token <= 128
and q_type == QuantType.per_1x32
and q_dtype_w == aiter.dtypes.fp4x2
and is_shuffled
):
return _make_k2_metadata(*args)
if (
expert == 33
and inter_dim == 512
and token == 512
and q_type == QuantType.per_1x32
and q_dtype_w == aiter.dtypes.fp4x2
and is_shuffled
):
return _make_exact_ck_metadata(
*args,
block_m=64,
kernel1=_KERNEL_33_512_STAGE1,
kernel2=_KERNEL_33_512_STAGE2,
)
md = orig_get_2stage_cfgs(*args)
if (
expert == 257
and inter_dim == 256
and token == 512
and q_type == QuantType.per_1x32
and q_dtype_w == aiter.dtypes.fp4x2
and is_shuffled
):
return _make_k1_large_metadata(md)
if (
expert == 33
and inter_dim == 2048
and token == 512
and q_type == QuantType.per_1x32
and q_dtype_w == aiter.dtypes.fp4x2
and is_shuffled
):
return _make_exact_ck_metadata(
*args,
block_m=32,
kernel1=_KERNEL_33_2048_STAGE1,
kernel2=_KERNEL_33_2048_STAGE2,
)
if md.run_1stage:
return md
stage1 = _repartial_with(md.stage1, use_non_temporal_load=True)
stage2 = _repartial_with(md.stage2, use_non_temporal_load=True)
return _fused_moe_mod.MOEMetadata(
stage1=stage1,
stage2=stage2,
block_m=md.block_m,
ksplit=md.ksplit,
run_1stage=md.run_1stage,
has_bias=md.has_bias,
use_non_temporal_load=True,
)
_fused_moe_mod.get_2stage_cfgs = wrapped_get_2stage_cfgs
_fused_moe_mod._oof_force_nt_sep256_257k2_33d512flydsl_33d2048bm64 = True
def _patch_qsort_threshold() -> None:
if getattr(_fused_moe_mod, "_oof_qsort_threshold", None) == _QSORT_THRESHOLD:
return
source = inspect.getsource(_fused_moe_mod.fused_moe_2stages)
source = source.replace(
"token_num_quant_moe_sort_switch = 1024",
f"token_num_quant_moe_sort_switch = {_QSORT_THRESHOLD}",
1,
)
source = source.replace(
" E, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)",
" E, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)\n"
" force_sep_qsort_257 = (\n"
f" token_num in {_FORCE_SEP_QSORT_257_TOKENS}\n"
" and E == 257\n"
" and inter_dim == 256\n"
" )",
1,
)
source = source.replace(
"if token_num <= token_num_quant_moe_sort_switch:",
"if token_num <= token_num_quant_moe_sort_switch and not force_sep_qsort_257:",
2,
)
source = source.replace("a2 = torch.empty(", "a2 = _oof_get_cached_empty(", 2)
_fused_moe_mod.__dict__["_oof_get_cached_empty"] = _get_cached_empty
namespace: dict[str, object] = {}
exec(source, _fused_moe_mod.__dict__, namespace)
_fused_moe_mod.fused_moe_2stages = namespace["fused_moe_2stages"]
_fused_moe_mod._oof_qsort_threshold = _QSORT_THRESHOLD
def _get_cached_empty(*args, **kwargs):
shape = args[0] if args else kwargs.get("size")
try:
shape = tuple(shape)
except TypeError:
return torch.empty(*args, **kwargs)
if len(shape) != 3 or shape[0] != 512:
return torch.empty(*args, **kwargs)
dtype = kwargs.get("dtype")
device = kwargs.get("device")
key = (str(device), shape, str(dtype))
cached = _A2_CACHE.get(key)
if cached is None:
cached = torch.empty(*args, **kwargs)
_A2_CACHE[key] = cached
return cached
def _get_cached_sort_buffers(
topk_ids: torch.Tensor,
num_experts: int,
model_dim: int,
moebuf_dtype: torch.dtype,
block_size: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] | None:
m, topk = topk_ids.shape
if m != 512:
return None
device = topk_ids.device
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 = (
str(device),
m,
topk,
int(num_experts),
int(model_dim),
str(moebuf_dtype),
int(block_size),
)
cached = _SORT_CACHE.get(key)
if cached is None:
cached = (
torch.empty(max_num_tokens_padded, dtype=aiter.dtypes.i32, device=device),
torch.empty(max_num_tokens_padded, dtype=aiter.dtypes.fp32, device=device),
torch.empty(max_num_m_blocks, dtype=aiter.dtypes.i32, device=device),
torch.empty(2, dtype=aiter.dtypes.i32, device=device),
torch.empty((m, model_dim), dtype=moebuf_dtype, device=device),
)
_SORT_CACHE[key] = cached
return cached
def _patch_sort_cache() -> None:
if getattr(_fused_moe_mod, "_oof_bs512_sort_cache", False):
return
orig_moe_sorting_impl = _fused_moe_mod._moe_sorting_impl
def wrapped_moe_sorting_impl(
topk_ids,
topk_weights,
num_experts,
model_dim,
moebuf_dtype,
block_size,
expert_mask,
num_local_tokens,
dispatch_policy,
use_opus,
):
cached = _get_cached_sort_buffers(
topk_ids, num_experts, model_dim, moebuf_dtype, block_size
)
if cached is None:
return orig_moe_sorting_impl(
topk_ids,
topk_weights,
num_experts,
model_dim,
moebuf_dtype,
block_size,
expert_mask,
num_local_tokens,
dispatch_policy,
use_opus,
)
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = cached
fwd_fn = aiter.moe_sorting_opus_fwd if use_opus else aiter.moe_sorting_fwd
fwd_fn(
topk_ids,
topk_weights,
sorted_ids,
sorted_weights,
sorted_expert_ids,
num_valid_ids,
moe_buf,
num_experts,
int(block_size),
expert_mask,
num_local_tokens,
dispatch_policy,
)
return sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf
_fused_moe_mod._moe_sorting_impl = wrapped_moe_sorting_impl
_fused_moe_mod._oof_bs512_sort_cache = True
def _patch_257_bs512_opus_sorting() -> None:
if getattr(_fused_moe_mod, "_oof_sorting_patch_257_bs512_opus", False):
return
def wrapped_moe_sorting(
topk_ids,
topk_weights,
num_experts,
model_dim,
moebuf_dtype,
block_size=_fused_moe_mod.BLOCK_SIZE_M,
expert_mask=None,
num_local_tokens=None,
dispatch_policy=0,
):
use_opus = num_experts == 257 and topk_ids.shape[0] == 512
return _fused_moe_mod._moe_sorting_impl(
topk_ids,
topk_weights,
num_experts,
model_dim,
moebuf_dtype,
block_size,
expert_mask,
num_local_tokens,
dispatch_policy,
use_opus=use_opus,
)
_fused_moe_mod.moe_sorting = wrapped_moe_sorting
_fused_moe_mod._oof_sorting_patch_257_bs512_opus = True
def _patch_33_bs512_opus_sorting() -> None:
if getattr(_fused_moe_mod, "_oof_sorting_patch_33_bs512_opus", False):
return
def wrapped_moe_sorting(
topk_ids,
topk_weights,
num_experts,
model_dim,
moebuf_dtype,
block_size=_fused_moe_mod.BLOCK_SIZE_M,
expert_mask=None,
num_local_tokens=None,
dispatch_policy=0,
):
use_opus = num_experts in (33, 257) and topk_ids.shape[0] == 512
return _fused_moe_mod._moe_sorting_impl(
topk_ids,
topk_weights,
num_experts,
model_dim,
moebuf_dtype,
block_size,
expert_mask,
num_local_tokens,
dispatch_policy,
use_opus=use_opus,
)
_fused_moe_mod.moe_sorting = wrapped_moe_sorting
_fused_moe_mod._oof_sorting_patch_33_bs512_opus = True
def _patch_cktile_stage1_out() -> None:
if getattr(_fused_moe_mod, "_oof_cktile_stage1_use_out", False):
return
def wrapped_cktile_moe_stage1(
hidden_states,
w1,
w2,
sorted_token_ids,
sorted_expert_ids,
num_valid_ids,
out,
topk,
block_m,
a1_scale,
w1_scale,
sorted_weights=None,
n_pad_zeros=0,
k_pad_zeros=0,
bias1=None,
activation=ActivationType.Silu,
split_k=1,
dtype=torch.bfloat16,
):
token_num = hidden_states.shape[0]
_, _, k1 = w1.shape
_, k2, n2 = w2.shape
d = n2 if k2 == k1 else n2 * 2
if w1.dtype is torch.uint32:
d *= 8
if out.shape != (token_num, topk, d) or out.dtype != dtype or out.device != hidden_states.device:
out = torch.empty((token_num, topk, d), dtype=dtype, device=hidden_states.device)
tmp_out = (
torch.zeros(
(token_num, topk, w1.shape[1]),
dtype=hidden_states.dtype,
device=out.device,
)
if split_k > 1
else out
)
aiter.moe_cktile2stages_gemm1(
hidden_states,
w1,
tmp_out,
sorted_token_ids,
sorted_expert_ids,
num_valid_ids,
topk,
n_pad_zeros,
k_pad_zeros,
sorted_weights,
a1_scale,
w1_scale,
bias1,
activation,
block_m,
split_k,
)
if split_k > 1:
if activation == ActivationType.Silu:
aiter.silu_and_mul(out, tmp_out)
else:
aiter.gelu_and_mul(out, tmp_out)
return out
_fused_moe_mod.cktile_moe_stage1 = wrapped_cktile_moe_stage1
_fused_moe_mod._oof_cktile_stage1_use_out = True
_patch_metadata()
_patch_qsort_threshold()
_patch_sort_cache()
_patch_257_bs512_opus_sorting()
_patch_33_bs512_opus_sorting()
_patch_cktile_stage1_out()
def _pick_block_size(config: dict) -> int | None:
total_experts = config["n_routed_experts"] + config["n_shared_experts"]
if total_experts == 33 and config["d_expert"] == 2048 and config["bs"] >= 512:
return 64
return None
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
hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
intermediate_pad = config["d_expert_pad"] - config["d_expert"]
use_shape7_exact_ck = (
config["n_routed_experts"] + config["n_shared_experts"] == 33
and config["d_expert"] == 2048
and config["bs"] == 512
)
output = _fused_moe_mod.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=None if use_shape7_exact_ck else _pick_block_size(config),
hidden_pad=hidden_pad,
intermediate_pad=intermediate_pad,
)
return output
scrolls · 620 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