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 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"""
|
||||||
|
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user