Improve engine customization interface (#15635)

This commit is contained in:
Lianmin Zheng
2025-12-22 14:24:16 -08:00
committed by GitHub
parent 34013d9d5a
commit 5e1a495c65
6 changed files with 195 additions and 188 deletions
+31 -14
View File
@@ -1,6 +1,7 @@
import atexit import atexit
import json import json
import multiprocessing import multiprocessing
import time
import warnings import warnings
from typing import Dict, List, Optional, Union from typing import Dict, List, Optional, Union
@@ -365,10 +366,18 @@ class Runtime:
def __init__( def __init__(
self, self,
log_level: str = "error", log_level: str = "error",
launch_timeout: float = 300.0,
*args, *args,
**kwargs, **kwargs,
): ):
"""See the arguments in server_args.py::ServerArgs""" """See the arguments in server_args.py::ServerArgs
Args:
log_level: Log level for the server.
timeout: Timeout in seconds for waiting for the server to start.
*args: Additional arguments passed to ServerArgs.
**kwargs: Additional keyword arguments passed to ServerArgs.
"""
# We delay the import of any `sglang.srt` components in `sglang.lang`, so users can run # We delay the import of any `sglang.srt` components in `sglang.lang`, so users can run
# client code without installing SRT server and its dependency if they want. # client code without installing SRT server and its dependency if they want.
from sglang.srt.entrypoints.http_server import launch_server from sglang.srt.entrypoints.http_server import launch_server
@@ -388,31 +397,39 @@ class Runtime:
# NOTE: We store pid instead of proc to fix some issues during __delete__ # NOTE: We store pid instead of proc to fix some issues during __delete__
self.pid = None self.pid = None
pipe_reader, pipe_writer = multiprocessing.Pipe(duplex=False)
ctx = multiprocessing.get_context("spawn") ctx = multiprocessing.get_context("spawn")
proc = ctx.Process( proc = ctx.Process(
target=launch_server, target=launch_server,
args=(self.server_args, pipe_writer), args=(self.server_args,),
) )
proc.start() proc.start()
pipe_writer.close()
self.pid = proc.pid self.pid = proc.pid
# Before python program terminates, call shutdown implicitly. Therefore, users don't have to explicitly call .shutdown() # Before python program terminates, call shutdown implicitly. Therefore, users don't have to explicitly call .shutdown()
atexit.register(self.shutdown) atexit.register(self.shutdown)
# TODO: remove this pipe_writer mechanism and use `/health_generate` instead. # Wait for server to be ready by polling /health_generate
try: start_time = time.time()
init_state = pipe_reader.recv() with requests.Session() as session:
except EOFError: while time.time() - start_time < launch_timeout:
init_state = "" try:
response = session.get(f"{self.url}/health_generate")
if response.status_code == 200:
break
except requests.RequestException:
pass
if init_state != "ready": if not proc.is_alive():
self.shutdown() self.shutdown()
raise RuntimeError( raise RuntimeError(
"Initialization failed. Please see the error messages above." "Initialization failed. Please see the error messages above."
) )
time.sleep(2)
else:
self.shutdown()
raise TimeoutError("Server failed to start within the timeout period.")
self.endpoint = RuntimeEndpoint(self.url) self.endpoint = RuntimeEndpoint(self.url)
+128 -112
View File
@@ -91,94 +91,25 @@ asyncio.set_event_loop_policy(uvloop.EventLoopPolicy())
_is_cuda = is_cuda() _is_cuda = is_cuda()
def _launch_subprocesses( def init_tokenizer_manager(
server_args: ServerArgs, port_args: Optional[PortArgs] = None server_args: ServerArgs,
) -> Tuple[TokenizerManager, TemplateManager, Dict, PortArgs]: port_args: PortArgs,
""" TokenizerManagerClass: Optional[TokenizerManager] = None,
Launch the TokenizerManager in the main process, the Scheduler in a subprocess, and the DetokenizerManager in another subprocess. ) -> Tuple[TokenizerManager, TemplateManager]:
""" # Launch tokenizer process
# Configure global environment TokenizerManagerClass = TokenizerManagerClass or TokenizerManager
configure_logger(server_args) tokenizer_manager = TokenizerManagerClass(server_args, port_args)
_set_envs_and_config(server_args)
server_args.check_server_args()
# Allocate ports for inter-process communications # Initialize templates
if port_args is None: template_manager = TemplateManager()
port_args = PortArgs.init_new(server_args) template_manager.initialize_templates(
logger.info(f"{server_args=}") tokenizer_manager=tokenizer_manager,
model_path=server_args.model_path,
# Launch scheduler processes chat_template=server_args.chat_template,
scheduler_procs, scheduler_pipe_readers = _launch_scheduler_processes( completion_template=server_args.completion_template,
server_args=server_args,
port_args=port_args,
) )
if server_args.node_rank >= 1: return tokenizer_manager, template_manager
# In multi-node cases, non-zero rank nodes do not need to run tokenizer or detokenizer,
# so they can just wait here.
for reader in scheduler_pipe_readers:
data = reader.recv()
assert data["status"] == "ready"
if os.getenv("SGLANG_BLOCK_NONZERO_RANK_CHILDREN") == "0":
# When using `Engine` as a Python API, we don't want to block here.
return None, None, None, port_args
launch_dummy_health_check_server(
server_args.host, server_args.port, server_args.enable_metrics
)
for proc in scheduler_procs:
proc.join()
logger.error(
f"Scheduler or DataParallelController {proc.pid} terminated with {proc.exitcode}"
)
return None, None, None, port_args
# Launch detokenizer process
detoken_proc = mp.Process(
target=run_detokenizer_process,
args=(
server_args,
port_args,
),
)
detoken_proc.start()
# Init tokenizer manager first, as the bootstrap server is initialized here
if server_args.tokenizer_worker_num == 1:
tokenizer_manager, template_manager = _init_tokenizer_manager(
server_args, port_args
)
else:
# Launch multi-tokenizer router
tokenizer_manager = MultiTokenizerRouter(server_args, port_args)
template_manager = None
# Wait for the model to finish loading
scheduler_infos = []
for i in range(len(scheduler_pipe_readers)):
try:
data = scheduler_pipe_readers[i].recv()
except EOFError:
logger.error(
f"Rank {i} scheduler is dead. Please check if there are relevant logs."
)
scheduler_procs[i].join()
logger.error(f"Exit code: {scheduler_procs[i].exitcode}")
raise
if data["status"] != "ready":
raise RuntimeError(
"Initialization failed. Please see the error messages above."
)
scheduler_infos.append(data)
# Get back some info from scheduler to tokenizer_manager
tokenizer_manager.max_req_input_len = scheduler_infos[0]["max_req_input_len"]
return tokenizer_manager, template_manager, scheduler_infos, port_args
class Engine(EngineBase): class Engine(EngineBase):
@@ -197,8 +128,10 @@ class Engine(EngineBase):
# Some fields to allow people to override the server args # Some fields to allow people to override the server args
# and launch processes for their private forks. # and launch processes for their private forks.
launch_subprocesses_func: Callable = staticmethod(_launch_subprocesses)
server_args_class: ServerArgs = ServerArgs server_args_class: ServerArgs = ServerArgs
init_tokenizer_manager_func: Callable = staticmethod(init_tokenizer_manager)
run_scheduler_process_func: Callable = staticmethod(run_scheduler_process)
run_detokenizer_process_func: Callable = staticmethod(run_detokenizer_process)
def __init__(self, **kwargs): def __init__(self, **kwargs):
""" """
@@ -224,7 +157,12 @@ class Engine(EngineBase):
# Launch subprocesses # Launch subprocesses
tokenizer_manager, template_manager, scheduler_infos, port_args = ( tokenizer_manager, template_manager, scheduler_infos, port_args = (
self.launch_subprocesses_func(server_args=server_args) _launch_subprocesses(
server_args=server_args,
init_tokenizer_manager_func=self.init_tokenizer_manager_func,
run_scheduler_process_func=self.run_scheduler_process_func,
run_detokenizer_process_func=self.run_detokenizer_process_func,
)
) )
self.tokenizer_manager = tokenizer_manager self.tokenizer_manager = tokenizer_manager
self.template_manager = template_manager self.template_manager = template_manager
@@ -873,32 +811,10 @@ def _set_envs_and_config(server_args: ServerArgs):
mp.set_start_method("spawn", force=True) mp.set_start_method("spawn", force=True)
def _init_tokenizer_manager(
server_args: ServerArgs,
port_args: PortArgs,
TokenizerManagerClass: Optional[TokenizerManager] = None,
) -> TokenizerManager:
# Launch tokenizer process
TokenizerManagerClass = TokenizerManagerClass or TokenizerManager
tokenizer_manager = TokenizerManagerClass(server_args, port_args)
# Initialize templates
template_manager = TemplateManager()
template_manager.initialize_templates(
tokenizer_manager=tokenizer_manager,
model_path=server_args.model_path,
chat_template=server_args.chat_template,
completion_template=server_args.completion_template,
)
return tokenizer_manager, template_manager
def _launch_scheduler_processes( def _launch_scheduler_processes(
server_args: ServerArgs, server_args: ServerArgs,
port_args: PortArgs, port_args: PortArgs,
run_scheduler_process_func: Callable = run_scheduler_process, run_scheduler_process_func: Callable,
run_data_parallel_controller_process_func: Callable = run_data_parallel_controller_process,
): ):
scheduler_procs = [] scheduler_procs = []
@@ -959,10 +875,110 @@ def _launch_scheduler_processes(
reader, writer = mp.Pipe(duplex=False) reader, writer = mp.Pipe(duplex=False)
scheduler_pipe_readers = [reader] scheduler_pipe_readers = [reader]
proc = mp.Process( proc = mp.Process(
target=run_data_parallel_controller_process_func, target=run_data_parallel_controller_process,
args=(server_args, port_args, writer), kwargs=dict(
server_args=server_args,
port_args=port_args,
pipe_writer=writer,
run_scheduler_process_func=run_scheduler_process_func,
),
) )
proc.start() proc.start()
scheduler_procs.append(proc) scheduler_procs.append(proc)
return scheduler_procs, scheduler_pipe_readers return scheduler_procs, scheduler_pipe_readers
def _launch_subprocesses(
server_args: ServerArgs,
init_tokenizer_manager_func: Callable,
run_scheduler_process_func: Callable,
run_detokenizer_process_func: Callable,
port_args: Optional[PortArgs] = None,
) -> Tuple[TokenizerManager, TemplateManager, Tuple[Dict], PortArgs]:
"""
Launch the TokenizerManager in the main process, the Scheduler in a subprocess, and the DetokenizerManager in another subprocess.
"""
# Configure global environment
configure_logger(server_args)
_set_envs_and_config(server_args)
server_args.check_server_args()
# Allocate ports for inter-process communications
if port_args is None:
port_args = PortArgs.init_new(server_args)
logger.info(f"{server_args=}")
# Launch scheduler processes
scheduler_procs, scheduler_pipe_readers = _launch_scheduler_processes(
server_args=server_args,
port_args=port_args,
run_scheduler_process_func=run_scheduler_process_func,
)
if server_args.node_rank >= 1:
# In multi-node cases, non-zero rank nodes do not need to run tokenizer or detokenizer,
# so they can just wait here.
for reader in scheduler_pipe_readers:
data = reader.recv()
assert data["status"] == "ready"
if os.getenv("SGLANG_BLOCK_NONZERO_RANK_CHILDREN") == "0":
# When using `Engine` as a Python API, we don't want to block here.
return None, None, None, port_args
launch_dummy_health_check_server(
server_args.host, server_args.port, server_args.enable_metrics
)
for proc in scheduler_procs:
proc.join()
logger.error(
f"Scheduler or DataParallelController {proc.pid} terminated with {proc.exitcode}"
)
return None, None, None, port_args
# Launch detokenizer process
detoken_proc = mp.Process(
target=run_detokenizer_process_func,
args=(
server_args,
port_args,
),
)
detoken_proc.start()
# Init tokenizer manager first, as the bootstrap server is initialized here
if server_args.tokenizer_worker_num == 1:
tokenizer_manager, template_manager = init_tokenizer_manager_func(
server_args, port_args
)
else:
# Launch multi-tokenizer router
tokenizer_manager = MultiTokenizerRouter(server_args, port_args)
template_manager = None
# Wait for the model to finish loading
scheduler_infos = []
for i in range(len(scheduler_pipe_readers)):
try:
data = scheduler_pipe_readers[i].recv()
except EOFError:
logger.error(
f"Rank {i} scheduler is dead. Please check if there are relevant logs."
)
scheduler_procs[i].join()
logger.error(f"Exit code: {scheduler_procs[i].exitcode}")
raise
if data["status"] != "ready":
raise RuntimeError(
"Initialization failed. Please see the error messages above."
)
scheduler_infos.append(data)
# Get back some info from scheduler to tokenizer_manager
tokenizer_manager.max_req_input_len = scheduler_infos[0]["max_req_input_len"]
return tokenizer_manager, template_manager, scheduler_infos, port_args
+3 -19
View File
@@ -6,7 +6,6 @@ Uses GrpcRequestManager for orchestration without tokenization.
import asyncio import asyncio
import dataclasses import dataclasses
import logging import logging
import multiprocessing as mp
import os import os
import signal import signal
import threading import threading
@@ -792,7 +791,7 @@ async def serve_grpc(
# Start warmup in a separate thread # Start warmup in a separate thread
warmup_thread = threading.Thread( warmup_thread = threading.Thread(
target=_wait_and_warmup_grpc, target=_wait_and_warmup_grpc,
args=(server_args, None, health_servicer), args=(server_args, health_servicer),
) )
warmup_thread.start() warmup_thread.start()
@@ -840,10 +839,7 @@ async def serve_grpc(
logger.info("All scheduler processes terminated") logger.info("All scheduler processes terminated")
def _execute_grpc_server_warmup( def _execute_grpc_server_warmup(server_args: ServerArgs):
server_args: ServerArgs,
pipe_finish_writer: Optional[mp.connection.Connection],
):
"""Execute warmup for gRPC server by checking health and sending test request.""" """Execute warmup for gRPC server by checking health and sending test request."""
try: try:
# Connect to the gRPC server # Connect to the gRPC server
@@ -874,8 +870,6 @@ def _execute_grpc_server_warmup(
if not success: if not success:
error_msg = f"gRPC server warmup failed: Could not connect to server after 120 seconds. Last error: {last_error}" error_msg = f"gRPC server warmup failed: Could not connect to server after 120 seconds. Last error: {last_error}"
logger.error(error_msg) logger.error(error_msg)
if pipe_finish_writer is not None:
pipe_finish_writer.send(error_msg)
channel.close() channel.close()
kill_process_tree(os.getpid()) kill_process_tree(os.getpid())
return False return False
@@ -938,8 +932,6 @@ def _execute_grpc_server_warmup(
except Exception as e: except Exception as e:
error_msg = f"gRPC warmup request failed: {e}" error_msg = f"gRPC warmup request failed: {e}"
logger.error(error_msg) logger.error(error_msg)
if pipe_finish_writer is not None:
pipe_finish_writer.send(error_msg)
channel.close() channel.close()
kill_process_tree(os.getpid()) kill_process_tree(os.getpid())
return False return False
@@ -966,8 +958,6 @@ def _execute_grpc_server_warmup(
except Exception as e: except Exception as e:
error_msg = f"gRPC warmup request failed: {e}" error_msg = f"gRPC warmup request failed: {e}"
logger.error(error_msg) logger.error(error_msg)
if pipe_finish_writer is not None:
pipe_finish_writer.send(error_msg)
channel.close() channel.close()
kill_process_tree(os.getpid()) kill_process_tree(os.getpid())
return False return False
@@ -980,8 +970,6 @@ def _execute_grpc_server_warmup(
f"gRPC warmup failed with exception: {e}\n{get_exception_traceback()}" f"gRPC warmup failed with exception: {e}\n{get_exception_traceback()}"
) )
logger.error(error_msg) logger.error(error_msg)
if pipe_finish_writer is not None:
pipe_finish_writer.send(error_msg)
try: try:
channel.close() channel.close()
except Exception: except Exception:
@@ -992,12 +980,11 @@ def _execute_grpc_server_warmup(
def _wait_and_warmup_grpc( def _wait_and_warmup_grpc(
server_args: ServerArgs, server_args: ServerArgs,
pipe_finish_writer: Optional[mp.connection.Connection],
health_servicer: Optional[SGLangHealthServicer] = None, health_servicer: Optional[SGLangHealthServicer] = None,
): ):
"""Wait for gRPC server to be ready and execute warmup.""" """Wait for gRPC server to be ready and execute warmup."""
if not server_args.skip_server_warmup: if not server_args.skip_server_warmup:
if not _execute_grpc_server_warmup(server_args, pipe_finish_writer): if not _execute_grpc_server_warmup(server_args):
return return
else: else:
logger.info("Skipping gRPC server warmup (skip_server_warmup=True)") logger.info("Skipping gRPC server warmup (skip_server_warmup=True)")
@@ -1007,6 +994,3 @@ def _wait_and_warmup_grpc(
health_servicer.set_serving() health_servicer.set_serving()
logger.info("The server is fired up and ready to roll!") logger.info("The server is fired up and ready to roll!")
if pipe_finish_writer is not None:
pipe_finish_writer.send("ready")
+28 -32
View File
@@ -20,7 +20,6 @@ This file implements HTTP APIs for the inference engine via fastapi.
import asyncio import asyncio
import dataclasses import dataclasses
import logging import logging
import multiprocessing
import os import os
import tempfile import tempfile
import threading import threading
@@ -53,7 +52,12 @@ from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import ORJSONResponse, Response, StreamingResponse from fastapi.responses import ORJSONResponse, Response, StreamingResponse
from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST, DisaggregationMode from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST, DisaggregationMode
from sglang.srt.entrypoints.engine import _launch_subprocesses from sglang.srt.entrypoints.engine import (
_launch_subprocesses,
init_tokenizer_manager,
run_detokenizer_process,
run_scheduler_process,
)
from sglang.srt.entrypoints.ollama.protocol import ( from sglang.srt.entrypoints.ollama.protocol import (
OllamaChatRequest, OllamaChatRequest,
OllamaGenerateRequest, OllamaGenerateRequest,
@@ -1463,10 +1467,7 @@ def _create_error_response(e):
MINIMUM_PNG_PICTURE_BASE64 = "iVBORw0KGgoAAAANSUhEUgAAACAAAAAgCAYAAABzenr0AAAACXBIWXMAAA7EAAAOxAGVKw4bAAAAbUlEQVRYhe3VsQ2AMAxE0Y/lIgNQULD/OqyCMgCihCKSG4yRuKuiNH6JLsoEbMACOGBcua9HOR7Y6w6swBwMy0qLTpkeI77qdEBpBFAHBBDAGH8WrwJKI4AAegUCfAKgEgpQDvh3CR3oQCuav58qlAw73kKCSgAAAABJRU5ErkJggg==" MINIMUM_PNG_PICTURE_BASE64 = "iVBORw0KGgoAAAANSUhEUgAAACAAAAAgCAYAAABzenr0AAAACXBIWXMAAA7EAAAOxAGVKw4bAAAAbUlEQVRYhe3VsQ2AMAxE0Y/lIgNQULD/OqyCMgCihCKSG4yRuKuiNH6JLsoEbMACOGBcua9HOR7Y6w6swBwMy0qLTpkeI77qdEBpBFAHBBDAGH8WrwJKI4AAegUCfAKgEgpQDvh3CR3oQCuav58qlAw73kKCSgAAAABJRU5ErkJggg=="
def _execute_server_warmup( def _execute_server_warmup(server_args: ServerArgs):
server_args: ServerArgs,
pipe_finish_writer: Optional[multiprocessing.connection.Connection],
):
headers = {} headers = {}
url = server_args.url() url = server_args.url()
if server_args.api_key: if server_args.api_key:
@@ -1486,8 +1487,6 @@ def _execute_server_warmup(
pass pass
if not success: if not success:
if pipe_finish_writer is not None:
pipe_finish_writer.send(last_traceback)
logger.error(f"Initialization failed. warmup error: {last_traceback}") logger.error(f"Initialization failed. warmup error: {last_traceback}")
kill_process_tree(os.getpid()) kill_process_tree(os.getpid())
return success return success
@@ -1607,8 +1606,6 @@ def _execute_server_warmup(
except Exception: except Exception:
last_traceback = get_exception_traceback() last_traceback = get_exception_traceback()
if pipe_finish_writer is not None:
pipe_finish_writer.send(last_traceback)
logger.error(f"Initialization failed. warmup error: {last_traceback}") logger.error(f"Initialization failed. warmup error: {last_traceback}")
kill_process_tree(os.getpid()) kill_process_tree(os.getpid())
return False return False
@@ -1620,7 +1617,6 @@ def _execute_server_warmup(
def _wait_and_warmup( def _wait_and_warmup(
server_args: ServerArgs, server_args: ServerArgs,
pipe_finish_writer: Optional[multiprocessing.connection.Connection] = None,
launch_callback: Optional[Callable[[], None]] = None, launch_callback: Optional[Callable[[], None]] = None,
execute_warmup_func: Callable = _execute_server_warmup, execute_warmup_func: Callable = _execute_server_warmup,
): ):
@@ -1629,10 +1625,7 @@ def _wait_and_warmup(
# Send a warmup request # Send a warmup request
if not server_args.skip_server_warmup: if not server_args.skip_server_warmup:
if not execute_warmup_func( if not execute_warmup_func(server_args):
server_args,
pipe_finish_writer,
):
return return
else: else:
_global_state.tokenizer_manager.server_status = ServerStatus.Up _global_state.tokenizer_manager.server_status = ServerStatus.Up
@@ -1640,9 +1633,6 @@ def _wait_and_warmup(
# The server is ready for requests # The server is ready for requests
logger.info("The server is fired up and ready to roll!") logger.info("The server is fired up and ready to roll!")
if pipe_finish_writer is not None:
pipe_finish_writer.send("ready")
if server_args.delete_ckpt_after_loading: if server_args.delete_ckpt_after_loading:
delete_directory(server_args.model_path) delete_directory(server_args.model_path)
@@ -1676,8 +1666,9 @@ def _wait_weights_ready():
def launch_server( def launch_server(
server_args: ServerArgs, server_args: ServerArgs,
pipe_finish_writer: Optional[multiprocessing.connection.Connection] = None, init_tokenizer_manager_func: Callable = init_tokenizer_manager,
launch_subprocesses_func: Callable = _launch_subprocesses, run_scheduler_process_func: Callable = run_scheduler_process,
run_detokenizer_process_func: Callable = run_detokenizer_process,
execute_warmup_func: Callable = _execute_server_warmup, execute_warmup_func: Callable = _execute_server_warmup,
launch_callback: Optional[Callable[[], None]] = None, launch_callback: Optional[Callable[[], None]] = None,
): ):
@@ -1696,23 +1687,27 @@ def launch_server(
1. The HTTP server, Engine, and TokenizerManager all run in the main process. 1. The HTTP server, Engine, and TokenizerManager all run in the main process.
2. Inter-process communication is done through IPC (each process uses a different port) via the ZMQ library. 2. Inter-process communication is done through IPC (each process uses a different port) via the ZMQ library.
""" """
# Launch subprocesses
tokenizer_manager, template_manager, scheduler_infos, port_args = ( tokenizer_manager, template_manager, scheduler_infos, port_args = (
launch_subprocesses_func(server_args=server_args) _launch_subprocesses(
server_args=server_args,
init_tokenizer_manager_func=init_tokenizer_manager_func,
run_scheduler_process_func=run_scheduler_process_func,
run_detokenizer_process_func=run_detokenizer_process_func,
)
) )
scheduler_info = scheduler_infos[0] # Parse info got from the schedulers
remote_instance_transfer_engine_info = None remote_instance_transfer_engine_info = (
if server_args.remote_instance_weight_loader_use_transfer_engine(): parse_remote_instance_transfer_engine_info_from_scheduler_infos(scheduler_infos)
remote_instance_transfer_engine_info = ( )
parse_remote_instance_transfer_engine_info_from_scheduler_infos(
scheduler_infos # Set global states
)
)
set_global_state( set_global_state(
_GlobalState( _GlobalState(
tokenizer_manager=tokenizer_manager, tokenizer_manager=tokenizer_manager,
template_manager=template_manager, template_manager=template_manager,
scheduler_info=scheduler_info, scheduler_info=scheduler_infos[0],
remote_instance_transfer_engine_info=remote_instance_transfer_engine_info, remote_instance_transfer_engine_info=remote_instance_transfer_engine_info,
) )
) )
@@ -1728,7 +1723,6 @@ def launch_server(
app.server_args = server_args app.server_args = server_args
app.warmup_thread_kwargs = dict( app.warmup_thread_kwargs = dict(
server_args=server_args, server_args=server_args,
pipe_finish_writer=pipe_finish_writer,
launch_callback=launch_callback, launch_callback=launch_callback,
execute_warmup_func=execute_warmup_func, execute_warmup_func=execute_warmup_func,
) )
@@ -1742,7 +1736,7 @@ def launch_server(
# for other worker processes to read. # for other worker processes to read.
app.is_single_tokenizer_mode = False app.is_single_tokenizer_mode = False
multi_tokenizer_args_shm = write_data_for_multi_tokenizer( multi_tokenizer_args_shm = write_data_for_multi_tokenizer(
port_args, server_args, scheduler_info port_args, server_args, scheduler_infos[0]
) )
try: try:
@@ -1751,6 +1745,7 @@ def launch_server(
# Listen for HTTP requests # Listen for HTTP requests
if server_args.tokenizer_worker_num == 1: if server_args.tokenizer_worker_num == 1:
# Default case, one tokenizer process
uvicorn.run( uvicorn.run(
app, app,
host=server_args.host, host=server_args.host,
@@ -1761,6 +1756,7 @@ def launch_server(
loop="uvloop", loop="uvloop",
) )
else: else:
# Multiple tokenizer and http processes
from uvicorn.config import LOGGING_CONFIG from uvicorn.config import LOGGING_CONFIG
LOGGING_CONFIG["loggers"]["sglang.srt.entrypoints.http_server"] = { LOGGING_CONFIG["loggers"]["sglang.srt.entrypoints.http_server"] = {
@@ -136,7 +136,7 @@ class SchedulerPPMixin:
# When the server is idle, self-check and re-init some states # When the server is idle, self-check and re-init some states
if server_is_idle: if server_is_idle:
self.check_during_pp_idle() self.self_check_during_idle()
@DynamicGradMode() @DynamicGradMode()
def event_loop_pp_disagg_prefill(self: Scheduler): def event_loop_pp_disagg_prefill(self: Scheduler):
@@ -312,7 +312,7 @@ class SchedulerPPMixin:
# When the server is idle, self-check and re-init some states # When the server is idle, self-check and re-init some states
if server_is_idle and len(self.disagg_prefill_inflight_queue) == 0: if server_is_idle and len(self.disagg_prefill_inflight_queue) == 0:
self.check_during_pp_idle() self.self_check_during_idle()
@DynamicGradMode() @DynamicGradMode()
def event_loop_pp_disagg_decode(self: Scheduler): def event_loop_pp_disagg_decode(self: Scheduler):
@@ -501,7 +501,7 @@ class SchedulerPPMixin:
queue_size += len(self.decode_offload_manager.ongoing_offload) queue_size += len(self.decode_offload_manager.ongoing_offload)
if server_is_idle and queue_size == 0: if server_is_idle and queue_size == 0:
self.check_during_pp_idle() self.self_check_during_idle()
def init_pp_loop_state(self: Scheduler): def init_pp_loop_state(self: Scheduler):
self.pp_loop_size: int = self.pp_size + self.server_args.pp_async_batch_depth self.pp_loop_size: int = self.pp_size + self.server_args.pp_async_batch_depth
@@ -700,12 +700,6 @@ class SchedulerPPMixin:
return predicted_size return predicted_size
def check_during_pp_idle(self: Scheduler):
self.check_memory()
self.check_tree_cache()
self.new_token_ratio = self.init_new_token_ratio
self.maybe_sleep_on_idle()
def process_bootstrapped_queue( def process_bootstrapped_queue(
self: Scheduler, bootstrapped_rids: Optional[List[str]] self: Scheduler, bootstrapped_rids: Optional[List[str]]
): ):
+2 -2
View File
@@ -47,5 +47,5 @@ It learns from [Copybara](https://github.com/google/copybara), a tool used at Go
default=ServerArgs.private_flag, default=ServerArgs.private_flag,
) )
``` ```
- Similarly, you can inherit `Engine` and override `launch_subprocesses_func`, `server_args_class`. - Similarly, you can inherit `Engine` and override its fields. You can override `server_args_class` to use your own ServerArgs,
- You can pass your own subprocesses launch functions to `launch_server.py::launch_server` override `init_tokenizer_manager_func` to use your own TokenizerManager, override `run_scheduler_process_func` to use your own scheduler.