Keep graph-pool borrows on their allocation stream (#39180)

Co-authored-by: cctry <17473714+cctry@users.noreply.github.com>
This commit is contained in:
cctry
2026-09-12 19:26:09 -07:00
committed by GitHub
co-authored by cctry
parent 18cc55dc0b
commit 206034e520
4 changed files with 85 additions and 11 deletions
+9 -5
View File
@@ -1107,9 +1107,15 @@ class Scheduler(
self.init_all_cuda_graphs()
model_runner = self.tp_worker.model_runner
with torch.get_device_module(model_runner.device).stream(
device_module = torch.get_device_module(model_runner.device)
self.schedule_stream = None if use_mlx() else device_module.Stream(priority=0)
# Match run_batch / _pp_launch_batch so warmup allocations stay reusable.
forward_stream = (
model_runner.forward_stream
):
if self.enable_overlap or self.ps.pp_size > 1 or use_mlx()
else self.schedule_stream
)
with device_module.stream(forward_stream):
if self.draft_worker is None:
model_runner.prewarm_sampling()
else:
@@ -1847,7 +1853,6 @@ class Scheduler(
def run_event_loop(self) -> None:
"""Run the scheduler's event loop.
Sets up the schedule stream and dispatches to the appropriate event loop.
The event loop blocks until shutdown.
"""
# Engine init (graph capture, warmups) is done; from here on any
@@ -1862,10 +1867,9 @@ class Scheduler(
dispatch_event_loop(self)
return
self.schedule_stream = self.device_module.Stream(priority=0)
if self.device == "cpu":
self.schedule_stream.synchronize = lambda: None # No-op for CPU
elif is_cuda() or _is_hip:
elif (is_cuda() or _is_hip) and (self.enable_overlap or self.ps.pp_size > 1):
# CUDA/HIP streams come from a fixed round-robin pool. Redraw if this
# stream aliases forward_stream, which would eliminate scheduler
# overlap. Only CUDA/HIP streams expose a ``cuda_stream`` handle;
@@ -45,6 +45,7 @@ class GraphPoolBorrowState:
active_user: Optional[str] = None
stub: Optional[BumpArenaStub] = None
mem_pool: Optional[torch.cuda.MemPool] = None
stream: Optional[torch.cuda.Stream] = None
disabled_reason: Optional[str] = None
static_runs: Optional[list[tuple[int, int]]] = None
check_pending: bool = False
@@ -279,6 +280,7 @@ def _teardown_borrow_pool() -> None:
torch.cuda.synchronize()
torch.empty(1, device="cuda")
state.mem_pool = None
state.stream = None
_PRECARVE_MIN_RUN_BYTES = 64 << 20
@@ -312,6 +314,8 @@ def _precarve_run_segments(runs: list[tuple[int, int]]) -> None:
def borrow_graph_pool(user: str) -> Iterator[None]:
"""Route this thread's torch allocations onto the graph pool's free runs.
All borrows must use the stream that first creates the borrow pool, so
the caching allocator can reuse its pre-carved segments.
Tensors allocated inside must be released before the next graph replay,
which rewrites their bytes; the next replay (or pool teardown) raises if
any are still referenced. An allocation no run can hold raises the
@@ -322,7 +326,13 @@ def borrow_graph_pool(user: str) -> Iterator[None]:
yield
return
with graph_pool_user_scope(user):
stream = torch.cuda.current_stream()
if state.mem_pool is not None:
if stream != state.stream:
raise RuntimeError(
"Graph-pool borrow must use the stream that created the borrow pool: "
f"expected {state.stream}, got {stream}"
)
# Return completed cross-stream frees to the cache. The allocator
# processes their events on a later allocation.
torch.empty(1, device="cuda")
@@ -346,6 +356,7 @@ def borrow_graph_pool(user: str) -> Iterator[None]:
# stream-ordered deferred frees remain allocator-managed. Capture
# retires it because capture changes the underlying free extents.
state.mem_pool = torch.cuda.MemPool(state.stub.allocator)
state.stream = stream
with torch.cuda.use_mem_pool(state.mem_pool):
_precarve_run_segments(runs)
# Only growth beyond the precarve is worth another log line.
@@ -17,6 +17,7 @@ maybe_stub_sgl_kernel()
from sglang.srt.configs.device_config import DeviceConfig
from sglang.srt.configs.load_config import LoadConfig, LoadFormat
from sglang.srt.configs.model_config import ModelImpl
from sglang.srt.managers.scheduler import Scheduler
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.model_executor.cuda_graph_config import Backend, CudaGraphConfig
from sglang.srt.model_executor.model_runner import ModelRunner
@@ -728,9 +729,9 @@ class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
reset_context()
self.addCleanup(reset_context)
def _scheduler(self, worker, trace, *, mode, draft_worker=None):
from sglang.srt.managers.scheduler import Scheduler
def _scheduler(
self, worker, trace, *, mode, draft_worker=None, enable_overlap=True, pp_size=1
):
# The schedule reads the mode from the bags, so the test states it by
# publishing a record rather than by standing one in.
reset_context()
@@ -739,6 +740,8 @@ class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
role="scheduler",
)
scheduler = Scheduler.__new__(Scheduler)
scheduler.enable_overlap = enable_overlap
scheduler.ps = SimpleNamespace(pp_size=pp_size)
scheduler.init_tp_model_worker = lambda: setattr(scheduler, "tp_worker", worker)
scheduler.maybe_init_draft_worker = lambda: setattr(
scheduler, "draft_worker", draft_worker
@@ -748,7 +751,9 @@ class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
scheduler.init_all_cuda_graphs = lambda: trace.append("capture")
return scheduler
def _run_startup(self, mode, *, use_draft_worker=False):
def _run_startup(
self, mode, *, use_draft_worker=False, enable_overlap=True, pp_size=1, mlx=False
):
trace = []
worker = _SchedulerWorker(trace, post_capture_active=True)
draft_worker = (
@@ -764,6 +769,14 @@ class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
trace,
mode=mode,
draft_worker=draft_worker,
enable_overlap=enable_overlap,
pp_size=pp_size,
)
schedule_stream = object()
expected_stream = (
worker.model_runner.forward_stream
if enable_overlap or pp_size > 1 or mlx
else schedule_stream
)
class StreamContext:
@@ -774,7 +787,7 @@ class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
trace.append("stream_exit")
def stream_context(stream):
self.assertIs(stream, worker.model_runner.forward_stream)
self.assertIs(stream, expected_stream)
return StreamContext()
def stop_after_startup():
@@ -783,6 +796,7 @@ class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
scheduler.spec_algorithm = SimpleNamespace(is_none=stop_after_startup)
with (
patch("sglang.srt.managers.scheduler.use_mlx", return_value=mlx),
patch(
"sglang.srt.managers.scheduler.get_exec",
return_value=SimpleNamespace(
@@ -794,12 +808,15 @@ class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
),
patch(
"sglang.srt.managers.scheduler.torch.get_device_module",
return_value=SimpleNamespace(stream=stream_context),
return_value=SimpleNamespace(
stream=stream_context, Stream=lambda priority: schedule_stream
),
),
self.assertRaisesRegex(RuntimeError, "stop after startup"),
):
scheduler.init_model_worker()
self.assertIs(scheduler.schedule_stream, None if mlx else schedule_stream)
worker.model_runner.post_capture_resize_kv_pool.assert_called_once_with(
draft_runners=(worker.model_runner,) if use_draft_worker else ()
)
@@ -819,6 +836,18 @@ class TestStartupWeightLoadSchedulerRouting(CustomTestCase):
],
)
def test_sampling_warmup_uses_the_serving_stream(self):
for enable_overlap, pp_size, mlx in (
(False, 1, False),
(False, 2, False),
(True, 1, False),
(False, 1, True),
):
with self.subTest(enable_overlap=enable_overlap, pp_size=pp_size, mlx=mlx):
self._run_startup(
"serial", enable_overlap=enable_overlap, pp_size=pp_size, mlx=mlx
)
def test_overlap_starts_before_capture_and_finalizes_after(self):
self.assertEqual(
self._run_startup("overlap"),
@@ -114,6 +114,7 @@ class TestGraphPoolBorrow(CustomTestCase):
patch.object(self.state, "mem_pool", None),
patch.object(pool.torch.cuda, "MemPool"),
patch.object(pool.torch.cuda, "use_mem_pool"),
patch.object(pool.torch.cuda, "current_stream"),
patch.object(pool.torch.cuda, "memory_snapshot", return_value=snapshot),
pool.borrow_graph_pool(user="test"),
):
@@ -127,15 +128,18 @@ class TestGraphPoolBorrow(CustomTestCase):
def test_high_cursor_keeps_reusable_cached_segments(self):
stub = MagicMock(cursor_bytes=600, freed_bytes=0)
mem_pool = MagicMock()
stream = object()
with (
patch.object(pool, "graph_pool_borrow_enabled", return_value=True),
patch.object(self.state, "stub", stub),
patch.object(self.state, "mem_pool", mem_pool),
patch.object(self.state, "stream", stream),
patch.object(self.state, "extents_total", 1000),
patch.object(pool, "_teardown_borrow_pool") as teardown,
patch.object(pool.torch, "empty"),
patch.object(pool.torch.cuda, "use_mem_pool"),
patch.object(pool.torch.cuda, "current_stream", return_value=stream),
):
with pool.borrow_graph_pool(user="test"):
pass
@@ -282,9 +286,35 @@ class TestGraphPoolBorrow(CustomTestCase):
self.assertEqual(c.data_ptr(), recycled_address)
del b, c
self.assertEqual(self.state.stream, torch.cuda.current_stream())
with (
torch.cuda.stream(stream),
self.assertRaisesRegex(
RuntimeError, "stream that created the borrow pool"
),
):
with pool.borrow_graph_pool(user="wrong stream"):
self.fail("cross-stream borrowing must fail before allocating")
self.assertIsNone(self.state.active_user)
with pool.borrow_graph_pool(user="same stream"):
reused = torch.empty(40 << 20, dtype=torch.uint8, device="cuda")
self.assertEqual(reused.data_ptr(), recycled_address)
del reused
# Captures retire the persistent borrow pool. Its storage aliases
# existing graph-pool runs, so the reserved footprint is unchanged.
pool._teardown_borrow_pool()
self.assertIsNone(self.state.stream)
with torch.cuda.stream(stream), pool.borrow_graph_pool(user="new pool"):
self.assertEqual(self.state.stream, stream)
reused = torch.empty(40 << 20, dtype=torch.uint8, device="cuda")
self.assertTrue(on_a_run(reused))
del reused
pool._teardown_borrow_pool()
with pool.graph_pool_replay_scope():
graph.replay()
torch.cuda.synchronize()
self.assertTrue(torch.equal(y, torch.ones_like(y)))
self.assertEqual(torch.cuda.memory_reserved(device_id), reserved_before)
del graph, y