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:
Yuhao Yang
2026-05-15 17:26:26 -07:00
committed by GitHub
co-authored by ybyang Shangming Cai
parent afc7c9f7f3
commit 5ba69f50fb
5 changed files with 308 additions and 32 deletions
+72 -12
View File
@@ -27,6 +27,7 @@ import multiprocessing as mp
import os import os
import random import random
import signal import signal
import tempfile
import threading import threading
import time import time
from typing import ( from typing import (
@@ -80,7 +81,10 @@ from sglang.srt.managers.io_struct import (
UpdateWeightsFromIPCReqInput, UpdateWeightsFromIPCReqInput,
UpdateWeightsFromTensorReqInput, 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.scheduler import run_scheduler_process
from sglang.srt.managers.template_detection import resolve_auto_parsers from sglang.srt.managers.template_detection import resolve_auto_parsers
from sglang.srt.managers.template_manager import TemplateManager from sglang.srt.managers.template_manager import TemplateManager
@@ -675,6 +679,63 @@ class Engine(EngineScoreMixin, EngineBase):
scheduler_procs, 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 @classmethod
def _launch_subprocesses( def _launch_subprocesses(
cls, cls,
@@ -776,16 +837,15 @@ class Engine(EngineScoreMixin, EngineBase):
None, None,
) )
# Launch detokenizer process # Launch detokenizer process(es) — optionally fronted by a router when
detoken_proc = mp.Process( # detokenizer_worker_num > 1.
target=run_detokenizer_process_func, detoken_procs, detoken_names = cls._launch_detokenizer_subprocesses(
args=( server_args=server_args,
server_args, port_args=port_args,
port_args, run_detokenizer_process_func=run_detokenizer_process_func,
),
) )
detoken_proc.start() for p in detoken_procs:
scheduler_init_result.all_child_pids.append(detoken_proc.pid) scheduler_init_result.all_child_pids.append(p.pid)
# Init tokenizer manager first, as the bootstrap server is initialized here # Init tokenizer manager first, as the bootstrap server is initialized here
if server_args.tokenizer_worker_num == 1: 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 # Note: RayEngine returns scheduler_procs=None as it uses Ray actors instead of mp.Process
processes = list(scheduler_procs or []) processes = list(scheduler_procs or [])
names = [f"scheduler_{i}" for i in range(len(processes))] names = [f"scheduler_{i}" for i in range(len(processes))]
processes.append(detoken_proc) processes.extend(detoken_procs)
names.append("detokenizer") names.extend(detoken_names)
subprocess_watchdog = SubprocessWatchdog( subprocess_watchdog = SubprocessWatchdog(
processes=processes, process_names=names processes=processes, process_names=names
) )
@@ -81,7 +81,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
port_args: PortArgs, port_args: PortArgs,
): ):
# Init inter-process communication # Init inter-process communication
self.init_ipc_channels(port_args) self.init_ipc_channels(port_args, server_args)
# Init tokenizer # Init tokenizer
self.init_tokenizer(server_args) self.init_tokenizer(server_args)
@@ -92,14 +92,18 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
# Init dispatcher # Init dispatcher
self.init_request_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) context = zmq.Context(2)
self.recv_from_scheduler = get_zmq_socket( self.recv_from_scheduler = get_zmq_socket(
context, zmq.PULL, port_args.detokenizer_ipc_name, True context, zmq.PULL, port_args.detokenizer_ipc_name, True
) )
self.send_to_tokenizer = get_zmq_socket( # In multi-tokenizer mode, results are pushed back to each TokenizerWorker
context, zmq.PUSH, port_args.tokenizer_ipc_name, False # 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): def init_tokenizer(self, server_args: ServerArgs):
if server_args.skip_tokenizer_init: if server_args.skip_tokenizer_init:
@@ -24,11 +24,14 @@ import logging
import multiprocessing as multiprocessing import multiprocessing as multiprocessing
import os import os
import pickle import pickle
import signal
import sys import sys
import threading import threading
import zlib
from multiprocessing import shared_memory 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 setproctitle
import zmq import zmq
import zmq.asyncio import zmq.asyncio
@@ -43,13 +46,18 @@ from sglang.srt.managers.io_struct import (
BatchStrOutput, BatchStrOutput,
BatchTokenIDOutput, BatchTokenIDOutput,
ContinueGenerationReqInput, ContinueGenerationReqInput,
FreezeGCReq,
PauseContinueBroadcast, PauseContinueBroadcast,
PauseGenerationReqInput, PauseGenerationReqInput,
TokenizerWorkerRegistration, TokenizerWorkerRegistration,
) )
from sglang.srt.managers.tokenizer_manager import TokenizerManager from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.server_args import PortArgs, ServerArgs 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.srt.utils.network import get_zmq_socket
from sglang.utils import get_exception_traceback 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) socket = get_zmq_socket(self._zmq_context, zmq.PUSH, ipc_name, False)
self._mapping[ipc_name] = socket 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: if ipc_name is None:
# Some unhandled cases # Some unhandled cases
logger.warning(f"IPC name is None, output type={type(output)}, skipping...") logger.warning(f"IPC name is None, output type={type(output)}, skipping...")
return return
if ipc_name not in self._mapping: 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) self._mapping[ipc_name].send_pyobj(output)
@@ -110,9 +118,7 @@ def _extract_field_by_index(
if isinstance(field, dict): if isinstance(field, dict):
new_field = {} new_field = {}
for k, v in field.items(): for k, v in field.items():
if len(v) <= index: new_field[k] = v[index] if len(v) > index else None
new_field[k] = None
new_field[k] = v[index]
return new_field return new_field
if check_length: if check_length:
@@ -196,11 +202,22 @@ def _handle_output_by_index(output, i):
output_hidden_states=_extract_field_by_index( output_hidden_states=_extract_field_by_index(
output, "output_hidden_states", i, check_length=False 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_idx=None,
placeholder_tokens_val=None, placeholder_tokens_val=None,
token_steps=_extract_field_by_index( token_steps=_extract_field_by_index(
output, "token_steps", i, check_length=False 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): elif isinstance(output, BatchEmbeddingOutput):
new_output = BatchEmbeddingOutput( new_output = BatchEmbeddingOutput(
@@ -310,14 +327,23 @@ class MultiHttpWorkerDetokenizerMixin:
if output is None: if output is None:
continue continue
assert isinstance( # Fan out the output back to the originating tokenizer worker(s).
recv_obj, BaseBatchReq # In multi-detokenizer mode the upstream MultiDetokenizerRouter may
), "for multi-http-worker, recv_obj must be BaseBatchReq" # forward either batched or single requests, so handle both shapes.
if isinstance(recv_obj, BaseBatchReq):
# Send data using the corresponding socket for i, ipc_name in enumerate(recv_obj.http_worker_ipcs):
for i, ipc_name in enumerate(recv_obj.http_worker_ipcs): new_output = _handle_output_by_index(output, i)
new_output = _handle_output_by_index(output, i) self.socket_mapping.send_output(
self.socket_mapping.send_output(ipc_name, new_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: class MultiTokenizerRouter:
@@ -415,6 +441,98 @@ class MultiTokenizerRouter:
self.socket_mapping.send_output(ipc_name, new_recv_obj) 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): class TokenizerWorker(TokenizerManager):
"""Tokenizer Worker in multi-http-worker mode""" """Tokenizer Worker in multi-http-worker mode"""
+14
View File
@@ -369,6 +369,7 @@ class ServerArgs:
tokenizer_mode: str = "auto" tokenizer_mode: str = "auto"
tokenizer_backend: str = "huggingface" tokenizer_backend: str = "huggingface"
tokenizer_worker_num: int = 1 tokenizer_worker_num: int = 1
detokenizer_worker_num: int = 1
skip_tokenizer_init: bool = False skip_tokenizer_init: bool = False
load_format: str = "auto" load_format: str = "auto"
model_loader_extra_config: str = "{}" model_loader_extra_config: str = "{}"
@@ -4186,6 +4187,12 @@ class ServerArgs:
f"(requested {self.tokenizer_worker_num})." f"(requested {self.tokenizer_worker_num})."
) )
self.tokenizer_worker_num = 1 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: if self.enable_tokenizer_batch_encode:
logger.warning( logger.warning(
@@ -4529,6 +4536,12 @@ class ServerArgs:
default=ServerArgs.tokenizer_worker_num, default=ServerArgs.tokenizer_worker_num,
help="The worker num of the tokenizer manager.", 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( parser.add_argument(
"--skip-tokenizer-init", "--skip-tokenizer-init",
action="store_true", action="store_true",
@@ -7275,6 +7288,7 @@ class ServerArgs:
) )
assert self.tokenizer_worker_num > 0, "Tokenizer worker num must >= 1" 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( self.validate_buckets_rule(
"--prompt-tokens-buckets", self.prompt_tokens_buckets "--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()