From 5ba69f50fb35d9c121b9cf7343f2d3776c580f3c Mon Sep 17 00:00:00 2001 From: Yuhao Yang <47235274+yhyang201@users.noreply.github.com> Date: Sat, 16 May 2026 08:26:26 +0800 Subject: [PATCH] Add multi-detokenizer support (#24944) Co-authored-by: ybyang <10629930+whybeyoung@users.noreply.github.com> Co-authored-by: Shangming Cai --- python/sglang/srt/entrypoints/engine.py | 84 ++++++++-- .../srt/managers/detokenizer_manager.py | 14 +- .../srt/managers/multi_tokenizer_mixin.py | 148 ++++++++++++++++-- python/sglang/srt/server_args.py | 14 ++ .../tokenizer/test_multi_detokenizer.py | 80 ++++++++++ 5 files changed, 308 insertions(+), 32 deletions(-) create mode 100644 test/registered/tokenizer/test_multi_detokenizer.py diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index f96445c31..1dde8bed8 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -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 ) diff --git a/python/sglang/srt/managers/detokenizer_manager.py b/python/sglang/srt/managers/detokenizer_manager.py index a4547bf36..a35e98167 100644 --- a/python/sglang/srt/managers/detokenizer_manager.py +++ b/python/sglang/srt/managers/detokenizer_manager.py @@ -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: diff --git a/python/sglang/srt/managers/multi_tokenizer_mixin.py b/python/sglang/srt/managers/multi_tokenizer_mixin.py index 658acc01c..baf25d332 100644 --- a/python/sglang/srt/managers/multi_tokenizer_mixin.py +++ b/python/sglang/srt/managers/multi_tokenizer_mixin.py @@ -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""" diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index a1dba0f72..5e48ad36d 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -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 ) diff --git a/test/registered/tokenizer/test_multi_detokenizer.py b/test/registered/tokenizer/test_multi_detokenizer.py new file mode 100644 index 000000000..09ca5e9b1 --- /dev/null +++ b/test/registered/tokenizer/test_multi_detokenizer.py @@ -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()