submission 754408
Jingze · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1451 lines, June 9 Researcher Reciprocity License v1.0.
submission_ck.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-754408?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:3df005cb95fe421d50ec694cc60ab6412fb63ce6ef4c3e307494c712c5185b63
license declaredunknown
license concludedunknown
authorsJingze
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
submission_ck.py1451 lines
from __future__ import annotations
import functools
import os
from typing import Any, Dict, Optional, Tuple
import aiter
import torch
from aiter import ActivationType, QuantType, dtypes, logger
from aiter import get_hip_quant as get_quant
from aiter.fused_moe import (
fp4_utils,
fused_dynamic_mxfp4_quant_moe_sort,
)
from aiter.ops.triton.moe.moe_op_mxfp4 import fused_moe_mxfp4
from aiter.ops.triton.moe.moe_op_mxfp4_silu_fused import fused_moe_mxfp4_silu
from aiter.ops.triton.utils.moe_config_utils import get_optimal_moe_config_func
from aiter.ops.triton.utils.types import torch_to_triton_dtype
from task import input_t, output_t
TOKEN_NUM_QUANT_MOE_SORT_SWITCH = 1024
# Keep this True to enable offline fixed-shape OPUS policy below.
_USE_OPUS_MOE_SORTING = True
_USE_TENSOR_CACHE = True
# Canonical FP4 shuffled CK kernel names used by this submission.
# These names must track block_m to avoid reusing a 32-tuned pair when block_m changes.
_FP4_CK_STAGE1_BY_BLOCK = {
32: "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
64: "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
128: "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
}
_FP4_CK_STAGE2_BY_BLOCK_INTER_SMALL = {
32: "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
64: "moe_ck2stages_gemm2_64x64x128x128_1x1_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
128: "moe_ck2stages_gemm2_64x128x128x128_1x1_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
}
_FP4_CK_STAGE2_BY_BLOCK_INTER_LARGE = {
32: "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
64: "moe_ck2stages_gemm2_256x64x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
128: "moe_ck2stages_gemm2_256x128x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
}
# Offline precomputed shape allowlist for OPUS moe sorting.
# Only these exact shapes will enable OPUS by default.
OPUS_SORTING_ENABLED_SHAPES = {
# benchmark: TP=8 (E=257, d_expert=256)
"M16_E257_H7168_I256_topk9",
"M128_E257_H7168_I256_topk9",
"M512_E257_H7168_I256_topk9",
# benchmark: TP=4 (E=33, d_expert=512)
"M16_E33_H7168_I512_topk9",
"M128_E33_H7168_I512_topk9",
"M512_E33_H7168_I512_topk9",
# benchmark: EP-on (E=33, d_expert=2048)
"M512_E33_H7168_I2048_topk9",
}
_BUFFER_CACHE: Dict[Tuple[Any, ...], torch.Tensor] = {}
_STATIC_TENSOR_CACHE: Dict[Tuple[Any, ...], Any] = {}
_LAST_BUFFER_CACHE_SHAPE_KEY: Optional[Tuple[Any, ...]] = None
_DEBUG_OUTPUT_COORDS = ((0, 0), (0, 2), (0, 3), (0, 4), (0, 5))
_DEBUG_PRINTED_CUSTOM_KERNEL_CALLS = 0
def _should_print_intermediates(config: Dict[str, Any]) -> bool:
override = config.get("print_intermediates")
if override is not None:
return bool(override)
env_value = os.getenv("SUBMISSION_CK_PRINT_INTERMEDIATES")
if env_value is not None:
return env_value not in ("0", "false", "False", "")
return True
def _get_print_intermediates_max_calls(config: Dict[str, Any]) -> int:
override = config.get("print_intermediates_max_calls")
if override is not None:
return int(override)
env_value = os.getenv("SUBMISSION_CK_PRINT_INTERMEDIATES_MAX_CALLS")
if env_value is not None:
return int(env_value)
return 1
def _should_print_this_custom_kernel_call(config: Dict[str, Any]) -> bool:
global _DEBUG_PRINTED_CUSTOM_KERNEL_CALLS
if not _should_print_intermediates(config):
return False
max_calls = _get_print_intermediates_max_calls(config)
if max_calls == 0:
return False
if max_calls > 0 and _DEBUG_PRINTED_CUSTOM_KERNEL_CALLS >= max_calls:
return False
_DEBUG_PRINTED_CUSTOM_KERNEL_CALLS += 1
return True
def _debug_view_tensor_for_summary(tensor: torch.Tensor) -> Tuple[torch.Tensor, str]:
dtype_label = str(tensor.dtype)
if tensor.dtype == dtypes.fp4x2:
return tensor.view(torch.uint8), f"{dtype_label}(storage_as_uint8)"
if tensor.dtype == dtypes.fp8_e8m0:
return tensor.view(torch.uint8), f"{dtype_label}(storage_as_uint8)"
return tensor, dtype_label
def _tensor_debug_summary(
tensor: Optional[torch.Tensor],
*,
max_items: int = 8,
coords_2d: Tuple[Tuple[int, int], ...] = _DEBUG_OUTPUT_COORDS,
) -> str:
if tensor is None:
return "None"
view_tensor, dtype = _debug_view_tensor_for_summary(tensor.detach())
shape = tuple(view_tensor.shape)
if view_tensor.numel() == 0:
return f"shape={shape} dtype={dtype} empty"
if view_tensor.ndim == 0:
if view_tensor.is_cuda:
view_tensor = view_tensor.cpu()
return f"shape={shape} dtype={dtype} value={view_tensor.item()}"
if view_tensor.is_cuda:
view_tensor = view_tensor.cpu()
pieces = [f"shape={shape}", f"dtype={dtype}"]
if torch.is_floating_point(view_tensor):
flat_float = view_tensor.float().reshape(-1)
pieces.append(f"min={flat_float.min().item():.6f}")
pieces.append(f"max={flat_float.max().item():.6f}")
pieces.append(f"mean={flat_float.mean().item():.6f}")
if view_tensor.ndim == 1:
pieces.append(f"head={view_tensor[:max_items].tolist()}")
elif view_tensor.ndim == 2:
coord_values = []
for row, col in coords_2d:
if row < view_tensor.shape[0] and col < view_tensor.shape[1]:
coord_values.append(f"({row},{col})={view_tensor[row, col].item()}")
if coord_values:
pieces.append("coords=[" + ", ".join(coord_values) + "]")
pieces.append(f"row0_head={view_tensor[0, :max_items].tolist()}")
else:
flattened = view_tensor.reshape(-1, view_tensor.shape[-1])
pieces.append(f"flat0_head={flattened[0, :max_items].tolist()}")
return " ".join(pieces)
def _debug_print_tensor(name: str, tensor: Optional[torch.Tensor]) -> None:
print(f"[submission_ck] {name}: {_tensor_debug_summary(tensor)}", flush=True)
def _debug_print_message(message: str) -> None:
print(f"[submission_ck] {message}", flush=True)
def _evict_cache_on_shape_change(shape_key: Tuple[Any, ...]) -> None:
"""Keep tensor cache only for the most recent runtime shape."""
global _LAST_BUFFER_CACHE_SHAPE_KEY
if _LAST_BUFFER_CACHE_SHAPE_KEY != shape_key:
_BUFFER_CACHE.clear()
_STATIC_TENSOR_CACHE.clear()
_LAST_BUFFER_CACHE_SHAPE_KEY = shape_key
QUANT_SORT_SWITCH_TOKENS_DEFAULT = {
"stage1": TOKEN_NUM_QUANT_MOE_SORT_SWITCH,
"stage2": TOKEN_NUM_QUANT_MOE_SORT_SWITCH,
}
# Offline fixed thresholds for fused quant+sort crossover.
# Keys follow _shape_key(token_num, e, model_dim, inter_dim, topk).
QUANT_SORT_SWITCH_TOKENS_BY_SHAPE: Dict[str, Dict[str, int]] = {
# benchmark: TP=8 (E=257, d_expert=256, topk=9)
"M16_E257_H7168_I256_topk9": {"stage1": 1536, "stage2": 1280},
"M128_E257_H7168_I256_topk9": {"stage1": 1536, "stage2": 1280},
"M512_E257_H7168_I256_topk9": {"stage1": 1536, "stage2": 1280},
# benchmark: TP=4 (E=33, d_expert=512, topk=9)
"M16_E33_H7168_I512_topk9": {"stage1": 1344, "stage2": 1088},
"M128_E33_H7168_I512_topk9": {"stage1": 1344, "stage2": 1088},
"M512_E33_H7168_I512_topk9": {"stage1": 1344, "stage2": 1088},
# benchmark: EP-on (E=33, d_expert=2048, topk=9)
"M512_E33_H7168_I2048_topk9": {"stage1": 1216, "stage2": 960},
}
# Offline fixed policy for where to apply routed expert weights.
# True: apply weights in stage1; False: apply weights in stage2.
DOWEIGHT_STAGE1_BY_SHAPE: Dict[str, bool] = {
# benchmark: TP=8 (E=257, d_expert=256, topk=9)
"M16_E257_H7168_I256_topk9": False,
"M128_E257_H7168_I256_topk9": False,
"M512_E257_H7168_I256_topk9": False,
# benchmark: TP=4 (E=33, d_expert=512, topk=9)
"M16_E33_H7168_I512_topk9": False,
"M128_E33_H7168_I512_topk9": False,
"M512_E33_H7168_I512_topk9": False,
# benchmark: EP-on (E=33, d_expert=2048, topk=9)
"M512_E33_H7168_I2048_topk9": False,
}
def _quant_sort_switch_tokens(
runtime_cfg: Dict[str, Any],
*,
stage: str,
token_num: int,
topk: int,
model_dim: int,
inter_dim: int,
num_experts: int,
) -> int:
"""Return token threshold for choosing fused quant+sort path."""
if stage not in ("stage1", "stage2"):
raise ValueError(f"invalid stage: {stage}")
# Highest priority: explicit per-stage override from runtime config.
stage_override_key = f"quant_sort_switch_tokens_{stage}"
stage_override = runtime_cfg.get(stage_override_key)
if isinstance(stage_override, int) and stage_override > 0:
return stage_override
# Second priority: generic override for both stages.
generic_override = runtime_cfg.get("quant_sort_switch_tokens")
if isinstance(generic_override, int) and generic_override > 0:
return generic_override
# Offline fixed lookup only: no online threshold computation.
shape_key = _shape_key(token_num, num_experts, model_dim, inter_dim, topk)
shape_cfg = QUANT_SORT_SWITCH_TOKENS_BY_SHAPE.get(shape_key)
if isinstance(shape_cfg, dict):
stage_value = shape_cfg.get(stage)
if isinstance(stage_value, int) and stage_value > 0:
return stage_value
return QUANT_SORT_SWITCH_TOKENS_DEFAULT[stage]
def _should_use_fused_quant_sort(
runtime_cfg: Dict[str, Any],
*,
stage: str,
token_num: int,
topk: int,
model_dim: int,
inter_dim: int,
num_experts: int,
) -> bool:
switch_tokens = _quant_sort_switch_tokens(
runtime_cfg,
stage=stage,
token_num=token_num,
topk=topk,
model_dim=model_dim,
inter_dim=inter_dim,
num_experts=num_experts,
)
return token_num <= switch_tokens
def _resolve_doweight_stage1(
runtime_cfg: Dict[str, Any],
ck_cfg: Dict[str, Any],
*,
token_num: int,
num_experts: int,
model_dim: int,
inter_dim: int,
topk: int,
) -> bool:
"""Resolve whether routed weights are multiplied in stage1 or stage2.
Priority:
1) runtime_cfg["doweight_stage1"] bool
2) offline shape table (default behavior)
3) ck_cfg fallback value
"""
runtime_override = runtime_cfg.get("doweight_stage1")
if isinstance(runtime_override, bool):
return runtime_override
shape_key = _shape_key(token_num, num_experts, model_dim, inter_dim, topk)
shape_value = DOWEIGHT_STAGE1_BY_SHAPE.get(shape_key)
if isinstance(shape_value, bool):
return shape_value
return bool(ck_cfg.get("doweight_stage1", False))
def _should_use_opus_moe_sorting(
runtime_cfg: Dict[str, Any],
token_num: int,
num_experts: int,
model_dim: int,
inter_dim: int,
topk: int,
topk_ids: torch.Tensor,
topk_weights: torch.Tensor,
) -> bool:
cfg_override = runtime_cfg.get("use_opus_moe_sorting")
if cfg_override is not None:
return bool(cfg_override)
if not _USE_OPUS_MOE_SORTING:
return False
if not hasattr(aiter, "moe_sorting_opus_fwd"):
return False
# Current OPUS binding accepts int32 topk ids and fp32 topk weights.
if topk_ids.dtype != dtypes.i32:
return False
if topk_weights.dtype != dtypes.fp32:
return False
# OPUS MP path does not support the mesh_byte_size==2 branch (topk >= 255 when tokens >= 512).
if topk >= 255:
return False
shape_key = _shape_key(token_num, num_experts, model_dim, inter_dim, topk)
cfg_shape_allowlist = runtime_cfg.get("opus_enabled_shapes")
if isinstance(cfg_shape_allowlist, (list, tuple, set)):
return shape_key in set(cfg_shape_allowlist)
return shape_key in OPUS_SORTING_ENABLED_SHAPES
def _make_cache_key(
name: str,
device: torch.device,
dtype: torch.dtype,
shape: Tuple[int, ...],
) -> Tuple[Any, ...]:
return (name, device.type, device.index, str(dtype), *shape)
def _get_cached_tensor(
cache_key: Tuple[Any, ...],
shape: Tuple[int, ...],
dtype: torch.dtype,
device: torch.device,
*,
enabled: bool,
) -> torch.Tensor:
if not enabled:
return torch.empty(shape, dtype=dtype, device=device)
cached = _BUFFER_CACHE.get(cache_key)
if cached is None:
cached = torch.empty(shape, dtype=dtype, device=device)
_BUFFER_CACHE[cache_key] = cached
return cached
def _get_cached_value(cache_key: Tuple[Any, ...], factory):
cached = _STATIC_TENSOR_CACHE.get(cache_key)
if cached is None:
cached = factory()
_STATIC_TENSOR_CACHE[cache_key] = cached
return cached
def _slice_expert_tensor(
tensor: Optional[torch.Tensor],
start_expert: int,
expert_count: int,
) -> Optional[torch.Tensor]:
if tensor is None:
return None
sliced = tensor[start_expert : start_expert + expert_count].contiguous()
if getattr(tensor, "is_shuffled", False):
sliced.is_shuffled = True
return sliced
def _slice_raw_scale_by_expert(
scale: Optional[torch.Tensor],
start_expert: int,
expert_count: int,
total_experts: int,
) -> Optional[torch.Tensor]:
if scale is None:
return None
if scale.shape[0] == total_experts:
return scale[start_expert : start_expert + expert_count].contiguous()
if scale.ndim == 2 and scale.shape[0] % total_experts == 0:
rows_per_expert = scale.shape[0] // total_experts
row_start = start_expert * rows_per_expert
row_end = (start_expert + expert_count) * rows_per_expert
return scale[row_start:row_end].contiguous()
raise ValueError(
f"cannot slice scale by expert: shape={tuple(scale.shape)}, total_experts={total_experts}"
)
def _make_branch_shuffled_scale(
raw_scale: Optional[torch.Tensor],
start_expert: int,
expert_count: int,
total_experts: int,
) -> Optional[torch.Tensor]:
branch_scale = _slice_raw_scale_by_expert(
raw_scale,
start_expert,
expert_count,
total_experts,
)
if branch_scale is None:
return None
if branch_scale.ndim > 2:
branch_scale = branch_scale.reshape(-1, branch_scale.shape[-1]).contiguous()
return fp4_utils.e8m0_shuffle(branch_scale)
def _get_cached_shared_branch_tensors(
*,
w1: torch.Tensor,
w2: torch.Tensor,
gate_up_weight_scale: Optional[torch.Tensor],
down_weight_scale: Optional[torch.Tensor],
n_routed_experts: int,
n_shared_experts: int,
total_experts: int,
) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor], Optional[torch.Tensor]]:
cache_key = (
"shared_branch_tensors",
w1.data_ptr(),
w2.data_ptr(),
None if gate_up_weight_scale is None else gate_up_weight_scale.data_ptr(),
None if down_weight_scale is None else down_weight_scale.data_ptr(),
n_routed_experts,
n_shared_experts,
total_experts,
)
def _factory():
shared_w1 = _slice_expert_tensor(w1, n_routed_experts, n_shared_experts)
shared_w2 = _slice_expert_tensor(w2, n_routed_experts, n_shared_experts)
shared_w1_scale = _slice_raw_scale_by_expert(
gate_up_weight_scale, n_routed_experts, n_shared_experts, total_experts
)
shared_w2_scale = _slice_raw_scale_by_expert(
down_weight_scale, n_routed_experts, n_shared_experts, total_experts
)
return shared_w1, shared_w2, shared_w1_scale, shared_w2_scale
return _get_cached_value(cache_key, _factory)
def _get_cached_shared_sorting(
*,
token_num: int,
n_shared_experts: int,
block_m: int,
device: torch.device,
use_opus: bool,
dispatch_policy: int,
) -> Tuple[
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
]:
cache_key = (
"shared_branch_sorting",
device.type,
device.index,
token_num,
n_shared_experts,
block_m,
use_opus,
dispatch_policy,
)
def _factory():
shared_ids = torch.arange(n_shared_experts, dtype=dtypes.i32, device=device).view(
1, n_shared_experts
).expand(token_num, -1).contiguous()
shared_weights = torch.ones(
(token_num, n_shared_experts), dtype=dtypes.fp32, device=device
)
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, _ = moe_sorting(
shared_ids,
shared_weights,
n_shared_experts,
0,
dtypes.fp32,
block_m,
expert_mask=None,
num_local_tokens=None,
dispatch_policy=dispatch_policy,
use_tensor_cache=False,
use_opus=use_opus,
)
num_tokens_post_padded = num_valid_ids[:1].contiguous()
return (
shared_ids,
shared_weights,
sorted_ids,
sorted_weights,
sorted_expert_ids,
num_valid_ids,
num_tokens_post_padded,
)
return _get_cached_value(cache_key, _factory)
def _as_triton_mx_scale(scale: torch.Tensor) -> torch.Tensor:
return scale.view(torch.uint8)
def _as_triton_mx_data(data: torch.Tensor) -> torch.Tensor:
return data.view(torch.uint8)
def _canonicalize_triton_weight_scale(
scale: torch.Tensor,
weight: torch.Tensor,
) -> torch.Tensor:
if scale.ndim == 3:
return _as_triton_mx_scale(scale)
if scale.ndim != 2:
raise ValueError(
f"unsupported shared Triton scale rank: scale_shape={tuple(scale.shape)}, weight_shape={tuple(weight.shape)}"
)
expert_count, out_dim, packed_k = weight.shape
scale_k = packed_k * 2 // 32
expected_rows = expert_count * out_dim
if scale.shape[0] != expected_rows:
raise ValueError(
f"cannot reshape shared Triton scale: scale_shape={tuple(scale.shape)}, weight_shape={tuple(weight.shape)}"
)
if scale.shape[1] != scale_k:
raise ValueError(
f"unexpected shared Triton scale K dimension: scale_shape={tuple(scale.shape)}, expected_scale_k={scale_k}"
)
return _as_triton_mx_scale(scale.reshape(expert_count, out_dim, scale_k))
def _run_shared_triton_dense(
*,
hidden_states: torch.Tensor,
shared_w1: torch.Tensor,
shared_w2: torch.Tensor,
shared_w1_scale: torch.Tensor,
shared_w2_scale: torch.Tensor,
shared_ids: torch.Tensor,
shared_weights: torch.Tensor,
shared_sorted_ids: torch.Tensor,
shared_sorted_expert_ids: torch.Tensor,
shared_num_tokens_post_padded: torch.Tensor,
token_num: int,
n_shared_experts: int,
inter_dim: int,
model_dim: int,
use_tensor_cache: bool,
) -> torch.Tensor:
moe_config = get_optimal_moe_config_func(hidden_states.dtype, use_mxfp4=True)(token_num)
compute_type = torch_to_triton_dtype[hidden_states.dtype]
a_scale = torch.ones((1,), dtype=torch.float32, device=hidden_states.device)
b1_scale = torch.ones((n_shared_experts,), dtype=torch.float32, device=hidden_states.device)
b2_scale = b1_scale
shared_w1_u8 = _as_triton_mx_data(shared_w1)
shared_w2_u8 = _as_triton_mx_data(shared_w2)
b1_mx_scale = _canonicalize_triton_weight_scale(shared_w1_scale, shared_w1)
b2_mx_scale = _canonicalize_triton_weight_scale(shared_w2_scale, shared_w2)
a1, a1_scale = fp4_utils.dynamic_mxfp4_quant(hidden_states)
a1 = _as_triton_mx_data(a1)
a1_scale = _as_triton_mx_scale(a1_scale)
stage1_out = _get_cached_tensor(
_make_cache_key(
"shared_triton_stage1_out",
hidden_states.device,
hidden_states.dtype,
(token_num * n_shared_experts, inter_dim),
),
(token_num * n_shared_experts, inter_dim),
hidden_states.dtype,
hidden_states.device,
enabled=use_tensor_cache,
)
fused_moe_mxfp4_silu(
a1,
shared_w1_u8,
stage1_out,
a_scale,
b1_scale,
a1_scale,
b1_mx_scale,
shared_weights,
shared_ids,
shared_sorted_ids,
shared_sorted_expert_ids,
shared_num_tokens_post_padded,
False,
n_shared_experts,
False,
False,
moe_config,
compute_type,
)
a2, a2_scale = fp4_utils.dynamic_mxfp4_quant(stage1_out)
a2 = _as_triton_mx_data(a2)
a2_scale = _as_triton_mx_scale(a2_scale)
stage2_out = _get_cached_tensor(
_make_cache_key(
"shared_triton_stage2_out",
hidden_states.device,
hidden_states.dtype,
(token_num, n_shared_experts, model_dim),
),
(token_num, n_shared_experts, model_dim),
hidden_states.dtype,
hidden_states.device,
enabled=use_tensor_cache,
)
fused_moe_mxfp4(
a2,
shared_w2_u8,
stage2_out,
a_scale,
b2_scale,
a2_scale,
b2_mx_scale,
shared_weights,
shared_ids,
shared_sorted_ids,
shared_sorted_expert_ids,
shared_num_tokens_post_padded,
False,
n_shared_experts,
False,
False,
moe_config,
compute_type,
)
if n_shared_experts == 1:
return stage2_out.view(token_num, model_dim)
return stage2_out.sum(dim=1)
# benchmarks:
# # TP=8
# - {"dhidden": 7168, "dexpert": 256, "nroutedexperts": 256, "nexpertspertoken": 8, "nsharedexperts": 1, "bs": 16, "seed": 9371}
# - {"dhidden": 7168, "dexpert": 256, "nroutedexperts": 256, "nexpertspertoken": 8, "nsharedexperts": 1, "bs": 128, "seed": 2291}
# - {"dhidden": 7168, "dexpert": 256, "nroutedexperts": 256, "nexpertspertoken": 8, "nsharedexperts": 1, "bs": 512, "seed": 81934}
# # TP=4
# - {"dhidden": 7168, "dexpert": 512, "nroutedexperts": 32, "nexpertspertoken": 8, "nsharedexperts": 1, "bs": 16, "seed": 2291}
# - {"dhidden": 7168, "dexpert": 512, "nroutedexperts": 32, "nexpertspertoken": 8, "nsharedexperts": 1, "bs": 128, "seed": 81934}
# - {"dhidden": 7168, "dexpert": 512, "nroutedexperts": 32, "nexpertspertoken": 8, "nsharedexperts": 1, "bs": 512, "seed": 81934}
# # EP on
# - {"dhidden": 7168, "dexpert": 2048, "nroutedexperts": 32, "nexpertspertoken": 8, "nsharedexperts": 1, "bs": 512, "seed": 81934}
# Manual kernel config table (CK).
#
# Runtime override priority:
# 1) config["kernel_manual_cfg"] or config["ck_manual_cfg"]
# 2) config["kernel_profile"] or config["ck_profile"]
# 3) shape profile in CK_MANUAL_CONFIGS
# 4) CK_MANUAL_CONFIGS["default"]
#
# Supported keys:
# - block_m: int
# - kernel_name1: str
# - kernel_name2: str
# - use_non_temporal_load: bool
# - doweight_stage1: bool
# - use_shuffled: bool
CK_MANUAL_CONFIGS: Dict[str, Dict[str, Any]] = {
"default": {
"block_m": 32,
"use_non_temporal_load": False,
"doweight_stage1": False,
"use_shuffled": True,
},
# TP=8, dexpert=256, E=257, topk=9
"M16_E257_H7168_I256_topk9": {
"block_m": 32,
},
"M128_E257_H7168_I256_topk9": {
"block_m": 32,
},
"M512_E257_H7168_I256_topk9": {
"block_m": 32,
},
# TP=4, dexpert=512, E=33, topk=9
# No exact tuned rows for E=33 in current CSVs; keep CK heuristic dispatch.
"M16_E33_H7168_I512_topk9": {
"block_m": 32,
},
"M128_E33_H7168_I512_topk9": {
"block_m": 32,
},
"M512_E33_H7168_I512_topk9": {
"block_m": 128,
},
# EP-on, dexpert=2048, E=33, topk=9
# No exact row found in current CSV; keep CK heuristic fallback.
"M512_E33_H7168_I2048_topk9": {
"block_m": 128,
},
}
def _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_tensor_cache,
):
device = topk_ids.device
M, topk = topk_ids.shape
max_num_tokens_padded = topk_ids.numel() + num_experts * block_size - topk
max_num_m_blocks = (max_num_tokens_padded + block_size - 1) // block_size
sorted_ids = _get_cached_tensor(
_make_cache_key(
"moe_sorting_sorted_ids",
device,
dtypes.i32,
(max_num_tokens_padded,),
),
(max_num_tokens_padded,),
dtypes.i32,
device,
enabled=use_tensor_cache,
)
sorted_weights = _get_cached_tensor(
_make_cache_key(
"moe_sorting_sorted_weights",
device,
dtypes.fp32,
(max_num_tokens_padded,),
),
(max_num_tokens_padded,),
dtypes.fp32,
device,
enabled=use_tensor_cache,
)
sorted_expert_ids = _get_cached_tensor(
_make_cache_key(
"moe_sorting_sorted_expert_ids",
device,
dtypes.i32,
(max_num_m_blocks,),
),
(max_num_m_blocks,),
dtypes.i32,
device,
enabled=use_tensor_cache,
)
num_valid_ids = _get_cached_tensor(
_make_cache_key("moe_sorting_num_valid_ids", device, dtypes.i32, (2,)),
(2,),
dtypes.i32,
device,
enabled=use_tensor_cache,
)
moe_buf = _get_cached_tensor(
_make_cache_key(
"moe_final_out",
device,
moebuf_dtype,
(M, model_dim),
),
(M, model_dim),
moebuf_dtype,
device,
enabled=use_tensor_cache,
)
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,
block_size,
expert_mask,
num_local_tokens,
dispatch_policy,
)
return sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf
def moe_sorting(
topk_ids,
topk_weights,
num_experts,
model_dim,
moebuf_dtype,
block_size,
expert_mask=None,
num_local_tokens=None,
dispatch_policy=0,
use_tensor_cache=True,
use_opus: Optional[bool] = None,
):
return _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_MOE_SORTING if use_opus is None else bool(use_opus)),
use_tensor_cache=use_tensor_cache,
)
# Lru cache will using hash to create key, which makes error when w1,w2 shape is symint.
# We can use torch.compile(dynamic=False) to avoid
@functools.lru_cache(maxsize=2048)
def get_inter_dim(w1_shape, w2_shape):
E, _, model_dim = w1_shape
E, model_dim, inter_dim = w2_shape
int4_war = model_dim // w1_shape[-1]
inter_dim *= int4_war
return E, model_dim, inter_dim
def _shape_key(token_num: int, e: int, model_dim: int, inter_dim: int, topk: int) -> str:
return f"M{token_num}_E{e}_H{model_dim}_I{inter_dim}_topk{topk}"
def _resolve_ck_cfg(
runtime_cfg: Dict[str, Any],
token_num: int,
e: int,
model_dim: int,
inter_dim: int,
topk: int,
) -> Dict[str, Any]:
cfg = CK_MANUAL_CONFIGS["default"].copy()
profile = runtime_cfg.get("kernel_profile", runtime_cfg.get("ck_profile"))
if isinstance(profile, str) and profile in CK_MANUAL_CONFIGS:
cfg.update(CK_MANUAL_CONFIGS[profile])
shape_profile = _shape_key(token_num, e, model_dim, inter_dim, topk)
if shape_profile in CK_MANUAL_CONFIGS:
cfg.update(CK_MANUAL_CONFIGS[shape_profile])
manual_cfg = runtime_cfg.get("kernel_manual_cfg", runtime_cfg.get("ck_manual_cfg"))
if isinstance(manual_cfg, dict):
cfg.update(manual_cfg)
if cfg["block_m"] <= 0:
raise ValueError(f"invalid block_m: {cfg['block_m']}")
return cfg
def _ensure_shuffled_flag(tensor: torch.Tensor) -> None:
# CK dispatch relies on this runtime flag for preshuffle path selection.
if not getattr(tensor, "is_shuffled", False):
tensor.is_shuffled = True
def _select_fp4_kernel_names_for_block_m(
block_m: int,
inter_dim: int,
kernel_name1: str,
kernel_name2: str,
*,
use_shuffled: bool,
w1_dtype: torch.dtype,
w2_dtype: torch.dtype,
allow_stage1_name_patch: bool,
) -> Tuple[str, str]:
"""Patch or fill CK kernel names when block_m changes on FP4 shuffled path.
We only touch names for the preshuffled FP4 path because this submission's
manual table currently hardcodes a 32-tuned pair for multiple shapes.
"""
if not use_shuffled:
return kernel_name1, kernel_name2
if w1_dtype != dtypes.fp4x2 or w2_dtype != dtypes.fp4x2:
return kernel_name1, kernel_name2
if block_m not in _FP4_CK_STAGE1_BY_BLOCK:
return kernel_name1, kernel_name2
expected_stage1 = _FP4_CK_STAGE1_BY_BLOCK[block_m]
stage2_table = (
_FP4_CK_STAGE2_BY_BLOCK_INTER_SMALL
if inter_dim <= 256
else _FP4_CK_STAGE2_BY_BLOCK_INTER_LARGE
)
expected_stage2 = stage2_table[block_m]
legacy_stage1_32 = _FP4_CK_STAGE1_BY_BLOCK[32]
legacy_stage2_32 = _FP4_CK_STAGE2_BY_BLOCK_INTER_SMALL[32]
# Fill empty stage2 name from the per-block defaults.
if not kernel_name2:
kernel_name2 = expected_stage2
# If user only changes block_m and leaves the old 32 stage2 name, rewrite safely.
if block_m != 32:
if kernel_name2 == legacy_stage2_32:
kernel_name2 = expected_stage2
# Stage1 FP4 named kernels do not match blockscale/per_1x128 dispatch.
if allow_stage1_name_patch:
if not kernel_name1:
kernel_name1 = expected_stage1
if block_m != 32 and kernel_name1 == legacy_stage1_32:
kernel_name1 = expected_stage1
return kernel_name1, kernel_name2
def _run_stage1(
a1: torch.Tensor,
w1: torch.Tensor,
w2: torch.Tensor,
sorted_ids: torch.Tensor,
sorted_expert_ids: torch.Tensor,
num_valid_ids: torch.Tensor,
out: Optional[torch.Tensor],
topk: int,
kernel_name: str,
w1_scale: Optional[torch.Tensor],
a1_scale: Optional[torch.Tensor],
block_m: int,
sorted_weights: Optional[torch.Tensor],
use_non_temporal_load: bool,
dst_type: Optional[torch.dtype],
):
assert out is not None
aiter.ck_moe_stage1_fwd(
a1,
w1,
w2,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
out,
topk,
kernelName=kernel_name,
w1_scale=w1_scale,
a1_scale=a1_scale,
block_m=block_m,
sorted_weights=sorted_weights,
quant_type=QuantType.per_1x32,
activation=ActivationType.Silu,
splitk=0,
use_non_temporal_load=use_non_temporal_load,
dst_type=dst_type,
)
return out
def _run_stage2(
a2: torch.Tensor,
w1: torch.Tensor,
w2: torch.Tensor,
sorted_ids: torch.Tensor,
sorted_expert_ids: torch.Tensor,
num_valid_ids: torch.Tensor,
out: torch.Tensor,
topk: int,
kernel_name: str,
w2_scale: Optional[torch.Tensor],
a2_scale: Optional[torch.Tensor],
block_m: int,
sorted_weights: Optional[torch.Tensor],
use_non_temporal_load: bool,
) -> torch.Tensor:
aiter.ck_moe_stage2_fwd(
a2,
w1,
w2,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
out,
topk,
kernelName=kernel_name,
w2_scale=w2_scale,
a2_scale=a2_scale,
block_m=block_m,
sorted_weights=sorted_weights,
quant_type=QuantType.per_1x32,
activation=ActivationType.Silu,
use_non_temporal_load=use_non_temporal_load,
)
return out
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
token_num, topk = topk_ids.shape
e, model_dim, inter_dim = get_inter_dim(
gate_up_weight_shuffled.shape, down_weight_shuffled.shape
)
ck_cfg = _resolve_ck_cfg(config, token_num, e, model_dim, inter_dim, topk)
use_shuffled = ck_cfg.get("use_shuffled", config.get("use_shuffled", True))
if use_shuffled:
w1 = gate_up_weight_shuffled
w2 = down_weight_shuffled
w1_scale = gate_up_weight_scale_shuffled
w2_scale = down_weight_scale_shuffled
_ensure_shuffled_flag(w1)
_ensure_shuffled_flag(w2)
else:
w1 = gate_up_weight
w2 = down_weight
w1_scale = gate_up_weight_scale
w2_scale = down_weight_scale
e, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)
stage1_quant_type = QuantType.per_1x32
quant_func_stage1 = get_quant(stage1_quant_type)
quant_func_stage2 = get_quant(QuantType.per_1x32)
# Keep cache entries only for one active shape.
_evict_cache_on_shape_change(
(
hidden_states.device,
hidden_states.dtype,
token_num,
topk,
e,
model_dim,
inter_dim,
)
)
block_m = ck_cfg["block_m"]
use_tensor_cache = config.get("use_tensor_cache", _USE_TENSOR_CACHE)
kernel_name1 = ck_cfg.get("kernel_name1", "")
kernel_name2 = ck_cfg.get("kernel_name2", "")
kernel_name1, kernel_name2 = _select_fp4_kernel_names_for_block_m(
block_m,
inter_dim,
kernel_name1,
kernel_name2,
use_shuffled=use_shuffled,
w1_dtype=w1.dtype,
w2_dtype=w2.dtype,
allow_stage1_name_patch=True,
)
use_non_temporal_load = ck_cfg.get("use_non_temporal_load", False)
doweight_stage1 = _resolve_doweight_stage1(
config,
ck_cfg,
token_num=token_num,
num_experts=e,
model_dim=model_dim,
inter_dim=inter_dim,
topk=topk,
)
moe_sorting_dispatch_policy = config.get("moe_sorting_dispatch_policy", 0)
debug_print_intermediates = _should_print_this_custom_kernel_call(config)
stage1_dst_type: Optional[torch.dtype] = None
w1_scale_fp8 = w1_scale.view(dtypes.fp8_e8m0) if w1.dtype == dtypes.fp4x2 else w1_scale
w2_scale_fp8 = w2_scale.view(dtypes.fp8_e8m0) if w2.dtype == dtypes.fp4x2 else w2_scale
def _run_one_group(
ids_group: torch.Tensor,
weights_group: torch.Tensor,
group_topk: int,
*,
group_num_experts: int,
group_w1: torch.Tensor,
group_w2: torch.Tensor,
group_w1_scale_fp8: Optional[torch.Tensor],
group_w2_scale_fp8: Optional[torch.Tensor],
group_kernel_name1: str,
group_name: str,
weights_are_unit: bool = False,
precomputed_sorting: Optional[
Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]
] = None,
output_buffer: Optional[torch.Tensor] = None,
zero_output_buffer: bool = True,
) -> torch.Tensor:
group_kernel_name2 = "" if weights_are_unit else kernel_name2
if debug_print_intermediates:
_debug_print_message(
f"group={group_name} topk={group_topk} num_experts={group_num_experts} weights_are_unit={weights_are_unit}"
)
_debug_print_tensor(f"{group_name}.ids_group", ids_group)
_debug_print_tensor(f"{group_name}.weights_group", weights_group)
if precomputed_sorting is None:
use_opus_sorting = _should_use_opus_moe_sorting(
config,
token_num,
group_num_experts,
model_dim,
inter_dim,
group_topk,
ids_group,
weights_group,
)
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_out = moe_sorting(
ids_group,
weights_group,
group_num_experts,
model_dim,
hidden_states.dtype,
block_m,
expert_mask=None,
num_local_tokens=None,
dispatch_policy=moe_sorting_dispatch_policy,
use_tensor_cache=use_tensor_cache,
use_opus=use_opus_sorting,
)
if output_buffer is not None:
moe_out = output_buffer
if zero_output_buffer:
moe_out.zero_()
else:
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids = precomputed_sorting
if output_buffer is None:
moe_out = _get_cached_tensor(
_make_cache_key(
f"moe_final_out_{group_name}",
hidden_states.device,
hidden_states.dtype,
(token_num, model_dim),
),
(token_num, model_dim),
hidden_states.dtype,
hidden_states.device,
enabled=use_tensor_cache,
)
else:
moe_out = output_buffer
if zero_output_buffer:
moe_out.zero_()
if debug_print_intermediates:
_debug_print_tensor(f"{group_name}.sorted_ids", sorted_ids)
_debug_print_tensor(f"{group_name}.sorted_weights", sorted_weights)
_debug_print_tensor(f"{group_name}.sorted_expert_ids", sorted_expert_ids)
_debug_print_tensor(f"{group_name}.num_valid_ids", num_valid_ids)
use_fused_stage1 = _should_use_fused_quant_sort(
config,
stage="stage1",
token_num=token_num,
topk=group_topk,
model_dim=model_dim,
inter_dim=inter_dim,
num_experts=group_num_experts,
)
if use_fused_stage1:
a1, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
hidden_states,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
topk=1,
block_size=block_m,
)
else:
a1, a1_scale = quant_func_stage1(
hidden_states,
scale=None,
quant_dtype=dtypes.fp4x2,
num_rows=None,
)
a1_scale = fp4_utils.moe_mxfp4_sort(
a1_scale,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
block_size=block_m,
)
if debug_print_intermediates:
_debug_print_message(f"group={group_name} use_fused_stage1={use_fused_stage1}")
_debug_print_tensor(f"{group_name}.a1", a1)
_debug_print_tensor(f"{group_name}.a1_scale", a1_scale)
stage1_out_dtype = hidden_states.dtype
stage1_out_buf = _get_cached_tensor(
_make_cache_key(
f"ck_stage1_out_{group_name}",
hidden_states.device,
stage1_out_dtype,
(token_num, group_topk, inter_dim),
),
(token_num, group_topk, inter_dim),
stage1_out_dtype,
hidden_states.device,
enabled=use_tensor_cache,
)
stage1_result = _run_stage1(
a1,
group_w1,
group_w2,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
stage1_out_buf,
group_topk,
group_kernel_name1,
w1_scale=group_w1_scale_fp8,
a1_scale=a1_scale,
block_m=block_m,
sorted_weights=(sorted_weights if doweight_stage1 and not weights_are_unit else None),
use_non_temporal_load=use_non_temporal_load,
dst_type=stage1_dst_type,
)
if debug_print_intermediates:
_debug_print_tensor(f"{group_name}.stage1_result", stage1_result)
a2_scale: Optional[torch.Tensor] = None
a2 = stage1_result.view(-1, inter_dim)
use_fused_stage2 = _should_use_fused_quant_sort(
config,
stage="stage2",
token_num=token_num,
topk=group_topk,
model_dim=model_dim,
inter_dim=inter_dim,
num_experts=group_num_experts,
)
if use_fused_stage2:
a2, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
a2,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
topk=group_topk,
block_size=block_m,
)
else:
a2, a2_scale = quant_func_stage2(
a2,
scale=a2_scale,
quant_dtype=dtypes.fp4x2,
num_rows=None,
num_rows_factor=group_topk,
)
a2_scale = fp4_utils.moe_mxfp4_sort(
a2_scale[: token_num * group_topk, :].view(token_num, group_topk, -1),
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
block_size=block_m,
)
a2 = a2.view(token_num, group_topk, -1)
if debug_print_intermediates:
_debug_print_message(f"group={group_name} use_fused_stage2={use_fused_stage2}")
_debug_print_tensor(f"{group_name}.a2", a2)
_debug_print_tensor(f"{group_name}.a2_scale", a2_scale)
_run_stage2(
a2,
group_w1,
group_w2,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
moe_out,
group_topk,
group_kernel_name2,
w2_scale=group_w2_scale_fp8,
a2_scale=a2_scale,
block_m=block_m,
sorted_weights=(None if doweight_stage1 or weights_are_unit else sorted_weights),
use_non_temporal_load=use_non_temporal_load,
)
if debug_print_intermediates:
_debug_print_tensor(f"{group_name}.moe_out", moe_out)
return moe_out
n_shared_experts = int(config.get("n_shared_experts", 0))
if n_shared_experts < 0 or n_shared_experts > topk:
raise ValueError(f"invalid n_shared_experts: {n_shared_experts}")
routed_topk = topk - n_shared_experts
n_routed_experts = int(config.get("n_routed_experts", e - n_shared_experts))
if n_routed_experts < 0 or n_routed_experts + n_shared_experts > e:
raise ValueError(
f"invalid expert partition: n_routed_experts={n_routed_experts}, n_shared_experts={n_shared_experts}, total_experts={e}"
)
split_shared_experts = bool(config.get("split_shared_experts", False))
if debug_print_intermediates:
_debug_print_message(
"custom_kernel "
f"token_num={token_num} topk={topk} num_experts={e} model_dim={model_dim} inter_dim={inter_dim} "
f"block_m={block_m} use_shuffled={use_shuffled} doweight_stage1={doweight_stage1} "
f"split_shared_experts={split_shared_experts} use_tensor_cache={use_tensor_cache}"
)
_debug_print_tensor("hidden_states", hidden_states)
_debug_print_tensor("topk_ids", topk_ids)
_debug_print_tensor("topk_weights", topk_weights)
_debug_print_message(
f"kernel_name1={kernel_name1 or '<empty>'} kernel_name2={kernel_name2 or '<empty>'}"
)
if n_shared_experts == 0 or not split_shared_experts:
moe_out = _run_one_group(
topk_ids,
topk_weights,
topk,
group_num_experts=e,
group_w1=w1,
group_w2=w2,
group_w1_scale_fp8=w1_scale_fp8,
group_w2_scale_fp8=w2_scale_fp8,
group_kernel_name1=kernel_name1,
group_name="all",
)
else:
moe_out: Optional[torch.Tensor] = None
if routed_topk > 0:
routed_out = _run_one_group(
topk_ids[:, :routed_topk].contiguous(),
topk_weights[:, :routed_topk].contiguous(),
routed_topk,
group_num_experts=e,
group_w1=w1,
group_w2=w2,
group_w1_scale_fp8=w1_scale_fp8,
group_w2_scale_fp8=w2_scale_fp8,
group_kernel_name1=kernel_name1,
group_name="routed",
)
moe_out = routed_out
shared_w1, shared_w2, shared_w1_scale, shared_w2_scale = _get_cached_shared_branch_tensors(
w1=gate_up_weight,
w2=down_weight,
gate_up_weight_scale=gate_up_weight_scale,
down_weight_scale=down_weight_scale,
n_routed_experts=n_routed_experts,
n_shared_experts=n_shared_experts,
total_experts=e,
)
shared_ids, shared_weights, shared_sorted_ids, shared_sorted_weights, shared_sorted_expert_ids, shared_num_valid_ids, shared_num_tokens_post_padded = _get_cached_shared_sorting(
token_num=token_num,
n_shared_experts=n_shared_experts,
block_m=get_optimal_moe_config_func(hidden_states.dtype, use_mxfp4=True)(token_num)["BLOCK_SIZE_M"],
device=hidden_states.device,
use_opus=False,
dispatch_policy=moe_sorting_dispatch_policy,
)
use_shared_triton = bool(config.get("use_shared_triton", False))
if use_shared_triton:
shared_out = _run_shared_triton_dense(
hidden_states=hidden_states,
shared_w1=shared_w1,
shared_w2=shared_w2,
shared_w1_scale=shared_w1_scale,
shared_w2_scale=shared_w2_scale,
shared_ids=shared_ids,
shared_weights=shared_weights,
shared_sorted_ids=shared_sorted_ids,
shared_sorted_expert_ids=shared_sorted_expert_ids,
shared_num_tokens_post_padded=shared_num_tokens_post_padded,
token_num=token_num,
n_shared_experts=n_shared_experts,
inter_dim=inter_dim,
model_dim=model_dim,
use_tensor_cache=use_tensor_cache,
)
else:
shared_out = _run_one_group(
shared_ids,
shared_weights,
n_shared_experts,
group_num_experts=n_shared_experts,
group_w1=shared_w1,
group_w2=shared_w2,
group_w1_scale_fp8=shared_w1_scale,
group_w2_scale_fp8=shared_w2_scale,
group_kernel_name1=kernel_name1,
group_name="shared",
weights_are_unit=True,
precomputed_sorting=(
shared_sorted_ids,
shared_sorted_weights,
shared_sorted_expert_ids,
shared_num_valid_ids,
),
)
if moe_out is None:
moe_out = shared_out
else:
moe_out.add_(shared_out)
d_hidden = config.get("d_hidden", hidden_states.shape[1])
final_out = moe_out[:, :d_hidden]
if debug_print_intermediates:
_debug_print_tensor("final_out", final_out)
return final_out
scrolls · 1451 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