Support gated launch to defer startup memory allocation (#35927)

This commit is contained in:
fzyzcjy
2026-08-24 20:19:45 +08:00
committed by GitHub
parent 3b24d8981b
commit c56cee0f80
5 changed files with 427 additions and 0 deletions
@@ -20,6 +20,7 @@ from sglang.srt.distributed import (
set_mscclpp_all_reduce,
set_torch_symm_mem_all_reduce,
)
from sglang.srt.distributed.gated_launch import maybe_wait_for_gated_launch
from sglang.srt.distributed.parallel_state import (
_tag_groups_for_flashinfer_allreduce_only,
)
@@ -132,6 +133,10 @@ def init_torch_distributed(
):
_prewarm_tp_lm_head_all_to_all()
maybe_wait_for_gated_launch(
host=server_args.host, port=server_args.gated_launch_port
)
pre_model_load_memory = get_available_gpu_memory(
device,
ps.gpu_id,
@@ -0,0 +1,96 @@
import logging
import threading
import time
from typing import Optional
import torch
import torch.distributed as dist
import uvicorn
from fastapi import FastAPI
from fastapi.responses import PlainTextResponse
from sglang.srt.distributed import get_world_group
logger = logging.getLogger(__name__)
POLL_INTERVAL_SECONDS = 1.0
LOG_INTERVAL_SECONDS = 10.0
_instance: Optional["_GatedLaunchServer"] = None
def maybe_wait_for_gated_launch(*, host: str, port: Optional[int]) -> None:
global _instance
if port is None or _instance is not None:
return
world_group = get_world_group()
_instance = _GatedLaunchServer()
if world_group.rank_in_group == 0:
_instance.serve(host=host, port=port)
logger.info(f"Gated launch waiting for activation. rank={world_group.rank}")
tic = time.perf_counter()
_wait_until_activated(world_group=world_group, server=_instance)
logger.info(f"Gated launch activated. elapsed={time.perf_counter() - tic:.2f} s")
def _wait_until_activated(*, world_group, server: "_GatedLaunchServer") -> None:
activated = torch.zeros(1, dtype=torch.int32)
started_at = time.perf_counter()
next_log_at = started_at + LOG_INTERVAL_SECONDS
while True:
activated[0] = int(server.activated)
if world_group.world_size > 1:
dist.broadcast(
activated,
src=world_group.ranks[0],
group=world_group.cpu_group,
)
if bool(activated[0]):
return
if (now := time.perf_counter()) >= next_log_at:
logger.info(
f"Gated launch still waiting for activation. "
f"rank={world_group.rank} elapsed={now - started_at:.0f} s"
)
next_log_at = now + LOG_INTERVAL_SECONDS
time.sleep(POLL_INTERVAL_SECONDS)
class _GatedLaunchServer:
def __init__(self):
self.activated = False
self._server: Optional[uvicorn.Server] = None
self._thread: Optional[threading.Thread] = None
def serve(self, *, host: str, port: int) -> None:
config = uvicorn.Config(
_build_app(self), host=host, port=port, log_level="warning"
)
self._server = uvicorn.Server(config)
self._thread = threading.Thread(target=self._server.run, daemon=True)
self._thread.start()
logger.info(f"Gated launch control server started on {host}:{port}")
def _build_app(server: _GatedLaunchServer) -> FastAPI:
app = FastAPI()
@app.get("/health")
def health():
return PlainTextResponse("OK")
@app.post("/gate/activate")
def activate():
server.activated = True
return PlainTextResponse("OK")
return app
+5
View File
@@ -1046,6 +1046,11 @@ class ServerArgs:
),
NS("parallel"),
] = None
gated_launch_port: A[
Optional[int],
"The port of the gated launch control server. When set, every rank blocks right after the distributed environment is initialized, before any sizable GPU allocation, until `POST /gate/activate` is sent to this port on the host of the first rank. This lets an external orchestrator defer the memory hungry part of startup to a safe window. Defaults to None, which disables the gate.",
NS("parallel"),
] = None
nnodes: A[int, "The number of nodes.", NS("parallel")] = 1
node_rank: A[int, "The node rank.", NS("parallel")] = 0
tp_size: A[