Support gated launch to defer startup memory allocation (#35927)
This commit is contained in:
@@ -20,6 +20,7 @@ from sglang.srt.distributed import (
|
|||||||
set_mscclpp_all_reduce,
|
set_mscclpp_all_reduce,
|
||||||
set_torch_symm_mem_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 (
|
from sglang.srt.distributed.parallel_state import (
|
||||||
_tag_groups_for_flashinfer_allreduce_only,
|
_tag_groups_for_flashinfer_allreduce_only,
|
||||||
)
|
)
|
||||||
@@ -132,6 +133,10 @@ def init_torch_distributed(
|
|||||||
):
|
):
|
||||||
_prewarm_tp_lm_head_all_to_all()
|
_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(
|
pre_model_load_memory = get_available_gpu_memory(
|
||||||
device,
|
device,
|
||||||
ps.gpu_id,
|
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
|
||||||
@@ -1046,6 +1046,11 @@ class ServerArgs:
|
|||||||
),
|
),
|
||||||
NS("parallel"),
|
NS("parallel"),
|
||||||
] = None
|
] = 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
|
nnodes: A[int, "The number of nodes.", NS("parallel")] = 1
|
||||||
node_rank: A[int, "The node rank.", NS("parallel")] = 0
|
node_rank: A[int, "The node rank.", NS("parallel")] = 0
|
||||||
tp_size: A[
|
tp_size: A[
|
||||||
|
|||||||
@@ -0,0 +1,124 @@
|
|||||||
|
import os
|
||||||
|
import subprocess
|
||||||
|
import time
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import psutil
|
||||||
|
import requests
|
||||||
|
|
||||||
|
from sglang.srt.utils.common import kill_process_tree
|
||||||
|
from sglang.srt.utils.network import get_open_port
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
|
||||||
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
DEFAULT_URL_FOR_TEST,
|
||||||
|
CustomTestCase,
|
||||||
|
)
|
||||||
|
|
||||||
|
register_cuda_ci(
|
||||||
|
est_time=180, stage="nightly", runner_config="1-gpu-large", nightly=True
|
||||||
|
)
|
||||||
|
|
||||||
|
MEM_FRACTION_STATIC = 0.6
|
||||||
|
GATED_MEMORY_CEILING_MB = 8 * 1024
|
||||||
|
SERVING_MEMORY_FLOOR_MB = 8 * 1024
|
||||||
|
|
||||||
|
|
||||||
|
class TestGatedLaunch(CustomTestCase):
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
_, host, port = cls.base_url.split(":")
|
||||||
|
cls.gate_port = get_open_port()
|
||||||
|
cls.gate_url = f"http:{host}:{cls.gate_port}"
|
||||||
|
|
||||||
|
command = [
|
||||||
|
"sglang",
|
||||||
|
"serve",
|
||||||
|
"--model-path",
|
||||||
|
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
|
||||||
|
"--host",
|
||||||
|
host[2:],
|
||||||
|
"--port",
|
||||||
|
port,
|
||||||
|
"--gated-launch-port",
|
||||||
|
str(cls.gate_port),
|
||||||
|
"--mem-fraction-static",
|
||||||
|
str(MEM_FRACTION_STATIC),
|
||||||
|
]
|
||||||
|
cls.process = subprocess.Popen(command, env=os.environ.copy())
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
if hasattr(cls, "process") and cls.process:
|
||||||
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
def test_gated_launch_defers_startup_until_activated(self):
|
||||||
|
"""The engine holds off every sizable allocation until it is activated."""
|
||||||
|
self._wait_for_health(self.gate_url, timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH)
|
||||||
|
|
||||||
|
with self.assertRaises(requests.exceptions.RequestException):
|
||||||
|
requests.get(f"{self.base_url}/health", timeout=5)
|
||||||
|
|
||||||
|
gated_memory_mb = self._device_memory_mb()
|
||||||
|
self.assertLess(gated_memory_mb, GATED_MEMORY_CEILING_MB)
|
||||||
|
|
||||||
|
for _ in range(2):
|
||||||
|
response = requests.post(f"{self.gate_url}/gate/activate", timeout=5)
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
|
||||||
|
self._wait_for_health(self.base_url, timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH)
|
||||||
|
|
||||||
|
response = requests.post(
|
||||||
|
f"{self.base_url}/generate",
|
||||||
|
json={
|
||||||
|
"text": "The capital of France is",
|
||||||
|
"sampling_params": {"max_new_tokens": 8, "temperature": 0},
|
||||||
|
},
|
||||||
|
timeout=60,
|
||||||
|
)
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
self.assertTrue(response.json()["text"])
|
||||||
|
|
||||||
|
self.assertGreater(self._device_memory_mb(), SERVING_MEMORY_FLOOR_MB)
|
||||||
|
|
||||||
|
def _wait_for_health(self, url: str, timeout: float) -> None:
|
||||||
|
deadline = time.perf_counter() + timeout
|
||||||
|
while time.perf_counter() < deadline:
|
||||||
|
self.assertIsNone(
|
||||||
|
self.process.poll(), msg=f"server died while waiting for {url}"
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
if requests.get(f"{url}/health", timeout=5).status_code == 200:
|
||||||
|
return
|
||||||
|
except requests.exceptions.RequestException:
|
||||||
|
pass
|
||||||
|
time.sleep(1)
|
||||||
|
self.fail(f"{url} did not become healthy within {timeout}s")
|
||||||
|
|
||||||
|
def _device_memory_mb(self) -> int:
|
||||||
|
parent = psutil.Process(self.process.pid)
|
||||||
|
pids = {parent.pid} | {child.pid for child in parent.children(recursive=True)}
|
||||||
|
|
||||||
|
output = subprocess.check_output(
|
||||||
|
[
|
||||||
|
"nvidia-smi",
|
||||||
|
"--query-compute-apps=pid,used_gpu_memory",
|
||||||
|
"--format=csv,noheader,nounits",
|
||||||
|
],
|
||||||
|
text=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
total_mb = 0
|
||||||
|
for line in output.splitlines():
|
||||||
|
if not line.strip():
|
||||||
|
continue
|
||||||
|
pid, used_mb = (field.strip() for field in line.split(","))
|
||||||
|
if int(pid) in pids:
|
||||||
|
total_mb += int(used_mb)
|
||||||
|
return total_mb
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,197 @@
|
|||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=30, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
import multiprocessing
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import requests
|
||||||
|
import torch.distributed as dist
|
||||||
|
|
||||||
|
from sglang.srt.distributed import gated_launch
|
||||||
|
from sglang.srt.distributed.gated_launch import (
|
||||||
|
POLL_INTERVAL_SECONDS,
|
||||||
|
_GatedLaunchServer,
|
||||||
|
_wait_until_activated,
|
||||||
|
maybe_wait_for_gated_launch,
|
||||||
|
)
|
||||||
|
from sglang.srt.utils.network import get_open_port
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
_JOIN_TIMEOUT_SECONDS = 120
|
||||||
|
|
||||||
|
|
||||||
|
def _fake_world_group(*, rank: int, world_size: int, cpu_group=None):
|
||||||
|
return SimpleNamespace(
|
||||||
|
rank=rank,
|
||||||
|
rank_in_group=rank,
|
||||||
|
ranks=list(range(world_size)),
|
||||||
|
world_size=world_size,
|
||||||
|
cpu_group=cpu_group,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _activate_over_http(base_url: str, delay: float):
|
||||||
|
time.sleep(delay)
|
||||||
|
deadline = time.perf_counter() + 30
|
||||||
|
while time.perf_counter() < deadline:
|
||||||
|
try:
|
||||||
|
if requests.post(f"{base_url}/gate/activate", timeout=1).status_code == 200:
|
||||||
|
return
|
||||||
|
except requests.exceptions.RequestException:
|
||||||
|
pass
|
||||||
|
time.sleep(0.1)
|
||||||
|
|
||||||
|
|
||||||
|
def _gated_launch_worker(
|
||||||
|
rank: int, dist_port: int, gate_port: int, activate_after: float, out
|
||||||
|
):
|
||||||
|
dist.init_process_group(
|
||||||
|
backend="gloo",
|
||||||
|
init_method=f"tcp://127.0.0.1:{dist_port}",
|
||||||
|
rank=rank,
|
||||||
|
world_size=2,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
if rank == 0:
|
||||||
|
threading.Thread(
|
||||||
|
target=_activate_over_http,
|
||||||
|
args=(f"http://127.0.0.1:{gate_port}", activate_after),
|
||||||
|
daemon=True,
|
||||||
|
).start()
|
||||||
|
|
||||||
|
world_group = _fake_world_group(
|
||||||
|
rank=rank, world_size=2, cpu_group=dist.group.WORLD
|
||||||
|
)
|
||||||
|
started_at = time.perf_counter()
|
||||||
|
with patch.object(gated_launch, "get_world_group", return_value=world_group):
|
||||||
|
maybe_wait_for_gated_launch(host="127.0.0.1", port=gate_port)
|
||||||
|
out.put((rank, time.perf_counter() - started_at))
|
||||||
|
finally:
|
||||||
|
dist.destroy_process_group()
|
||||||
|
|
||||||
|
|
||||||
|
class TestGatedLaunchServer(CustomTestCase):
|
||||||
|
def setUp(self):
|
||||||
|
self.server = _GatedLaunchServer()
|
||||||
|
self.port = get_open_port()
|
||||||
|
self.server.serve(host="127.0.0.1", port=self.port)
|
||||||
|
self.base_url = f"http://127.0.0.1:{self.port}"
|
||||||
|
self._wait_until_listening()
|
||||||
|
|
||||||
|
def _wait_until_listening(self):
|
||||||
|
deadline = time.perf_counter() + 30
|
||||||
|
while time.perf_counter() < deadline:
|
||||||
|
try:
|
||||||
|
requests.get(f"{self.base_url}/health", timeout=1)
|
||||||
|
return
|
||||||
|
except requests.exceptions.RequestException:
|
||||||
|
time.sleep(0.1)
|
||||||
|
self.fail(f"control server did not start listening on port {self.port}")
|
||||||
|
|
||||||
|
def test_health_is_served_while_the_gate_is_still_closed(self):
|
||||||
|
"""The control port answers before activation so a caller can find it."""
|
||||||
|
response = requests.get(f"{self.base_url}/health", timeout=5)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
self.assertFalse(self.server.activated)
|
||||||
|
|
||||||
|
def test_activate_flips_the_flag_and_stays_successful_when_repeated(self):
|
||||||
|
"""A retried activation still succeeds instead of erroring or toggling back."""
|
||||||
|
first = requests.post(f"{self.base_url}/gate/activate", timeout=5)
|
||||||
|
self.assertEqual(first.status_code, 200)
|
||||||
|
self.assertTrue(self.server.activated)
|
||||||
|
|
||||||
|
second = requests.post(f"{self.base_url}/gate/activate", timeout=5)
|
||||||
|
self.assertEqual(second.status_code, 200)
|
||||||
|
self.assertTrue(self.server.activated)
|
||||||
|
|
||||||
|
def test_activating_one_server_leaves_another_one_closed(self):
|
||||||
|
"""The route acts on its own server instead of process wide state."""
|
||||||
|
other = _GatedLaunchServer()
|
||||||
|
other_port = get_open_port()
|
||||||
|
other.serve(host="127.0.0.1", port=other_port)
|
||||||
|
|
||||||
|
requests.post(f"{self.base_url}/gate/activate", timeout=5)
|
||||||
|
|
||||||
|
self.assertTrue(self.server.activated)
|
||||||
|
self.assertFalse(other.activated)
|
||||||
|
|
||||||
|
|
||||||
|
class TestWaitUntilActivated(CustomTestCase):
|
||||||
|
def test_single_rank_keeps_polling_until_the_flag_is_set(self):
|
||||||
|
"""A lone rank leaves the gate only after its own flag flips."""
|
||||||
|
server = _GatedLaunchServer()
|
||||||
|
activate_after = 2 * POLL_INTERVAL_SECONDS
|
||||||
|
threading.Timer(
|
||||||
|
activate_after, lambda: setattr(server, "activated", True)
|
||||||
|
).start()
|
||||||
|
|
||||||
|
started_at = time.perf_counter()
|
||||||
|
_wait_until_activated(
|
||||||
|
world_group=_fake_world_group(rank=0, world_size=1), server=server
|
||||||
|
)
|
||||||
|
elapsed = time.perf_counter() - started_at
|
||||||
|
|
||||||
|
self.assertGreaterEqual(elapsed, activate_after)
|
||||||
|
|
||||||
|
def test_second_rank_learns_about_activation_through_the_cpu_group(self):
|
||||||
|
"""Driven through maybe_wait_for_gated_launch: the rank without the control server is released by the gloo broadcast."""
|
||||||
|
context = multiprocessing.get_context("spawn")
|
||||||
|
out = context.Queue()
|
||||||
|
dist_port = get_open_port()
|
||||||
|
gate_port = get_open_port()
|
||||||
|
activate_after = 2 * POLL_INTERVAL_SECONDS
|
||||||
|
|
||||||
|
processes = [
|
||||||
|
context.Process(
|
||||||
|
target=_gated_launch_worker,
|
||||||
|
args=(rank, dist_port, gate_port, activate_after, out),
|
||||||
|
)
|
||||||
|
for rank in range(2)
|
||||||
|
]
|
||||||
|
for process in processes:
|
||||||
|
process.start()
|
||||||
|
|
||||||
|
elapsed_by_rank = {}
|
||||||
|
for _ in processes:
|
||||||
|
rank, elapsed = out.get(timeout=_JOIN_TIMEOUT_SECONDS)
|
||||||
|
elapsed_by_rank[rank] = elapsed
|
||||||
|
|
||||||
|
for process in processes:
|
||||||
|
process.join(timeout=_JOIN_TIMEOUT_SECONDS)
|
||||||
|
for process in processes:
|
||||||
|
self.assertEqual(process.exitcode, 0)
|
||||||
|
self.assertEqual(sorted(elapsed_by_rank), [0, 1])
|
||||||
|
self.assertGreaterEqual(elapsed_by_rank[1], activate_after)
|
||||||
|
|
||||||
|
|
||||||
|
class TestMaybeWaitForGatedLaunch(CustomTestCase):
|
||||||
|
def setUp(self):
|
||||||
|
self.addCleanup(setattr, gated_launch, "_instance", None)
|
||||||
|
gated_launch._instance = None
|
||||||
|
|
||||||
|
def test_an_unset_port_leaves_the_startup_path_untouched(self):
|
||||||
|
"""Without the flag the gate never reaches the distributed environment."""
|
||||||
|
with patch.object(gated_launch, "get_world_group") as get_world_group:
|
||||||
|
maybe_wait_for_gated_launch(host="127.0.0.1", port=None)
|
||||||
|
|
||||||
|
get_world_group.assert_not_called()
|
||||||
|
self.assertIsNone(gated_launch._instance)
|
||||||
|
|
||||||
|
def test_a_second_call_in_the_same_process_does_not_gate_again(self):
|
||||||
|
"""A draft worker re-entering the init path must not wait a second time."""
|
||||||
|
gated_launch._instance = _GatedLaunchServer()
|
||||||
|
|
||||||
|
with patch.object(gated_launch, "get_world_group") as get_world_group:
|
||||||
|
maybe_wait_for_gated_launch(host="127.0.0.1", port=get_open_port())
|
||||||
|
|
||||||
|
get_world_group.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user