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
@@ -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