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

View File

@@ -1,6 +1,7 @@
import atexit
import json
import multiprocessing
import time
import warnings
from typing import Dict, List, Optional, Union
@@ -365,10 +366,18 @@ class Runtime:
def __init__(
self,
log_level: str = "error",
launch_timeout: float = 300.0,
*args,
**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
# client code without installing SRT server and its dependency if they want.
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__
self.pid = None
pipe_reader, pipe_writer = multiprocessing.Pipe(duplex=False)
ctx = multiprocessing.get_context("spawn")
proc = ctx.Process(
target=launch_server,
args=(self.server_args, pipe_writer),
args=(self.server_args,),
)
proc.start()
pipe_writer.close()
self.pid = proc.pid
# Before python program terminates, call shutdown implicitly. Therefore, users don't have to explicitly call .shutdown()
atexit.register(self.shutdown)
# TODO: remove this pipe_writer mechanism and use `/health_generate` instead.
try:
init_state = pipe_reader.recv()
except EOFError:
init_state = ""
# Wait for server to be ready by polling /health_generate
start_time = time.time()
with requests.Session() as session:
while time.time() - start_time < launch_timeout:
try:
response = session.get(f"{self.url}/health_generate")
if response.status_code == 200:
break
except requests.RequestException:
pass
if init_state != "ready":
self.shutdown()
raise RuntimeError(
"Initialization failed. Please see the error messages above."
)
if not proc.is_alive():
self.shutdown()
raise RuntimeError(
"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)

View File

@@ -91,94 +91,25 @@ asyncio.set_event_loop_policy(uvloop.EventLoopPolicy())
_is_cuda = is_cuda()
def _launch_subprocesses(
server_args: ServerArgs, port_args: Optional[PortArgs] = None
) -> Tuple[TokenizerManager, TemplateManager, 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()
def init_tokenizer_manager(
server_args: ServerArgs,
port_args: PortArgs,
TokenizerManagerClass: Optional[TokenizerManager] = None,
) -> Tuple[TokenizerManager, TemplateManager]:
# Launch tokenizer process
TokenizerManagerClass = TokenizerManagerClass or TokenizerManager
tokenizer_manager = TokenizerManagerClass(server_args, port_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,
# 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,
)
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,
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
return tokenizer_manager, template_manager
class Engine(EngineBase):
@@ -197,8 +128,10 @@ class Engine(EngineBase):
# Some fields to allow people to override the server args
# and launch processes for their private forks.
launch_subprocesses_func: Callable = staticmethod(_launch_subprocesses)
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):
"""
@@ -224,7 +157,12 @@ class Engine(EngineBase):
# Launch subprocesses
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.template_manager = template_manager
@@ -873,32 +811,10 @@ def _set_envs_and_config(server_args: ServerArgs):
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(
server_args: ServerArgs,
port_args: PortArgs,
run_scheduler_process_func: Callable = run_scheduler_process,
run_data_parallel_controller_process_func: Callable = run_data_parallel_controller_process,
run_scheduler_process_func: Callable,
):
scheduler_procs = []
@@ -959,10 +875,110 @@ def _launch_scheduler_processes(
reader, writer = mp.Pipe(duplex=False)
scheduler_pipe_readers = [reader]
proc = mp.Process(
target=run_data_parallel_controller_process_func,
args=(server_args, port_args, writer),
target=run_data_parallel_controller_process,
kwargs=dict(
server_args=server_args,
port_args=port_args,
pipe_writer=writer,
run_scheduler_process_func=run_scheduler_process_func,
),
)
proc.start()
scheduler_procs.append(proc)
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

View File

@@ -6,7 +6,6 @@ Uses GrpcRequestManager for orchestration without tokenization.
import asyncio
import dataclasses
import logging
import multiprocessing as mp
import os
import signal
import threading
@@ -792,7 +791,7 @@ async def serve_grpc(
# Start warmup in a separate thread
warmup_thread = threading.Thread(
target=_wait_and_warmup_grpc,
args=(server_args, None, health_servicer),
args=(server_args, health_servicer),
)
warmup_thread.start()
@@ -840,10 +839,7 @@ async def serve_grpc(
logger.info("All scheduler processes terminated")
def _execute_grpc_server_warmup(
server_args: ServerArgs,
pipe_finish_writer: Optional[mp.connection.Connection],
):
def _execute_grpc_server_warmup(server_args: ServerArgs):
"""Execute warmup for gRPC server by checking health and sending test request."""
try:
# Connect to the gRPC server
@@ -874,8 +870,6 @@ def _execute_grpc_server_warmup(
if not success:
error_msg = f"gRPC server warmup failed: Could not connect to server after 120 seconds. Last error: {last_error}"
logger.error(error_msg)
if pipe_finish_writer is not None:
pipe_finish_writer.send(error_msg)
channel.close()
kill_process_tree(os.getpid())
return False
@@ -938,8 +932,6 @@ def _execute_grpc_server_warmup(
except Exception as e:
error_msg = f"gRPC warmup request failed: {e}"
logger.error(error_msg)
if pipe_finish_writer is not None:
pipe_finish_writer.send(error_msg)
channel.close()
kill_process_tree(os.getpid())
return False
@@ -966,8 +958,6 @@ def _execute_grpc_server_warmup(
except Exception as e:
error_msg = f"gRPC warmup request failed: {e}"
logger.error(error_msg)
if pipe_finish_writer is not None:
pipe_finish_writer.send(error_msg)
channel.close()
kill_process_tree(os.getpid())
return False
@@ -980,8 +970,6 @@ def _execute_grpc_server_warmup(
f"gRPC warmup failed with exception: {e}\n{get_exception_traceback()}"
)
logger.error(error_msg)
if pipe_finish_writer is not None:
pipe_finish_writer.send(error_msg)
try:
channel.close()
except Exception:
@@ -992,12 +980,11 @@ def _execute_grpc_server_warmup(
def _wait_and_warmup_grpc(
server_args: ServerArgs,
pipe_finish_writer: Optional[mp.connection.Connection],
health_servicer: Optional[SGLangHealthServicer] = None,
):
"""Wait for gRPC server to be ready and execute 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
else:
logger.info("Skipping gRPC server warmup (skip_server_warmup=True)")
@@ -1007,6 +994,3 @@ def _wait_and_warmup_grpc(
health_servicer.set_serving()
logger.info("The server is fired up and ready to roll!")
if pipe_finish_writer is not None:
pipe_finish_writer.send("ready")

View File

@@ -20,7 +20,6 @@ This file implements HTTP APIs for the inference engine via fastapi.
import asyncio
import dataclasses
import logging
import multiprocessing
import os
import tempfile
import threading
@@ -53,7 +52,12 @@ from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import ORJSONResponse, Response, StreamingResponse
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 (
OllamaChatRequest,
OllamaGenerateRequest,
@@ -1463,10 +1467,7 @@ def _create_error_response(e):
MINIMUM_PNG_PICTURE_BASE64 = "iVBORw0KGgoAAAANSUhEUgAAACAAAAAgCAYAAABzenr0AAAACXBIWXMAAA7EAAAOxAGVKw4bAAAAbUlEQVRYhe3VsQ2AMAxE0Y/lIgNQULD/OqyCMgCihCKSG4yRuKuiNH6JLsoEbMACOGBcua9HOR7Y6w6swBwMy0qLTpkeI77qdEBpBFAHBBDAGH8WrwJKI4AAegUCfAKgEgpQDvh3CR3oQCuav58qlAw73kKCSgAAAABJRU5ErkJggg=="
def _execute_server_warmup(
server_args: ServerArgs,
pipe_finish_writer: Optional[multiprocessing.connection.Connection],
):
def _execute_server_warmup(server_args: ServerArgs):
headers = {}
url = server_args.url()
if server_args.api_key:
@@ -1486,8 +1487,6 @@ def _execute_server_warmup(
pass
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}")
kill_process_tree(os.getpid())
return success
@@ -1607,8 +1606,6 @@ def _execute_server_warmup(
except Exception:
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}")
kill_process_tree(os.getpid())
return False
@@ -1620,7 +1617,6 @@ def _execute_server_warmup(
def _wait_and_warmup(
server_args: ServerArgs,
pipe_finish_writer: Optional[multiprocessing.connection.Connection] = None,
launch_callback: Optional[Callable[[], None]] = None,
execute_warmup_func: Callable = _execute_server_warmup,
):
@@ -1629,10 +1625,7 @@ def _wait_and_warmup(
# Send a warmup request
if not server_args.skip_server_warmup:
if not execute_warmup_func(
server_args,
pipe_finish_writer,
):
if not execute_warmup_func(server_args):
return
else:
_global_state.tokenizer_manager.server_status = ServerStatus.Up
@@ -1640,9 +1633,6 @@ def _wait_and_warmup(
# The server is ready for requests
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:
delete_directory(server_args.model_path)
@@ -1676,8 +1666,9 @@ def _wait_weights_ready():
def launch_server(
server_args: ServerArgs,
pipe_finish_writer: Optional[multiprocessing.connection.Connection] = None,
launch_subprocesses_func: Callable = _launch_subprocesses,
init_tokenizer_manager_func: Callable = init_tokenizer_manager,
run_scheduler_process_func: Callable = run_scheduler_process,
run_detokenizer_process_func: Callable = run_detokenizer_process,
execute_warmup_func: Callable = _execute_server_warmup,
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.
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 = (
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]
remote_instance_transfer_engine_info = None
if server_args.remote_instance_weight_loader_use_transfer_engine():
remote_instance_transfer_engine_info = (
parse_remote_instance_transfer_engine_info_from_scheduler_infos(
scheduler_infos
)
)
# Parse info got from the schedulers
remote_instance_transfer_engine_info = (
parse_remote_instance_transfer_engine_info_from_scheduler_infos(scheduler_infos)
)
# Set global states
set_global_state(
_GlobalState(
tokenizer_manager=tokenizer_manager,
template_manager=template_manager,
scheduler_info=scheduler_info,
scheduler_info=scheduler_infos[0],
remote_instance_transfer_engine_info=remote_instance_transfer_engine_info,
)
)
@@ -1728,7 +1723,6 @@ def launch_server(
app.server_args = server_args
app.warmup_thread_kwargs = dict(
server_args=server_args,
pipe_finish_writer=pipe_finish_writer,
launch_callback=launch_callback,
execute_warmup_func=execute_warmup_func,
)
@@ -1742,7 +1736,7 @@ def launch_server(
# for other worker processes to read.
app.is_single_tokenizer_mode = False
multi_tokenizer_args_shm = write_data_for_multi_tokenizer(
port_args, server_args, scheduler_info
port_args, server_args, scheduler_infos[0]
)
try:
@@ -1751,6 +1745,7 @@ def launch_server(
# Listen for HTTP requests
if server_args.tokenizer_worker_num == 1:
# Default case, one tokenizer process
uvicorn.run(
app,
host=server_args.host,
@@ -1761,6 +1756,7 @@ def launch_server(
loop="uvloop",
)
else:
# Multiple tokenizer and http processes
from uvicorn.config import LOGGING_CONFIG
LOGGING_CONFIG["loggers"]["sglang.srt.entrypoints.http_server"] = {

View File

@@ -136,7 +136,7 @@ class SchedulerPPMixin:
# When the server is idle, self-check and re-init some states
if server_is_idle:
self.check_during_pp_idle()
self.self_check_during_idle()
@DynamicGradMode()
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
if server_is_idle and len(self.disagg_prefill_inflight_queue) == 0:
self.check_during_pp_idle()
self.self_check_during_idle()
@DynamicGradMode()
def event_loop_pp_disagg_decode(self: Scheduler):
@@ -501,7 +501,7 @@ class SchedulerPPMixin:
queue_size += len(self.decode_offload_manager.ongoing_offload)
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):
self.pp_loop_size: int = self.pp_size + self.server_args.pp_async_batch_depth
@@ -700,12 +700,6 @@ class SchedulerPPMixin:
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(
self: Scheduler, bootstrapped_rids: Optional[List[str]]
):