Introduce CUDA graph debug mode with breakable CUDA graph (#19102)

Co-authored-by: Cheng Wan <54331508+ch-wan@users.noreply.github.com>
Co-authored-by: Cheng Wan <chwan@rice.edu>
Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
cctry
2026-04-11 00:36:56 -07:00
committed by GitHub
co-authored by Cheng Wan Cheng Wan Claude Opus 4.6
parent d11da2403c
commit f855a0bde6
9 changed files with 896 additions and 7 deletions
@@ -0,0 +1,139 @@
# Breakable CUDA Graph
## Motivation
Standard CUDA graphs capture an entire forward pass as a single, opaque graph. This is great for performance, but creates two problems:
1. **Debugging is hard.** When something goes wrong inside a captured graph (wrong outputs, numerical mismatches, crashes), there is no way to step through the operations or insert print statements because the graph replays as a monolithic unit.
2. **Some ops are incompatible.** Certain operations — dynamic control flow, host-device synchronization, JIT compilation, or ops that change behavior across iterations — cannot be captured into a CUDA graph at all. Today, the only workaround is to disable CUDA graphs entirely, which sacrifices the kernel launch overhead savings for the rest of the model.
**Breakable CUDA Graph** solves both problems by allowing graph breaks to be inserted at specific points. The computation is split into multiple captured graph segments with eager (non-graph) execution in between. This preserves most of the CUDA graph performance benefit while allowing targeted operations to run outside the graph.
## Usage
### Debug Mode: Run Everything Eagerly
The simplest use case is debugging. The `--debug-cuda-graph` flag wraps the entire decode forward pass in a graph break, so every operation runs eagerly while still going through the full CUDA graph capture/replay code path. This lets you debug CUDA graph issues without changing model code.
```bash
python -m sglang.launch_server \
--model meta-llama/Llama-3.1-8B-Instruct \
--debug-cuda-graph
```
This mode is intended for debugging only — it eliminates the performance benefit of CUDA graphs since every op runs eagerly.
### Selective Graph Breaks in Model Code
For production use, you can mark specific functions as "non-graphable" using the `@eager_on_graph` decorator. During CUDA graph capture, these functions run eagerly between captured graph segments. Outside of capture, they behave normally.
```python
from sglang.srt.model_executor.breakable_cuda_graph.breakable_cuda_graph import eager_on_graph
@eager_on_graph(enable=True)
def my_dynamic_op(x):
# This op is incompatible with CUDA graph capture
return some_dynamic_operation(x)
```
You can also insert a bare graph break (no computation) using the `break_graph()` helper:
```python
from sglang.srt.model_executor.breakable_cuda_graph.breakable_cuda_graph import break_graph
def forward(self, x):
x = self.layer1(x)
break_graph() # force a segment split here
x = self.layer2(x)
return x
```
To enable breakable CUDA graph at the environment level (without debug mode), set the environment variable:
```bash
export SGLANG_USE_BREAKABLE_CUDA_GRAPH=1
python -m sglang.launch_server \
--model meta-llama/Llama-3.1-8B-Instruct
```
### Server Args
| Argument | Default | Description |
|---|---|---|
| `--debug-cuda-graph` | `False` | Enable debug/eager mode. Wraps the entire forward pass in a graph break so every op runs eagerly through the capture/replay path. |
| `SGLANG_USE_BREAKABLE_CUDA_GRAPH` | `0` | Environment variable. Enables breakable CUDA graph without debug mode. Required for `@eager_on_graph` decorators to take effect. |
## How It Works
### Capture
Breakable CUDA graph extends PyTorch's `torch.cuda.CUDAGraph` by splitting a single capture into multiple segments separated by graph breaks.
During capture, the flow is:
```
Begin capture (segment 1)
... graphable ops ...
@eager_on_graph function encountered:
1. End current capture segment
2. Run the function eagerly (allocates output tensors)
3. Record the function for later replay
4. Begin new capture segment
... more graphable ops ...
End capture (segment N)
```
Each segment is independently instantiated as a CUDA graph executable. The non-graph functions and their argument references are stored for replay.
### Replay
During replay:
```
For each segment i:
1. Launch CUDA graph segment i
2. Run the recorded non-graph function i eagerly
Launch final CUDA graph segment
```
The non-graph functions are re-invoked with the same tensor references as capture time. Since these references point to the CUDA graph's static input/output buffers, they see updated values on each replay.
### Output Writeback
When a non-graph function produces output during replay, the result must be written back into the same tensor buffers that downstream graph segments reference. The mechanism handles:
- **Plain tensors**: In-place `copy_()` into the original buffer.
- **Structured outputs** (dataclasses, objects with tensor attributes): Tensor fields are copied in-place; non-tensor fields are replaced.
- **Dicts of tensors**: Tensor values are copied in-place; non-tensor values are replaced.
### Stream Fork/Join Tracking
Some models fork work onto secondary CUDA streams (e.g., for overlapped computation). Breakable CUDA graph hooks `torch.cuda.Stream.wait_stream` to track which streams are forked from the capture stream. When a graph break occurs, all forked streams are automatically joined back before ending the capture segment, and re-forked after beginning the next segment.
## Compatibility
- **NVIDIA CUDA only.** Breakable CUDA graph is not supported on ROCm/HIP or other non-CUDA platforms. On unsupported platforms, `--debug-cuda-graph` is automatically disabled with a warning.
- **Requires `cuda-python`.** The `cuda.bindings` package must be installed (`pip install cuda-python`).
- **Not compatible with memory saver mode.** Cannot be used together with `SGLANG_MEMORY_SAVER_CUDA_GRAPH`.
## Performance
When no graph breaks are inserted, breakable CUDA graph has minimal overhead compared to standard CUDA graph — the capture/replay path is nearly identical.
Each graph break adds:
- One `cudaGraphLaunch` call (to replay the segment before the break)
- One eager Python function call
- One `cudaStreamBeginCapture` / `cudaStreamEndCapture` pair during capture
For typical use cases with a small number of graph breaks, the overhead is negligible compared to the saved kernel launch overhead from the captured segments.
## Code Reference
| File | Description |
|---|---|
| `python/sglang/srt/model_executor/breakable_cuda_graph/breakable_cuda_graph.py` | Core implementation: `eager_on_graph`, `BreakableCUDAGraph`, `BreakableCUDAGraphCapture` |
| `python/sglang/srt/model_executor/breakable_cuda_graph/cuda_utils.py` | CUDA runtime binding utilities |
| `python/sglang/srt/model_executor/cuda_graph_runner.py` | Integration with the main CUDA graph runner |
| `python/sglang/srt/server_args.py` | `--debug-cuda-graph` flag and environment variable handling |
| `python/sglang/srt/environ.py` | `SGLANG_USE_BREAKABLE_CUDA_GRAPH` environment variable definition |
+1
View File
@@ -63,6 +63,7 @@ Its core features include:
advanced_features/dp_for_multi_modal_encoder.md
advanced_features/cuda_graph_for_multi_modal_encoder.md
advanced_features/piecewise_cuda_graph.md
advanced_features/breakable_cuda_graph.md
advanced_features/sgl_model_gateway.md
advanced_features/deterministic_inference.md
advanced_features/observability.md
+3
View File
@@ -473,6 +473,9 @@ class Envs:
SGLANG_MAMBA_CONV_DTYPE = EnvStr("bfloat16")
SGLANG_MAMBA_SSM_DTYPE = EnvStr(None)
# Breakable CUDA Graph
SGLANG_USE_BREAKABLE_CUDA_GRAPH = EnvBool(False)
# Release & Resume Memory
SGLANG_MEMORY_SAVER_CUDA_GRAPH = EnvBool(False)
@@ -0,0 +1,352 @@
# Copyright 2023-2026 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
import logging
import threading
from contextvars import ContextVar
from typing import Any, Callable, NamedTuple
import torch
try:
from cuda.bindings import runtime as rt
except ImportError:
rt = None
from sglang.srt.model_executor.breakable_cuda_graph.cuda_utils import checkCudaErrors
logger = logging.getLogger(__name__)
__all__ = [
"eager_on_graph",
"BreakableCUDAGraph",
"BreakableCUDAGraphCapture",
"break_graph",
]
def _check_cuda_bindings():
if rt is None:
raise ImportError(
"Breakable CUDA graph requires the 'cuda-python' package. "
"Install it with: pip install cuda-python"
)
class GraphBreakInfo(NamedTuple):
# python function breaking the graph
func: Callable
# output of the function (must be a tensor so we keep them)
output: Any
# raw handle after capture or raw exec handle after instantiate
graph_handle: Any
_captured_graphs_var: ContextVar[list[GraphBreakInfo] | None] = ContextVar(
"captured_graphs", default=None
)
_current_stream_var: ContextVar[torch.cuda.Stream | None] = ContextVar(
"current_stream", default=None
)
_forked_streams_var: ContextVar[set[torch.cuda.Stream] | None] = ContextVar(
"forked_streams", default=None
)
def get_current_stream(device: torch.device | None = None) -> torch.cuda.Stream:
stream = _current_stream_var.get()
if stream is None:
return torch.cuda.current_stream(device)
return stream
def _capture_status(stream_ptr: int) -> "rt.cudaStreamCaptureStatus":
_check_cuda_bindings()
status, *_ = checkCudaErrors(rt.cudaStreamGetCaptureInfo(stream_ptr))
return status
def _is_capturing(stream_ptr: int) -> bool:
_check_cuda_bindings()
return (
_capture_status(stream_ptr)
== rt.cudaStreamCaptureStatus.cudaStreamCaptureStatusActive
)
# hook wait_stream to track forks/joins during breakable capture.
_original_wait_stream: Callable | None = None
_hook_lock = threading.Lock()
_hook_refcount = 0
def _hooked_wait_stream(self: torch.cuda.Stream, other: torch.cuda.Stream):
assert _original_wait_stream is not None
forked = _forked_streams_var.get()
if forked is None:
_original_wait_stream(self, other)
return
capturing = _current_stream_var.get()
if capturing is None:
_original_wait_stream(self, other)
return
cap_ptr = capturing.cuda_stream
is_self_cap = self is capturing or self.cuda_stream == cap_ptr
is_other_cap = other is capturing or other.cuda_stream == cap_ptr
if is_self_cap and not is_other_cap:
# Join: capturing_stream.wait_stream(other).
# other might not be part of the capture because we join it in the last segment
# skip the wait to avoid cuda error
if (
_capture_status(other.cuda_stream)
!= rt.cudaStreamCaptureStatus.cudaStreamCaptureStatusActive
):
return
_original_wait_stream(self, other)
forked.discard(other)
elif is_other_cap and not is_self_cap:
# Fork: other.wait_stream(capturing_stream).
_original_wait_stream(self, other)
forked.add(self)
else:
_original_wait_stream(self, other)
def _install_wait_stream_hook():
global _original_wait_stream, _hook_refcount
with _hook_lock:
if _hook_refcount == 0:
_original_wait_stream = torch.cuda.Stream.wait_stream
torch.cuda.Stream.wait_stream = _hooked_wait_stream # type: ignore[assignment]
_hook_refcount += 1
def _uninstall_wait_stream_hook():
global _original_wait_stream, _hook_refcount
with _hook_lock:
_hook_refcount -= 1
if _hook_refcount == 0:
assert _original_wait_stream is not None, "wait_stream hook not installed"
torch.cuda.Stream.wait_stream = _original_wait_stream # type: ignore[assignment]
_original_wait_stream = None
def _end_capture_segment(stream: torch.cuda.Stream):
"""End a capture segment, auto-joining any forked streams first."""
# Join forked streams that are still part of this capture.
forked = _forked_streams_var.get()
if forked:
assert _original_wait_stream is not None
for s in forked:
if _is_capturing(s.cuda_stream):
_original_wait_stream(stream, s)
forked.clear()
graph = checkCudaErrors(rt.cudaStreamEndCapture(stream.cuda_stream))
assert graph is not None
return graph
def _begin_capture_segment(stream: torch.cuda.Stream):
checkCudaErrors(
rt.cudaStreamBeginCapture(
stream.cuda_stream,
rt.cudaStreamCaptureMode.cudaStreamCaptureModeGlobal,
)
)
def _instantiate_graph(graph_ptr: int) -> int:
graph_exec = checkCudaErrors(
rt.cudaGraphInstantiateWithFlags(
graph_ptr,
rt.cudaGraphInstantiateFlags.cudaGraphInstantiateFlagAutoFreeOnLaunch,
)
)
assert graph_exec is not None
checkCudaErrors(rt.cudaGraphDestroy(graph_ptr))
return graph_exec
def _destroy_graph_exec(graph_exec_ptr: int) -> None:
checkCudaErrors(rt.cudaGraphExecDestroy(graph_exec_ptr))
def _replay_graph(graph_exec_ptr: int, stream_ptr: int) -> None:
checkCudaErrors(rt.cudaGraphLaunch(graph_exec_ptr, stream_ptr))
def _copy_output(dst: Any, src: Any) -> Any:
"""Copy src output into dst in-place where possible.
Handles plain tensors, dataclass/object with tensor attributes,
and dicts of tensors. Returns dst if in-place copy succeeded,
otherwise returns src.
"""
if torch.is_tensor(dst) and torch.is_tensor(src):
dst.copy_(src)
return dst
# Handle objects with __dict__ (dataclasses, regular objects)
if hasattr(dst, "__dict__") and hasattr(src, "__dict__"):
for key, src_val in src.__dict__.items():
dst_val = getattr(dst, key, None)
if torch.is_tensor(dst_val) and torch.is_tensor(src_val):
dst_val.copy_(src_val)
else:
setattr(dst, key, src_val)
return dst
# Handle dicts of tensors
if isinstance(dst, dict) and isinstance(src, dict):
for key, src_val in src.items():
dst_val = dst.get(key)
if torch.is_tensor(dst_val) and torch.is_tensor(src_val):
dst_val.copy_(src_val)
else:
dst[key] = src_val
return dst
return src
def eager_on_graph(enable: bool):
def decorator(inner: Callable):
if not enable:
return inner
def wrapper(*args, **kwargs):
stream = get_current_stream()
if not _is_capturing(stream.cuda_stream):
return inner(*args, **kwargs)
last_graph = _end_capture_segment(stream)
logger.debug(f"Break graph due to function: {inner.__name__}")
# run the function once to allocate the output tensor captured by later graphs
output = inner(*args, **kwargs)
# Store the callable and its arguments so replay can re-invoke with
# the same argument *references* (which point to CUDA graph input
# buffers whose contents are updated before replay).
captured_inner = inner
captured_args = args
captured_kwargs = kwargs
captured_output = output
def replay_fn():
new_out = captured_inner(*captured_args, **captured_kwargs)
return _copy_output(captured_output, new_out)
captured_graphs = _captured_graphs_var.get()
assert (
captured_graphs is not None
), "eager_on_graph wrapper called outside of BreakableCUDAGraphCapture"
captured_graphs.append(GraphBreakInfo(replay_fn, output, last_graph))
_begin_capture_segment(stream)
return output
return wrapper
return decorator
class BreakableCUDAGraph(torch.cuda.CUDAGraph):
def __new__(cls) -> "BreakableCUDAGraph":
return super().__new__(cls, True)
def capture_begin(self, pool=None, capture_error_mode: str = "global") -> None:
_check_cuda_bindings()
super().capture_begin(pool, capture_error_mode)
stream = get_current_stream()
# torch graph will not record any operation but only for compatibility
_end_capture_segment(stream)
_begin_capture_segment(stream)
def capture_end(self):
stream = get_current_stream()
self.last_graph = _end_capture_segment(stream)
self.last_graph_exec = _instantiate_graph(self.last_graph)
breaks = _captured_graphs_var.get()
self._exec = []
if breaks:
for replay_fn, output, handle in breaks:
graph_exec = _instantiate_graph(handle)
self._exec.append(GraphBreakInfo(replay_fn, output, graph_exec))
# start a dummy capture so torch's capture_end() can finalize
_begin_capture_segment(stream)
super().capture_end()
def replay(self):
stream = torch.cuda.current_stream()
token = _current_stream_var.set(stream)
try:
if not self._exec:
_replay_graph(self.last_graph_exec, stream.cuda_stream)
return
for func, _, handle in self._exec:
_replay_graph(handle, stream.cuda_stream)
func()
_replay_graph(self.last_graph_exec, stream.cuda_stream)
finally:
_current_stream_var.reset(token)
def __del__(self):
try:
if hasattr(self, "_exec"):
for _, _, handle in self._exec:
_destroy_graph_exec(handle)
if hasattr(self, "last_graph_exec"):
_destroy_graph_exec(self.last_graph_exec)
except Exception:
pass
class BreakableCUDAGraphCapture(torch.cuda.graph):
def __init__(
self,
cuda_graph: BreakableCUDAGraph,
pool=None,
stream: torch.cuda.Stream | None = None,
capture_error_mode: str = "global",
):
super().__init__(
cuda_graph, pool=pool, stream=stream, capture_error_mode=capture_error_mode
)
self._stream = stream
assert isinstance(
cuda_graph, BreakableCUDAGraph
), "cuda_graph must be a BreakableCUDAGraph"
def __enter__(self):
_install_wait_stream_hook()
self._breaks_token = _captured_graphs_var.set([])
self._stream_token = _current_stream_var.set(self._stream)
self._forked_streams_token = _forked_streams_var.set(set())
return super().__enter__()
def __exit__(self, *args: object):
super().__exit__(*args)
_current_stream_var.reset(self._stream_token)
_captured_graphs_var.reset(self._breaks_token)
_forked_streams_var.reset(self._forked_streams_token)
_uninstall_wait_stream_hook()
@eager_on_graph(True)
def break_graph():
"""Insert a graph break. The @eager_on_graph decorator does the actual
segment split; this function body intentionally does nothing."""
pass
@@ -0,0 +1,53 @@
# Copyright 2023-2026 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""CUDA runtime binding utilities."""
try:
from cuda.bindings import runtime as rt
except ImportError:
rt = None
def _cudaGetErrorString(error):
if rt is None:
return "<cuda.bindings not available>"
err, msg = rt.cudaGetErrorString(error)
if err != rt.cudaError_t.cudaSuccess:
return "<unknown>"
if isinstance(msg, bytes):
return msg.decode("utf-8", "replace")
return str(msg)
def checkCudaErrors(result):
if rt is None:
raise RuntimeError(
"cuda.bindings is not available. "
"Install it with: pip install cuda-python"
)
if rt is None:
raise RuntimeError(
"cuda.bindings is not available. "
"Install it with: pip install cuda-python"
)
if result[0] != rt.cudaError_t.cudaSuccess:
raise RuntimeError(
f"CUDA error {int(result[0])}({_cudaGetErrorString(result[0])})"
)
if len(result) == 1:
return None
elif len(result) == 2:
return result[1]
else:
return result[1:]
@@ -41,6 +41,7 @@ from sglang.srt.distributed.parallel_state import (
set_pdmux_status,
)
from sglang.srt.dllm.config import DllmConfig
from sglang.srt.environ import envs
from sglang.srt.layers.attention.nsa.utils import is_nsa_enable_prefill_cp
from sglang.srt.layers.dp_attention import (
DpPaddingMode,
@@ -88,6 +89,13 @@ except ImportError:
_is_hip = is_hip()
if not _is_hip:
from sglang.srt.model_executor.breakable_cuda_graph.breakable_cuda_graph import (
BreakableCUDAGraph,
BreakableCUDAGraphCapture,
eager_on_graph,
)
logger = logging.getLogger(__name__)
if TYPE_CHECKING:
@@ -833,20 +841,43 @@ class CudaGraphRunner:
self._post_process_after_profile(prof)
def _capture_graph(self, graph, pool, stream, run_once_fn):
if self.model_runner.server_args.debug_cuda_graph:
assert (
envs.SGLANG_USE_BREAKABLE_CUDA_GRAPH.get()
), "Breakable CUDA graph is not enabled in debug mode"
memory_saver_adapter = TorchMemorySaverAdapter.create(
enable=self.model_runner.server_args.enable_memory_saver
and get_bool_env_var("SGLANG_MEMORY_SAVER_CUDA_GRAPH")
)
graph_fn = (
partial(memory_saver_adapter.cuda_graph, tag=GPU_MEMORY_TYPE_CUDA_GRAPH)
if memory_saver_adapter.enabled
else self.device_module.graph
)
with graph_fn(cuda_graph=graph, pool=pool, stream=stream):
out = run_once_fn()
if envs.SGLANG_USE_BREAKABLE_CUDA_GRAPH.get():
if memory_saver_adapter.enabled:
raise NotImplementedError(
"Breakable CUDA graph is not compatible with memory saver mode"
)
graph_ctx = BreakableCUDAGraphCapture
else:
graph_ctx = (
partial(memory_saver_adapter.cuda_graph, tag=GPU_MEMORY_TYPE_CUDA_GRAPH)
if memory_saver_adapter.enabled
else self.device_module.graph
)
if self.model_runner.server_args.debug_cuda_graph:
captured_fn = eager_on_graph(True)(run_once_fn)
else:
captured_fn = run_once_fn
with graph_ctx(cuda_graph=graph, pool=pool, stream=stream):
out = captured_fn()
return out
def _create_device_graph(self):
if envs.SGLANG_USE_BREAKABLE_CUDA_GRAPH.get():
if _is_hip:
raise RuntimeError("Breakable CUDA graph is not supported on ROCm/HIP")
return BreakableCUDAGraph()
return torch.cuda.CUDAGraph()
def capture_one_batch_size(
+25
View File
@@ -622,6 +622,7 @@ class ServerArgs:
disable_cuda_graph_padding: bool = False
enable_profile_cuda_graph: bool = False
enable_cudagraph_gc: bool = False
debug_cuda_graph: bool = False
enable_layerwise_nvtx_marker: bool = False
enable_nccl_nvls: bool = False
enable_symm_mem: bool = False
@@ -1167,6 +1168,9 @@ class ServerArgs:
# 17. Context parallel
if self.attn_cp_size > 1:
self.disable_piecewise_cuda_graph = True
# 18. CUDA Graph debug mode
if self.debug_cuda_graph:
self.disable_piecewise_cuda_graph = True
def _handle_gpu_memory_settings(self, gpu_mem):
"""
@@ -3629,6 +3633,19 @@ class ServerArgs:
envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.set(
"1" if self.enable_deterministic_inference else "0"
)
if self.debug_cuda_graph:
if not is_cuda():
logger.warning(
"--debug-cuda-graph is not supported on non CUDA devices. "
"Disabling breakable CUDA graph."
)
self.debug_cuda_graph = False
else:
envs.SGLANG_USE_BREAKABLE_CUDA_GRAPH.set("1")
logger.warning(
"Debug mode for CUDA graph is enabled via breakable CUDA graph. "
"All operations will run eagerly through the graph capture/replay path."
)
def _handle_cache_compatibility(self):
if self.enable_hierarchical_cache and self.disable_radix_cache:
@@ -5650,6 +5667,14 @@ class ServerArgs:
action="store_true",
help="Enable garbage collection during CUDA graph capture. If disabled (default), GC is frozen during capture to speed up the process.",
)
parser.add_argument(
"--debug-cuda-graph",
action="store_true",
help="Enable debug/eager mode for CUDA graph using breakable CUDA graph. "
"When enabled, graph breaks are inserted so every operation runs eagerly "
"while still going through the CUDA graph capture / replay path. "
"Useful for debugging CUDA graph capture / replay issues.",
)
parser.add_argument(
"--enable-layerwise-nvtx-marker",
action="store_true",
@@ -0,0 +1,285 @@
"""Unit tests for the breakable CUDA graph mechanism.
Tests the core capture/replay logic with simple tensor operations,
verifying that graph breaks work correctly and outputs are properly
propagated across segments.
"""
import unittest
import torch
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
# CI Registration
register_cuda_ci(est_time=30, suite="stage-b-test-1-gpu-small")
def _skip_if_no_cuda(test_func):
return unittest.skipUnless(torch.cuda.is_available(), "CUDA not available")(
test_func
)
def _skip_if_no_cuda_bindings(test_func):
try:
from cuda.bindings import runtime as rt # noqa: F401
return test_func
except ImportError:
return unittest.skip("cuda-python not installed")(test_func)
class TestBreakableCUDAGraphBasic(CustomTestCase):
"""Test basic breakable CUDA graph capture and replay."""
@classmethod
def setUpClass(cls):
if not torch.cuda.is_available():
raise unittest.SkipTest("CUDA not available")
try:
from cuda.bindings import runtime # noqa: F401
except ImportError:
raise unittest.SkipTest("cuda-python not installed")
from sglang.srt.model_executor.breakable_cuda_graph.breakable_cuda_graph import (
BreakableCUDAGraph,
BreakableCUDAGraphCapture,
eager_on_graph,
)
cls.BreakableCUDAGraph = BreakableCUDAGraph
cls.BreakableCUDAGraphCapture = BreakableCUDAGraphCapture
cls.eager_on_graph = staticmethod(eager_on_graph)
cls.device = torch.device("cuda:0")
def test_no_break_capture_replay(self):
"""Capture and replay without any graph breaks should work like normal CUDA graph."""
x = torch.zeros(4, device=self.device)
y = torch.zeros(4, device=self.device)
graph = self.BreakableCUDAGraph()
stream = torch.cuda.Stream(self.device)
with self.BreakableCUDAGraphCapture(graph, stream=stream):
y.copy_(x + 1.0)
# Replay with new input
x.fill_(5.0)
graph.replay()
torch.cuda.synchronize()
self.assertTrue(torch.allclose(y, torch.full((4,), 6.0, device=self.device)))
def test_single_break(self):
"""A single graph break should split capture into two segments."""
x = torch.zeros(4, device=self.device)
intermediate = torch.zeros(4, device=self.device)
y = torch.zeros(4, device=self.device)
@self.eager_on_graph(enable=True)
def eager_op(src):
return src * 2.0
graph = self.BreakableCUDAGraph()
stream = torch.cuda.Stream(self.device)
with self.BreakableCUDAGraphCapture(graph, stream=stream):
intermediate.copy_(x + 1.0)
broken = eager_op(intermediate)
y.copy_(broken + 3.0)
# Replay with new input
x.fill_(10.0)
graph.replay()
torch.cuda.synchronize()
# x=10 -> intermediate=11 -> eager: 11*2=22 -> y=22+3=25
self.assertTrue(torch.allclose(y, torch.full((4,), 25.0, device=self.device)))
def test_multiple_breaks(self):
"""Multiple graph breaks should produce correct chained results."""
x = torch.zeros(4, device=self.device)
y = torch.zeros(4, device=self.device)
@self.eager_on_graph(enable=True)
def add_one(src):
return src + 1.0
@self.eager_on_graph(enable=True)
def double(src):
return src * 2.0
graph = self.BreakableCUDAGraph()
stream = torch.cuda.Stream(self.device)
with self.BreakableCUDAGraphCapture(graph, stream=stream):
t1 = x + 1.0 # graph segment 1
t2 = add_one(t1) # break 1: eager
t3 = t2 + 1.0 # graph segment 2
t4 = double(t3) # break 2: eager
y.copy_(t4) # graph segment 3
# Replay: x=5 -> +1=6 -> add_one=7 -> +1=8 -> double=16
x.fill_(5.0)
graph.replay()
torch.cuda.synchronize()
self.assertTrue(torch.allclose(y, torch.full((4,), 16.0, device=self.device)))
def test_eager_on_graph_disabled(self):
"""@eager_on_graph(enable=False) should be a no-op passthrough."""
@self.eager_on_graph(enable=False)
def my_fn(x):
return x + 1.0
# Should just be the original function
t = torch.tensor([1.0, 2.0], device=self.device)
result = my_fn(t)
self.assertTrue(
torch.allclose(result, torch.tensor([2.0, 3.0], device=self.device))
)
def test_eager_on_graph_outside_capture(self):
"""@eager_on_graph called outside capture should run the function directly."""
@self.eager_on_graph(enable=True)
def my_fn(x):
return x + 1.0
t = torch.tensor([1.0, 2.0], device=self.device)
result = my_fn(t)
self.assertTrue(
torch.allclose(result, torch.tensor([2.0, 3.0], device=self.device))
)
def test_replay_updates_output(self):
"""Replay should produce different results when input buffers change."""
x = torch.zeros(4, device=self.device)
y = torch.zeros(4, device=self.device)
@self.eager_on_graph(enable=True)
def scale(src):
return src * 3.0
graph = self.BreakableCUDAGraph()
stream = torch.cuda.Stream(self.device)
with self.BreakableCUDAGraphCapture(graph, stream=stream):
t = x + 1.0
t2 = scale(t)
y.copy_(t2)
# First replay: x=0 -> 0+1=1 -> 1*3=3
graph.replay()
torch.cuda.synchronize()
self.assertTrue(torch.allclose(y, torch.full((4,), 3.0, device=self.device)))
# Second replay: x=10 -> 10+1=11 -> 11*3=33
x.fill_(10.0)
graph.replay()
torch.cuda.synchronize()
self.assertTrue(torch.allclose(y, torch.full((4,), 33.0, device=self.device)))
class TestCopyOutput(CustomTestCase):
"""Test the _copy_output helper for structured output writeback."""
@classmethod
def setUpClass(cls):
if not torch.cuda.is_available():
raise unittest.SkipTest("CUDA not available")
try:
from cuda.bindings import runtime # noqa: F401
except ImportError:
raise unittest.SkipTest("cuda-python not installed")
from sglang.srt.model_executor.breakable_cuda_graph.breakable_cuda_graph import (
_copy_output,
)
cls._copy_output = staticmethod(_copy_output)
cls.device = torch.device("cuda:0")
def test_tensor_copy(self):
dst = torch.zeros(4, device=self.device)
src = torch.ones(4, device=self.device) * 5.0
result = self._copy_output(dst, src)
self.assertIs(result, dst)
self.assertTrue(torch.allclose(dst, src))
def test_dict_copy(self):
dst = {
"a": torch.zeros(4, device=self.device),
"b": torch.zeros(4, device=self.device),
}
src = {
"a": torch.ones(4, device=self.device),
"b": torch.ones(4, device=self.device) * 2.0,
}
result = self._copy_output(dst, src)
self.assertIs(result, dst)
self.assertTrue(torch.allclose(dst["a"], torch.ones(4, device=self.device)))
self.assertTrue(
torch.allclose(dst["b"], torch.ones(4, device=self.device) * 2.0)
)
def test_object_copy(self):
class FakeOutput:
def __init__(self, t, label):
self.tensor = t
self.label = label
dst = FakeOutput(torch.zeros(4, device=self.device), "old")
src = FakeOutput(torch.ones(4, device=self.device) * 3.0, "new")
result = self._copy_output(dst, src)
self.assertIs(result, dst)
self.assertTrue(
torch.allclose(dst.tensor, torch.ones(4, device=self.device) * 3.0)
)
self.assertEqual(dst.label, "new")
def test_non_tensor_fallback(self):
result = self._copy_output(42, 99)
self.assertEqual(result, 99)
class TestBreakGraphHelper(CustomTestCase):
"""Test the break_graph() convenience function."""
@classmethod
def setUpClass(cls):
if not torch.cuda.is_available():
raise unittest.SkipTest("CUDA not available")
try:
from cuda.bindings import runtime # noqa: F401
except ImportError:
raise unittest.SkipTest("cuda-python not installed")
from sglang.srt.model_executor.breakable_cuda_graph.breakable_cuda_graph import (
BreakableCUDAGraph,
BreakableCUDAGraphCapture,
break_graph,
)
cls.BreakableCUDAGraph = BreakableCUDAGraph
cls.BreakableCUDAGraphCapture = BreakableCUDAGraphCapture
cls.break_graph = staticmethod(break_graph)
cls.device = torch.device("cuda:0")
def test_break_graph_inserts_segment(self):
"""break_graph() should insert a graph break even though it does nothing."""
x = torch.zeros(4, device=self.device)
y = torch.zeros(4, device=self.device)
graph = self.BreakableCUDAGraph()
stream = torch.cuda.Stream(self.device)
with self.BreakableCUDAGraphCapture(graph, stream=stream):
t = x + 1.0
self.break_graph()
y.copy_(t + 2.0)
x.fill_(10.0)
graph.replay()
torch.cuda.synchronize()
# x=10 -> +1=11 -> break -> +2=13
self.assertTrue(torch.allclose(y, torch.full((4,), 13.0, device=self.device)))
if __name__ == "__main__":
unittest.main()