Scope graph-pool borrowing to the runtime and reduce fragmentation (#39177)

Co-authored-by: cctry <17473714+cctry@users.noreply.github.com>
This commit is contained in:
cctry
2026-09-12 11:45:17 -07:00
committed by GitHub
co-authored by cctry
parent 2784a86062
commit 7ae4af8187
3 changed files with 270 additions and 89 deletions
@@ -21,6 +21,8 @@ from __future__ import annotations
import logging import logging
from contextlib import contextmanager from contextlib import contextmanager
from dataclasses import dataclass, field
from functools import cache
from typing import Any, Iterator, Optional from typing import Any, Iterator, Optional
import torch import torch
@@ -31,21 +33,50 @@ from sglang.srt.utils import is_cuda
from sglang.srt.utils.cuda_vmm_utils import BumpArenaStub from sglang.srt.utils.cuda_vmm_utils import BumpArenaStub
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_active_graph_pool_user: Optional[str] = None
_borrow_stub: Optional[BumpArenaStub] = None
_borrow_mem_pool: Optional[torch.cuda.MemPool] = None
_borrow_disabled_reason: Optional[str] = None
_borrow_static_runs: Optional[list[tuple[int, int]]] = None
_borrow_extents_total = 0
_largest_logged_graph_pool_borrow = 0
_MIB = 1 << 20
_CAPTURE_STREAM_NAME = "cuda_graph_capture" _CAPTURE_STREAM_NAME = "cuda_graph_capture"
@dataclass(eq=False)
class GraphPoolBorrowState:
"""Mutable borrowing state and free-run cache for one runtime."""
active_user: Optional[str] = None
stub: Optional[BumpArenaStub] = None
mem_pool: Optional[torch.cuda.MemPool] = None
disabled_reason: Optional[str] = None
static_runs: Optional[list[tuple[int, int]]] = None
check_pending: bool = False
extents_total: int = 0
largest_logged_borrow: int = 0
snapshot_free_runs: Any = field(
default_factory=lambda: cache(find_free_graph_pool_runs), init=False, repr=False
)
def _get_graph_pool_borrow_state() -> GraphPoolBorrowState:
resources = get_resources()
if resources.graph_pool_borrow is None:
resources.graph_pool_borrow = GraphPoolBorrowState()
return resources.graph_pool_borrow
def _log_graph_pool_borrow_capacity(runs: list[tuple[int, int]]) -> None:
total_bytes = sum(nbytes for _, nbytes in runs)
largest_bytes = max((nbytes for _, nbytes in runs), default=0)
logger.info(
"Graph pool borrow capacity: %.1f MiB available, largest contiguous "
"region %.1f MiB, %d regions",
total_bytes / _MIB,
largest_bytes / _MIB,
len(runs),
)
def disable_graph_pool_borrow(reason: str) -> None: def disable_graph_pool_borrow(reason: str) -> None:
"""Disable borrowing when graph storage is managed outside the shared pool.""" """Disable borrowing when graph storage is managed outside the shared pool."""
global _borrow_disabled_reason _get_graph_pool_borrow_state().disabled_reason = reason
_borrow_disabled_reason = reason
_teardown_borrow_pool() _teardown_borrow_pool()
logger.info("Graph pool borrow disabled: %s", reason) logger.info("Graph pool borrow disabled: %s", reason)
@@ -57,17 +88,13 @@ def set_graph_pool_borrow_runs(runs: list[tuple[int, int]]) -> None:
remain stable for the process lifetime. Registering an empty list disables remain stable for the process lifetime. Registering an empty list disables
borrowing. borrowing.
""" """
global _borrow_static_runs state = _get_graph_pool_borrow_state()
static_runs = sorted(runs, key=lambda run: run[1], reverse=True)[ static_runs = sorted(runs, key=lambda run: run[1], reverse=True)[
: BumpArenaStub.MAX_EXTENTS : BumpArenaStub.MAX_EXTENTS
] ]
_teardown_borrow_pool() _teardown_borrow_pool()
_borrow_static_runs = static_runs state.static_runs = static_runs
logger.info( _log_graph_pool_borrow_capacity(state.static_runs)
"Graph pool borrow runs pinned: runs=%d free=%d",
len(_borrow_static_runs),
sum(nbytes for _, nbytes in _borrow_static_runs),
)
def get_global_graph_memory_pool() -> Optional[Any]: def get_global_graph_memory_pool() -> Optional[Any]:
@@ -134,31 +161,32 @@ class GraphPoolPrecarve:
def graph_pool_borrow_enabled() -> bool: def graph_pool_borrow_enabled() -> bool:
state = _get_graph_pool_borrow_state()
if ( if (
_borrow_disabled_reason is not None state.disabled_reason is not None
or not envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.get() or not envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.get()
or not is_cuda() or not is_cuda()
): ):
return False return False
if _borrow_static_runs is not None: if state.static_runs is not None:
return len(_borrow_static_runs) > 0 return len(state.static_runs) > 0
return get_global_graph_memory_pool() is not None return get_global_graph_memory_pool() is not None
@contextmanager @contextmanager
def graph_pool_user_scope(user: str) -> Iterator[None]: def graph_pool_user_scope(user: str) -> Iterator[None]:
global _active_graph_pool_user state = _get_graph_pool_borrow_state()
# Graph replay silently overwrites aliases of its allocator-free blocks. # Graph replay silently overwrites aliases of its allocator-free blocks.
if _active_graph_pool_user is not None: if state.active_user is not None:
raise RuntimeError( raise RuntimeError(
f"graph pool already has live user {_active_graph_pool_user!r}; " f"graph pool already has live user {state.active_user!r}; "
f"cannot use it for {user!r}" f"cannot use it for {user!r}"
) )
_active_graph_pool_user = user state.active_user = user
try: try:
yield yield
finally: finally:
_active_graph_pool_user = None state.active_user = None
@contextmanager @contextmanager
@@ -167,6 +195,8 @@ def graph_pool_replay_scope() -> Iterator[None]:
yield yield
return return
with graph_pool_user_scope("CUDA graph"): with graph_pool_user_scope("CUDA graph"):
if _get_graph_pool_borrow_state().check_pending:
_raise_on_live_borrows("graph replay")
yield yield
@@ -205,68 +235,126 @@ def find_free_graph_pool_runs(pool_id: Any) -> list[tuple[int, int]]:
return runs return runs
def graph_pool_borrow_largest_run() -> int:
"""Largest contiguous free extent a single borrow can occupy, in bytes."""
if not graph_pool_borrow_enabled():
return 0
state = _get_graph_pool_borrow_state()
if state.static_runs is not None:
return state.static_runs[0][1]
runs = state.snapshot_free_runs(get_global_graph_memory_pool())
return runs[0][1] if runs else 0
def _raise_on_live_borrows(event: str) -> None:
"""Reject borrows that graph replay or pool teardown would overwrite."""
state = _get_graph_pool_borrow_state()
state.check_pending = False
if state.mem_pool is None:
return
live = sum(
block["size"]
for segment in torch.cuda.memory_snapshot(
state.mem_pool.id, include_traces=False
)
for block in segment["blocks"]
if block["state"] == "active_allocated"
)
if live:
raise RuntimeError(
f"Graph-pool borrow leak at {event}: {live} bytes still referenced"
)
def _teardown_borrow_pool() -> None: def _teardown_borrow_pool() -> None:
"""Retire the persistent borrow pool after draining deferred frees.""" """Retire the persistent borrow pool after draining deferred frees."""
global _borrow_mem_pool state = _get_graph_pool_borrow_state()
if _borrow_mem_pool is None: state.snapshot_free_runs.cache_clear()
if state.mem_pool is None:
return return
if state.check_pending:
_raise_on_live_borrows("borrow pool teardown")
# Borrowed blocks that saw cross-stream use can remain in event limbo. # Borrowed blocks that saw cross-stream use can remain in event limbo.
# Synchronize, then drive allocator event processing before dropping the pool. # Synchronize, then drive allocator event processing before dropping the pool.
torch.cuda.synchronize() torch.cuda.synchronize()
torch.empty(1, device="cuda") torch.empty(1, device="cuda")
_borrow_mem_pool = None state.mem_pool = None
_PRECARVE_MIN_RUN_BYTES = 64 << 20
# Left uncarved per run so 2 MiB small-pool segments keep a home.
_PRECARVE_SMALL_RESERVE_BYTES = 32 << 20
# The caching allocator rounds large segment requests up to 2 MiB.
_PRECARVE_GRANULARITY = 2 << 20
def _precarve_run_segments(runs: list[tuple[int, int]]) -> None:
"""Seed coalescible segments on the stream that will allocate borrows."""
for _, run_bytes in runs:
seed = (
(run_bytes - _PRECARVE_SMALL_RESERVE_BYTES) // _PRECARVE_GRANULARITY
) * _PRECARVE_GRANULARITY
if seed >= _PRECARVE_MIN_RUN_BYTES:
torch.empty(seed, dtype=torch.uint8, device="cuda")
@contextmanager @contextmanager
def borrow_graph_pool(user: str) -> Iterator[None]: def borrow_graph_pool(user: str) -> Iterator[None]:
"""Route this thread's torch allocations onto the graph pool's free runs. """Route this thread's torch allocations onto the graph pool's free runs.
Tensors allocated inside must not survive the current step: the next graph Tensors allocated inside must be released before the next graph replay,
replay may overwrite their bytes. An allocation no run can hold raises the which rewrites their bytes; the next replay (or pool teardown) raises if
any are still referenced. An allocation no run can hold raises the
allocator's normal OOM. This is a no-op while borrowing is disabled. allocator's normal OOM. This is a no-op while borrowing is disabled.
""" """
global _borrow_stub, _borrow_mem_pool, _borrow_extents_total state = _get_graph_pool_borrow_state()
global _largest_logged_graph_pool_borrow
if not graph_pool_borrow_enabled(): if not graph_pool_borrow_enabled():
yield yield
return return
with graph_pool_user_scope(user): with graph_pool_user_scope(user):
if _borrow_mem_pool is not None: if state.mem_pool is not None:
# Return completed cross-stream frees to the cache. The allocator # Return completed cross-stream frees to the cache. The allocator
# processes their events on a later allocation. # processes their events on a later allocation.
torch.empty(1, device="cuda") torch.empty(1, device="cuda")
if _borrow_stub.freed_bytes: if state.stub.freed_bytes:
# Rebuild if the caching allocator released an arena segment. # The bump arena cannot reuse segments returned by empty_cache().
# A high cursor alone means the cache owns reusable segments;
# rebuilding discards them and can turn the next large borrow
# into a fragmented cold-allocation OOM.
_teardown_borrow_pool() _teardown_borrow_pool()
if _borrow_mem_pool is None: if state.mem_pool is None:
if _borrow_stub is None: if state.stub is None:
_borrow_stub = BumpArenaStub() state.stub = BumpArenaStub()
# Runs are sorted largest first, so first fit would carve every # Runs are sorted largest first, so first fit would carve every
# small allocation out of the run a probability matrix needs. # small allocation out of the run a probability matrix needs.
_borrow_stub.set_best_fit(True) state.stub.set_best_fit(True)
if _borrow_static_runs is not None: if state.static_runs is not None:
runs = _borrow_static_runs runs = state.static_runs
else: else:
runs = find_free_graph_pool_runs(get_global_graph_memory_pool())[ runs = state.snapshot_free_runs(get_global_graph_memory_pool())[
: BumpArenaStub.MAX_EXTENTS : BumpArenaStub.MAX_EXTENTS
] ]
_borrow_stub.set_extents(runs) state.stub.set_extents(runs)
# Keep one caching layer across borrows so normal block reuse and # Keep one caching layer across borrows so normal block reuse and
# stream-ordered deferred frees remain allocator-managed. Capture # stream-ordered deferred frees remain allocator-managed. Capture
# retires it because capture changes the underlying free extents. # retires it because capture changes the underlying free extents.
_borrow_mem_pool = torch.cuda.MemPool(_borrow_stub.allocator) state.mem_pool = torch.cuda.MemPool(state.stub.allocator)
_borrow_extents_total = sum(run_bytes for _, run_bytes in runs) with torch.cuda.use_mem_pool(state.mem_pool):
_precarve_run_segments(runs)
# Only growth beyond the precarve is worth another log line.
state.largest_logged_borrow = state.stub.cursor_bytes
state.extents_total = sum(run_bytes for _, run_bytes in runs)
_log_graph_pool_borrow_capacity(runs)
logger.info( logger.info(
"Graph pool borrow extents: runs=%d free=%d", "Graph pool borrow pre-carved: %.1f MiB",
len(runs), state.stub.cursor_bytes / _MIB,
_borrow_extents_total,
) )
with torch.cuda.use_mem_pool(_borrow_mem_pool): with torch.cuda.use_mem_pool(state.mem_pool):
yield yield
consumed = _borrow_stub.cursor_bytes state.check_pending = True
if consumed > _largest_logged_graph_pool_borrow: consumed = state.stub.cursor_bytes
logger.info("Graph pool borrow: consumed=%d", consumed) if consumed > state.largest_logged_borrow:
_largest_logged_graph_pool_borrow = consumed logger.info(
"Graph pool borrow high-water mark: %.1f MiB used of %.1f MiB "
"available",
consumed / _MIB,
state.extents_total / _MIB,
)
state.largest_logged_borrow = consumed
+2
View File
@@ -58,6 +58,7 @@ from typing import TYPE_CHECKING, Any, Dict, Optional
import msgspec import msgspec
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.model_executor.runner_utils.pool import GraphPoolBorrowState
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -607,6 +608,7 @@ class Resources(_FlagGroupBase):
# CUDA graph memory pool shared across the prefill and decode graph # CUDA graph memory pool shared across the prefill and decode graph
# backends (created lazily by model_executor.runner_utils.pool). # backends (created lazily by model_executor.runner_utils.pool).
graph_memory_pool: Any = None graph_memory_pool: Any = None
graph_pool_borrow: GraphPoolBorrowState | None = None
# EPLB: per-process recorder and the publish-once location metadata # EPLB: per-process recorder and the publish-once location metadata
# (owning accessors live in sglang.srt.eplb). # (owning accessors live in sglang.srt.eplb).
expert_distribution_recorder: Any = None expert_distribution_recorder: Any = None
@@ -1,5 +1,6 @@
"""CUDA graph-pool borrowing allocator and lifetime regression tests.""" """CUDA graph-pool borrowing allocator and lifetime regression tests."""
import contextlib
import unittest import unittest
from contextlib import contextmanager from contextlib import contextmanager
from types import SimpleNamespace from types import SimpleNamespace
@@ -17,32 +18,23 @@ from sglang.srt.speculative.dflash_worker_v2 import DFlashWorkerV2
from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=13, stage="base-b", runner_config="1-gpu-small") register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-small")
class TestGraphPoolBorrow(CustomTestCase): class TestGraphPoolBorrow(CustomTestCase):
def setUp(self): def setUp(self):
super().setUp() super().setUp()
self._reset_borrow_state() self.state = pool.GraphPoolBorrowState()
# ExitStack rather than enterContext, which is 3.11+.
stack = contextlib.ExitStack()
self.addCleanup(stack.close)
stack.enter_context(pool.get_resources().override(graph_pool_borrow=self.state))
def tearDown(self): def tearDown(self):
try: if torch.cuda.is_available():
if torch.cuda.is_available(): pool._teardown_borrow_pool()
pool._teardown_borrow_pool() torch.cuda.synchronize()
torch.cuda.synchronize() torch.cuda.empty_cache()
torch.cuda.empty_cache()
finally:
self._reset_borrow_state()
@staticmethod
def _reset_borrow_state():
pool._active_graph_pool_user = None
pool._borrow_stub = None
pool._borrow_mem_pool = None
pool._borrow_disabled_reason = None
pool._borrow_static_runs = None
pool._borrow_extents_total = 0
pool._largest_logged_graph_pool_borrow = 0
def test_mixed_segment_runs_exclude_live_blocks(self): def test_mixed_segment_runs_exclude_live_blocks(self):
"""A mixed segment's free runs are borrowable, but a returned run must """A mixed segment's free runs are borrowable, but a returned run must
@@ -63,6 +55,27 @@ class TestGraphPoolBorrow(CustomTestCase):
runs = pool.find_free_graph_pool_runs((0, 1)) runs = pool.find_free_graph_pool_runs((0, 1))
self.assertEqual(sorted(runs), [(0x1000, 4096), (0x3000, 4096)]) self.assertEqual(sorted(runs), [(0x1000, 4096), (0x3000, 4096)])
def test_free_run_snapshot_lifetime_follows_borrow_state(self):
runs = [{"blocks": [{"state": "inactive", "address": 0x1000, "size": 8192}]}]
with (
envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.override(True),
patch.object(pool, "is_cuda", return_value=True),
patch.object(pool, "get_global_graph_memory_pool", return_value=(1, 2)),
patch.object(
pool.torch.cuda, "memory_snapshot", side_effect=[runs, [], runs]
) as snapshot,
):
self.assertEqual(pool.graph_pool_borrow_largest_run(), 8192)
self.assertEqual(pool.graph_pool_borrow_largest_run(), 8192)
snapshot.assert_called_once_with((1, 2), include_traces=False)
pool._teardown_borrow_pool()
self.assertEqual(pool.graph_pool_borrow_largest_run(), 0)
self.assertEqual(snapshot.call_count, 2)
with pool.get_resources().override(graph_pool_borrow=None):
self.assertEqual(pool.graph_pool_borrow_largest_run(), 8192)
self.assertEqual(pool.graph_pool_borrow_largest_run(), 0)
self.assertEqual(snapshot.call_count, 3)
def test_graph_replay_fails_during_active_pool_borrow(self): def test_graph_replay_fails_during_active_pool_borrow(self):
graph = Mock() graph = Mock()
backend = object.__new__(FullCudaGraphBackend) backend = object.__new__(FullCudaGraphBackend)
@@ -78,10 +91,8 @@ class TestGraphPoolBorrow(CustomTestCase):
with ( with (
envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.override(True), envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.override(True),
patch.object(pool, "get_global_graph_memory_pool", return_value=(1, 2)), patch.object(pool, "get_global_graph_memory_pool", return_value=(1, 2)),
patch.object( patch.object(self.state, "stub", MagicMock(cursor_bytes=0, freed_bytes=0)),
pool, "_borrow_stub", MagicMock(cursor_bytes=0, freed_bytes=0) patch.object(self.state, "mem_pool", None),
),
patch.object(pool, "_borrow_mem_pool", None),
patch.object(pool.torch.cuda, "MemPool"), patch.object(pool.torch.cuda, "MemPool"),
patch.object(pool.torch.cuda, "use_mem_pool"), patch.object(pool.torch.cuda, "use_mem_pool"),
patch.object(pool.torch.cuda, "memory_snapshot", return_value=snapshot), patch.object(pool.torch.cuda, "memory_snapshot", return_value=snapshot),
@@ -100,9 +111,9 @@ class TestGraphPoolBorrow(CustomTestCase):
with ( with (
patch.object(pool, "graph_pool_borrow_enabled", return_value=True), patch.object(pool, "graph_pool_borrow_enabled", return_value=True),
patch.object(pool, "_borrow_stub", stub), patch.object(self.state, "stub", stub),
patch.object(pool, "_borrow_mem_pool", mem_pool), patch.object(self.state, "mem_pool", mem_pool),
patch.object(pool, "_borrow_extents_total", 1000), patch.object(self.state, "extents_total", 1000),
patch.object(pool, "_teardown_borrow_pool") as teardown, patch.object(pool, "_teardown_borrow_pool") as teardown,
patch.object(pool.torch, "empty"), patch.object(pool.torch, "empty"),
patch.object(pool.torch.cuda, "use_mem_pool"), patch.object(pool.torch.cuda, "use_mem_pool"),
@@ -126,7 +137,7 @@ class TestGraphPoolBorrow(CustomTestCase):
runs = [(0x1000, 4096), (0x2000, 8192)] runs = [(0x1000, 4096), (0x2000, 8192)]
def reset_static_runs(): def reset_static_runs():
pool._borrow_static_runs = None self.state.static_runs = None
with patch.object( with patch.object(
pool, "_teardown_borrow_pool", side_effect=reset_static_runs pool, "_teardown_borrow_pool", side_effect=reset_static_runs
@@ -134,7 +145,7 @@ class TestGraphPoolBorrow(CustomTestCase):
pool.set_graph_pool_borrow_runs(runs) pool.set_graph_pool_borrow_runs(runs)
teardown.assert_called_once_with() teardown.assert_called_once_with()
self.assertEqual(pool._borrow_static_runs, [(0x2000, 8192), (0x1000, 4096)]) self.assertEqual(self.state.static_runs, [(0x2000, 8192), (0x1000, 4096)])
def test_eagle_non_greedy_probabilities_do_not_borrow_graph_pool(self): def test_eagle_non_greedy_probabilities_do_not_borrow_graph_pool(self):
def fake_sampling(**kwargs): def fake_sampling(**kwargs):
@@ -221,7 +232,6 @@ class TestGraphPoolBorrow(CustomTestCase):
with ( with (
envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.override(True), envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.override(True),
patch.object(pool, "get_global_graph_memory_pool", return_value=handle), patch.object(pool, "get_global_graph_memory_pool", return_value=handle),
patch.object(pool, "_borrow_mem_pool", None),
): ):
runs = pool.find_free_graph_pool_runs(handle) runs = pool.find_free_graph_pool_runs(handle)
self.assertGreaterEqual(len(runs), 2) self.assertGreaterEqual(len(runs), 2)
@@ -260,6 +270,86 @@ class TestGraphPoolBorrow(CustomTestCase):
self.assertEqual(torch.cuda.memory_reserved(device_id), reserved_before) self.assertEqual(torch.cuda.memory_reserved(device_id), reserved_before)
del graph, y del graph, y
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
def test_borrow_recovers_from_arena_fragmentation(self):
handle = torch.cuda.graph_pool_handle()
graph = torch.cuda.CUDAGraph()
seed = torch.zeros(8, device="cuda")
stream = torch.cuda.Stream()
with (
torch.cuda.stream(stream),
torch.cuda.graph(graph, pool=handle, stream=stream),
):
transient = torch.empty(200 << 20, dtype=torch.uint8, device="cuda")
keep = seed + 1
del transient
torch.cuda.synchronize()
address, run_bytes = pool.find_free_graph_pool_runs(handle)[0]
self.assertEqual(run_bytes, 200 << 20)
with (
envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.override(True),
patch.object(pool, "get_global_graph_memory_pool", return_value=handle),
):
# Unseeded 24/32/40 MiB segments strand 192 MiB before the 42 MiB request.
for rows in (1500, 2000, 2500, 2600):
with self.subTest(rows=rows), pool.borrow_graph_pool(user="test"):
first = torch.empty(
(rows, 4096), dtype=torch.float32, device="cuda"
)
second = torch.empty(
(rows, 4096), dtype=torch.float32, device="cuda"
)
self.assertTrue(
all(
address <= tensor.data_ptr()
and tensor.data_ptr() + tensor.nbytes <= address + run_bytes
for tensor in (first, second)
)
)
del first, second
pool._teardown_borrow_pool()
del graph, keep
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
def test_replay_raises_when_borrowed_tensor_is_still_referenced(self):
"""Reject borrowed tensors that replay would silently overwrite."""
handle = torch.cuda.graph_pool_handle()
graph = torch.cuda.CUDAGraph()
seed = torch.zeros(8, device="cuda")
stream = torch.cuda.Stream()
with (
torch.cuda.stream(stream),
torch.cuda.graph(graph, pool=handle, stream=stream),
):
transient = torch.empty(48 << 20, dtype=torch.uint8, device="cuda")
keep = seed + 1
del transient
torch.cuda.synchronize()
with (
envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.override(True),
patch.object(pool, "get_global_graph_memory_pool", return_value=handle),
):
with pool.borrow_graph_pool(user="leaky"):
leaked = torch.empty(1 << 20, device="cuda")
with self.assertRaisesRegex(
RuntimeError,
f"graph replay: {leaked.nbytes} bytes",
):
with pool.graph_pool_replay_scope():
pass
# A fresh borrow re-arms the replay check after releasing the leak.
del leaked
with pool.borrow_graph_pool(user="clean"):
released = torch.empty(1 << 20, device="cuda")
del released
with pool.graph_pool_replay_scope():
pass
pool._teardown_borrow_pool()
del graph, keep
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA") @unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
def test_static_borrow_runs_serve_without_a_pool_snapshot(self): def test_static_borrow_runs_serve_without_a_pool_snapshot(self):
"""Fixed extents serve borrows without consulting the shared pool.""" """Fixed extents serve borrows without consulting the shared pool."""
@@ -281,8 +371,6 @@ class TestGraphPoolBorrow(CustomTestCase):
with ( with (
envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.override(True), envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.override(True),
patch.object(pool, "get_global_graph_memory_pool", return_value=None), patch.object(pool, "get_global_graph_memory_pool", return_value=None),
patch.object(pool, "_borrow_static_runs", None),
patch.object(pool, "_borrow_mem_pool", None),
): ):
pool.set_graph_pool_borrow_runs(runs) pool.set_graph_pool_borrow_runs(runs)
self.assertTrue(pool.graph_pool_borrow_enabled()) self.assertTrue(pool.graph_pool_borrow_enabled())
@@ -295,6 +383,7 @@ class TestGraphPoolBorrow(CustomTestCase):
) )
) )
del borrowed del borrowed
pool._teardown_borrow_pool()
del graph, y del graph, y
@@ -352,18 +441,20 @@ class TestGraphPoolBorrow(CustomTestCase):
with ( with (
envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.override(True), envs.SGLANG_ENABLE_GRAPH_POOL_BORROW.override(True),
patch.object(pool, "get_global_graph_memory_pool", return_value=handle), patch.object(pool, "get_global_graph_memory_pool", return_value=handle),
patch.object(pool, "_borrow_mem_pool", None),
): ):
for _ in range(3): for _ in range(3):
with pool.borrow_graph_pool(user="test"): with pool.borrow_graph_pool(user="test"):
borrowed = torch.empty(16 << 20, dtype=torch.uint8, device="cuda") borrowed = torch.empty(16 << 20, dtype=torch.uint8, device="cuda")
# Stream-keyed segments allow side-stream copies, not allocations.
sink = torch.empty_like(borrowed)
with torch.cuda.stream(side): with torch.cuda.stream(side):
widened = borrowed.to(torch.int32) sink.copy_(borrowed)
borrowed.record_stream(side) borrowed.record_stream(side)
del borrowed, widened del borrowed, sink
# Regression: this used to fail with "Trying to free a pointer not # Regression: this used to fail with "Trying to free a pointer not
# allocated here" after a deferred free was re-issued too early. # allocated here" after a deferred free was re-issued too early.
torch.cuda.empty_cache() torch.cuda.empty_cache()
pool._teardown_borrow_pool()
del graph, y del graph, y