[dpc]: unify DP controller load balancing and simplify dispatch logic (#16258)

This commit is contained in:
Ratish P
2026-01-11 12:38:03 +08:00
committed by GitHub
parent cc25f9df50
commit c0248d6f37
4 changed files with 32 additions and 78 deletions
+1 -1
View File
@@ -205,7 +205,7 @@ Please consult the documentation below and [server_args.py](https://github.com/s
| Argument | Description | Defaults | Options | | Argument | Description | Defaults | Options |
| --- | --- | --- | --- | | --- | --- | --- | --- |
| `--data-parallel-size`<br>`--dp-size` | The data parallelism size. | `1` | Type: int | | `--data-parallel-size`<br>`--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 | | `--load-watch-interval` | The interval of load watching in seconds. | `0.1` | Type: float |
## Multi-node distributed serving ## Multi-node distributed serving
@@ -153,7 +153,7 @@ click [Server Arguments](https://docs.sglang.io/advanced_features/server_argumen
| Argument | Defaults | Options | A2 | A3 | | Argument | Defaults | Options | A2 | A3 |
|----------------------------------------|---------------|-------------------------------------------------------------|:----------------------------------------:|:----------------------------------------:| |----------------------------------------|---------------|-------------------------------------------------------------|:----------------------------------------:|:----------------------------------------:|
| `--data-parallel-size`<br/>`--dp-size` | `1` | Type: int | **<span style="color: green;">√</span>** | **<span style="color: green;">√</span>** | | `--data-parallel-size`<br/>`--dp-size` | `1` | Type: int | **<span style="color: green;">√</span>** | **<span style="color: green;">√</span>** |
| `--load-balance-method` | `round_robin` | `round_robin`,<br/> `shortest_queue`,<br/> `minimum_tokens` | **<span style="color: green;">√</span>** | **<span style="color: green;">√</span>** | | `--load-balance-method` | `round_robin` | `round_robin`,<br/> `total_requests`,<br/> `total_tokens` | **<span style="color: green;">√</span>** | **<span style="color: green;">√</span>** |
| `--prefill-round-robin-balance` | `False` | bool flag<br/> (set to enable) | **<span style="color: green;">√</span>** | **<span style="color: green;">√</span>** | | `--prefill-round-robin-balance` | `False` | bool flag<br/> (set to enable) | **<span style="color: green;">√</span>** | **<span style="color: green;">√</span>** |
## Multi-node distributed serving ## Multi-node distributed serving
@@ -19,7 +19,6 @@ import multiprocessing as mp
import signal import signal
import threading import threading
import time import time
from collections import deque
from enum import Enum, auto from enum import Enum, auto
from typing import Callable, List, Optional from typing import Callable, List, Optional
@@ -72,8 +71,8 @@ class LoadBalanceMethod(Enum):
ROUND_ROBIN = auto() ROUND_ROBIN = auto()
FOLLOW_BOOTSTRAP_ROOM = auto() FOLLOW_BOOTSTRAP_ROOM = auto()
SHORTEST_QUEUE = auto() TOTAL_REQUESTS = auto()
MINIMUM_TOKENS = auto() TOTAL_TOKENS = auto()
@classmethod @classmethod
def from_str(cls, method: str): def from_str(cls, method: str):
@@ -86,63 +85,31 @@ class LoadBalanceMethod(Enum):
class DPBudget: class DPBudget:
def __init__(self, dp_size: int): def __init__(self, dp_size: int):
# TODO: support minimum tokens method
self.budget_queue = deque()
self.dp_size = dp_size self.dp_size = dp_size
self.ts_tic = 0.0 self.total_requests = [0] * dp_size
self.pending_loads = {} self.total_tokens = [0] * dp_size
# 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): def update_budget(self, load_update: WatchLoadUpdateReq):
"""Update the budget queue.""" """Update the budget."""
# Update budget queue together for load updating from the same round.
for load in load_update.loads: for load in load_update.loads:
if abs(load.ts_tic - self.ts_tic) > self.tic_window: self.total_requests[load.dp_rank] = load.num_reqs
logger.debug(f"Proceed to next round: {self.ts_tic=} {load.ts_tic=}") self.total_tokens[load.dp_rank] = load.num_tokens
self.pending_loads.clear()
self.ts_tic = load.ts_tic
self.pending_loads[load.dp_rank] = load
if len(self.pending_loads) < self.dp_size: def dispatch(self, method: LoadBalanceMethod):
logger.debug(f"Waiting for all DP ranks: {len(self.pending_loads)=}") if method == LoadBalanceMethod.TOTAL_REQUESTS:
return target_rank = self.total_requests.index(min(self.total_requests))
elif method == LoadBalanceMethod.TOTAL_TOKENS:
self.update_budget_count = (self.update_budget_count + 1) % self.update_interval # Use total_requests as a tie-breaker when total_tokens are equal
if self.update_budget_count: target_rank = min(
return range(self.dp_size),
key=lambda i: (self.total_tokens[i], self.total_requests[i]),
# 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)
) )
for idx in min_indices: else:
num_reqs[idx] = second_min_load return None
def dispatch(self): # Increment the load of that worker by one as a heuristic
if not self.budget_queue: self.total_requests[target_rank] += 1
self.budget_queue.extend(range(self.dp_size)) return target_rank
return self.budget_queue.popleft()
class DataParallelController: class DataParallelController:
@@ -177,8 +144,8 @@ class DataParallelController:
dispatch_lookup = { dispatch_lookup = {
LoadBalanceMethod.ROUND_ROBIN: self.round_robin_scheduler, LoadBalanceMethod.ROUND_ROBIN: self.round_robin_scheduler,
LoadBalanceMethod.FOLLOW_BOOTSTRAP_ROOM: self.follow_bootstrap_room_scheduler, LoadBalanceMethod.FOLLOW_BOOTSTRAP_ROOM: self.follow_bootstrap_room_scheduler,
LoadBalanceMethod.SHORTEST_QUEUE: self.shortest_queue_scheduler, LoadBalanceMethod.TOTAL_REQUESTS: self.total_requests_scheduler,
LoadBalanceMethod.MINIMUM_TOKENS: self.minimum_tokens_scheduler, LoadBalanceMethod.TOTAL_TOKENS: self.total_tokens_scheduler,
} }
self.dispatching = dispatch_lookup[self.load_balance_method] self.dispatching = dispatch_lookup[self.load_balance_method]
@@ -536,30 +503,17 @@ class DataParallelController:
target_rank = req.bootstrap_room % len(self.workers) target_rank = req.bootstrap_room % len(self.workers)
self.workers[target_rank].send_pyobj(req) 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): if self.maybe_external_dp_rank_routing(req):
return return
target_worker = self.dp_budget.dispatch() target_worker = self.dp_budget.dispatch(LoadBalanceMethod.TOTAL_REQUESTS)
if target_worker is None: self.workers[target_worker].send_pyobj(req)
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)
def minimum_tokens_scheduler(self, req): def total_tokens_scheduler(self, req: Req):
if self.maybe_external_dp_rank_routing(req): if self.maybe_external_dp_rank_routing(req):
return return
target_worker = self.dp_budget.dispatch(LoadBalanceMethod.TOTAL_TOKENS)
logger.warning( self.workers[target_worker].send_pyobj(req)
"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)
def event_loop(self): def event_loop(self):
while True: while True:
+2 -2
View File
@@ -3327,8 +3327,8 @@ class ServerArgs:
"auto", "auto",
"round_robin", "round_robin",
"follow_bootstrap_room", "follow_bootstrap_room",
"shortest_queue", "total_requests",
"minimum_tokens", "total_tokens",
], ],
) )
parser.add_argument( parser.add_argument(