From 716e670d3dc0404961756353d95ef219a9aab595 Mon Sep 17 00:00:00 2001 From: Jimmy Shong <69131491+Jiminator@users.noreply.github.com> Date: Sat, 30 May 2026 00:27:20 -0700 Subject: [PATCH] [bugfix]: size CuteDSL MoE allgather buffers for the worst-case forward (#26696) --- .../moe/moe_runner/flashinfer_cutedsl.py | 26 ++---- .../srt/layers/quantization/modelopt_quant.py | 2 +- python/sglang/srt/server_args.py | 90 ++++++++++++------- .../unit/server_args/test_server_args.py | 41 +++++++++ 4 files changed, 109 insertions(+), 50 deletions(-) diff --git a/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutedsl.py b/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutedsl.py index 6b16841a5..b78da24c9 100644 --- a/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutedsl.py +++ b/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutedsl.py @@ -221,17 +221,9 @@ def resolve_cutedsl_standard_scales( return w1_alpha, fc2_input_scale, w2_alpha, used_input_scale -def ensure_cutedsl_wrapper(layer: torch.nn.Module, num_tokens: int = 0) -> None: +def ensure_cutedsl_wrapper(layer: torch.nn.Module) -> None: """Lazily create CuteDslMoEWrapper and resolve scales on first forward. - Args: - layer: The FusedMoE layer module. - num_tokens: Current token count entering the MoE layer. Used as - the buffer size for the non-a2a (allgather) path, where the - autotune dummy run passes req_to_token_pool.size * dp_size — - the worst-case post-allgather batch. For the a2a path this - is ignored in favour of the dispatcher's workspace limit. - The wrapper is created lazily (not in __init__ / create_weights) because it depends on final weight shapes and EP configuration. The wrapper's CUDA-graph buffers are allocated inside CuteDslMoEWrapper.__init__, which @@ -259,20 +251,18 @@ def ensure_cutedsl_wrapper(layer: torch.nn.Module, num_tokens: int = 0) -> None: ) server_args = get_global_server_args() - use_cuda_graph = server_args is not None and not server_args.disable_cuda_graph + use_cuda_graph = not server_args.disable_cuda_graph - # Buffer size must cover the worst-case token count the MoE layer can see. - # - A2A path: dispatch returns tensors flattened from - # [ep_size, max_tokens_per_rank, ...]. - # - Standard allgather path: dp_size * max local tokens per rank. + # Size the wrapper's CUDA-graph buffers for the largest number of tokens a + # single forward can route through this layer. dispatcher = getattr(layer, "dispatcher", None) if hasattr(dispatcher, "max_num_tokens"): + # A2A path: bounded by the dispatcher's own workspace limit. max_num_tokens = dispatcher.max_num_tokens * getattr(dispatcher, "ep_size", 1) else: - # Standard allgather path: num_tokens from the first forward is - # req_to_token_pool.size * dp_size (the autotune dummy run's batch), - # which is the worst-case post-allgather token count. - max_num_tokens = max(num_tokens, 1) + # Standard allgather path: the MoE sees up to dp_size local forwards + # gathered together, so scale the per-rank forward bound by dp_size. + max_num_tokens = server_args.dp_size * server_args.cutedsl_moe_max_num_tokens() top_k = layer.top_k if layer.top_k is not None else layer.moe_runner_config.top_k # inference_mode(False) ensures the wrapper's pre-allocated CUDA-graph # buffers are normal tensors. This call typically happens inside diff --git a/python/sglang/srt/layers/quantization/modelopt_quant.py b/python/sglang/srt/layers/quantization/modelopt_quant.py index adfb8443a..c91643e83 100755 --- a/python/sglang/srt/layers/quantization/modelopt_quant.py +++ b/python/sglang/srt/layers/quantization/modelopt_quant.py @@ -2200,7 +2200,7 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): # v2 standard path (a2a=none/flashinfer): uses CuteDslMoEWrapper # with [Up, Gate] interleaved weights and MMA blockscales. - ensure_cutedsl_wrapper(layer, dispatch_output.hidden_states.shape[0]) + ensure_cutedsl_wrapper(layer) w1_alpha, fc2_input_scale, w2_alpha = layer._cutedsl_scales quant_info = CuteDslFp4MoeQuantInfo( w13_weight=layer.w13_weight, diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 015125e96..247d113d8 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -991,6 +991,9 @@ class ServerArgs: handle_speculative_decoding(self) + # Validate the CuteDSL A2A token budget now that num_tokens_per_bs is final. + self._validate_cutedsl_a2a_token_budget() + # Handle model loading format. self._handle_load_format() @@ -3328,6 +3331,62 @@ class ServerArgs: self.ep_size == 1 ), "FP8/MXFP8 Cutlass MoE is only supported with ep_size == 1" + def cutedsl_moe_max_num_tokens(self) -> int: + """Largest number of tokens a single forward routes through a CuteDSL + MoE layer on one (DP) rank. Single source of truth for both the + standard-allgather wrapper buffers and the FlashInfer A2A dispatcher + budget. Max over the prefill (max_prefill_tokens), piecewise-prefill + capture (piecewise_cuda_graph_max_tokens), and decode/verify + (cuda_graph_max_bs * num_tokens_per_bs) bounds; num_tokens_per_bs is + speculative_num_draft_tokens under speculative decoding, else 1. + """ + if self.speculative_algorithm: + num_tokens_per_bs = self.speculative_num_draft_tokens or 1 + else: + num_tokens_per_bs = 1 + prefill_tokens = self.max_prefill_tokens + if not self.disable_piecewise_cuda_graph: + prefill_tokens = max( + prefill_tokens, self.piecewise_cuda_graph_max_tokens or 0 + ) + decode_tokens = (self.cuda_graph_max_bs or 0) * num_tokens_per_bs + return max(prefill_tokens, decode_tokens) + + def _validate_cutedsl_a2a_token_budget(self): + """Fail fast if the FlashInfer A2A dispatcher workspace cannot cover the + largest CuteDSL MoE forward. Runs after speculative decoding is resolved + so cutedsl_moe_max_num_tokens() sees the final num_tokens_per_bs.""" + if not ( + self.moe_a2a_backend == "flashinfer" + and self.moe_runner_backend == "flashinfer_cutedsl" + and self.max_prefill_tokens > 0 + and self.disaggregation_mode != "decode" + ): + return + required_tokens = self.cutedsl_moe_max_num_tokens() + max_dispatch_tokens_per_rank = get_int_env_var( + "SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK", 1024 + ) + max_cutedsl_tokens = max_dispatch_tokens_per_rank * self.ep_size + if max_cutedsl_tokens < required_tokens: + required_per_rank = (required_tokens + self.ep_size - 1) // self.ep_size + raise ValueError( + "FlashInfer MoE A2A with flashinfer_cutedsl requires " + "SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK * " + "ep_size to cover the largest CuteDSL MoE forward " + f"({required_tokens} tokens). Otherwise the FlashInfer " + "dispatcher can crash at runtime with " + "`ValueError: num_tokens (...) exceeds max_num_tokens (...)`. " + "Current values: " + f"SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK=" + f"{max_dispatch_tokens_per_rank}, ep_size={self.ep_size}, " + f"capacity={max_cutedsl_tokens}, required={required_tokens}. " + f"Set `export " + f"SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK=" + f"{required_per_rank}` or lower the relevant limit " + f"(e.g. --max-prefill-tokens) to <= {max_cutedsl_tokens}." + ) + def _handle_a2a_moe(self): if self.enable_deepep_waterfill and self.moe_a2a_backend != "deepep": logger.warning( @@ -3423,37 +3482,6 @@ class ServerArgs: "flashinfer_cutlass", "flashinfer_cutedsl", ], "Flashinfer MoE A2A is only supported with flashinfer_cutlass or flashinfer_cutedsl moe runner backend" - if ( - self.moe_runner_backend == "flashinfer_cutedsl" - and self.max_prefill_tokens is not None - and self.max_prefill_tokens > 0 - and self.disaggregation_mode != "decode" - ): - max_dispatch_tokens_per_rank = get_int_env_var( - "SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK", 1024 - ) - max_cutedsl_tokens = max_dispatch_tokens_per_rank * self.ep_size - if max_cutedsl_tokens < self.max_prefill_tokens: - required_per_rank = ( - self.max_prefill_tokens + self.ep_size - 1 - ) // self.ep_size - raise ValueError( - "FlashInfer MoE A2A with flashinfer_cutedsl requires " - "SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK * " - "ep_size to cover --max-prefill-tokens. Otherwise the " - "FlashInfer dispatcher can crash at runtime with " - "`ValueError: num_tokens (...) exceeds max_num_tokens (...)` " - "when a local DP rank schedules too many prefill tokens. " - "Current values: " - f"SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK=" - f"{max_dispatch_tokens_per_rank}, ep_size={self.ep_size}, " - f"capacity={max_cutedsl_tokens}, " - f"max_prefill_tokens={self.max_prefill_tokens}. " - f"Set `export " - f"SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK=" - f"{required_per_rank}` or lower `--max-prefill-tokens` " - f"to <= {max_cutedsl_tokens}." - ) if self.moe_a2a_backend == "mori": self.ep_size = self.tp_size diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 3ec7956ae..b04faa8fb 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -618,5 +618,46 @@ class TestPrefillOnlyDisableKvCache(unittest.TestCase): ServerArgs(**self._base_kwargs(kv_cache_dtype="fp4_e2m1")) +class TestCutedslMoeMaxNumTokens(unittest.TestCase): + """The shared CuteDSL MoE per-forward token bound. Fields are set directly + to exercise the math independently of __post_init__ resolution.""" + + def _args(self, **overrides): + server_args = ServerArgs(model_path="dummy") + fields = dict( + speculative_algorithm=None, + speculative_num_draft_tokens=None, + max_prefill_tokens=16384, + disable_piecewise_cuda_graph=False, + piecewise_cuda_graph_max_tokens=2048, + cuda_graph_max_bs=512, + ) + fields.update(overrides) + for key, value in fields.items(): + setattr(server_args, key, value) + return server_args + + def test_prefill_dominates_in_default_config(self): + self.assertEqual(self._args().cutedsl_moe_max_num_tokens(), 16384) + + def test_speculative_decoding_scales_decode_bound(self): + # decode bound 512 * 8 dominates the small prefill/piecewise bounds + args = self._args( + max_prefill_tokens=512, + piecewise_cuda_graph_max_tokens=512, + speculative_algorithm="EAGLE", + speculative_num_draft_tokens=8, + ) + self.assertEqual(args.cutedsl_moe_max_num_tokens(), 4096) + + def test_piecewise_bound_excluded_when_disabled(self): + args = self._args( + max_prefill_tokens=512, + disable_piecewise_cuda_graph=True, + cuda_graph_max_bs=64, + ) + self.assertEqual(args.cutedsl_moe_max_num_tokens(), 512) + + if __name__ == "__main__": unittest.main()