diff --git a/docs/advanced_features/server_arguments.md b/docs/advanced_features/server_arguments.md index c2914f789..38ad385cf 100644 --- a/docs/advanced_features/server_arguments.md +++ b/docs/advanced_features/server_arguments.md @@ -205,7 +205,7 @@ Please consult the documentation below and [server_args.py](https://github.com/s | Argument | Description | Defaults | Options | | --- | --- | --- | --- | | `--data-parallel-size`
`--dp-size` | The data parallelism size. | `1` | Type: int | -| `--load-balance-method` | The load balancing strategy for data parallelism. The Minimum Token algorithm can only be used when DP attention is applied. This algorithm performs load balancing based on the real-time token load of the DP workers. | `auto` | `auto`, `round_robin`, `follow_bootstrap_room`, `shortest_queue`, `minimum_tokens` | +| `--load-balance-method` | The load balancing strategy for data parallelism. The `total_tokens` algorithm can only be used when DP attention is applied. This algorithm performs load balancing based on the real-time token load of the DP workers. | `auto` | `auto`, `round_robin`, `follow_bootstrap_room`, `total_requests`, `total_tokens` | | `--load-watch-interval` | The interval of load watching in seconds. | `0.1` | Type: float | ## Multi-node distributed serving diff --git a/docs/platforms/ascend_npu_support_features.md b/docs/platforms/ascend_npu_support_features.md index 0a0060d8b..ee233a9f4 100644 --- a/docs/platforms/ascend_npu_support_features.md +++ b/docs/platforms/ascend_npu_support_features.md @@ -153,7 +153,7 @@ click [Server Arguments](https://docs.sglang.io/advanced_features/server_argumen | Argument | Defaults | Options | A2 | A3 | |----------------------------------------|---------------|-------------------------------------------------------------|:----------------------------------------:|:----------------------------------------:| | `--data-parallel-size`
`--dp-size` | `1` | Type: int | **** | **** | -| `--load-balance-method` | `round_robin` | `round_robin`,
`shortest_queue`,
`minimum_tokens` | **** | **** | +| `--load-balance-method` | `round_robin` | `round_robin`,
`total_requests`,
`total_tokens` | **** | **** | | `--prefill-round-robin-balance` | `False` | bool flag
(set to enable) | **** | **** | ## Multi-node distributed serving diff --git a/python/sglang/srt/managers/data_parallel_controller.py b/python/sglang/srt/managers/data_parallel_controller.py index 9790649a9..eb51352fd 100644 --- a/python/sglang/srt/managers/data_parallel_controller.py +++ b/python/sglang/srt/managers/data_parallel_controller.py @@ -19,7 +19,6 @@ import multiprocessing as mp import signal import threading import time -from collections import deque from enum import Enum, auto from typing import Callable, List, Optional @@ -72,8 +71,8 @@ class LoadBalanceMethod(Enum): ROUND_ROBIN = auto() FOLLOW_BOOTSTRAP_ROOM = auto() - SHORTEST_QUEUE = auto() - MINIMUM_TOKENS = auto() + TOTAL_REQUESTS = auto() + TOTAL_TOKENS = auto() @classmethod def from_str(cls, method: str): @@ -86,63 +85,31 @@ class LoadBalanceMethod(Enum): class DPBudget: 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() + self.total_requests = [0] * dp_size + self.total_tokens = [0] * dp_size def update_budget(self, load_update: WatchLoadUpdateReq): - """Update the budget queue.""" - # Update budget queue together for load updating from the same round. + """Update the budget.""" 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 + self.total_requests[load.dp_rank] = load.num_reqs + self.total_tokens[load.dp_rank] = load.num_tokens - 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 - - max_num_reqs = max(num_reqs) - if all(x == max_num_reqs for x in num_reqs): - return - - while any(x != num_reqs[0] for x in num_reqs): - min_load = min(num_reqs) - 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( - [dp_rank for dp_rank in min_indices] * (second_min_load - min_load) + def dispatch(self, method: LoadBalanceMethod): + if method == LoadBalanceMethod.TOTAL_REQUESTS: + target_rank = self.total_requests.index(min(self.total_requests)) + elif method == LoadBalanceMethod.TOTAL_TOKENS: + # Use total_requests as a tie-breaker when total_tokens are equal + target_rank = min( + range(self.dp_size), + key=lambda i: (self.total_tokens[i], self.total_requests[i]), ) - for idx in min_indices: - num_reqs[idx] = second_min_load + else: + return None - def dispatch(self): - if not self.budget_queue: - self.budget_queue.extend(range(self.dp_size)) - - return self.budget_queue.popleft() + # Increment the load of that worker by one as a heuristic + self.total_requests[target_rank] += 1 + return target_rank class DataParallelController: @@ -177,8 +144,8 @@ class DataParallelController: dispatch_lookup = { LoadBalanceMethod.ROUND_ROBIN: self.round_robin_scheduler, LoadBalanceMethod.FOLLOW_BOOTSTRAP_ROOM: self.follow_bootstrap_room_scheduler, - LoadBalanceMethod.SHORTEST_QUEUE: self.shortest_queue_scheduler, - LoadBalanceMethod.MINIMUM_TOKENS: self.minimum_tokens_scheduler, + LoadBalanceMethod.TOTAL_REQUESTS: self.total_requests_scheduler, + LoadBalanceMethod.TOTAL_TOKENS: self.total_tokens_scheduler, } self.dispatching = dispatch_lookup[self.load_balance_method] @@ -536,30 +503,17 @@ class DataParallelController: target_rank = req.bootstrap_room % len(self.workers) self.workers[target_rank].send_pyobj(req) - def shortest_queue_scheduler(self, req): + def total_requests_scheduler(self, req: Req): if self.maybe_external_dp_rank_routing(req): return - target_worker = self.dp_budget.dispatch() - if target_worker is None: - if self.server_args.disaggregation_mode == "null": - self.round_robin_scheduler(req) - else: - self.follow_bootstrap_room_scheduler(req) - else: - self.workers[target_worker].send_pyobj(req) + target_worker = self.dp_budget.dispatch(LoadBalanceMethod.TOTAL_REQUESTS) + self.workers[target_worker].send_pyobj(req) - def minimum_tokens_scheduler(self, req): + def total_tokens_scheduler(self, req: Req): if self.maybe_external_dp_rank_routing(req): return - - logger.warning( - "The 'minimum_tokens' load balancing method is deprecated for now and will introduced later." - "Fall back to 'round_robin_scheduler'" - ) - if self.server_args.disaggregation_mode == "null": - self.round_robin_scheduler(req) - else: - self.follow_bootstrap_room_scheduler(req) + target_worker = self.dp_budget.dispatch(LoadBalanceMethod.TOTAL_TOKENS) + self.workers[target_worker].send_pyobj(req) def event_loop(self): while True: diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 4451a08dd..6802f2011 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -3327,8 +3327,8 @@ class ServerArgs: "auto", "round_robin", "follow_bootstrap_room", - "shortest_queue", - "minimum_tokens", + "total_requests", + "total_tokens", ], ) parser.add_argument(