support MPServer and embedded server for granian to enable muti tokenizer worker (#28573)

This commit is contained in:
Rain Jiang
2026-06-18 17:59:12 -07:00
committed by GitHub
parent ea407df4b0
commit ef01618dfb
5 changed files with 131 additions and 128 deletions
+90 -92
View File
@@ -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,50 +2179,76 @@ 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)
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() server.serve()
@@ -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,6 +2417,22 @@ def _setup_and_run_http_server(
"SSL refresh will be disabled." "SSL refresh will be disabled."
) )
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,
tokenizer_worker_num=server_args.tokenizer_worker_num,
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,
)
else:
uvicorn.run( uvicorn.run(
"sglang.srt.entrypoints.http_server:app", "sglang.srt.entrypoints.http_server:app",
host=server_args.host, host=server_args.host,
-3
View File
@@ -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
-6
View File
@@ -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)