From f855a0bde6c9db10212183261be1f355977d9ea3 Mon Sep 17 00:00:00 2001 From: cctry Date: Sat, 11 Apr 2026 00:36:56 -0700 Subject: [PATCH] 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 Co-authored-by: Claude Opus 4.6 --- .../advanced_features/breakable_cuda_graph.md | 139 +++++++ docs/index.rst | 1 + python/sglang/srt/environ.py | 3 + .../breakable_cuda_graph/__init__.py | 0 .../breakable_cuda_graph.py | 352 ++++++++++++++++++ .../breakable_cuda_graph/cuda_utils.py | 53 +++ .../srt/model_executor/cuda_graph_runner.py | 45 ++- python/sglang/srt/server_args.py | 25 ++ .../cuda_graph/test_breakable_cuda_graph.py | 285 ++++++++++++++ 9 files changed, 896 insertions(+), 7 deletions(-) create mode 100644 docs/advanced_features/breakable_cuda_graph.md create mode 100644 python/sglang/srt/model_executor/breakable_cuda_graph/__init__.py create mode 100644 python/sglang/srt/model_executor/breakable_cuda_graph/breakable_cuda_graph.py create mode 100644 python/sglang/srt/model_executor/breakable_cuda_graph/cuda_utils.py create mode 100644 test/registered/cuda_graph/test_breakable_cuda_graph.py diff --git a/docs/advanced_features/breakable_cuda_graph.md b/docs/advanced_features/breakable_cuda_graph.md new file mode 100644 index 000000000..4fb2c090c --- /dev/null +++ b/docs/advanced_features/breakable_cuda_graph.md @@ -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 | diff --git a/docs/index.rst b/docs/index.rst index 1a4defe9f..2b327b6c9 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -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 diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index e29aa6d12..78babcb7a 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -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) diff --git a/python/sglang/srt/model_executor/breakable_cuda_graph/__init__.py b/python/sglang/srt/model_executor/breakable_cuda_graph/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/python/sglang/srt/model_executor/breakable_cuda_graph/breakable_cuda_graph.py b/python/sglang/srt/model_executor/breakable_cuda_graph/breakable_cuda_graph.py new file mode 100644 index 000000000..adad2b172 --- /dev/null +++ b/python/sglang/srt/model_executor/breakable_cuda_graph/breakable_cuda_graph.py @@ -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 diff --git a/python/sglang/srt/model_executor/breakable_cuda_graph/cuda_utils.py b/python/sglang/srt/model_executor/breakable_cuda_graph/cuda_utils.py new file mode 100644 index 000000000..c95471012 --- /dev/null +++ b/python/sglang/srt/model_executor/breakable_cuda_graph/cuda_utils.py @@ -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 "" + err, msg = rt.cudaGetErrorString(error) + if err != rt.cudaError_t.cudaSuccess: + return "" + 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:] diff --git a/python/sglang/srt/model_executor/cuda_graph_runner.py b/python/sglang/srt/model_executor/cuda_graph_runner.py index b4e21bb3b..381cb800f 100644 --- a/python/sglang/srt/model_executor/cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/cuda_graph_runner.py @@ -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( diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 93eba2b1e..31fb5b535 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -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", diff --git a/test/registered/cuda_graph/test_breakable_cuda_graph.py b/test/registered/cuda_graph/test_breakable_cuda_graph.py new file mode 100644 index 000000000..4354e90ef --- /dev/null +++ b/test/registered/cuda_graph/test_breakable_cuda_graph.py @@ -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()