From 6fdcb9934cee42bdb1c5247d323d21b554d58da4 Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Fri, 19 Jun 2026 13:23:26 -0700 Subject: [PATCH] fix(runner): size eager static buffers for prefill budget and MLP-sync autotune (#28677) --- .../srt/model_executor/runner/eager_runner.py | 12 +++++++++++- python/sglang/srt/server_args.py | 16 ++++++++++++++++ 2 files changed, 27 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/model_executor/runner/eager_runner.py b/python/sglang/srt/model_executor/runner/eager_runner.py index 6a8ffd370..d93af75a3 100644 --- a/python/sglang/srt/model_executor/runner/eager_runner.py +++ b/python/sglang/srt/model_executor/runner/eager_runner.py @@ -43,6 +43,7 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo set_tc_piecewise_forward_context, ) from sglang.srt.utils import is_hip, require_mlp_tp_gather +from sglang.srt.utils.common import ceil_align, require_mlp_sync logger = logging.getLogger(__name__) @@ -94,12 +95,21 @@ class EagerRunner(BaseRunner): # Frozen-KV MTP expands the draft batch by topk on the bs axis # (expand_for_topk_draft) before the eager fallback. max_bs *= sa.speculative_eagle_topk + # Mirror prepare_mlp_sync_batch padding so the registry holds what load_batch copies. + if require_mlp_sync(sa): + from sglang.srt.layers.utils.cp_utils import get_cp_padding_align_size + + max_bs = ceil_align(max_bs, self.attn_tp_size) + max_bs = ceil_align(max_bs, get_cp_padding_align_size()) prefill_ceiling = ( - sa.chunked_prefill_size + sa.max_prefill_buffer_tokens() if sa.chunked_prefill_size and sa.chunked_prefill_size > 0 else mr.max_total_num_tokens ) max_num_token = max(prefill_ceiling, max_bs * num_tokens_per_bs) + if require_mlp_sync(sa): + max_num_token = ceil_align(max_num_token, self.attn_tp_size) + max_num_token = ceil_align(max_num_token, get_cp_padding_align_size()) self._eager_max_bs = max_bs self._eager_num_tokens_per_bs = num_tokens_per_bs is_encoder_decoder = mr.model_config.is_encoder_decoder diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index e491f2c8a..21420e265 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -22,6 +22,7 @@ import importlib import importlib.util import json import logging +import math import os import random import socket @@ -3673,6 +3674,21 @@ class ServerArgs: decode_tokens = decode_max_bs * num_tokens_per_bs return max(prefill_tokens, decode_tokens) + def max_prefill_buffer_tokens(self) -> int: + """Prefill-buffer ceiling: chunked_prefill_size, except PP dynamic + chunking can grow chunks toward max_prefill_tokens and probe at 1.25x.""" + chunked = ( + self.chunked_prefill_size + if self.chunked_prefill_size and self.chunked_prefill_size > 0 + else 0 + ) + tokens = chunked + if self.enable_dynamic_chunking and self.pp_size > 1 and chunked: + tokens = max( + tokens, self.max_prefill_tokens or 0, math.ceil(chunked * 1.25) + ) + return 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