diff --git a/python/sglang/srt/layers/moe/moe_runner/aiter.py b/python/sglang/srt/layers/moe/moe_runner/aiter.py index 8fa06487d..07020e62b 100644 --- a/python/sglang/srt/layers/moe/moe_runner/aiter.py +++ b/python/sglang/srt/layers/moe/moe_runner/aiter.py @@ -2,6 +2,7 @@ from __future__ import annotations import functools import inspect +import logging from dataclasses import dataclass from enum import Enum from typing import TYPE_CHECKING, Any, Optional, Union @@ -36,6 +37,9 @@ if TYPE_CHECKING: ) +logger = logging.getLogger(__name__) + + class AiterQuantType(str, Enum): NONE = "No" PER_TOKEN = "per_Token" @@ -123,6 +127,98 @@ def _aiter_fused_moe_supports_no_combine() -> bool: return "no_combine" in inspect.signature(fused_moe).parameters +_RECV_BOUND_LOGGED: set[int] = set() +_RECV_BOUND_WARNED = False + + +def _warn_recv_bound_unavailable() -> None: + global _RECV_BOUND_WARNED + if not _RECV_BOUND_WARNED: + _RECV_BOUND_WARNED = True + logger.warning( + "SGLANG_MORI_RECV_BOUND is set but the per-rank DP token counts do " + "not cover every mori sender, so the receive fan-in is unknown; " + "leaving the receive buffer unbounded." + ) + + +def _mori_decode_recv_bound(recv_rows: int, topk: int) -> int: + """Live rows mori's receive buffer can hold in decode, or 0 for "do not bound". + + Worst case fan-in is every rank routing all of its tokens to this one, so + `sum(per-rank tokens) * topk`, where topk already includes the fused shared + expert. The per-rank counts come from the DP sync, so this is the fan-in for + the batch actually being run rather than an upper bound over all batches. + + That is only sound because enabling this gate also makes + `require_mlp_tp_gather()` true for mori, which gives every rank the same + cuda-graph bucket. The value is baked into a captured graph and has to hold + for every later replay; with per-rank buckets a rank on a narrow tier could + be handed rows by a peer on a wider one, and the only bound valid under that + is the widest tier's -- 4-16x looser than the batch being run, which costs + more in expert-GEMM tiles (M 32/64 -> 128) than the trim saves. + + Two cases stay unbounded, because a bound below the real fan-in silently + drops rows from the all-to-all -- wrong output rather than an error: + + * Prefill, whose per-rank counts are uneven and not knowable here. + * Anything that leaves the per-rank counts unpopulated, or where the EP world + is wider than the DP world so the counts do not cover every sender. + """ + if not get_bool_env_var("SGLANG_MORI_RECV_BOUND", "false"): + return 0 + + from sglang.srt.layers.dp_attention import ( + get_dp_global_num_tokens, + get_is_extend_in_batch, + ) + from sglang.srt.runtime_context import get_parallel + + if get_is_extend_in_batch(): + return 0 + + per_rank_tokens = get_dp_global_num_tokens() + ep_size = get_parallel().moe_ep_size + if not per_rank_tokens or len(per_rank_tokens) < ep_size: + # Either the DP sync did not publish counts, or they do not cover every + # mori sender. Both mean the fan-in is unknown here. + _warn_recv_bound_unavailable() + return 0 + + max_tokens = sum(per_rank_tokens) + bound = max_tokens * topk + # Never grow the tensor, and nothing to do when there is nothing to trim. + if not 0 < bound < recv_rows: + return 0 + + # One INFO line the first time it engages, so an inert bound is not mistaken + # for an active one in the results. Per-tier values go to DEBUG: capture + # visits every tier, and at INFO on every rank that is dozens of lines. + if get_parallel().tp_rank == 0 and bound not in _RECV_BOUND_LOGGED: + first = not _RECV_BOUND_LOGGED + _RECV_BOUND_LOGGED.add(bound) + if first: + logger.info( + "mori recv bound active: %d rows -> %d for this tier " + "(dp_tokens=%d ep=%d topk=%d); per-tier values at DEBUG", + recv_rows, + bound, + max_tokens, + get_parallel().moe_ep_size, + topk, + ) + else: + logger.debug( + "mori recv bound: %d rows -> %d (dp_tokens=%d ep=%d topk=%d)", + recv_rows, + bound, + max_tokens, + get_parallel().moe_ep_size, + topk, + ) + return bound + + class AiterRunnerCore(MoeRunnerCore): def run( self, @@ -322,6 +418,10 @@ def _pre_permute_deepep_to_aiter( # reads [0, totalRecvTokenNum), so the truncated result needs no # padding back. mori_max = get_int_env_var("SGLANG_MORI_MOE_MAX_INPUT_TOKENS", 0) + if mori_max <= 0: + mori_max = _mori_decode_recv_bound( + hidden_states.shape[0], topk_ids.shape[-1] + ) if mori_max > 0: hidden_states = hidden_states[:mori_max] if a1_scale is not None: diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index 28058bd26..3ac8ace65 100644 --- a/python/sglang/srt/utils/common.py +++ b/python/sglang/srt/utils/common.py @@ -3753,6 +3753,19 @@ def require_mlp_tp_gather(): # reuse this flag's DP-sync bookkeeping (uniform global_num_tokens + # max-based graph bucket). See #30432 re: the misleading flag name. return True + elif get_moe_a2a_backend().is_mori() and get_bool_env_var( + "SGLANG_MORI_RECV_BOUND", "false" + ): + # Same bookkeeping, for the same reason. Bounding mori's receive + # buffer means baking a fan-in size into a captured graph, and the + # fan-in depends on what the *peers* send. Without a DP-synchronized + # bucket every rank buckets its own batch, so a rank on a narrow tier + # can be handed rows by a peer on a wider one; the only bound valid + # under that is the widest tier's, which is 4-16x looser than the + # batch actually being run and costs more in expert-GEMM tiles than + # the trim saves. With uniform buckets the per-tier fan-in is exact. + # Scoped to the opt-in gate so the default path is untouched. + return True else: return ( get_parallel().moe_dense_tp_size