support MPServer and embedded server for granian to enable muti tokenizer worker (#28573)
This commit is contained in:
@@ -207,32 +207,6 @@ def get_global_state() -> _GlobalState:
|
|||||||
return _global_state
|
return _global_state
|
||||||
|
|
||||||
|
|
||||||
async def _init_granian_worker() -> ServerArgs:
|
|
||||||
main_pid = get_main_process_id()
|
|
||||||
port_args, server_args, scheduler_info = read_from_shared_memory(
|
|
||||||
f"multi_tokenizer_args_{main_pid}"
|
|
||||||
)
|
|
||||||
|
|
||||||
tokenizer_manager = TokenizerManager(server_args, port_args)
|
|
||||||
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,
|
|
||||||
)
|
|
||||||
tokenizer_manager.max_req_input_len = scheduler_info["max_req_input_len"]
|
|
||||||
|
|
||||||
set_global_state(
|
|
||||||
_GlobalState(
|
|
||||||
tokenizer_manager=tokenizer_manager,
|
|
||||||
template_manager=template_manager,
|
|
||||||
scheduler_info=scheduler_info,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
return server_args
|
|
||||||
|
|
||||||
|
|
||||||
async def init_multi_tokenizer() -> ServerArgs:
|
async def init_multi_tokenizer() -> ServerArgs:
|
||||||
"""
|
"""
|
||||||
Initialization function for multi-process tokenizer mode.
|
Initialization function for multi-process tokenizer mode.
|
||||||
@@ -290,10 +264,6 @@ async def lifespan(fast_api_app: FastAPI):
|
|||||||
server_args = fast_api_app.server_args
|
server_args = fast_api_app.server_args
|
||||||
warmup_thread_kwargs = fast_api_app.warmup_thread_kwargs
|
warmup_thread_kwargs = fast_api_app.warmup_thread_kwargs
|
||||||
thread_label = "Tokenizer"
|
thread_label = "Tokenizer"
|
||||||
elif envs.SGLANG_GRANIAN_PARENT_PID.get() is not None:
|
|
||||||
server_args = await _init_granian_worker()
|
|
||||||
warmup_thread_kwargs = dict(server_args=server_args)
|
|
||||||
thread_label = "Tokenizer"
|
|
||||||
else:
|
else:
|
||||||
# Initialize multi-tokenizer support for worker processes
|
# Initialize multi-tokenizer support for worker processes
|
||||||
server_args = await init_multi_tokenizer()
|
server_args = await init_multi_tokenizer()
|
||||||
@@ -2209,51 +2179,77 @@ def _wait_weights_ready():
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _close_main_process_sockets():
|
def _run_granian_server(
|
||||||
"""Close the main process's ZMQ sockets before spawning Granian workers.
|
host,
|
||||||
|
port,
|
||||||
|
log_level,
|
||||||
|
tokenizer_worker_num=1,
|
||||||
|
ssl_certfile=None,
|
||||||
|
ssl_keyfile=None,
|
||||||
|
ssl_ca_certs=None,
|
||||||
|
ssl_keyfile_password=None,
|
||||||
|
ssl_verify=False, # MTls is not supported
|
||||||
|
backlog=2048,
|
||||||
|
backpressure=2048,
|
||||||
|
):
|
||||||
|
"""Serve the in-process ASGI app with Granian (embedded mode) over HTTP/2.
|
||||||
|
|
||||||
Granian workers create their own TokenizerManager with fresh ZMQ sockets.
|
Unlike Granian's default multi-process server, the embedded server runs a
|
||||||
The main process must release its sockets first to avoid binding conflicts
|
single worker as an asyncio task inside the current process. It therefore
|
||||||
on the same IPC addresses.
|
serves the live ``app`` object directly and reuses the already-initialized
|
||||||
|
global state (tokenizer manager, templates, ...) through the normal
|
||||||
|
single-tokenizer lifespan path -- no shared memory or worker re-init needed.
|
||||||
|
The event loop is uvloop. The default backlog and backpressure values are set
|
||||||
|
exactly like uvicorn's defaults.
|
||||||
"""
|
"""
|
||||||
if _global_state is None or _global_state.tokenizer_manager is None:
|
import signal
|
||||||
return
|
|
||||||
tm = _global_state.tokenizer_manager
|
|
||||||
for attr in ("recv_from_detokenizer", "send_to_scheduler"):
|
|
||||||
sock = getattr(tm, attr, None)
|
|
||||||
if sock is None:
|
|
||||||
continue
|
|
||||||
inner = getattr(sock, "socket", None)
|
|
||||||
if inner is not None:
|
|
||||||
inner.close()
|
|
||||||
elif hasattr(sock, "close"):
|
|
||||||
sock.close()
|
|
||||||
setattr(tm, attr, None)
|
|
||||||
|
|
||||||
|
|
||||||
def _run_granian_server(server_args: ServerArgs):
|
|
||||||
"""Launch Granian with HTTP/2 support"""
|
|
||||||
from granian import Granian
|
from granian import Granian
|
||||||
from granian.constants import HTTPModes, Interfaces, Loops
|
from granian.constants import HTTPModes, Interfaces, Loops
|
||||||
|
from granian.server.embed import Server as GranianEmbeddedServer
|
||||||
|
|
||||||
|
Server = GranianEmbeddedServer if tokenizer_worker_num == 1 else Granian
|
||||||
|
target = (
|
||||||
|
app if tokenizer_worker_num == 1 else "sglang.srt.entrypoints.http_server:app"
|
||||||
|
)
|
||||||
granian_kwargs = dict(
|
granian_kwargs = dict(
|
||||||
target="sglang.srt.entrypoints.http_server:app",
|
target=target,
|
||||||
address=server_args.host,
|
address=host,
|
||||||
port=server_args.port,
|
port=port,
|
||||||
interface=Interfaces.ASGI,
|
interface=Interfaces.ASGI,
|
||||||
http=HTTPModes.auto,
|
http=HTTPModes.auto,
|
||||||
loop=Loops.uvloop,
|
log_level=log_level,
|
||||||
log_level=server_args.log_level_http or server_args.log_level or "info",
|
ssl_cert=ssl_certfile,
|
||||||
workers=1,
|
ssl_key=ssl_keyfile,
|
||||||
|
ssl_key_password=ssl_keyfile_password,
|
||||||
|
ssl_ca=ssl_ca_certs,
|
||||||
|
ssl_client_verify=ssl_verify,
|
||||||
|
backlog=backlog,
|
||||||
|
backpressure=backpressure,
|
||||||
)
|
)
|
||||||
|
|
||||||
ssl_enabled = server_args.ssl_certfile and server_args.ssl_keyfile
|
if tokenizer_worker_num > 1:
|
||||||
if ssl_enabled:
|
granian_kwargs["workers"] = tokenizer_worker_num
|
||||||
granian_kwargs["ssl_cert"] = server_args.ssl_certfile
|
granian_kwargs["loop"] = Loops.uvloop
|
||||||
granian_kwargs["ssl_key"] = server_args.ssl_keyfile
|
|
||||||
|
|
||||||
server = Granian(**granian_kwargs)
|
server = Server(**granian_kwargs)
|
||||||
server.serve()
|
|
||||||
|
if tokenizer_worker_num == 1:
|
||||||
|
|
||||||
|
async def serve():
|
||||||
|
# The embedded server does not install its own signal handlers, so wire
|
||||||
|
# SIGINT/SIGTERM to a graceful stop, mirroring uvicorn's behavior.
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
for sig in (signal.SIGINT, signal.SIGTERM):
|
||||||
|
try:
|
||||||
|
loop.add_signal_handler(sig, server.stop)
|
||||||
|
except (NotImplementedError, ValueError):
|
||||||
|
pass
|
||||||
|
await server.serve()
|
||||||
|
|
||||||
|
uvloop.run(serve())
|
||||||
|
else:
|
||||||
|
server.serve()
|
||||||
|
|
||||||
|
|
||||||
def _setup_and_run_http_server(
|
def _setup_and_run_http_server(
|
||||||
@@ -2286,35 +2282,6 @@ def _setup_and_run_http_server(
|
|||||||
if server_args.enable_metrics:
|
if server_args.enable_metrics:
|
||||||
add_prometheus_track_response_middleware(app)
|
add_prometheus_track_response_middleware(app)
|
||||||
|
|
||||||
# Use Granian for HTTP/2 server
|
|
||||||
if server_args.enable_http2:
|
|
||||||
# Reuse the multi-tokenizer shared memory mechanism to pass
|
|
||||||
# init args (port_args, server_args, scheduler_info) to
|
|
||||||
# Granian workers, which are independent processes.
|
|
||||||
multi_tokenizer_args_shm = write_data_for_multi_tokenizer(
|
|
||||||
port_args, server_args, scheduler_infos[0]
|
|
||||||
)
|
|
||||||
try:
|
|
||||||
if server_args.ssl_certfile:
|
|
||||||
logger.info(
|
|
||||||
f"SSL enabled: certfile={server_args.ssl_certfile}, "
|
|
||||||
f"keyfile={server_args.ssl_keyfile}"
|
|
||||||
)
|
|
||||||
logger.info(
|
|
||||||
f"Starting Granian HTTP/2 server on "
|
|
||||||
f"{server_args.host}:{server_args.port}"
|
|
||||||
)
|
|
||||||
# Propagate the main process PID via os.environ so Granian
|
|
||||||
# workers (forked or spawned) can locate the shared memory
|
|
||||||
# segment created above.
|
|
||||||
envs.SGLANG_GRANIAN_PARENT_PID.set(os.getpid())
|
|
||||||
_close_main_process_sockets()
|
|
||||||
_run_granian_server(server_args)
|
|
||||||
finally:
|
|
||||||
if multi_tokenizer_args_shm is not None:
|
|
||||||
multi_tokenizer_args_shm.unlink()
|
|
||||||
return
|
|
||||||
|
|
||||||
# Pass additional arguments to the lifespan function.
|
# Pass additional arguments to the lifespan function.
|
||||||
# They will be used for additional initialization setups.
|
# They will be used for additional initialization setups.
|
||||||
if server_args.tokenizer_worker_num == 1:
|
if server_args.tokenizer_worker_num == 1:
|
||||||
@@ -2366,7 +2333,22 @@ def _setup_and_run_http_server(
|
|||||||
|
|
||||||
# Listen for HTTP requests
|
# Listen for HTTP requests
|
||||||
if server_args.tokenizer_worker_num == 1:
|
if server_args.tokenizer_worker_num == 1:
|
||||||
if server_args.enable_ssl_refresh:
|
if server_args.enable_http2:
|
||||||
|
logger.info(
|
||||||
|
f"Starting embedded Granian HTTP/2 server on "
|
||||||
|
f"{server_args.host}:{server_args.port}"
|
||||||
|
)
|
||||||
|
_run_granian_server(
|
||||||
|
host=server_args.host,
|
||||||
|
port=server_args.port,
|
||||||
|
log_level=server_args.log_level_http or server_args.log_level,
|
||||||
|
ssl_certfile=server_args.ssl_certfile,
|
||||||
|
ssl_keyfile=server_args.ssl_keyfile,
|
||||||
|
ssl_ca_certs=server_args.ssl_ca_certs,
|
||||||
|
ssl_keyfile_password=server_args.ssl_keyfile_password,
|
||||||
|
ssl_verify=False, # No MTLS supported for now.
|
||||||
|
)
|
||||||
|
elif server_args.enable_ssl_refresh:
|
||||||
# Use Config/Server API for access to the SSLContext.
|
# Use Config/Server API for access to the SSLContext.
|
||||||
config = uvicorn.Config(
|
config = uvicorn.Config(
|
||||||
app,
|
app,
|
||||||
@@ -2435,21 +2417,37 @@ def _setup_and_run_http_server(
|
|||||||
"SSL refresh will be disabled."
|
"SSL refresh will be disabled."
|
||||||
)
|
)
|
||||||
|
|
||||||
uvicorn.run(
|
if server_args.enable_http2:
|
||||||
"sglang.srt.entrypoints.http_server:app",
|
logger.info(
|
||||||
host=server_args.host,
|
f"Starting embedded Granian HTTP/2 server on "
|
||||||
port=server_args.port,
|
f"{server_args.host}:{server_args.port}"
|
||||||
root_path=server_args.fastapi_root_path,
|
)
|
||||||
log_level=server_args.log_level_http or server_args.log_level,
|
_run_granian_server(
|
||||||
timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(),
|
host=server_args.host,
|
||||||
timeout_worker_healthcheck=envs.SGLANG_UVICORN_WORKER_HEALTHCHECK_TIMEOUT.get(),
|
port=server_args.port,
|
||||||
loop="uvloop",
|
log_level=server_args.log_level_http or server_args.log_level,
|
||||||
workers=server_args.tokenizer_worker_num,
|
tokenizer_worker_num=server_args.tokenizer_worker_num,
|
||||||
ssl_keyfile=server_args.ssl_keyfile,
|
ssl_certfile=server_args.ssl_certfile,
|
||||||
ssl_certfile=server_args.ssl_certfile,
|
ssl_keyfile=server_args.ssl_keyfile,
|
||||||
ssl_ca_certs=server_args.ssl_ca_certs,
|
ssl_ca_certs=server_args.ssl_ca_certs,
|
||||||
ssl_keyfile_password=server_args.ssl_keyfile_password,
|
ssl_keyfile_password=server_args.ssl_keyfile_password,
|
||||||
)
|
)
|
||||||
|
else:
|
||||||
|
uvicorn.run(
|
||||||
|
"sglang.srt.entrypoints.http_server:app",
|
||||||
|
host=server_args.host,
|
||||||
|
port=server_args.port,
|
||||||
|
root_path=server_args.fastapi_root_path,
|
||||||
|
log_level=server_args.log_level_http or server_args.log_level,
|
||||||
|
timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(),
|
||||||
|
timeout_worker_healthcheck=envs.SGLANG_UVICORN_WORKER_HEALTHCHECK_TIMEOUT.get(),
|
||||||
|
loop="uvloop",
|
||||||
|
workers=server_args.tokenizer_worker_num,
|
||||||
|
ssl_keyfile=server_args.ssl_keyfile,
|
||||||
|
ssl_certfile=server_args.ssl_certfile,
|
||||||
|
ssl_ca_certs=server_args.ssl_ca_certs,
|
||||||
|
ssl_keyfile_password=server_args.ssl_keyfile_password,
|
||||||
|
)
|
||||||
finally:
|
finally:
|
||||||
if server_args.tokenizer_worker_num > 1:
|
if server_args.tokenizer_worker_num > 1:
|
||||||
if multi_tokenizer_args_shm is not None:
|
if multi_tokenizer_args_shm is not None:
|
||||||
|
|||||||
@@ -708,9 +708,6 @@ class Envs:
|
|||||||
# too short when many workers cold-start and load tokenizers in parallel.
|
# too short when many workers cold-start and load tokenizers in parallel.
|
||||||
SGLANG_UVICORN_WORKER_HEALTHCHECK_TIMEOUT = EnvInt(10)
|
SGLANG_UVICORN_WORKER_HEALTHCHECK_TIMEOUT = EnvInt(10)
|
||||||
|
|
||||||
# HTTP/2 Server
|
|
||||||
SGLANG_GRANIAN_PARENT_PID = EnvInt(None)
|
|
||||||
|
|
||||||
# Health Check
|
# Health Check
|
||||||
SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION = EnvBool(True)
|
SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION = EnvBool(True)
|
||||||
|
|
||||||
|
|||||||
@@ -684,16 +684,9 @@ async def print_exception_wrapper(func):
|
|||||||
|
|
||||||
|
|
||||||
def get_main_process_id() -> int:
|
def get_main_process_id() -> int:
|
||||||
"""Get the main process ID.
|
|
||||||
|
|
||||||
Supports override via SGLANG_GRANIAN_PARENT_PID for workers whose
|
|
||||||
multiprocessing parent PID differs from the shared-memory owner.
|
|
||||||
"""
|
"""
|
||||||
from sglang.srt.environ import envs
|
Get the main process ID.
|
||||||
|
"""
|
||||||
override = envs.SGLANG_GRANIAN_PARENT_PID.get()
|
|
||||||
if override is not None:
|
|
||||||
return override
|
|
||||||
return multiprocessing.current_process()._parent_pid
|
return multiprocessing.current_process()._parent_pid
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1178,12 +1178,6 @@ class ServerArgs:
|
|||||||
"Use Uvicorn (the default) or handle certificate rotation externally."
|
"Use Uvicorn (the default) or handle certificate rotation externally."
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.tokenizer_worker_num > 1:
|
|
||||||
raise ValueError(
|
|
||||||
"--enable-http2 does not yet support --tokenizer-worker-num > 1. "
|
|
||||||
"Multi-worker HTTP/2 support will be added in a future release."
|
|
||||||
)
|
|
||||||
|
|
||||||
def _handle_multimodal(self):
|
def _handle_multimodal(self):
|
||||||
"""Validate mm_process_config structure before model loading."""
|
"""Validate mm_process_config structure before model loading."""
|
||||||
if self.mm_process_config is not None:
|
if self.mm_process_config is not None:
|
||||||
|
|||||||
@@ -27,8 +27,8 @@ try:
|
|||||||
except ImportError:
|
except ImportError:
|
||||||
_HAS_GRANIAN = False
|
_HAS_GRANIAN = False
|
||||||
|
|
||||||
register_cuda_ci(est_time=52, stage="base-b", runner_config="1-gpu-small")
|
register_cuda_ci(est_time=150, stage="base-b", runner_config="1-gpu-small")
|
||||||
register_amd_ci(est_time=52, suite="stage-b-test-1-gpu-small-amd")
|
register_amd_ci(est_time=150, suite="stage-b-test-1-gpu-small-amd")
|
||||||
|
|
||||||
|
|
||||||
@unittest.skipUnless(_HAS_GRANIAN, "granian not installed (pip install sglang[http2])")
|
@unittest.skipUnless(_HAS_GRANIAN, "granian not installed (pip install sglang[http2])")
|
||||||
@@ -109,5 +109,26 @@ class TestHTTP2Server(CustomTestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(_HAS_GRANIAN, "granian not installed (pip install sglang[http2])")
|
||||||
|
class TestHTTP2ServerMultiTokenizer(TestHTTP2Server):
|
||||||
|
"""Same checks as TestHTTP2Server but with multiple tokenizer workers.
|
||||||
|
|
||||||
|
With --tokenizer-worker-num > 1 the HTTP/2 server is served by Granian's
|
||||||
|
multi-process server (instead of the single-process embedded server), so
|
||||||
|
this exercises the multi-worker code path.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model = DEFAULT_SMALL_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=["--enable-http2", "--tokenizer-worker-num", "2"],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main(verbosity=3)
|
unittest.main(verbosity=3)
|
||||||
|
|||||||
Reference in New Issue
Block a user