Revert "[AMD] [GLM5] Fuse shared-expert append into aiter grouped-topk (skip per-layer append kernel)" (#35105)
This commit is contained in:
@@ -144,78 +144,6 @@ _RENORMALIZE_SUM_EPSILON = 1e-20
|
|||||||
# an accuracy run before becoming the default.
|
# an accuracy run before becoming the default.
|
||||||
_skip_hip_pad_mask = get_bool_env_var("SGLANG_MORI_NO_PAD_MASK", "False")
|
_skip_hip_pad_mask = get_bool_env_var("SGLANG_MORI_NO_PAD_MASK", "False")
|
||||||
|
|
||||||
# ATOM-style shared-expert fusion for the aiter grouped-topk path: keep a
|
|
||||||
# persistent topk buffer whose shared-expert columns are pre-populated once (NOT
|
|
||||||
# related to the LLM prefill phase — this applies to both prefill and decode), and
|
|
||||||
# let the aiter kernel write only the routed columns (via row stride). This removes
|
|
||||||
# the per-layer _fused_append_shared_experts kernel.
|
|
||||||
# Auto-enabled (no env) when: non-EP aiter path (moe_ep_size == 1) and
|
|
||||||
# num_fused_shared_experts > 0. Bit-identical to the plain append: the shared column
|
|
||||||
# is filled with the same constant scale_factor the aiter append writes. The
|
|
||||||
# persistent buffer is sized to the max prefill batch (chunked-prefill-size); for
|
|
||||||
# token counts above it we fall back to the plain path (condition mirrored in
|
|
||||||
# _post_process_topk_ids). Mirrors ATOM's init_aiter_topK_meta_data.
|
|
||||||
# Hard upper bound on the persistent buffer (safety cap when chunked prefill is
|
|
||||||
# disabled / unexpectedly huge); the buffer only costs ~[MAX, topk+n_shared] * 8B.
|
|
||||||
_AITER_TOPK_FUSE_SHARED_MAX_TOKENS_CAP = 131072
|
|
||||||
_aiter_topk_fuse_shared_max_tokens_cache = None
|
|
||||||
_aiter_topk_fuse_shared_bufs: dict = {}
|
|
||||||
|
|
||||||
|
|
||||||
def _get_aiter_topk_fuse_shared_max_tokens() -> int:
|
|
||||||
"""Max per-forward token count the persistent buffer must cover. Sized to the
|
|
||||||
largest prefill batch (chunked-prefill-size / max-prefill-tokens); decode is
|
|
||||||
always tiny (bs * num_tokens_per_bs). Above this we fall back to the plain
|
|
||||||
path, so this only bounds the fast-path coverage, not correctness. Cached
|
|
||||||
(server args are fixed after startup)."""
|
|
||||||
global _aiter_topk_fuse_shared_max_tokens_cache
|
|
||||||
if _aiter_topk_fuse_shared_max_tokens_cache is None:
|
|
||||||
from sglang.srt.runtime_context import get_server_args
|
|
||||||
|
|
||||||
try:
|
|
||||||
sa = get_server_args()
|
|
||||||
except ValueError:
|
|
||||||
# Global server args not published yet (e.g. a unit test or offline
|
|
||||||
# init that reaches the aiter grouped-topk path before startup).
|
|
||||||
# Degrade gracefully instead of crashing -- this value only bounds
|
|
||||||
# the fast-path coverage, not correctness (see docstring). Return the
|
|
||||||
# safety cap WITHOUT caching, so a later call (once args are set)
|
|
||||||
# still computes and caches the real value.
|
|
||||||
return _AITER_TOPK_FUSE_SHARED_MAX_TOKENS_CAP
|
|
||||||
cps = getattr(sa, "chunked_prefill_size", None) or 0
|
|
||||||
mpt = getattr(sa, "max_prefill_tokens", None) or 0
|
|
||||||
m = max(int(cps), int(mpt), 8192) # 8192 floor for tiny configs
|
|
||||||
if int(cps) <= 0 and int(mpt) <= 0:
|
|
||||||
# chunked prefill disabled -> use the safety cap
|
|
||||||
m = _AITER_TOPK_FUSE_SHARED_MAX_TOKENS_CAP
|
|
||||||
_aiter_topk_fuse_shared_max_tokens_cache = min(
|
|
||||||
m, _AITER_TOPK_FUSE_SHARED_MAX_TOKENS_CAP
|
|
||||||
)
|
|
||||||
return _aiter_topk_fuse_shared_max_tokens_cache
|
|
||||||
|
|
||||||
|
|
||||||
def _get_aiter_topk_fuse_shared_buf(
|
|
||||||
topk_routed: int, n_shared: int, num_experts: int, shared_weight: float, device
|
|
||||||
):
|
|
||||||
"""Persistent [MAX, topk_routed + n_shared] weight/id buffers whose shared
|
|
||||||
columns are pre-filled once (id = num_experts + i, weight = shared_weight).
|
|
||||||
Fixed max size (>= max prefill batch) so the tensor address is stable across
|
|
||||||
CUDA-graph replays."""
|
|
||||||
key = (topk_routed, n_shared, num_experts, float(shared_weight), str(device))
|
|
||||||
buf = _aiter_topk_fuse_shared_bufs.get(key)
|
|
||||||
if buf is None:
|
|
||||||
total = topk_routed + n_shared
|
|
||||||
M = _get_aiter_topk_fuse_shared_max_tokens()
|
|
||||||
w = torch.empty((M, total), dtype=torch.float32, device=device)
|
|
||||||
ids = torch.empty((M, total), dtype=torch.int32, device=device)
|
|
||||||
ids[:, topk_routed:] = torch.arange(
|
|
||||||
num_experts, num_experts + n_shared, dtype=torch.int32, device=device
|
|
||||||
).unsqueeze(0)
|
|
||||||
w[:, topk_routed:] = shared_weight
|
|
||||||
buf = (w, ids)
|
|
||||||
_aiter_topk_fuse_shared_bufs[key] = buf
|
|
||||||
return buf
|
|
||||||
|
|
||||||
|
|
||||||
if _is_cuda:
|
if _is_cuda:
|
||||||
try:
|
try:
|
||||||
@@ -1518,7 +1446,6 @@ def biased_grouped_topk_gpu(
|
|||||||
num_fused_shared_experts: int = 0,
|
num_fused_shared_experts: int = 0,
|
||||||
routed_scaling_factor: Optional[float] = None,
|
routed_scaling_factor: Optional[float] = None,
|
||||||
apply_routed_scaling_factor_on_output: Optional[bool] = False,
|
apply_routed_scaling_factor_on_output: Optional[bool] = False,
|
||||||
fused_shared_experts_scaling_factor: Optional[float] = None,
|
|
||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
num_tokens = gating_output.shape[0]
|
num_tokens = gating_output.shape[0]
|
||||||
num_experts = gating_output.shape[1]
|
num_experts = gating_output.shape[1]
|
||||||
@@ -1644,34 +1571,8 @@ def biased_grouped_topk_gpu(
|
|||||||
assert (
|
assert (
|
||||||
hidden_states.shape[0] == gating_output.shape[0]
|
hidden_states.shape[0] == gating_output.shape[0]
|
||||||
), f"Number of tokens mismatch: hidden_states.shape[0] = {hidden_states.shape[0]}, gating_output.shape[0] = {gating_output.shape[0]}"
|
), f"Number of tokens mismatch: hidden_states.shape[0] = {hidden_states.shape[0]}, gating_output.shape[0] = {gating_output.shape[0]}"
|
||||||
_shared_fuse = (
|
topk_weights = torch.empty((token, topk), dtype=torch.float32, device=device)
|
||||||
num_fused_shared_experts > 0
|
topk_ids = torch.empty((token, topk), dtype=torch.int32, device=device)
|
||||||
and get_parallel().moe_ep_size == 1
|
|
||||||
and token <= _get_aiter_topk_fuse_shared_max_tokens()
|
|
||||||
)
|
|
||||||
if _shared_fuse:
|
|
||||||
# Persistent buffer with pre-populated shared columns. The weight is the
|
|
||||||
# constant scale_factor the aiter _post_process append writes (default
|
|
||||||
# 1.0), so this is bit-identical for any shared-expert scaling. aiter
|
|
||||||
# writes only the routed columns [:, :topk] via row stride; the shared
|
|
||||||
# column stays intact. If token exceeds the buffer size we drop to the
|
|
||||||
# plain path below (and _post_process appends shared experts as usual) —
|
|
||||||
# the same condition is mirrored there so the two paths stay consistent.
|
|
||||||
_shared_w = (
|
|
||||||
1.0
|
|
||||||
if fused_shared_experts_scaling_factor is None
|
|
||||||
else fused_shared_experts_scaling_factor
|
|
||||||
)
|
|
||||||
full_w, full_ids = _get_aiter_topk_fuse_shared_buf(
|
|
||||||
topk, num_fused_shared_experts, num_experts, _shared_w, device
|
|
||||||
)
|
|
||||||
topk_weights = full_w[:token, :topk]
|
|
||||||
topk_ids = full_ids[:token, :topk]
|
|
||||||
else:
|
|
||||||
topk_weights = torch.empty(
|
|
||||||
(token, topk), dtype=torch.float32, device=device
|
|
||||||
)
|
|
||||||
topk_ids = torch.empty((token, topk), dtype=torch.int32, device=device)
|
|
||||||
aiter_biased_grouped_topk(
|
aiter_biased_grouped_topk(
|
||||||
gating_output,
|
gating_output,
|
||||||
correction_bias.to(dtype=gating_output.dtype),
|
correction_bias.to(dtype=gating_output.dtype),
|
||||||
@@ -1682,10 +1583,6 @@ def biased_grouped_topk_gpu(
|
|||||||
renormalize,
|
renormalize,
|
||||||
routed_scaling_factor if routed_scaling_factor is not None else 1.0,
|
routed_scaling_factor if routed_scaling_factor is not None else 1.0,
|
||||||
)
|
)
|
||||||
if _shared_fuse:
|
|
||||||
# Return the full [token, topk + n_shared] view (routed just written,
|
|
||||||
# shared pre-populated). _post_process_topk_ids skips its append.
|
|
||||||
return full_w[:token], full_ids[:token]
|
|
||||||
return topk_weights, topk_ids
|
return topk_weights, topk_ids
|
||||||
elif _is_musa and (
|
elif _is_musa and (
|
||||||
gating_output.shape[1] // num_expert_group <= 32
|
gating_output.shape[1] // num_expert_group <= 32
|
||||||
@@ -2139,43 +2036,25 @@ def _post_process_topk_ids(
|
|||||||
num_local_routed,
|
num_local_routed,
|
||||||
)
|
)
|
||||||
elif _aiter_append:
|
elif _aiter_append:
|
||||||
# Detect whether the shared experts were already fused/appended in
|
M, N = router_logits.shape
|
||||||
# biased_grouped_topk_gpu via the persistent pre-populated topk buffer.
|
scale_factor = (
|
||||||
# When fused, topk_ids already has the full width (routed + shared), i.e.
|
1.0
|
||||||
# topk_ids.shape[1] == topk_config.top_k; when the fast path fell back to
|
if fused_shared_experts_scaling_factor is None
|
||||||
# the plain buffer (e.g. token > MAX during a large prefill), only the
|
else fused_shared_experts_scaling_factor
|
||||||
# routed columns are present and we must append the shared experts here.
|
|
||||||
# Checking the tensor shape is robust to the exact fast-path conditions
|
|
||||||
# (no fragile mirroring of biased_grouped_topk_gpu's _shared_fuse check).
|
|
||||||
_shared_fused_in_topk = (
|
|
||||||
num_fused_shared_experts > 0
|
|
||||||
and get_parallel().moe_ep_size == 1
|
|
||||||
and topk_ids.shape[1] == topk_config.top_k
|
|
||||||
)
|
)
|
||||||
if _shared_fused_in_topk:
|
|
||||||
# Shared experts were already appended in biased_grouped_topk_gpu via
|
|
||||||
# the persistent pre-populated topk buffer; nothing to do here.
|
|
||||||
pass
|
|
||||||
else:
|
|
||||||
M, N = router_logits.shape
|
|
||||||
scale_factor = (
|
|
||||||
1.0
|
|
||||||
if fused_shared_experts_scaling_factor is None
|
|
||||||
else fused_shared_experts_scaling_factor
|
|
||||||
)
|
|
||||||
|
|
||||||
# Lazy import to avoid circular-import issues
|
# Lazy import to avoid circular-import issues
|
||||||
from sglang.kernels.ops.moe.fused_moe_triton_kernels import (
|
from sglang.kernels.ops.moe.fused_moe_triton_kernels import (
|
||||||
fused_append_shared_experts,
|
fused_append_shared_experts,
|
||||||
)
|
)
|
||||||
|
|
||||||
topk_ids, topk_weights = fused_append_shared_experts(
|
topk_ids, topk_weights = fused_append_shared_experts(
|
||||||
topk_ids,
|
topk_ids,
|
||||||
topk_weights,
|
topk_weights,
|
||||||
num_fused_shared_experts,
|
num_fused_shared_experts,
|
||||||
scale_factor,
|
scale_factor,
|
||||||
N, # base id for shared experts
|
N, # base id for shared experts
|
||||||
)
|
)
|
||||||
|
|
||||||
elif use_per_rank_shared_slots:
|
elif use_per_rank_shared_slots:
|
||||||
# DeepEP/MegaMOE: remap to per-rank shared-slot layout where each
|
# DeepEP/MegaMOE: remap to per-rank shared-slot layout where each
|
||||||
@@ -2270,7 +2149,6 @@ def select_experts(
|
|||||||
num_fused_shared_experts=num_fused_shared_experts,
|
num_fused_shared_experts=num_fused_shared_experts,
|
||||||
routed_scaling_factor=routed_scaling_factor,
|
routed_scaling_factor=routed_scaling_factor,
|
||||||
apply_routed_scaling_factor_on_output=apply_routed_scaling_factor_on_output,
|
apply_routed_scaling_factor_on_output=apply_routed_scaling_factor_on_output,
|
||||||
fused_shared_experts_scaling_factor=topk_config.fused_shared_experts_scaling_factor,
|
|
||||||
)
|
)
|
||||||
elif torch_native and custom_routing_function is None:
|
elif torch_native and custom_routing_function is None:
|
||||||
assert (
|
assert (
|
||||||
|
|||||||
Reference in New Issue
Block a user