[4/N]DP refactor: support watching mode get_load and shortest queue strategy (#10201)
This commit is contained in:
@@ -21,6 +21,7 @@ import struct
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from collections import deque
|
||||
from enum import Enum, auto
|
||||
from multiprocessing import shared_memory
|
||||
from typing import Dict, List
|
||||
@@ -34,6 +35,7 @@ from sglang.srt.managers.io_struct import (
|
||||
BlockReqInput,
|
||||
TokenizedEmbeddingReqInput,
|
||||
TokenizedGenerateReqInput,
|
||||
WatchLoadUpdateReq,
|
||||
)
|
||||
from sglang.srt.managers.schedule_batch import Req
|
||||
from sglang.srt.managers.scheduler import run_scheduler_process
|
||||
@@ -46,7 +48,7 @@ from sglang.srt.utils import (
|
||||
get_zmq_socket,
|
||||
kill_itself_when_parent_died,
|
||||
)
|
||||
from sglang.utils import get_exception_traceback
|
||||
from sglang.utils import TypeBasedDispatcher, get_exception_traceback
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -67,6 +69,42 @@ class LoadBalanceMethod(Enum):
|
||||
raise ValueError(f"Invalid load balance method: {method}") from exc
|
||||
|
||||
|
||||
class DPBudget:
|
||||
def __init__(self):
|
||||
# TODO: support minimum tokens method
|
||||
self.budget_queue = deque()
|
||||
|
||||
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()
|
||||
|
||||
num_reqs = [load.num_reqs for load in loads]
|
||||
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 = [i for i, 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)
|
||||
)
|
||||
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
|
||||
|
||||
|
||||
class DataParallelController:
|
||||
"""A controller that dispatches requests to multiple data parallel workers."""
|
||||
|
||||
@@ -104,6 +142,9 @@ class DataParallelController:
|
||||
}
|
||||
self.dispatching = dispatch_lookup[self.load_balance_method]
|
||||
|
||||
# Load balance budget
|
||||
self.dp_budget = DPBudget()
|
||||
|
||||
# Launch data parallel workers
|
||||
self.scheduler_procs = []
|
||||
self.workers: List[zmq.Socket] = [None] * server_args.dp_size
|
||||
@@ -127,6 +168,31 @@ class DataParallelController:
|
||||
|
||||
self.max_req_input_len = None
|
||||
|
||||
self.init_dispatcher()
|
||||
|
||||
def send_to_all_workers(self, obj):
|
||||
for worker in self.workers:
|
||||
worker.send_pyobj(obj)
|
||||
|
||||
def send_control_message(self, obj):
|
||||
# Send control messages to first worker of tp group
|
||||
for worker in self.workers[:: self.control_message_step]:
|
||||
worker.send_pyobj(obj)
|
||||
|
||||
def handle_load_update_req(self, obj):
|
||||
self.dp_budget.update_budget(obj)
|
||||
|
||||
def init_dispatcher(self):
|
||||
self._request_dispatcher = TypeBasedDispatcher(
|
||||
[
|
||||
(TokenizedGenerateReqInput, self.dispatching),
|
||||
(TokenizedEmbeddingReqInput, self.dispatching),
|
||||
(BlockReqInput, self.send_to_all_workers),
|
||||
(WatchLoadUpdateReq, self.handle_load_update_req),
|
||||
]
|
||||
)
|
||||
self._request_dispatcher.add_fallback_fn(self.send_control_message)
|
||||
|
||||
def launch_dp_schedulers(self, server_args, port_args):
|
||||
base_gpu_id = 0
|
||||
|
||||
@@ -291,10 +357,14 @@ class DataParallelController:
|
||||
else:
|
||||
self.workers[req.bootstrap_room % len(self.workers)].send_pyobj(req)
|
||||
|
||||
def shortest_queue_scheduler(self, input_requests):
|
||||
def shortest_queue_scheduler(self, req):
|
||||
if self.maybe_external_dp_rank_routing(req):
|
||||
return
|
||||
raise NotImplementedError()
|
||||
target_worker = self.dp_budget.dispatch()
|
||||
if target_worker is None:
|
||||
self.round_robin_scheduler(req)
|
||||
else:
|
||||
self.workers[target_worker].send_pyobj(req)
|
||||
|
||||
def minimum_tokens_scheduler(self, req):
|
||||
if self.maybe_external_dp_rank_routing(req):
|
||||
@@ -333,22 +403,7 @@ class DataParallelController:
|
||||
recv_req = self.recv_from_tokenizer.recv_pyobj(zmq.NOBLOCK)
|
||||
except zmq.ZMQError:
|
||||
break
|
||||
|
||||
if isinstance(
|
||||
recv_req,
|
||||
(
|
||||
TokenizedGenerateReqInput,
|
||||
TokenizedEmbeddingReqInput,
|
||||
),
|
||||
):
|
||||
self.dispatching(recv_req)
|
||||
elif isinstance(recv_req, BlockReqInput):
|
||||
for worker in self.workers:
|
||||
worker.send_pyobj(recv_req)
|
||||
else:
|
||||
# Send other control messages to first worker of tp group
|
||||
for worker in self.workers[:: self.control_message_step]:
|
||||
worker.send_pyobj(recv_req)
|
||||
self._request_dispatcher(recv_req)
|
||||
|
||||
|
||||
def run_data_parallel_controller_process(
|
||||
|
||||
Reference in New Issue
Block a user