Add multi-detokenizer support (#24944)
Co-authored-by: ybyang <10629930+whybeyoung@users.noreply.github.com> Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
co-authored by
ybyang
Shangming Cai
parent
afc7c9f7f3
commit
5ba69f50fb
@@ -27,6 +27,7 @@ import multiprocessing as mp
|
||||
import os
|
||||
import random
|
||||
import signal
|
||||
import tempfile
|
||||
import threading
|
||||
import time
|
||||
from typing import (
|
||||
@@ -80,7 +81,10 @@ from sglang.srt.managers.io_struct import (
|
||||
UpdateWeightsFromIPCReqInput,
|
||||
UpdateWeightsFromTensorReqInput,
|
||||
)
|
||||
from sglang.srt.managers.multi_tokenizer_mixin import MultiTokenizerRouter
|
||||
from sglang.srt.managers.multi_tokenizer_mixin import (
|
||||
MultiTokenizerRouter,
|
||||
run_multi_detokenizer_router_process,
|
||||
)
|
||||
from sglang.srt.managers.scheduler import run_scheduler_process
|
||||
from sglang.srt.managers.template_detection import resolve_auto_parsers
|
||||
from sglang.srt.managers.template_manager import TemplateManager
|
||||
@@ -675,6 +679,63 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
scheduler_procs,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _launch_detokenizer_subprocesses(
|
||||
cls,
|
||||
server_args: ServerArgs,
|
||||
port_args: PortArgs,
|
||||
run_detokenizer_process_func: Callable,
|
||||
) -> Tuple[List[mp.Process], List[str]]:
|
||||
"""Launch detokenizer worker(s).
|
||||
|
||||
- When ``detokenizer_worker_num == 1``: a single detokenizer process listens on
|
||||
``port_args.detokenizer_ipc_name`` (the original behavior).
|
||||
- When ``detokenizer_worker_num > 1``: each detokenizer worker gets its own
|
||||
private IPC socket, and a ``MultiDetokenizerRouter`` process owns the
|
||||
original ``port_args.detokenizer_ipc_name`` and fans out to them.
|
||||
|
||||
Returns (processes, names) for SubprocessWatchdog.
|
||||
"""
|
||||
processes: List[mp.Process] = []
|
||||
names: List[str] = []
|
||||
|
||||
if server_args.detokenizer_worker_num <= 1:
|
||||
proc = mp.Process(
|
||||
target=run_detokenizer_process_func,
|
||||
args=(server_args, port_args),
|
||||
)
|
||||
proc.start()
|
||||
processes.append(proc)
|
||||
names.append("detokenizer")
|
||||
return processes, names
|
||||
|
||||
router_ipc_name = port_args.detokenizer_ipc_name
|
||||
worker_ipc_names: List[str] = []
|
||||
try:
|
||||
for i in range(server_args.detokenizer_worker_num):
|
||||
worker_ipc = f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}"
|
||||
port_args.detokenizer_ipc_name = worker_ipc
|
||||
proc = mp.Process(
|
||||
target=run_detokenizer_process_func,
|
||||
args=(server_args, port_args),
|
||||
)
|
||||
proc.start()
|
||||
processes.append(proc)
|
||||
names.append(f"detokenizer_{i}")
|
||||
worker_ipc_names.append(worker_ipc)
|
||||
finally:
|
||||
port_args.detokenizer_ipc_name = router_ipc_name
|
||||
|
||||
router_proc = mp.Process(
|
||||
target=run_multi_detokenizer_router_process,
|
||||
args=(worker_ipc_names, server_args, port_args),
|
||||
)
|
||||
router_proc.start()
|
||||
processes.append(router_proc)
|
||||
names.append("detokenizer_router")
|
||||
|
||||
return processes, names
|
||||
|
||||
@classmethod
|
||||
def _launch_subprocesses(
|
||||
cls,
|
||||
@@ -776,16 +837,15 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
None,
|
||||
)
|
||||
|
||||
# Launch detokenizer process
|
||||
detoken_proc = mp.Process(
|
||||
target=run_detokenizer_process_func,
|
||||
args=(
|
||||
server_args,
|
||||
port_args,
|
||||
),
|
||||
# Launch detokenizer process(es) — optionally fronted by a router when
|
||||
# detokenizer_worker_num > 1.
|
||||
detoken_procs, detoken_names = cls._launch_detokenizer_subprocesses(
|
||||
server_args=server_args,
|
||||
port_args=port_args,
|
||||
run_detokenizer_process_func=run_detokenizer_process_func,
|
||||
)
|
||||
detoken_proc.start()
|
||||
scheduler_init_result.all_child_pids.append(detoken_proc.pid)
|
||||
for p in detoken_procs:
|
||||
scheduler_init_result.all_child_pids.append(p.pid)
|
||||
|
||||
# Init tokenizer manager first, as the bootstrap server is initialized here
|
||||
if server_args.tokenizer_worker_num == 1:
|
||||
@@ -809,8 +869,8 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
# Note: RayEngine returns scheduler_procs=None as it uses Ray actors instead of mp.Process
|
||||
processes = list(scheduler_procs or [])
|
||||
names = [f"scheduler_{i}" for i in range(len(processes))]
|
||||
processes.append(detoken_proc)
|
||||
names.append("detokenizer")
|
||||
processes.extend(detoken_procs)
|
||||
names.extend(detoken_names)
|
||||
subprocess_watchdog = SubprocessWatchdog(
|
||||
processes=processes, process_names=names
|
||||
)
|
||||
|
||||
@@ -81,7 +81,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
||||
port_args: PortArgs,
|
||||
):
|
||||
# Init inter-process communication
|
||||
self.init_ipc_channels(port_args)
|
||||
self.init_ipc_channels(port_args, server_args)
|
||||
|
||||
# Init tokenizer
|
||||
self.init_tokenizer(server_args)
|
||||
@@ -92,14 +92,18 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
||||
# Init dispatcher
|
||||
self.init_request_dispatcher()
|
||||
|
||||
def init_ipc_channels(self, port_args: PortArgs):
|
||||
def init_ipc_channels(self, port_args: PortArgs, server_args: ServerArgs):
|
||||
context = zmq.Context(2)
|
||||
self.recv_from_scheduler = get_zmq_socket(
|
||||
context, zmq.PULL, port_args.detokenizer_ipc_name, True
|
||||
)
|
||||
self.send_to_tokenizer = get_zmq_socket(
|
||||
context, zmq.PUSH, port_args.tokenizer_ipc_name, False
|
||||
)
|
||||
# In multi-tokenizer mode, results are pushed back to each TokenizerWorker
|
||||
# directly via SocketMapping inside multi_http_worker_event_loop, so the
|
||||
# single send_to_tokenizer socket is unused.
|
||||
if server_args.tokenizer_worker_num == 1:
|
||||
self.send_to_tokenizer = get_zmq_socket(
|
||||
context, zmq.PUSH, port_args.tokenizer_ipc_name, False
|
||||
)
|
||||
|
||||
def init_tokenizer(self, server_args: ServerArgs):
|
||||
if server_args.skip_tokenizer_init:
|
||||
|
||||
@@ -24,11 +24,14 @@ import logging
|
||||
import multiprocessing as multiprocessing
|
||||
import os
|
||||
import pickle
|
||||
import signal
|
||||
import sys
|
||||
import threading
|
||||
import zlib
|
||||
from multiprocessing import shared_memory
|
||||
from typing import TYPE_CHECKING, Any, Dict, Optional, Union
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||
|
||||
import psutil
|
||||
import setproctitle
|
||||
import zmq
|
||||
import zmq.asyncio
|
||||
@@ -43,13 +46,18 @@ from sglang.srt.managers.io_struct import (
|
||||
BatchStrOutput,
|
||||
BatchTokenIDOutput,
|
||||
ContinueGenerationReqInput,
|
||||
FreezeGCReq,
|
||||
PauseContinueBroadcast,
|
||||
PauseGenerationReqInput,
|
||||
TokenizerWorkerRegistration,
|
||||
)
|
||||
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||
from sglang.srt.utils import kill_process_tree
|
||||
from sglang.srt.utils import (
|
||||
configure_logger,
|
||||
kill_itself_when_parent_died,
|
||||
kill_process_tree,
|
||||
)
|
||||
from sglang.srt.utils.network import get_zmq_socket
|
||||
from sglang.utils import get_exception_traceback
|
||||
|
||||
@@ -78,14 +86,14 @@ class SocketMapping:
|
||||
socket = get_zmq_socket(self._zmq_context, zmq.PUSH, ipc_name, False)
|
||||
self._mapping[ipc_name] = socket
|
||||
|
||||
def send_output(self, ipc_name: str, output: Any):
|
||||
def send_output(self, ipc_name: str, output: Any, is_tokenizer: bool = False):
|
||||
if ipc_name is None:
|
||||
# Some unhandled cases
|
||||
logger.warning(f"IPC name is None, output type={type(output)}, skipping...")
|
||||
return
|
||||
|
||||
if ipc_name not in self._mapping:
|
||||
self._register_ipc_mapping(ipc_name, is_tokenizer=False)
|
||||
self._register_ipc_mapping(ipc_name, is_tokenizer=is_tokenizer)
|
||||
self._mapping[ipc_name].send_pyobj(output)
|
||||
|
||||
|
||||
@@ -110,9 +118,7 @@ def _extract_field_by_index(
|
||||
if isinstance(field, dict):
|
||||
new_field = {}
|
||||
for k, v in field.items():
|
||||
if len(v) <= index:
|
||||
new_field[k] = None
|
||||
new_field[k] = v[index]
|
||||
new_field[k] = v[index] if len(v) > index else None
|
||||
return new_field
|
||||
|
||||
if check_length:
|
||||
@@ -196,11 +202,22 @@ def _handle_output_by_index(output, i):
|
||||
output_hidden_states=_extract_field_by_index(
|
||||
output, "output_hidden_states", i, check_length=False
|
||||
),
|
||||
routed_experts=_extract_field_by_index(
|
||||
output, "routed_experts", i, check_length=False
|
||||
),
|
||||
indexer_topk=_extract_field_by_index(
|
||||
output, "indexer_topk", i, check_length=False
|
||||
),
|
||||
retraction_counts=_extract_field_by_index(output, "retraction_counts", i),
|
||||
placeholder_tokens_idx=None,
|
||||
placeholder_tokens_val=None,
|
||||
token_steps=_extract_field_by_index(
|
||||
output, "token_steps", i, check_length=False
|
||||
),
|
||||
customized_info=_extract_field_by_index(
|
||||
output, "customized_info", i, check_length=False
|
||||
),
|
||||
dp_ranks=_extract_field_by_index(output, "dp_ranks", i, check_length=False),
|
||||
)
|
||||
elif isinstance(output, BatchEmbeddingOutput):
|
||||
new_output = BatchEmbeddingOutput(
|
||||
@@ -310,14 +327,23 @@ class MultiHttpWorkerDetokenizerMixin:
|
||||
if output is None:
|
||||
continue
|
||||
|
||||
assert isinstance(
|
||||
recv_obj, BaseBatchReq
|
||||
), "for multi-http-worker, recv_obj must be BaseBatchReq"
|
||||
|
||||
# Send data using the corresponding socket
|
||||
for i, ipc_name in enumerate(recv_obj.http_worker_ipcs):
|
||||
new_output = _handle_output_by_index(output, i)
|
||||
self.socket_mapping.send_output(ipc_name, new_output)
|
||||
# Fan out the output back to the originating tokenizer worker(s).
|
||||
# In multi-detokenizer mode the upstream MultiDetokenizerRouter may
|
||||
# forward either batched or single requests, so handle both shapes.
|
||||
if isinstance(recv_obj, BaseBatchReq):
|
||||
for i, ipc_name in enumerate(recv_obj.http_worker_ipcs):
|
||||
new_output = _handle_output_by_index(output, i)
|
||||
self.socket_mapping.send_output(
|
||||
ipc_name, new_output, is_tokenizer=True
|
||||
)
|
||||
elif isinstance(recv_obj, BaseReq):
|
||||
self.socket_mapping.send_output(
|
||||
recv_obj.http_worker_ipc, output, is_tokenizer=True
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"multi_http_worker_event_loop got unexpected req type {type(recv_obj)}"
|
||||
)
|
||||
|
||||
|
||||
class MultiTokenizerRouter:
|
||||
@@ -415,6 +441,98 @@ class MultiTokenizerRouter:
|
||||
self.socket_mapping.send_output(ipc_name, new_recv_obj)
|
||||
|
||||
|
||||
class MultiDetokenizerRouter:
|
||||
"""Route scheduler outputs to one of N DetokenizerManager workers.
|
||||
|
||||
Each request is pinned to a worker by hashing its ``http_worker_ipc`` with
|
||||
``zlib.crc32`` (deterministic across runs), so all outputs of the same rid
|
||||
always land on the same detokenizer and ``decode_status`` stays consistent.
|
||||
"""
|
||||
|
||||
def __init__(self, ipc_name_list: List[str], port_args: PortArgs):
|
||||
self.ipc_name_list = ipc_name_list
|
||||
self.num_workers = len(ipc_name_list)
|
||||
self.socket_mapping = SocketMapping()
|
||||
context = zmq.Context(2)
|
||||
self.recv_from_scheduler = get_zmq_socket(
|
||||
context, zmq.PULL, port_args.detokenizer_ipc_name, True
|
||||
)
|
||||
|
||||
def _pick(self, key: str) -> str:
|
||||
return self.ipc_name_list[zlib.crc32(key.encode()) % self.num_workers]
|
||||
|
||||
def _send(self, ipc_name: str, obj: Any) -> None:
|
||||
self.socket_mapping.send_output(ipc_name, obj, is_tokenizer=False)
|
||||
|
||||
def event_loop(self):
|
||||
while True:
|
||||
recv_obj = self.recv_from_scheduler.recv_pyobj()
|
||||
|
||||
# FreezeGCReq must freeze every detokenizer process.
|
||||
if isinstance(recv_obj, FreezeGCReq):
|
||||
for ipc in self.ipc_name_list:
|
||||
self._send(ipc, recv_obj)
|
||||
continue
|
||||
|
||||
# Single request: route by its own http_worker_ipc.
|
||||
if isinstance(recv_obj, BaseReq):
|
||||
assert (
|
||||
recv_obj.http_worker_ipc is not None
|
||||
), f"Single req {recv_obj.rid=} missing http_worker_ipc"
|
||||
self._send(self._pick(recv_obj.http_worker_ipc), recv_obj)
|
||||
continue
|
||||
|
||||
# Batch request.
|
||||
if isinstance(recv_obj, BaseBatchReq):
|
||||
# Idle/no-op batch (rids=[]): broadcast to all detokenizers
|
||||
if not recv_obj.rids:
|
||||
for ipc in self.ipc_name_list:
|
||||
self._send(ipc, recv_obj)
|
||||
continue
|
||||
|
||||
ipcs = recv_obj.http_worker_ipcs
|
||||
assert (
|
||||
ipcs is not None
|
||||
and len(ipcs) == len(recv_obj.rids)
|
||||
and all(x is not None for x in ipcs)
|
||||
), f"Batch req {recv_obj.rids=} has invalid http_worker_ipcs"
|
||||
|
||||
# Split per-item and route each by its own ipc.
|
||||
for i, ipc_key in enumerate(ipcs):
|
||||
one = _handle_output_by_index(recv_obj, i)
|
||||
if one is recv_obj:
|
||||
raise TypeError(f"Cannot split {type(recv_obj)}")
|
||||
one.http_worker_ipcs = [ipc_key]
|
||||
self._send(self._pick(ipc_key), one)
|
||||
continue
|
||||
|
||||
raise ValueError(
|
||||
f"MultiDetokenizerRouter got unsupported type {type(recv_obj)}"
|
||||
)
|
||||
|
||||
|
||||
def run_multi_detokenizer_router_process(
|
||||
ipc_name_list: List[str],
|
||||
server_args: ServerArgs,
|
||||
port_args: PortArgs,
|
||||
):
|
||||
kill_itself_when_parent_died()
|
||||
setproctitle.setproctitle("sglang::detokenizer_router")
|
||||
configure_logger(server_args)
|
||||
parent_process = psutil.Process().parent()
|
||||
|
||||
router = None
|
||||
try:
|
||||
router = MultiDetokenizerRouter(ipc_name_list, port_args)
|
||||
router.event_loop()
|
||||
except Exception:
|
||||
traceback = get_exception_traceback()
|
||||
logger.error(f"MultiDetokenizerRouter hit an exception: {traceback}")
|
||||
if router is not None:
|
||||
router.socket_mapping.clear_all_sockets()
|
||||
parent_process.send_signal(signal.SIGQUIT)
|
||||
|
||||
|
||||
class TokenizerWorker(TokenizerManager):
|
||||
"""Tokenizer Worker in multi-http-worker mode"""
|
||||
|
||||
|
||||
@@ -369,6 +369,7 @@ class ServerArgs:
|
||||
tokenizer_mode: str = "auto"
|
||||
tokenizer_backend: str = "huggingface"
|
||||
tokenizer_worker_num: int = 1
|
||||
detokenizer_worker_num: int = 1
|
||||
skip_tokenizer_init: bool = False
|
||||
load_format: str = "auto"
|
||||
model_loader_extra_config: str = "{}"
|
||||
@@ -4186,6 +4187,12 @@ class ServerArgs:
|
||||
f"(requested {self.tokenizer_worker_num})."
|
||||
)
|
||||
self.tokenizer_worker_num = 1
|
||||
if self.detokenizer_worker_num != 1:
|
||||
logger.warning(
|
||||
"skip_tokenizer_init=True disables detokenizer workers; forcing detokenizer_worker_num=1 "
|
||||
f"(requested {self.detokenizer_worker_num})."
|
||||
)
|
||||
self.detokenizer_worker_num = 1
|
||||
|
||||
if self.enable_tokenizer_batch_encode:
|
||||
logger.warning(
|
||||
@@ -4529,6 +4536,12 @@ class ServerArgs:
|
||||
default=ServerArgs.tokenizer_worker_num,
|
||||
help="The worker num of the tokenizer manager.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--detokenizer-worker-num",
|
||||
type=int,
|
||||
default=ServerArgs.detokenizer_worker_num,
|
||||
help="The worker num of the detokenizer manager.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--skip-tokenizer-init",
|
||||
action="store_true",
|
||||
@@ -7275,6 +7288,7 @@ class ServerArgs:
|
||||
)
|
||||
|
||||
assert self.tokenizer_worker_num > 0, "Tokenizer worker num must >= 1"
|
||||
assert self.detokenizer_worker_num > 0, "Detokenizer worker num must >= 1"
|
||||
self.validate_buckets_rule(
|
||||
"--prompt-tokens-buckets", self.prompt_tokens_buckets
|
||||
)
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
import unittest
|
||||
|
||||
from sglang.srt.utils import kill_process_tree
|
||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||
from sglang.test.kits.eval_accuracy_kit import MMLUMixin
|
||||
from sglang.test.test_utils import (
|
||||
DEFAULT_MODEL_NAME_FOR_TEST,
|
||||
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
DEFAULT_URL_FOR_TEST,
|
||||
CustomTestCase,
|
||||
auto_config_device,
|
||||
get_benchmark_args,
|
||||
is_in_amd_ci,
|
||||
is_in_ci,
|
||||
popen_launch_server,
|
||||
run_benchmark,
|
||||
write_github_step_summary,
|
||||
)
|
||||
|
||||
register_cuda_ci(est_time=211, suite="stage-b-test-1-gpu-large")
|
||||
register_amd_ci(est_time=345, suite="stage-b-test-1-gpu-small-amd")
|
||||
|
||||
|
||||
class TestMultiDetokenizer(CustomTestCase, MMLUMixin):
|
||||
mmlu_score_threshold = 0.65
|
||||
mmlu_num_examples = 64
|
||||
mmlu_num_threads = 32
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
cls.model = DEFAULT_MODEL_NAME_FOR_TEST
|
||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||
cls.process = popen_launch_server(
|
||||
cls.model,
|
||||
cls.base_url,
|
||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
other_args=[
|
||||
"--tokenizer-worker-num",
|
||||
8,
|
||||
"--detokenizer-worker-num",
|
||||
4,
|
||||
"--mem-fraction-static",
|
||||
0.7,
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
kill_process_tree(cls.process.pid)
|
||||
|
||||
def test_multi_detokenizer_ttft(self):
|
||||
args = get_benchmark_args(
|
||||
base_url=self.base_url,
|
||||
dataset_name="random",
|
||||
dataset_path="",
|
||||
tokenizer=None,
|
||||
num_prompts=100,
|
||||
random_input_len=4096,
|
||||
random_output_len=2048,
|
||||
sharegpt_context_len=None,
|
||||
request_rate=1,
|
||||
disable_stream=False,
|
||||
disable_ignore_eos=False,
|
||||
seed=0,
|
||||
device=auto_config_device(),
|
||||
lora_name=None,
|
||||
)
|
||||
res = run_benchmark(args)
|
||||
if is_in_ci():
|
||||
write_github_step_summary(
|
||||
f"### test_multi_detokenizer_ttft\n"
|
||||
f"median_e2e_latency_ms: {res['median_e2e_latency_ms']:.2f} ms\n"
|
||||
)
|
||||
self.assertLess(res["median_e2e_latency_ms"], 11000)
|
||||
self.assertLess(res["median_ttft_ms"], 130 if is_in_amd_ci() else 86)
|
||||
self.assertLess(res["median_itl_ms"], 10)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user