[Experimental] Breakable Piecewise Cuda Graph (#22218)
Signed-off-by: Oasis-Git <ayw.sirius19@gmail.com> Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.6
parent
b3b03369a5
commit
60bbb800db
@@ -23,6 +23,12 @@ from torch import nn
|
|||||||
|
|
||||||
from sglang.srt.compilation.compilation_config import register_split_op
|
from sglang.srt.compilation.compilation_config import register_split_op
|
||||||
from sglang.srt.compilation.piecewise_context_manager import get_forward_context
|
from sglang.srt.compilation.piecewise_context_manager import get_forward_context
|
||||||
|
from sglang.srt.model_executor.breakable_cuda_graph.breakable_cuda_graph import (
|
||||||
|
eager_on_graph,
|
||||||
|
)
|
||||||
|
from sglang.srt.model_executor.breakable_cuda_graph.context import (
|
||||||
|
is_in_breakable_cuda_graph,
|
||||||
|
)
|
||||||
from sglang.srt.utils.custom_op import register_custom_op
|
from sglang.srt.utils.custom_op import register_custom_op
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -119,6 +125,11 @@ class RadixAttention(nn.Module):
|
|||||||
output = q.new_empty((q.shape[0], self.tp_q_head_num * self.v_head_dim))
|
output = q.new_empty((q.shape[0], self.tp_q_head_num * self.v_head_dim))
|
||||||
else:
|
else:
|
||||||
output = torch.empty_like(q)
|
output = torch.empty_like(q)
|
||||||
|
if is_in_breakable_cuda_graph():
|
||||||
|
bcg_unified_attention_with_output(
|
||||||
|
q, k, v, output, save_kv_cache, self.layer_id, **kwargs
|
||||||
|
)
|
||||||
|
else:
|
||||||
unified_attention_with_output(
|
unified_attention_with_output(
|
||||||
q, k, v, output, save_kv_cache, self.layer_id, **kwargs
|
q, k, v, output, save_kv_cache, self.layer_id, **kwargs
|
||||||
)
|
)
|
||||||
@@ -197,3 +208,6 @@ def unified_attention_with_output(
|
|||||||
|
|
||||||
output[:real_num_tokens].view(ret.shape).copy_(ret)
|
output[:real_num_tokens].view(ret.shape).copy_(ret)
|
||||||
return
|
return
|
||||||
|
|
||||||
|
|
||||||
|
bcg_unified_attention_with_output = eager_on_graph(True)(unified_attention_with_output)
|
||||||
|
|||||||
@@ -11,10 +11,21 @@
|
|||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
# ==============================================================================
|
# ==============================================================================
|
||||||
|
"""Breakable CUDA Graph: capture a region as a sequence of
|
||||||
|
``torch.cuda.CUDAGraph`` segments separated by eager break points.
|
||||||
|
|
||||||
|
Each segment is a real ``torch.cuda.CUDAGraph``. Its destructor calls
|
||||||
|
``releasePool`` on the shared mempool, so the pool's ``use_count`` tracks how
|
||||||
|
many segments are alive; the pool stays pinned as long as any segment graph
|
||||||
|
is alive. This lets ``weak_ref_tensor`` views of intermediate pool-allocated
|
||||||
|
tensors remain valid across replays — we don't need Python-managed bridge
|
||||||
|
buffers to keep break-point tensors at stable addresses.
|
||||||
|
"""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
import threading
|
import threading
|
||||||
from contextvars import ContextVar
|
from contextvars import ContextVar
|
||||||
from typing import Any, Callable, NamedTuple
|
from typing import Any, Callable
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -43,17 +54,11 @@ def _check_cuda_bindings():
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class GraphBreakInfo(NamedTuple):
|
# Active BreakableCUDAGraphCapture context for the currently-capturing thread.
|
||||||
# python function breaking the graph
|
# eager_on_graph's wrapper uses this to split the current torch.cuda.CUDAGraph
|
||||||
func: Callable
|
# at break points.
|
||||||
# output of the function (must be a tensor so we keep them)
|
_current_capture_var: ContextVar["BreakableCUDAGraphCapture | None"] = ContextVar(
|
||||||
output: Any
|
"current_capture", default=None
|
||||||
# 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_var: ContextVar[torch.cuda.Stream | None] = ContextVar(
|
||||||
"current_stream", default=None
|
"current_stream", default=None
|
||||||
@@ -84,7 +89,10 @@ def _is_capturing(stream_ptr: int) -> bool:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
# hook wait_stream to track forks/joins during breakable capture.
|
# Hook torch.cuda.Stream.wait_stream to track side-stream forks/joins that happen
|
||||||
|
# during breakable capture. We need this because capture_end() on a torch
|
||||||
|
# CUDAGraph fails if there are still side streams participating in the capture
|
||||||
|
# — so before ending each segment we auto-join any forked-but-not-rejoined streams.
|
||||||
_original_wait_stream: Callable | None = None
|
_original_wait_stream: Callable | None = None
|
||||||
_hook_lock = threading.Lock()
|
_hook_lock = threading.Lock()
|
||||||
_hook_refcount = 0
|
_hook_refcount = 0
|
||||||
@@ -106,9 +114,6 @@ def _hooked_wait_stream(self: torch.cuda.Stream, other: torch.cuda.Stream):
|
|||||||
is_other_cap = other is capturing or other.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:
|
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 (
|
if (
|
||||||
_capture_status(other.cuda_stream)
|
_capture_status(other.cuda_stream)
|
||||||
!= rt.cudaStreamCaptureStatus.cudaStreamCaptureStatusActive
|
!= rt.cudaStreamCaptureStatus.cudaStreamCaptureStatusActive
|
||||||
@@ -117,7 +122,6 @@ def _hooked_wait_stream(self: torch.cuda.Stream, other: torch.cuda.Stream):
|
|||||||
_original_wait_stream(self, other)
|
_original_wait_stream(self, other)
|
||||||
forked.discard(other)
|
forked.discard(other)
|
||||||
elif is_other_cap and not is_self_cap:
|
elif is_other_cap and not is_self_cap:
|
||||||
# Fork: other.wait_stream(capturing_stream).
|
|
||||||
_original_wait_stream(self, other)
|
_original_wait_stream(self, other)
|
||||||
forked.add(self)
|
forked.add(self)
|
||||||
else:
|
else:
|
||||||
@@ -143,49 +147,20 @@ def _uninstall_wait_stream_hook():
|
|||||||
_original_wait_stream = None
|
_original_wait_stream = None
|
||||||
|
|
||||||
|
|
||||||
def _end_capture_segment(stream: torch.cuda.Stream):
|
def _weak_ref_if_tensor(x):
|
||||||
"""End a capture segment, auto-joining any forked streams first."""
|
"""Return a weak-ref tensor view (shared storage, no refcount) for tensors;
|
||||||
# Join forked streams that are still part of this capture.
|
pass-through for non-tensors. Weak-ref'ing captured args lets the shared
|
||||||
forked = _forked_streams_var.get()
|
mempool reclaim per-layer intermediates between segments — storage stays
|
||||||
if forked:
|
alive for each segment CUDAGraph's lifetime via its pool use_count.
|
||||||
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))
|
``weak_ref_tensors`` is imported lazily: the module hard-raises on
|
||||||
assert graph is not None
|
non-CUDA/NPU platforms, and we only reach this code during an active
|
||||||
return graph
|
BCG capture (which can't happen on CPU-only runners anyway)."""
|
||||||
|
if torch.is_tensor(x):
|
||||||
|
from sglang.srt.compilation.weak_ref_tensor import weak_ref_tensors
|
||||||
|
|
||||||
|
return weak_ref_tensors(x)
|
||||||
def _begin_capture_segment(stream: torch.cuda.Stream):
|
return x
|
||||||
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:
|
def _copy_output(dst: Any, src: Any) -> Any:
|
||||||
@@ -199,7 +174,6 @@ def _copy_output(dst: Any, src: Any) -> Any:
|
|||||||
dst.copy_(src)
|
dst.copy_(src)
|
||||||
return dst
|
return dst
|
||||||
|
|
||||||
# Handle objects with __dict__ (dataclasses, regular objects)
|
|
||||||
if hasattr(dst, "__dict__") and hasattr(src, "__dict__"):
|
if hasattr(dst, "__dict__") and hasattr(src, "__dict__"):
|
||||||
for key, src_val in src.__dict__.items():
|
for key, src_val in src.__dict__.items():
|
||||||
dst_val = getattr(dst, key, None)
|
dst_val = getattr(dst, key, None)
|
||||||
@@ -209,7 +183,6 @@ def _copy_output(dst: Any, src: Any) -> Any:
|
|||||||
setattr(dst, key, src_val)
|
setattr(dst, key, src_val)
|
||||||
return dst
|
return dst
|
||||||
|
|
||||||
# Handle dicts of tensors
|
|
||||||
if isinstance(dst, dict) and isinstance(src, dict):
|
if isinstance(dst, dict) and isinstance(src, dict):
|
||||||
for key, src_val in src.items():
|
for key, src_val in src.items():
|
||||||
dst_val = dst.get(key)
|
dst_val = dst.get(key)
|
||||||
@@ -228,32 +201,35 @@ def eager_on_graph(enable: bool):
|
|||||||
return inner
|
return inner
|
||||||
|
|
||||||
def wrapper(*args, **kwargs):
|
def wrapper(*args, **kwargs):
|
||||||
stream = get_current_stream()
|
capture = _current_capture_var.get()
|
||||||
if not _is_capturing(stream.cuda_stream):
|
if capture is None:
|
||||||
return inner(*args, **kwargs)
|
return inner(*args, **kwargs)
|
||||||
last_graph = _end_capture_segment(stream)
|
|
||||||
logger.debug(f"Break graph due to function: {inner.__name__}")
|
logger.debug("Break graph due to function: %s", inner.__name__)
|
||||||
# run the function once to allocate the output tensor captured by later graphs
|
|
||||||
|
# End the segment that captured up to this break point.
|
||||||
|
capture._end_current_segment()
|
||||||
|
|
||||||
|
# Run the eager function once so it allocates its outputs and
|
||||||
|
# writes real data into them.
|
||||||
output = inner(*args, **kwargs)
|
output = inner(*args, **kwargs)
|
||||||
|
|
||||||
# Store the callable and its arguments so replay can re-invoke with
|
# Weak-ref the closure state. Storage lives with the segment
|
||||||
# the same argument *references* (which point to CUDA graph input
|
# CUDAGraphs' mempool pin; Python refs don't need to prevent
|
||||||
# buffers whose contents are updated before replay).
|
# pool reuse across layers.
|
||||||
captured_inner = inner
|
captured_inner = inner
|
||||||
captured_args = args
|
captured_args = tuple(_weak_ref_if_tensor(a) for a in args)
|
||||||
captured_kwargs = kwargs
|
captured_kwargs = {k: _weak_ref_if_tensor(v) for k, v in kwargs.items()}
|
||||||
captured_output = output
|
captured_output = _weak_ref_if_tensor(output)
|
||||||
|
|
||||||
def replay_fn():
|
def replay_fn():
|
||||||
new_out = captured_inner(*captured_args, **captured_kwargs)
|
new_out = captured_inner(*captured_args, **captured_kwargs)
|
||||||
return _copy_output(captured_output, new_out)
|
return _copy_output(captured_output, new_out)
|
||||||
|
|
||||||
captured_graphs = _captured_graphs_var.get()
|
capture.cuda_graph._break_fns.append(replay_fn)
|
||||||
assert (
|
|
||||||
captured_graphs is not None
|
# Start a fresh CUDAGraph segment for the remainder of the forward.
|
||||||
), "eager_on_graph wrapper called outside of BreakableCUDAGraphCapture"
|
capture._begin_new_segment()
|
||||||
captured_graphs.append(GraphBreakInfo(replay_fn, output, last_graph))
|
|
||||||
_begin_capture_segment(stream)
|
|
||||||
return output
|
return output
|
||||||
|
|
||||||
return wrapper
|
return wrapper
|
||||||
@@ -261,60 +237,37 @@ def eager_on_graph(enable: bool):
|
|||||||
return decorator
|
return decorator
|
||||||
|
|
||||||
|
|
||||||
class BreakableCUDAGraph(torch.cuda.CUDAGraph):
|
class BreakableCUDAGraph:
|
||||||
|
"""Container holding one ``torch.cuda.CUDAGraph`` per segment plus an
|
||||||
|
eager break function between consecutive segments."""
|
||||||
|
|
||||||
def __new__(cls) -> "BreakableCUDAGraph":
|
def __init__(self) -> None:
|
||||||
return super().__new__(cls, True)
|
self._segments: list[torch.cuda.CUDAGraph] = []
|
||||||
|
self._break_fns: list[Callable[[], Any]] = []
|
||||||
|
|
||||||
def capture_begin(self, pool=None, capture_error_mode: str = "global") -> None:
|
def replay(self) -> 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()
|
stream = torch.cuda.current_stream()
|
||||||
token = _current_stream_var.set(stream)
|
token = _current_stream_var.set(stream)
|
||||||
try:
|
try:
|
||||||
if not self._exec:
|
for i, seg in enumerate(self._segments):
|
||||||
_replay_graph(self.last_graph_exec, stream.cuda_stream)
|
seg.replay()
|
||||||
return
|
if i < len(self._break_fns):
|
||||||
for func, _, handle in self._exec:
|
self._break_fns[i]()
|
||||||
_replay_graph(handle, stream.cuda_stream)
|
|
||||||
func()
|
|
||||||
_replay_graph(self.last_graph_exec, stream.cuda_stream)
|
|
||||||
finally:
|
finally:
|
||||||
_current_stream_var.reset(token)
|
_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:
|
||||||
|
"""Context manager that captures the enclosed code as one or more
|
||||||
|
``torch.cuda.CUDAGraph`` segments separated by eager break points.
|
||||||
|
|
||||||
|
Each segment shares the supplied ``pool`` (``MempoolId_t`` tuple) so
|
||||||
|
pool-allocated intermediates can be reused across segments. While any
|
||||||
|
segment is alive, its ``beginAllocateToPool`` call keeps the mempool's
|
||||||
|
``use_count`` > 0, which makes ``weak_ref_tensor`` of segment-allocated
|
||||||
|
tensors safe across subsequent replays.
|
||||||
|
"""
|
||||||
|
|
||||||
class BreakableCUDAGraphCapture(torch.cuda.graph):
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
cuda_graph: BreakableCUDAGraph,
|
cuda_graph: BreakableCUDAGraph,
|
||||||
@@ -322,31 +275,66 @@ class BreakableCUDAGraphCapture(torch.cuda.graph):
|
|||||||
stream: torch.cuda.Stream | None = None,
|
stream: torch.cuda.Stream | None = None,
|
||||||
capture_error_mode: str = "global",
|
capture_error_mode: str = "global",
|
||||||
):
|
):
|
||||||
super().__init__(
|
|
||||||
cuda_graph, pool=pool, stream=stream, capture_error_mode=capture_error_mode
|
|
||||||
)
|
|
||||||
self._stream = stream
|
|
||||||
assert isinstance(
|
assert isinstance(
|
||||||
cuda_graph, BreakableCUDAGraph
|
cuda_graph, BreakableCUDAGraph
|
||||||
), "cuda_graph must be a BreakableCUDAGraph"
|
), "cuda_graph must be a BreakableCUDAGraph"
|
||||||
|
self.cuda_graph = cuda_graph
|
||||||
|
self._pool = pool if pool is not None else (0, 0)
|
||||||
|
self._stream = stream
|
||||||
|
self._capture_error_mode = capture_error_mode
|
||||||
|
self._stream_ctx = None
|
||||||
|
self._capture_token = None
|
||||||
|
self._stream_token = None
|
||||||
|
self._forked_token = None
|
||||||
|
|
||||||
def __enter__(self):
|
def __enter__(self):
|
||||||
_install_wait_stream_hook()
|
_install_wait_stream_hook()
|
||||||
self._breaks_token = _captured_graphs_var.set([])
|
if self._stream is not None:
|
||||||
self._stream_token = _current_stream_var.set(self._stream)
|
self._stream_ctx = torch.cuda.stream(self._stream)
|
||||||
self._forked_streams_token = _forked_streams_var.set(set())
|
self._stream_ctx.__enter__()
|
||||||
return super().__enter__()
|
self._capture_token = _current_capture_var.set(self)
|
||||||
|
self._stream_token = _current_stream_var.set(
|
||||||
|
self._stream or torch.cuda.current_stream()
|
||||||
|
)
|
||||||
|
self._forked_token = _forked_streams_var.set(set())
|
||||||
|
self._begin_new_segment()
|
||||||
|
return self
|
||||||
|
|
||||||
def __exit__(self, *args: object):
|
def __exit__(self, *args: object):
|
||||||
super().__exit__(*args)
|
try:
|
||||||
|
self._end_current_segment()
|
||||||
|
finally:
|
||||||
|
_forked_streams_var.reset(self._forked_token)
|
||||||
_current_stream_var.reset(self._stream_token)
|
_current_stream_var.reset(self._stream_token)
|
||||||
_captured_graphs_var.reset(self._breaks_token)
|
_current_capture_var.reset(self._capture_token)
|
||||||
_forked_streams_var.reset(self._forked_streams_token)
|
if self._stream_ctx is not None:
|
||||||
|
self._stream_ctx.__exit__(*args)
|
||||||
|
self._stream_ctx = None
|
||||||
_uninstall_wait_stream_hook()
|
_uninstall_wait_stream_hook()
|
||||||
|
return False
|
||||||
|
|
||||||
|
def _begin_new_segment(self) -> None:
|
||||||
|
graph = torch.cuda.CUDAGraph()
|
||||||
|
graph.capture_begin(
|
||||||
|
pool=self._pool, capture_error_mode=self._capture_error_mode
|
||||||
|
)
|
||||||
|
self.cuda_graph._segments.append(graph)
|
||||||
|
|
||||||
|
def _end_current_segment(self) -> None:
|
||||||
|
# Auto-join any side streams forked during this segment but not joined.
|
||||||
|
main_stream = get_current_stream()
|
||||||
|
forked = _forked_streams_var.get()
|
||||||
|
if forked:
|
||||||
|
assert _original_wait_stream is not None
|
||||||
|
for side in list(forked):
|
||||||
|
if _is_capturing(side.cuda_stream):
|
||||||
|
_original_wait_stream(main_stream, side)
|
||||||
|
forked.clear()
|
||||||
|
self.cuda_graph._segments[-1].capture_end()
|
||||||
|
|
||||||
|
|
||||||
@eager_on_graph(True)
|
@eager_on_graph(True)
|
||||||
def break_graph():
|
def break_graph() -> None:
|
||||||
"""Insert a graph break. The @eager_on_graph decorator does the actual
|
"""Insert a graph break. The @eager_on_graph decorator does the actual
|
||||||
segment split; this function body intentionally does nothing."""
|
segment split; this function body intentionally does nothing."""
|
||||||
pass
|
pass
|
||||||
|
|||||||
@@ -0,0 +1,39 @@
|
|||||||
|
# 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.
|
||||||
|
# ==============================================================================
|
||||||
|
"""Runtime state for the breakable CUDA graph (BCG) runner.
|
||||||
|
|
||||||
|
Kept intentionally separate from ``compilation/piecewise_context_manager.py``:
|
||||||
|
BCG no longer inherits from the torch.compile-based PCG path, so its
|
||||||
|
capture/replay lifecycle is managed on its own.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from contextlib import contextmanager
|
||||||
|
|
||||||
|
_in_breakable_cuda_graph = False
|
||||||
|
|
||||||
|
|
||||||
|
def is_in_breakable_cuda_graph() -> bool:
|
||||||
|
return _in_breakable_cuda_graph
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def enable_breakable_cuda_graph():
|
||||||
|
global _in_breakable_cuda_graph
|
||||||
|
_in_breakable_cuda_graph = True
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
_in_breakable_cuda_graph = False
|
||||||
@@ -0,0 +1,402 @@
|
|||||||
|
# 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.
|
||||||
|
# ==============================================================================
|
||||||
|
"""Breakable CUDA graph (BCG) runner.
|
||||||
|
|
||||||
|
Captures the model forward as a sequence of ``torch.cuda.CUDAGraph`` segments
|
||||||
|
split at attention layers. Functionally parallel to the torch.compile-based
|
||||||
|
PCG runner but does not depend on torch.compile or FX graph splitting — graph
|
||||||
|
breaks are inserted eagerly via :func:`eager_on_graph` decorated callables
|
||||||
|
(radix attention for dense models, mamba for hybrid models).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import bisect
|
||||||
|
import logging
|
||||||
|
from typing import TYPE_CHECKING, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import tqdm
|
||||||
|
|
||||||
|
from sglang.srt.compilation.piecewise_context_manager import set_forward_context
|
||||||
|
from sglang.srt.distributed import get_tensor_model_parallel_rank
|
||||||
|
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||||
|
set_graph_pool_id,
|
||||||
|
)
|
||||||
|
from sglang.srt.distributed.parallel_state import graph_capture
|
||||||
|
from sglang.srt.layers.dp_attention import (
|
||||||
|
set_dp_buffer_len,
|
||||||
|
set_is_extend_in_batch,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||||
|
from sglang.srt.layers.pooler import EmbeddingPoolerOutput
|
||||||
|
from sglang.srt.model_executor.breakable_cuda_graph.breakable_cuda_graph import (
|
||||||
|
BreakableCUDAGraph,
|
||||||
|
BreakableCUDAGraphCapture,
|
||||||
|
)
|
||||||
|
from sglang.srt.model_executor.breakable_cuda_graph.context import (
|
||||||
|
enable_breakable_cuda_graph,
|
||||||
|
)
|
||||||
|
from sglang.srt.model_executor.cuda_graph_runner import (
|
||||||
|
get_global_graph_memory_pool,
|
||||||
|
set_global_graph_memory_pool,
|
||||||
|
)
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import (
|
||||||
|
PPProxyTensors,
|
||||||
|
)
|
||||||
|
from sglang.srt.model_executor.piecewise_cuda_graph_runner import (
|
||||||
|
PiecewiseCudaGraphRunner,
|
||||||
|
freeze_gc,
|
||||||
|
)
|
||||||
|
from sglang.srt.utils import get_available_gpu_memory, log_info_on_rank0
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
|
|
||||||
|
|
||||||
|
class BreakableCudaGraphRunner:
|
||||||
|
"""Breakable CUDA graph runner.
|
||||||
|
|
||||||
|
Captures the model forward as a series of ``torch.cuda.CUDAGraph`` segments
|
||||||
|
with graph breaks at attention layers. Simpler than the torch.compile-based
|
||||||
|
PCG runner: no FX tracing, no compiled-kernel fusion — just segment-level
|
||||||
|
graph capture of the eager kernel stream.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# replay_prepare shares its buffer-population logic with the PCG runner —
|
||||||
|
# bind the method here without inheriting. __init__, capture, and replay
|
||||||
|
# diverge enough that inheritance would obscure more than it saves.
|
||||||
|
replay_prepare = PiecewiseCudaGraphRunner.replay_prepare
|
||||||
|
|
||||||
|
def __init__(self, model_runner: ModelRunner):
|
||||||
|
self.model_runner = model_runner
|
||||||
|
self.device = model_runner.device
|
||||||
|
self.device_module = torch.get_device_module(self.device)
|
||||||
|
self.graphs = {}
|
||||||
|
self.output_buffers = {}
|
||||||
|
|
||||||
|
self.quant_config = getattr(model_runner.model, "quant_config", None)
|
||||||
|
self.is_multimodal = model_runner.is_multimodal
|
||||||
|
# Read by the shared replay_prepare (bound from PiecewiseCudaGraphRunner).
|
||||||
|
self.capture_return_pooled_hidden_states = not model_runner.is_generation
|
||||||
|
|
||||||
|
# Capture sizes
|
||||||
|
capture_tokens = model_runner.server_args.piecewise_cuda_graph_tokens
|
||||||
|
assert capture_tokens is not None
|
||||||
|
self.capture_num_tokens = sorted(capture_tokens)
|
||||||
|
self.max_num_tokens = (
|
||||||
|
max(self.capture_num_tokens) if self.capture_num_tokens else 8192
|
||||||
|
)
|
||||||
|
self.max_bs = model_runner.req_to_token_pool.size
|
||||||
|
|
||||||
|
log_info_on_rank0(
|
||||||
|
logger,
|
||||||
|
f"[BCG] Capture num tokens: {self.capture_num_tokens}",
|
||||||
|
)
|
||||||
|
|
||||||
|
self._init_buffers(model_runner)
|
||||||
|
|
||||||
|
self.attention_layers = model_runner.attention_layers
|
||||||
|
self.moe_layers = model_runner.moe_layers
|
||||||
|
self.moe_fusions = model_runner.moe_fusions
|
||||||
|
|
||||||
|
with torch.device(self.device):
|
||||||
|
self.static_seq_lens = torch.zeros((self.max_bs,), dtype=torch.int64)
|
||||||
|
self.static_extend_seq_lens = torch.zeros((self.max_bs,), dtype=torch.int64)
|
||||||
|
self.static_extend_prefix_lens = torch.zeros(
|
||||||
|
(self.max_bs,), dtype=torch.int64
|
||||||
|
)
|
||||||
|
self.static_extend_start_loc = torch.zeros(
|
||||||
|
(self.max_bs,), dtype=torch.int64
|
||||||
|
)
|
||||||
|
self.static_req_pool_indices = torch.zeros(
|
||||||
|
(self.max_bs,), dtype=torch.int64
|
||||||
|
)
|
||||||
|
self.static_orig_seq_lens = torch.zeros((self.max_bs,), dtype=torch.int64)
|
||||||
|
|
||||||
|
# Memory pool
|
||||||
|
if get_global_graph_memory_pool() is None:
|
||||||
|
set_global_graph_memory_pool(self.device_module.graph_pool_handle())
|
||||||
|
set_graph_pool_id(get_global_graph_memory_pool())
|
||||||
|
|
||||||
|
# Warmup then capture
|
||||||
|
self._warmup()
|
||||||
|
self.device_module.synchronize()
|
||||||
|
self.model_runner.tp_group.barrier()
|
||||||
|
self._capture_all()
|
||||||
|
|
||||||
|
self.raw_num_tokens = 0
|
||||||
|
|
||||||
|
def _init_buffers(self, model_runner):
|
||||||
|
"""Initialize input buffers."""
|
||||||
|
from sglang.srt.model_executor.piecewise_cuda_graph_runner import (
|
||||||
|
PrefillInputBuffers,
|
||||||
|
)
|
||||||
|
from sglang.srt.utils import is_npu
|
||||||
|
|
||||||
|
with torch.device(self.device):
|
||||||
|
input_ids = torch.zeros((self.max_num_tokens,), dtype=torch.int64)
|
||||||
|
out_cache_loc = torch.zeros(
|
||||||
|
(self.max_num_tokens,),
|
||||||
|
dtype=torch.int64 if not is_npu() else torch.int32,
|
||||||
|
)
|
||||||
|
out_cache_loc_swa = (
|
||||||
|
torch.zeros((self.max_num_tokens,), dtype=torch.int64)
|
||||||
|
if model_runner.is_hybrid_swa
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
positions = torch.zeros((self.max_num_tokens,), dtype=torch.int64)
|
||||||
|
if self.is_multimodal:
|
||||||
|
input_embeds = torch.zeros(
|
||||||
|
(self.max_num_tokens, model_runner.model_config.hidden_size),
|
||||||
|
dtype=model_runner.dtype,
|
||||||
|
)
|
||||||
|
mrope_positions = torch.zeros(
|
||||||
|
(3, self.max_num_tokens), dtype=torch.int64
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
input_embeds = None
|
||||||
|
mrope_positions = None
|
||||||
|
|
||||||
|
self.buffers = PrefillInputBuffers(
|
||||||
|
input_ids=input_ids,
|
||||||
|
out_cache_loc=out_cache_loc,
|
||||||
|
out_cache_loc_swa=out_cache_loc_swa,
|
||||||
|
mamba_track_indices=None,
|
||||||
|
mamba_track_mask=None,
|
||||||
|
mamba_track_seqlens=None,
|
||||||
|
positions=positions,
|
||||||
|
input_embeds=input_embeds,
|
||||||
|
mrope_positions=mrope_positions,
|
||||||
|
)
|
||||||
|
self.buffers.share_buffers()
|
||||||
|
|
||||||
|
def _run_forward(self, forward_batch, num_tokens):
|
||||||
|
"""Run model forward with proper context."""
|
||||||
|
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
|
||||||
|
set_dp_buffer_len(None, num_tokens, forward_batch.dp_padding_mode.is_max_len())
|
||||||
|
set_is_extend_in_batch(False)
|
||||||
|
|
||||||
|
with set_forward_context(
|
||||||
|
forward_batch,
|
||||||
|
self.attention_layers,
|
||||||
|
self.quant_config,
|
||||||
|
self.moe_layers,
|
||||||
|
self.moe_fusions,
|
||||||
|
):
|
||||||
|
output = self.model_runner.model.forward(
|
||||||
|
forward_batch.input_ids,
|
||||||
|
forward_batch.positions,
|
||||||
|
forward_batch,
|
||||||
|
)
|
||||||
|
return output
|
||||||
|
|
||||||
|
def _build_capture_forward_batch(self, num_tokens):
|
||||||
|
"""Build a ForwardBatch for capture using static buffers for stable addresses."""
|
||||||
|
from sglang.srt.layers.dp_attention import DpPaddingMode
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import (
|
||||||
|
CaptureHiddenMode,
|
||||||
|
ForwardBatch,
|
||||||
|
ForwardMode,
|
||||||
|
)
|
||||||
|
|
||||||
|
buffers = self.buffers
|
||||||
|
bs = 1
|
||||||
|
self.static_seq_lens[:bs].fill_(num_tokens)
|
||||||
|
self.static_extend_seq_lens[:bs].fill_(num_tokens)
|
||||||
|
self.static_extend_prefix_lens[:bs].zero_()
|
||||||
|
self.static_extend_start_loc[:bs].zero_()
|
||||||
|
self.static_req_pool_indices[:bs].copy_(torch.arange(bs, device=self.device))
|
||||||
|
self.static_orig_seq_lens[:bs].fill_(num_tokens)
|
||||||
|
|
||||||
|
return ForwardBatch(
|
||||||
|
forward_mode=ForwardMode.EXTEND,
|
||||||
|
batch_size=bs,
|
||||||
|
input_ids=buffers.input_ids[:num_tokens],
|
||||||
|
input_embeds=(
|
||||||
|
buffers.input_embeds[:num_tokens] if self.is_multimodal else None
|
||||||
|
),
|
||||||
|
req_pool_indices=self.static_req_pool_indices[:bs],
|
||||||
|
seq_lens=self.static_seq_lens[:bs],
|
||||||
|
next_token_logits_buffer=None,
|
||||||
|
orig_seq_lens=self.static_orig_seq_lens[:bs],
|
||||||
|
seq_lens_cpu=torch.tensor([num_tokens], device="cpu"),
|
||||||
|
req_to_token_pool=self.model_runner.req_to_token_pool,
|
||||||
|
token_to_kv_pool=self.model_runner.token_to_kv_pool,
|
||||||
|
attn_backend=self.model_runner.attn_backend,
|
||||||
|
out_cache_loc=buffers.out_cache_loc[:num_tokens],
|
||||||
|
out_cache_loc_swa=(
|
||||||
|
buffers.out_cache_loc_swa[:num_tokens]
|
||||||
|
if buffers.out_cache_loc_swa is not None
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
seq_lens_sum=num_tokens,
|
||||||
|
mamba_track_indices=None,
|
||||||
|
mamba_track_mask=None,
|
||||||
|
mamba_track_seqlens=None,
|
||||||
|
encoder_lens=None,
|
||||||
|
return_logprob=False,
|
||||||
|
extend_num_tokens=num_tokens,
|
||||||
|
extend_seq_lens=self.static_extend_seq_lens[:bs],
|
||||||
|
extend_prefix_lens=self.static_extend_prefix_lens[:bs],
|
||||||
|
extend_start_loc=self.static_extend_start_loc[:bs],
|
||||||
|
extend_prefix_lens_cpu=torch.tensor([0], device="cpu"),
|
||||||
|
extend_seq_lens_cpu=torch.tensor([num_tokens], device="cpu"),
|
||||||
|
extend_logprob_start_lens_cpu=torch.tensor([num_tokens], device="cpu"),
|
||||||
|
positions=buffers.positions[:num_tokens],
|
||||||
|
global_num_tokens_gpu=None,
|
||||||
|
global_num_tokens_for_logprob_gpu=None,
|
||||||
|
dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(),
|
||||||
|
global_dp_buffer_len=None,
|
||||||
|
mrope_positions=(
|
||||||
|
buffers.mrope_positions[:, :num_tokens] if self.is_multimodal else None
|
||||||
|
),
|
||||||
|
spec_algorithm=None,
|
||||||
|
spec_info=None,
|
||||||
|
capture_hidden_mode=CaptureHiddenMode.NULL,
|
||||||
|
num_token_non_padded=None,
|
||||||
|
global_forward_mode=ForwardMode.EXTEND,
|
||||||
|
lora_ids=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _warmup(self):
|
||||||
|
"""Warmup the model with a forward pass."""
|
||||||
|
num_tokens = self.capture_num_tokens[0]
|
||||||
|
forward_batch = self._build_capture_forward_batch(num_tokens)
|
||||||
|
self.model_runner.attn_backend.init_forward_metadata(forward_batch)
|
||||||
|
self._run_forward(forward_batch, num_tokens)
|
||||||
|
|
||||||
|
def _capture_all(self):
|
||||||
|
"""Capture breakable CUDA graphs for all token sizes."""
|
||||||
|
with freeze_gc(
|
||||||
|
self.model_runner.server_args.enable_cudagraph_gc
|
||||||
|
), graph_capture() as graph_capture_context, enable_breakable_cuda_graph():
|
||||||
|
stream = graph_capture_context.stream
|
||||||
|
pool = get_global_graph_memory_pool()
|
||||||
|
|
||||||
|
capture_range = (
|
||||||
|
tqdm.tqdm(list(reversed(self.capture_num_tokens)))
|
||||||
|
if get_tensor_model_parallel_rank() == 0
|
||||||
|
else reversed(self.capture_num_tokens)
|
||||||
|
)
|
||||||
|
for num_tokens in capture_range:
|
||||||
|
if get_tensor_model_parallel_rank() == 0:
|
||||||
|
avail_mem = get_available_gpu_memory(
|
||||||
|
self.model_runner.device,
|
||||||
|
self.model_runner.gpu_id,
|
||||||
|
empty_cache=False,
|
||||||
|
)
|
||||||
|
capture_range.set_description(
|
||||||
|
f"[BCG] Capturing ({num_tokens=} {avail_mem=:.2f} GB)"
|
||||||
|
)
|
||||||
|
|
||||||
|
graph, output = self._capture_one(num_tokens, pool, stream)
|
||||||
|
self.graphs[num_tokens] = graph
|
||||||
|
self.output_buffers[num_tokens] = output
|
||||||
|
|
||||||
|
def can_run(self, forward_batch: "ForwardBatch"):
|
||||||
|
# BCG graphs are captured with batch_size=1 (see _build_capture_forward_batch);
|
||||||
|
# the captured logits-gather / sampler path yields bs=1 outputs. Multi-req
|
||||||
|
# prefill would silently return wrong-shaped logits, corrupting downstream
|
||||||
|
# output_ids and breaking the subsequent decode step. Reject here so the
|
||||||
|
# caller falls back to the eager extend path.
|
||||||
|
if forward_batch.batch_size > 1:
|
||||||
|
return False
|
||||||
|
if forward_batch.input_embeds is not None:
|
||||||
|
return False
|
||||||
|
if forward_batch.replace_embeds is not None:
|
||||||
|
return False
|
||||||
|
num_tokens = len(forward_batch.input_ids)
|
||||||
|
if forward_batch.return_logprob:
|
||||||
|
for start_len, seq_len in zip(
|
||||||
|
forward_batch.extend_logprob_start_lens_cpu,
|
||||||
|
forward_batch.extend_seq_lens_cpu,
|
||||||
|
):
|
||||||
|
if start_len is not None and start_len < seq_len:
|
||||||
|
return False
|
||||||
|
return num_tokens <= self.max_num_tokens
|
||||||
|
|
||||||
|
def _capture_one(self, num_tokens, pool, stream):
|
||||||
|
"""Capture a breakable CUDA graph for one token size."""
|
||||||
|
forward_batch = self._build_capture_forward_batch(num_tokens)
|
||||||
|
self.model_runner.attn_backend.init_forward_metadata(forward_batch)
|
||||||
|
|
||||||
|
def run_once():
|
||||||
|
return self._run_forward(forward_batch, num_tokens)
|
||||||
|
|
||||||
|
for _ in range(2):
|
||||||
|
self.device_module.synchronize()
|
||||||
|
self.model_runner.tp_group.barrier()
|
||||||
|
run_once()
|
||||||
|
|
||||||
|
graph = BreakableCUDAGraph()
|
||||||
|
with BreakableCUDAGraphCapture(cuda_graph=graph, pool=pool, stream=stream):
|
||||||
|
output = run_once()
|
||||||
|
|
||||||
|
return graph, output
|
||||||
|
|
||||||
|
def replay(
|
||||||
|
self,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
**kwargs,
|
||||||
|
) -> Union[LogitsProcessorOutput, PPProxyTensors, EmbeddingPoolerOutput]:
|
||||||
|
num_tokens = len(forward_batch.input_ids)
|
||||||
|
index = bisect.bisect_left(self.capture_num_tokens, num_tokens)
|
||||||
|
static_num_tokens = self.capture_num_tokens[index]
|
||||||
|
|
||||||
|
with enable_breakable_cuda_graph():
|
||||||
|
static_forward_batch = self.replay_prepare(forward_batch, **kwargs)
|
||||||
|
bs = forward_batch.batch_size
|
||||||
|
|
||||||
|
# Update static buffers used by graph segments (esp. logits processor).
|
||||||
|
# The graph reads from these addresses — they must have serving-time values.
|
||||||
|
self.static_seq_lens[:bs].copy_(forward_batch.seq_lens)
|
||||||
|
self.static_extend_seq_lens[:bs].copy_(forward_batch.extend_seq_lens)
|
||||||
|
self.static_extend_prefix_lens[:bs].copy_(forward_batch.extend_prefix_lens)
|
||||||
|
self.static_extend_start_loc[:bs].copy_(forward_batch.extend_start_loc)
|
||||||
|
self.static_req_pool_indices[:bs].copy_(forward_batch.req_pool_indices)
|
||||||
|
if forward_batch.orig_seq_lens is not None:
|
||||||
|
self.static_orig_seq_lens[:bs].copy_(forward_batch.orig_seq_lens)
|
||||||
|
|
||||||
|
# Set forward context and replay
|
||||||
|
self.model_runner.attn_backend.init_forward_metadata(forward_batch)
|
||||||
|
with set_forward_context(
|
||||||
|
static_forward_batch,
|
||||||
|
self.attention_layers,
|
||||||
|
self.quant_config,
|
||||||
|
self.moe_layers,
|
||||||
|
self.moe_fusions,
|
||||||
|
):
|
||||||
|
self.graphs[static_num_tokens].replay()
|
||||||
|
|
||||||
|
output = self.output_buffers[static_num_tokens]
|
||||||
|
if isinstance(output, LogitsProcessorOutput):
|
||||||
|
return LogitsProcessorOutput(
|
||||||
|
next_token_logits=output.next_token_logits[: self.raw_num_tokens],
|
||||||
|
hidden_states=(
|
||||||
|
output.hidden_states[: self.raw_num_tokens]
|
||||||
|
if output.hidden_states is not None
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
)
|
||||||
|
elif isinstance(output, EmbeddingPoolerOutput):
|
||||||
|
return output
|
||||||
|
else:
|
||||||
|
assert isinstance(output, PPProxyTensors)
|
||||||
|
raise NotImplementedError(
|
||||||
|
"PPProxyTensors is not supported in BreakableCudaGraphRunner."
|
||||||
|
)
|
||||||
@@ -124,6 +124,9 @@ from sglang.srt.lora.lora_registry import LoRARef
|
|||||||
from sglang.srt.managers.schedule_batch import sanity_check_mm_pad_shift_value
|
from sglang.srt.managers.schedule_batch import sanity_check_mm_pad_shift_value
|
||||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||||
|
from sglang.srt.model_executor.breakable_cuda_graph_runner import (
|
||||||
|
BreakableCudaGraphRunner,
|
||||||
|
)
|
||||||
from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner
|
from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner
|
||||||
from sglang.srt.model_executor.cuda_graph_runner import (
|
from sglang.srt.model_executor.cuda_graph_runner import (
|
||||||
CudaGraphRunner,
|
CudaGraphRunner,
|
||||||
@@ -2726,6 +2729,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
f"Capture piecewise CUDA graph begin. avail mem={before_mem:.2f} GB"
|
f"Capture piecewise CUDA graph begin. avail mem={before_mem:.2f} GB"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if self.server_args.enable_breakable_cuda_graph:
|
||||||
|
# Experimental feature
|
||||||
|
self.piecewise_cuda_graph_runner = BreakableCudaGraphRunner(self)
|
||||||
|
else:
|
||||||
self.piecewise_cuda_graph_runner = PiecewiseCudaGraphRunner(self)
|
self.piecewise_cuda_graph_runner = PiecewiseCudaGraphRunner(self)
|
||||||
|
|
||||||
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
||||||
|
|||||||
@@ -59,6 +59,12 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
|||||||
ParallelLMHead,
|
ParallelLMHead,
|
||||||
VocabParallelEmbedding,
|
VocabParallelEmbedding,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.model_executor.breakable_cuda_graph.breakable_cuda_graph import (
|
||||||
|
eager_on_graph,
|
||||||
|
)
|
||||||
|
from sglang.srt.model_executor.breakable_cuda_graph.context import (
|
||||||
|
is_in_breakable_cuda_graph,
|
||||||
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||||
from sglang.srt.model_loader.weight_utils import (
|
from sglang.srt.model_loader.weight_utils import (
|
||||||
default_weight_loader,
|
default_weight_loader,
|
||||||
@@ -429,6 +435,11 @@ class NemotronHMambaDecoderLayer(nn.Module):
|
|||||||
else:
|
else:
|
||||||
hidden_states, residual = self.norm(hidden_states, residual)
|
hidden_states, residual = self.norm(hidden_states, residual)
|
||||||
|
|
||||||
|
if is_in_breakable_cuda_graph():
|
||||||
|
output = torch.empty_like(hidden_states)
|
||||||
|
breakable_nemotron_mamba2_with_output(hidden_states, output, self.layer_id)
|
||||||
|
return output, residual
|
||||||
|
|
||||||
if is_in_piecewise_cuda_graph():
|
if is_in_piecewise_cuda_graph():
|
||||||
output = torch.empty_like(hidden_states)
|
output = torch.empty_like(hidden_states)
|
||||||
nemotron_mamba2_with_output(hidden_states, output, self.layer_id)
|
nemotron_mamba2_with_output(hidden_states, output, self.layer_id)
|
||||||
@@ -912,4 +923,8 @@ def nemotron_mamba2_with_output(
|
|||||||
|
|
||||||
# Copy result back; output may be larger (padded) so only fill actual tokens
|
# Copy result back; output may be larger (padded) so only fill actual tokens
|
||||||
output[:num_actual_tokens].view(ret.shape).copy_(ret)
|
output[:num_actual_tokens].view(ret.shape).copy_(ret)
|
||||||
return
|
|
||||||
|
|
||||||
|
breakable_nemotron_mamba2_with_output = eager_on_graph(True)(
|
||||||
|
nemotron_mamba2_with_output
|
||||||
|
)
|
||||||
|
|||||||
@@ -624,6 +624,7 @@ class ServerArgs:
|
|||||||
cuda_graph_bs: Optional[List[int]] = None
|
cuda_graph_bs: Optional[List[int]] = None
|
||||||
disable_cuda_graph: bool = False
|
disable_cuda_graph: bool = False
|
||||||
disable_cuda_graph_padding: bool = False
|
disable_cuda_graph_padding: bool = False
|
||||||
|
enable_breakable_cuda_graph: bool = False
|
||||||
enable_profile_cuda_graph: bool = False
|
enable_profile_cuda_graph: bool = False
|
||||||
enable_cudagraph_gc: bool = False
|
enable_cudagraph_gc: bool = False
|
||||||
debug_cuda_graph: bool = False
|
debug_cuda_graph: bool = False
|
||||||
@@ -5852,6 +5853,11 @@ class ServerArgs:
|
|||||||
action="store_true",
|
action="store_true",
|
||||||
help="Disable cuda graph when padding is needed. Still uses cuda graph when padding is not needed.",
|
help="Disable cuda graph when padding is needed. Still uses cuda graph when padding is not needed.",
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--enable-breakable-cuda-graph",
|
||||||
|
action="store_true",
|
||||||
|
help="Use breakable CUDA graph for piecewise capture instead of torch.compile-based splitting.",
|
||||||
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--enable-profile-cuda-graph",
|
"--enable-profile-cuda-graph",
|
||||||
action="store_true",
|
action="store_true",
|
||||||
|
|||||||
+53
-7
@@ -1,19 +1,29 @@
|
|||||||
"""Unit tests for the breakable CUDA graph mechanism.
|
"""Tests for the breakable CUDA graph (BCG) runner.
|
||||||
|
|
||||||
Tests the core capture/replay logic with simple tensor operations,
|
Two test classes:
|
||||||
verifying that graph breaks work correctly and outputs are properly
|
- ``TestBreakableCUDAGraphBasic`` / ``TestCopyOutput`` / ``TestBreakGraphHelper``:
|
||||||
propagated across segments.
|
unit tests for the core capture / replay mechanism (simple tensor ops).
|
||||||
|
- ``TestBreakableCudaGraph``: integration test — spin up Qwen3-8B with
|
||||||
|
``--enable-breakable-cuda-graph`` and check mgsm_en accuracy.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.utils import kill_process_tree
|
||||||
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.run_eval import run_eval
|
||||||
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
DEFAULT_URL_FOR_TEST,
|
||||||
|
CustomTestCase,
|
||||||
|
SimpleNamespace,
|
||||||
|
popen_launch_server,
|
||||||
|
)
|
||||||
|
|
||||||
# CI Registration
|
# CI Registration — large suite to fit the integration test's server startup.
|
||||||
register_cuda_ci(est_time=30, suite="stage-b-test-1-gpu-small")
|
register_cuda_ci(est_time=130, suite="stage-b-test-1-gpu-large")
|
||||||
|
|
||||||
|
|
||||||
def _skip_if_no_cuda(test_func):
|
def _skip_if_no_cuda(test_func):
|
||||||
@@ -281,5 +291,41 @@ class TestBreakGraphHelper(CustomTestCase):
|
|||||||
self.assertTrue(torch.allclose(y, torch.full((4,), 13.0, device=self.device)))
|
self.assertTrue(torch.allclose(y, torch.full((4,), 13.0, device=self.device)))
|
||||||
|
|
||||||
|
|
||||||
|
class TestBreakableCudaGraph(CustomTestCase):
|
||||||
|
"""Integration: Qwen3-8B with --enable-breakable-cuda-graph on mgsm_en."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model = "Qwen/Qwen3-8B"
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=[
|
||||||
|
"--enable-breakable-cuda-graph",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
def test_gsm8k_accuracy(self):
|
||||||
|
args = SimpleNamespace(
|
||||||
|
base_url=self.base_url,
|
||||||
|
model=self.model,
|
||||||
|
eval_name="mgsm_en",
|
||||||
|
num_examples=1319,
|
||||||
|
num_threads=1024,
|
||||||
|
)
|
||||||
|
|
||||||
|
metrics = run_eval(args)
|
||||||
|
score = metrics["score"]
|
||||||
|
print(f"mgsm_en accuracy with breakable CUDA graph: {score:.3f}")
|
||||||
|
|
||||||
|
self.assertGreaterEqual(score, 0.80)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
Reference in New Issue
Block a user