DP: support piggyback server load report (#11469)
Signed-off-by: Chang Huaixin (OpenAnolis) <changhuaixin@linux.alibaba.com>
This commit is contained in:
@@ -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")
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user