Improve engine customization interface (#15635)
This commit is contained in:
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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")
|
|
||||||
|
|||||||
@@ -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]]
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
Reference in New Issue
Block a user