From 0c39730b18f9f2f383e26deacc632f60f9a3051f Mon Sep 17 00:00:00 2001 From: Huaixin Chang <61184708+changhuaixin@users.noreply.github.com> Date: Thu, 25 Dec 2025 11:35:05 +0800 Subject: [PATCH] DP: support piggyback server load report (#11469) Signed-off-by: Chang Huaixin (OpenAnolis) --- docs/references/environment_variables.md | 1 + python/sglang/srt/environ.py | 1 + .../srt/managers/data_parallel_controller.py | 54 ++++++++++++++----- .../srt/managers/detokenizer_manager.py | 8 ++- python/sglang/srt/managers/io_struct.py | 9 ++++ python/sglang/srt/managers/scheduler.py | 4 +- .../srt/managers/scheduler_metrics_mixin.py | 1 + .../scheduler_output_processor_mixin.py | 17 +++++- .../sglang/srt/managers/tokenizer_manager.py | 29 ++++------ python/sglang/srt/server_args.py | 7 --- 10 files changed, 86 insertions(+), 45 deletions(-) diff --git a/docs/references/environment_variables.md b/docs/references/environment_variables.md index a379c2cd6..b53f3e1b2 100644 --- a/docs/references/environment_variables.md +++ b/docs/references/environment_variables.md @@ -32,6 +32,7 @@ SGLang supports various environment variables that can be used to configure its | `SGLANG_DISABLE_CONSECUTIVE_PREFILL_OVERLAP` | Disable overlap schedule for consecutive prefill batches | `false` | | `SGLANG_SCHEDULER_MAX_RECV_PER_POLL` | Set the maximum number of requests per poll, with a negative value indicating no limit | `-1` | | `SGLANG_DISABLE_FA4_WARMUP` | Disable Flash Attention 4 warmup passes (set to `1`, `true`, `yes`, or `on` to disable) | `false` | +| `SGLANG_DATA_PARALLEL_BUDGET_INTERVAL` | Interval for DPBudget updates | `1` | | `SGLANG_SCHEDULER_RECV_SKIPPER_WEIGHT_DEFAULT` | Default weight value for scheduler recv skipper counter (used when forward mode doesn't match specific modes). Only active when `--scheduler-recv-interval > 1`. The counter accumulates weights and triggers request polling when reaching the interval threshold. | `1000` | | `SGLANG_SCHEDULER_RECV_SKIPPER_WEIGHT_DECODE` | Weight increment for decode forward mode in scheduler recv skipper. Works with `--scheduler-recv-interval` to control polling frequency during decode phase. | `1` | | `SGLANG_SCHEDULER_RECV_SKIPPER_WEIGHT_VERIFY` | Weight increment for target verify forward mode in scheduler recv skipper. Works with `--scheduler-recv-interval` to control polling frequency during verification phase. | `1` | diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 61954412c..e443170f8 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -204,6 +204,7 @@ class Envs: SGLANG_DYNAMIC_CHUNKING_SMOOTH_FACTOR = EnvFloat(0.75) SGLANG_SCHEDULER_SKIP_ALL_GATHER = EnvBool(False) SGLANG_SCHEDULER_DECREASE_PREFILL_IDLE = EnvBool(False) + SGLANG_DATA_PARALLEL_BUDGET_INTERVAL = EnvInt(1) # Test: pd-disaggregation SGLANG_TEST_PD_DISAGG_BACKEND = EnvStr("mooncake") diff --git a/python/sglang/srt/managers/data_parallel_controller.py b/python/sglang/srt/managers/data_parallel_controller.py index df37be6cf..de79ec056 100644 --- a/python/sglang/srt/managers/data_parallel_controller.py +++ b/python/sglang/srt/managers/data_parallel_controller.py @@ -84,18 +84,40 @@ class LoadBalanceMethod(Enum): class DPBudget: - def __init__(self): + def __init__(self, dp_size: int): # TODO: support minimum tokens method self.budget_queue = deque() + self.dp_size = dp_size + self.ts_tic = 0.0 + self.pending_loads = {} + # Set time window to 2ms + self.tic_window = 0.002 + self.update_budget_count = 0 + self.update_interval = envs.SGLANG_DATA_PARALLEL_BUDGET_INTERVAL.get() def update_budget(self, load_update: WatchLoadUpdateReq): - """Update the budget queue. - Use num_reqs instead of num_waiting_reqs to balance decode running batch. - """ - loads = load_update.loads - self.budget_queue.clear() + """Update the budget queue.""" + # Update budget queue together for load updating from the same round. + for load in load_update.loads: + if abs(load.ts_tic - self.ts_tic) > self.tic_window: + logger.debug(f"Proceed to next round: {self.ts_tic=} {load.ts_tic=}") + self.pending_loads.clear() + self.ts_tic = load.ts_tic + self.pending_loads[load.dp_rank] = load - num_reqs = [load.num_reqs for load in loads] + if len(self.pending_loads) < self.dp_size: + logger.debug(f"Waiting for all DP ranks: {len(self.pending_loads)=}") + return + + self.update_budget_count = (self.update_budget_count + 1) % self.update_interval + if self.update_budget_count: + return + + # Ready to update budget_queue. + self.budget_queue.clear() + num_reqs = [0] * self.dp_size + for dp_rank, load in self.pending_loads.items(): + num_reqs[dp_rank] = load.num_reqs if not num_reqs: return @@ -105,18 +127,21 @@ class DPBudget: while any(x != num_reqs[0] for x in num_reqs): min_load = min(num_reqs) - min_indices = [i for i, x in enumerate(num_reqs) if x == min_load] + min_indices = [ + dp_rank for dp_rank, x in enumerate(num_reqs) if x == min_load + ] second_min_load = min(x for x in num_reqs if x > min_load) self.budget_queue.extend( - [loads[i].dp_rank for i in min_indices] * (second_min_load - min_load) + [dp_rank for dp_rank in min_indices] * (second_min_load - min_load) ) for idx in min_indices: num_reqs[idx] = second_min_load def dispatch(self): - if self.budget_queue: - return self.budget_queue.popleft() - return None + if not self.budget_queue: + self.budget_queue.extend(range(self.dp_size)) + + return self.budget_queue.popleft() class DataParallelController: @@ -157,7 +182,7 @@ class DataParallelController: self.dispatching = dispatch_lookup[self.load_balance_method] # Load balance budget - self.dp_budget = DPBudget() + self.dp_budget = DPBudget(server_args.dp_size) # To protect changing env vars to set CUDA_VISIBLE_DEVICES. self.env_lock = threading.Lock() @@ -502,7 +527,8 @@ class DataParallelController: assert ( req.bootstrap_room is not None ), "req.bootstrap_room should not be None. Do not send requests directly to prefill or decode instances, but send to the router instead." - self.workers[req.bootstrap_room % len(self.workers)].send_pyobj(req) + target_rank = req.bootstrap_room % len(self.workers) + self.workers[target_rank].send_pyobj(req) def decode_round_robin_scheduler(self, req: Req): if self.maybe_external_dp_rank_routing(req): diff --git a/python/sglang/srt/managers/detokenizer_manager.py b/python/sglang/srt/managers/detokenizer_manager.py index 60fb1a9c2..0e64a80b5 100644 --- a/python/sglang/srt/managers/detokenizer_manager.py +++ b/python/sglang/srt/managers/detokenizer_manager.py @@ -294,7 +294,12 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin): return output_routed_experts def handle_batch_token_id_out(self, recv_obj: BatchTokenIDOutput): - output_strs = self._decode_batch_token_id_output(recv_obj) + # If handling idle batch, set output_strs to []. + output_strs = ( + self._decode_batch_token_id_output(recv_obj) + if len(recv_obj.rids) > 0 + else [] + ) output_routed_experts = self._extract_routed_experts(recv_obj) return BatchStrOutput( @@ -331,6 +336,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin): forward_entry_time=recv_obj.forward_entry_time, prefill_launch_delay=recv_obj.prefill_launch_delay, prefill_launch_latency=recv_obj.prefill_launch_latency, + load=recv_obj.load, prefill_finished_ts=recv_obj.prefill_finished_ts, ) diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index c592f0cb5..02131604d 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -16,6 +16,8 @@ The definition of objects transferred between different processes (TokenizerManager, DetokenizerManager, Scheduler). """ +from __future__ import annotations + import copy import uuid from abc import ABC @@ -975,6 +977,9 @@ class BatchTokenIDOutput( # The trainer step id. Used to know which step's weights are used for sampling. token_steps: List[List[int]] = None + # Load for DP balance + load: GetLoadReqOutput = None + @dataclass class BatchMultimodalDecodeReq(BaseBatchReq): @@ -1057,6 +1062,9 @@ class BatchStrOutput( # The trainer step id. Used to know which step's weights are used for sampling. token_steps: List[List[int]] = None + # Load for DP balance + load: GetLoadReqOutput = None + @dataclass class BatchMultimodalOutput(BaseBatchReq): @@ -1642,6 +1650,7 @@ class GetLoadReqOutput(BaseReq): num_reqs: int num_waiting_reqs: int num_tokens: int + ts_tic: float @dataclass diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index d1bb2c2ed..141d95d5d 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2324,9 +2324,7 @@ class Scheduler( elif batch.forward_mode.is_prebuilt(): self.process_batch_result_prebuilt(batch) elif batch.forward_mode.is_idle(): - if self.enable_overlap: - if result.copy_done is not None: - result.copy_done.synchronize() + self.process_batch_result_idle(batch, result) self.log_batch_result_stats(batch, result) self.maybe_send_health_check_signal() diff --git a/python/sglang/srt/managers/scheduler_metrics_mixin.py b/python/sglang/srt/managers/scheduler_metrics_mixin.py index 22437b4b3..96a0020a4 100644 --- a/python/sglang/srt/managers/scheduler_metrics_mixin.py +++ b/python/sglang/srt/managers/scheduler_metrics_mixin.py @@ -539,6 +539,7 @@ class SchedulerMetricsMixin: num_reqs=len(self.running_batch.reqs) + num_waiting_reqs, num_waiting_reqs=num_waiting_reqs, num_tokens=num_tokens, + ts_tic=time.perf_counter(), ) @contextmanager diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py index 7b68cae46..450e06bb3 100644 --- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py +++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py @@ -299,6 +299,18 @@ class SchedulerOutputProcessorMixin: return predict_tokens + def process_batch_result_idle( + self: Scheduler, + batch: ScheduleBatch, + result: GenerationBatchResult, + ): + if result.copy_done is not None: + result.copy_done.synchronize() + + self.stream_output_generation( + batch.reqs, batch.return_logprob, is_idle_batch=True + ) + def process_batch_result_dllm( self: Scheduler, batch: ScheduleBatch, @@ -790,6 +802,7 @@ class SchedulerOutputProcessorMixin: reqs: List[Req], return_logprob: bool, skip_req: Optional[Req] = None, + is_idle_batch: bool = False, ): rids = [] http_worker_ipcs = [] @@ -810,6 +823,7 @@ class SchedulerOutputProcessorMixin: spec_accepted_tokens = [] retraction_counts = [] output_hidden_states = None + load = self.get_load() output_routed_experts = None queue_times = [] @@ -1018,7 +1032,7 @@ class SchedulerOutputProcessorMixin: req.log_time_stats() # Send to detokenizer - if rids: + if reqs or is_idle_batch: if self.model_config.is_multimodal_gen: return @@ -1062,6 +1076,7 @@ class SchedulerOutputProcessorMixin: placeholder_tokens_idx=None, placeholder_tokens_val=None, retraction_counts=retraction_counts, + load=load, ) ) diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 0c20bbe7e..9be7d4381 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -57,7 +57,6 @@ from sglang.srt.managers.io_struct import ( EmbeddingReqInput, FreezeGCReq, GenerateReqInput, - GetLoadReqInput, HealthCheckOutput, LoadLoRAAdapterReqInput, OpenSessionReqOutput, @@ -1379,9 +1378,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi self.asyncio_tasks.add( loop.create_task(print_exception_wrapper(self.sigterm_watchdog)) ) - self.asyncio_tasks.add( - loop.create_task(print_exception_wrapper(self.watch_load_thread)) - ) def dump_requests_before_crash(self): if self.crash_dump_performed: @@ -1614,6 +1610,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi "output_ids": output_token_ids, "meta_info": meta_info, } + elif isinstance(recv_obj, BatchTokenIDOutput): is_stream = getattr(state.obj, "stream", False) if self.server_args.stream_output and is_stream: @@ -1667,6 +1664,15 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi if self.crash_dump_folder and state.finished and state.obj.log_metrics: self.record_request_for_crash_dump(state, out_dict) + # When skip_tokenizer_init is enabled, tokensizer_manager receives + # BatchTokenIDOutput. + if self.server_args.dp_size > 1 and ( + isinstance(recv_obj, BatchStrOutput) + or isinstance(recv_obj, BatchTokenIDOutput) + ): + load_update_req = WatchLoadUpdateReq(loads=[recv_obj.load]) + self.send_to_scheduler.send_pyobj(load_update_req) + def add_logprob_to_meta_info( self, meta_info: dict, @@ -2077,21 +2083,6 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi logprobs[token_id] = logprob return logprobs - async def watch_load_thread(self): - # Only for dp_controller when dp_size > 1 - if ( - self.server_args.dp_size == 1 - or self.server_args.load_balance_method == "round_robin" - or self.server_args.load_balance_method == "decode_round_robin" - ): - return - - while True: - await asyncio.sleep(self.server_args.load_watch_interval) - loads = await self.get_load_communicator(GetLoadReqInput()) - load_udpate_req = WatchLoadUpdateReq(loads=loads) - self.send_to_scheduler.send_pyobj(load_udpate_req) - async def _resolve_lora_path(self, obj: Union[GenerateReqInput, EmbeddingReqInput]): if isinstance(obj.lora_path, str): unique_lora_paths = set([obj.lora_path]) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 60157c888..ce61d4a67 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -377,7 +377,6 @@ class ServerArgs: # Data parallelism dp_size: int = 1 load_balance_method: str = "round_robin" - load_watch_interval: float = 0.1 # FIXME: remove this after dp rank scheduling is fully supported with PD-Disaggregation prefill_round_robin_balance: bool = False @@ -3145,12 +3144,6 @@ class ServerArgs: "minimum_tokens", ], ) - parser.add_argument( - "--load-watch-interval", - type=float, - default=ServerArgs.load_watch_interval, - help="The interval of load watching in seconds.", - ) parser.add_argument( "--prefill-round-robin-balance", default=ServerArgs.prefill_round_robin_balance,