[BCG] Share output buffers across capture sizes + typed ShapeKey (#27857)
This commit is contained in:
@@ -24,6 +24,7 @@ from sglang.srt.constants import GPU_MEMORY_TYPE_CUDA_GRAPH
|
|||||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||||
set_graph_pool_id,
|
set_graph_pool_id,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.model_executor.runner.shape_key import ShapeKey
|
||||||
from sglang.srt.model_executor.runner_backend.base_cuda_graph_backend import (
|
from sglang.srt.model_executor.runner_backend.base_cuda_graph_backend import (
|
||||||
BaseCudaGraphBackend,
|
BaseCudaGraphBackend,
|
||||||
)
|
)
|
||||||
@@ -75,7 +76,7 @@ class NPUCudaGraphBackend(BaseCudaGraphBackend):
|
|||||||
|
|
||||||
def capture_one(
|
def capture_one(
|
||||||
self,
|
self,
|
||||||
shape_key: Any,
|
shape_key: ShapeKey,
|
||||||
forward_fn: Callable[[], Any],
|
forward_fn: Callable[[], Any],
|
||||||
dummies: Optional[Any] = None,
|
dummies: Optional[Any] = None,
|
||||||
post_warmup_hook: Optional[Callable[[], None]] = None,
|
post_warmup_hook: Optional[Callable[[], None]] = None,
|
||||||
@@ -121,7 +122,7 @@ class NPUCudaGraphBackend(BaseCudaGraphBackend):
|
|||||||
self._graphs[shape_key] = graph
|
self._graphs[shape_key] = graph
|
||||||
self._outputs[shape_key] = out
|
self._outputs[shape_key] = out
|
||||||
|
|
||||||
def can_run(self, forward_batch: ForwardBatch, shape_key: Any) -> bool:
|
def can_run(self, forward_batch: ForwardBatch, shape_key: ShapeKey) -> bool:
|
||||||
return shape_key in self._graphs
|
return shape_key in self._graphs
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
@@ -130,7 +131,7 @@ class NPUCudaGraphBackend(BaseCudaGraphBackend):
|
|||||||
|
|
||||||
def replay(
|
def replay(
|
||||||
self,
|
self,
|
||||||
shape_key: Any,
|
shape_key: ShapeKey,
|
||||||
static_forward_batch: ForwardBatch,
|
static_forward_batch: ForwardBatch,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> Any:
|
) -> Any:
|
||||||
@@ -139,7 +140,7 @@ class NPUCudaGraphBackend(BaseCudaGraphBackend):
|
|||||||
|
|
||||||
def replay_with_input_update(
|
def replay_with_input_update(
|
||||||
self,
|
self,
|
||||||
shape_key: Any,
|
shape_key: ShapeKey,
|
||||||
seq_lens: Any,
|
seq_lens: Any,
|
||||||
attr_name: str = None,
|
attr_name: str = None,
|
||||||
attr_type: Any = None,
|
attr_type: Any = None,
|
||||||
|
|||||||
@@ -25,12 +25,10 @@ from sglang.srt.model_executor.runner.base_cuda_graph_runner import ( # noqa: F
|
|||||||
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
|
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
|
||||||
DecodeCudaGraphRunner,
|
DecodeCudaGraphRunner,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import ( # noqa: F401
|
|
||||||
_make_graph_key as _default_make_graph_key,
|
|
||||||
)
|
|
||||||
from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import ( # noqa: F401
|
from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import ( # noqa: F401
|
||||||
PrefillCudaGraphRunner,
|
PrefillCudaGraphRunner,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.model_executor.runner.shape_key import ShapeKey # noqa: F401
|
||||||
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( # noqa: F401
|
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( # noqa: F401
|
||||||
TC_PIECEWISE_CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
TC_PIECEWISE_CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -73,6 +73,7 @@ from sglang.srt.model_executor.runner.base_cuda_graph_runner import (
|
|||||||
freeze_gc,
|
freeze_gc,
|
||||||
get_batch_sizes_to_capture,
|
get_batch_sizes_to_capture,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.model_executor.runner.shape_key import ShapeKey
|
||||||
from sglang.srt.model_executor.runner_backend.breakable_cuda_graph_backend import (
|
from sglang.srt.model_executor.runner_backend.breakable_cuda_graph_backend import (
|
||||||
BreakableCudaGraphBackend,
|
BreakableCudaGraphBackend,
|
||||||
)
|
)
|
||||||
@@ -114,18 +115,6 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
|
|
||||||
|
|
||||||
def _make_graph_key(bs, stream_idx=None, variant_label=None):
|
|
||||||
"""Build a graph dict key from batch size, stream index, and lora variant.
|
|
||||||
|
|
||||||
Standalone function so speculative runners (which don't subclass
|
|
||||||
DecodeCudaGraphRunner) can use the same key encoding.
|
|
||||||
"""
|
|
||||||
key = bs if stream_idx is None else f"{stream_idx}_{bs}"
|
|
||||||
if variant_label is not None:
|
|
||||||
key = f"{variant_label}_{key}"
|
|
||||||
return key
|
|
||||||
|
|
||||||
|
|
||||||
def build_replay_fb_view(
|
def build_replay_fb_view(
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
buffers: DecodeInputBuffers,
|
buffers: DecodeInputBuffers,
|
||||||
@@ -499,7 +488,11 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
return torch.int64
|
return torch.int64
|
||||||
|
|
||||||
def _make_graph_key(self, bs, stream_idx=None, variant_label=None):
|
def _make_graph_key(self, bs, stream_idx=None, variant_label=None):
|
||||||
return _make_graph_key(bs, stream_idx, variant_label)
|
return ShapeKey(
|
||||||
|
size=bs,
|
||||||
|
stream_idx=stream_idx,
|
||||||
|
variant_label=variant_label,
|
||||||
|
)
|
||||||
|
|
||||||
def _resolve_lora_variant(self, forward_batch: ForwardBatch):
|
def _resolve_lora_variant(self, forward_batch: ForwardBatch):
|
||||||
if not getattr(self, "record_nolora_graph", False):
|
if not getattr(self, "record_nolora_graph", False):
|
||||||
|
|||||||
@@ -58,6 +58,7 @@ from sglang.srt.model_executor.runner.base_cuda_graph_runner import (
|
|||||||
BaseCudaGraphRunner,
|
BaseCudaGraphRunner,
|
||||||
freeze_gc,
|
freeze_gc,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.model_executor.runner.shape_key import ShapeKey
|
||||||
from sglang.srt.model_executor.runner_backend.breakable_cuda_graph_backend import (
|
from sglang.srt.model_executor.runner_backend.breakable_cuda_graph_backend import (
|
||||||
BreakableCudaGraphBackend,
|
BreakableCudaGraphBackend,
|
||||||
)
|
)
|
||||||
@@ -658,7 +659,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
else:
|
else:
|
||||||
post_warmup_hook = getattr(attn_backend, "on_after_cuda_graph_warmup", None)
|
post_warmup_hook = getattr(attn_backend, "on_after_cuda_graph_warmup", None)
|
||||||
self.backend.capture_one(
|
self.backend.capture_one(
|
||||||
num_tokens,
|
ShapeKey(size=num_tokens),
|
||||||
run_once,
|
run_once,
|
||||||
dummies=None,
|
dummies=None,
|
||||||
post_warmup_hook=post_warmup_hook,
|
post_warmup_hook=post_warmup_hook,
|
||||||
@@ -820,7 +821,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
# model.forward eagerly with the live multi-req
|
# model.forward eagerly with the live multi-req
|
||||||
# static_forward_batch. The outer's logits_processor /
|
# static_forward_batch. The outer's logits_processor /
|
||||||
# pooler then runs on top with live multi-req metadata.
|
# pooler then runs on top with live multi-req metadata.
|
||||||
shape_key = self._static_num_tokens
|
shape_key = ShapeKey(size=self._static_num_tokens)
|
||||||
|
|
||||||
def replay_layer_forward(*args, **layer_kwargs):
|
def replay_layer_forward(*args, **layer_kwargs):
|
||||||
return self.backend.replay(
|
return self.backend.replay(
|
||||||
|
|||||||
@@ -0,0 +1,23 @@
|
|||||||
|
"""ShapeKey — typed identifier for one captured CUDA-graph shape."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class ShapeKey:
|
||||||
|
"""Identifies one captured CUDA-graph shape across all runners.
|
||||||
|
|
||||||
|
size: the per-phase capture size — what the runner iterates over.
|
||||||
|
- prefill: num_tokens
|
||||||
|
- decode: bs
|
||||||
|
stream_idx: pdmux stream index, or None for single-stream runners.
|
||||||
|
variant_label: LoRA-variant label ("lora" / "nolora"), or None
|
||||||
|
for runners that don't record per-variant graphs.
|
||||||
|
"""
|
||||||
|
|
||||||
|
size: int
|
||||||
|
stream_idx: Optional[int] = None
|
||||||
|
variant_label: Optional[str] = None
|
||||||
@@ -7,6 +7,8 @@ from typing import TYPE_CHECKING, Any, Callable, Iterator, Optional
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.model_executor.runner.shape_key import ShapeKey
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
|
||||||
@@ -43,14 +45,14 @@ class BaseCudaGraphBackend(ABC):
|
|||||||
@abstractmethod
|
@abstractmethod
|
||||||
def capture_one(
|
def capture_one(
|
||||||
self,
|
self,
|
||||||
shape_key: Any,
|
shape_key: ShapeKey,
|
||||||
forward_fn,
|
forward_fn,
|
||||||
dummies: Optional[Any] = None,
|
dummies: Optional[Any] = None,
|
||||||
post_warmup_hook: Optional[Callable[[], None]] = None,
|
post_warmup_hook: Optional[Callable[[], None]] = None,
|
||||||
) -> None: ...
|
) -> None: ...
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def can_run(self, forward_batch: ForwardBatch, shape_key: Any) -> bool: ...
|
def can_run(self, forward_batch: ForwardBatch, shape_key: ShapeKey) -> bool: ...
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def replay_session(self) -> Iterator[None]: ...
|
def replay_session(self) -> Iterator[None]: ...
|
||||||
@@ -58,7 +60,7 @@ class BaseCudaGraphBackend(ABC):
|
|||||||
@abstractmethod
|
@abstractmethod
|
||||||
def replay(
|
def replay(
|
||||||
self,
|
self,
|
||||||
shape_key: Any,
|
shape_key: ShapeKey,
|
||||||
static_forward_batch: ForwardBatch,
|
static_forward_batch: ForwardBatch,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> Any: ...
|
) -> Any: ...
|
||||||
|
|||||||
@@ -13,6 +13,8 @@ import torch
|
|||||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||||
set_graph_pool_id,
|
set_graph_pool_id,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import PPProxyTensors
|
||||||
|
from sglang.srt.model_executor.runner.shape_key import ShapeKey
|
||||||
from sglang.srt.model_executor.runner_backend.base_cuda_graph_backend import (
|
from sglang.srt.model_executor.runner_backend.base_cuda_graph_backend import (
|
||||||
BaseCudaGraphBackend,
|
BaseCudaGraphBackend,
|
||||||
)
|
)
|
||||||
@@ -53,6 +55,7 @@ class BreakableCudaGraphBackend(BaseCudaGraphBackend):
|
|||||||
self._tp_group = cuda_graph_runner.model_runner.tp_group
|
self._tp_group = cuda_graph_runner.model_runner.tp_group
|
||||||
self._capture_stream: Optional[torch.cuda.Stream] = None
|
self._capture_stream: Optional[torch.cuda.Stream] = None
|
||||||
self._debug_eager = debug_eager
|
self._debug_eager = debug_eager
|
||||||
|
self._shared_output_buffer: Optional[Any] = None
|
||||||
self._memory_saver_adapter: Optional[Any] = TorchMemorySaverAdapter.create(
|
self._memory_saver_adapter: Optional[Any] = TorchMemorySaverAdapter.create(
|
||||||
enable=enable_memory_saver
|
enable=enable_memory_saver
|
||||||
and get_bool_env_var("SGLANG_MEMORY_SAVER_CUDA_GRAPH")
|
and get_bool_env_var("SGLANG_MEMORY_SAVER_CUDA_GRAPH")
|
||||||
@@ -71,6 +74,7 @@ class BreakableCudaGraphBackend(BaseCudaGraphBackend):
|
|||||||
self._pool = self._device_module.graph_pool_handle()
|
self._pool = self._device_module.graph_pool_handle()
|
||||||
set_graph_pool_id(self._pool)
|
set_graph_pool_id(self._pool)
|
||||||
self._capture_stream = stream
|
self._capture_stream = stream
|
||||||
|
self._shared_output_buffer = None
|
||||||
try:
|
try:
|
||||||
with self.replay_session():
|
with self.replay_session():
|
||||||
yield
|
yield
|
||||||
@@ -79,7 +83,7 @@ class BreakableCudaGraphBackend(BaseCudaGraphBackend):
|
|||||||
|
|
||||||
def capture_one(
|
def capture_one(
|
||||||
self,
|
self,
|
||||||
shape_key: Any,
|
shape_key: ShapeKey,
|
||||||
forward_fn: Callable[[], Any],
|
forward_fn: Callable[[], Any],
|
||||||
dummies: Optional[Any] = None,
|
dummies: Optional[Any] = None,
|
||||||
post_warmup_hook: Optional[Callable[[], None]] = None,
|
post_warmup_hook: Optional[Callable[[], None]] = None,
|
||||||
@@ -95,16 +99,81 @@ class BreakableCudaGraphBackend(BaseCudaGraphBackend):
|
|||||||
captured_fn = (
|
captured_fn = (
|
||||||
eager_on_graph(True)(forward_fn) if self._debug_eager else forward_fn
|
eager_on_graph(True)(forward_fn) if self._debug_eager else forward_fn
|
||||||
)
|
)
|
||||||
|
size = shape_key.size
|
||||||
with BreakableCUDAGraphCapture(
|
with BreakableCUDAGraphCapture(
|
||||||
cuda_graph=graph,
|
cuda_graph=graph,
|
||||||
pool=self._pool,
|
pool=self._pool,
|
||||||
stream=self._capture_stream,
|
stream=self._capture_stream,
|
||||||
):
|
):
|
||||||
out = captured_fn()
|
out = captured_fn()
|
||||||
self._graphs[shape_key] = graph
|
if self._shared_output_buffer is not None:
|
||||||
self._outputs[shape_key] = out
|
self._copy_output_to_buffer(out, self._shared_output_buffer, size)
|
||||||
|
|
||||||
def can_run(self, forward_batch: ForwardBatch, shape_key: Any) -> bool:
|
if self._shared_output_buffer is None:
|
||||||
|
self._shared_output_buffer = out
|
||||||
|
stored = self._slice_output(out, size)
|
||||||
|
else:
|
||||||
|
stored = self._slice_output(self._shared_output_buffer, size)
|
||||||
|
|
||||||
|
self._graphs[shape_key] = graph
|
||||||
|
self._outputs[shape_key] = stored
|
||||||
|
|
||||||
|
def _slice_output(self, output: Any, num_tokens: int) -> Any:
|
||||||
|
if output is None:
|
||||||
|
return None
|
||||||
|
if torch.is_tensor(output):
|
||||||
|
return output[:num_tokens]
|
||||||
|
if isinstance(output, PPProxyTensors):
|
||||||
|
return output[:num_tokens]
|
||||||
|
if isinstance(output, tuple):
|
||||||
|
return tuple(self._slice_output(item, num_tokens) for item in output)
|
||||||
|
if isinstance(output, list):
|
||||||
|
return [self._slice_output(item, num_tokens) for item in output]
|
||||||
|
raise TypeError(f"Unsupported BCG output type: {type(output)}")
|
||||||
|
|
||||||
|
def _copy_output_to_buffer(
|
||||||
|
self, output: Any, output_buffer: Any, num_tokens: int
|
||||||
|
) -> None:
|
||||||
|
if output is None or output_buffer is None:
|
||||||
|
if output is None and output_buffer is None:
|
||||||
|
return
|
||||||
|
raise ValueError(
|
||||||
|
"BCG output structure changed between capture sizes: "
|
||||||
|
f"{type(output)} vs {type(output_buffer)}"
|
||||||
|
)
|
||||||
|
if torch.is_tensor(output) and torch.is_tensor(output_buffer):
|
||||||
|
output_buffer[:num_tokens].copy_(output[:num_tokens])
|
||||||
|
return
|
||||||
|
if isinstance(output, PPProxyTensors) and isinstance(
|
||||||
|
output_buffer, PPProxyTensors
|
||||||
|
):
|
||||||
|
if output.tensors.keys() != output_buffer.tensors.keys():
|
||||||
|
raise ValueError(
|
||||||
|
"BCG output proxy structure changed between capture sizes: "
|
||||||
|
f"{output.tensors.keys()} != {output_buffer.tensors.keys()}"
|
||||||
|
)
|
||||||
|
for key, tensor in output.tensors.items():
|
||||||
|
self._copy_output_to_buffer(
|
||||||
|
tensor, output_buffer.tensors[key], num_tokens
|
||||||
|
)
|
||||||
|
return
|
||||||
|
if isinstance(output, (list, tuple)) and isinstance(
|
||||||
|
output_buffer, type(output)
|
||||||
|
):
|
||||||
|
if len(output) != len(output_buffer):
|
||||||
|
raise ValueError(
|
||||||
|
"BCG output sequence structure changed between capture sizes: "
|
||||||
|
f"{len(output)} != {len(output_buffer)}"
|
||||||
|
)
|
||||||
|
for item, buffer in zip(output, output_buffer):
|
||||||
|
self._copy_output_to_buffer(item, buffer, num_tokens)
|
||||||
|
return
|
||||||
|
raise TypeError(
|
||||||
|
"Unsupported BCG output buffer pair: "
|
||||||
|
f"{type(output)} vs {type(output_buffer)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
def can_run(self, forward_batch: ForwardBatch, shape_key: ShapeKey) -> bool:
|
||||||
return shape_key in self._graphs
|
return shape_key in self._graphs
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
@@ -114,7 +183,7 @@ class BreakableCudaGraphBackend(BaseCudaGraphBackend):
|
|||||||
|
|
||||||
def replay(
|
def replay(
|
||||||
self,
|
self,
|
||||||
shape_key: Any,
|
shape_key: ShapeKey,
|
||||||
static_forward_batch: ForwardBatch,
|
static_forward_batch: ForwardBatch,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> Any:
|
) -> Any:
|
||||||
@@ -125,3 +194,4 @@ class BreakableCudaGraphBackend(BaseCudaGraphBackend):
|
|||||||
self._graphs.clear()
|
self._graphs.clear()
|
||||||
self._outputs.clear()
|
self._outputs.clear()
|
||||||
self._pool = None
|
self._pool = None
|
||||||
|
self._shared_output_buffer = None
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from sglang.srt.constants import GPU_MEMORY_TYPE_CUDA_GRAPH
|
|||||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||||
set_graph_pool_id,
|
set_graph_pool_id,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.model_executor.runner.shape_key import ShapeKey
|
||||||
from sglang.srt.model_executor.runner_backend.base_cuda_graph_backend import (
|
from sglang.srt.model_executor.runner_backend.base_cuda_graph_backend import (
|
||||||
BaseCudaGraphBackend,
|
BaseCudaGraphBackend,
|
||||||
)
|
)
|
||||||
@@ -62,7 +63,7 @@ class FullCudaGraphBackend(BaseCudaGraphBackend):
|
|||||||
|
|
||||||
def capture_one(
|
def capture_one(
|
||||||
self,
|
self,
|
||||||
shape_key: Any,
|
shape_key: ShapeKey,
|
||||||
forward_fn: Callable[[], Any],
|
forward_fn: Callable[[], Any],
|
||||||
dummies: Optional[Any] = None,
|
dummies: Optional[Any] = None,
|
||||||
post_warmup_hook: Optional[Callable[[], None]] = None,
|
post_warmup_hook: Optional[Callable[[], None]] = None,
|
||||||
@@ -96,7 +97,7 @@ class FullCudaGraphBackend(BaseCudaGraphBackend):
|
|||||||
self._graphs[shape_key] = graph
|
self._graphs[shape_key] = graph
|
||||||
self._outputs[shape_key] = out
|
self._outputs[shape_key] = out
|
||||||
|
|
||||||
def can_run(self, forward_batch: ForwardBatch, shape_key: Any) -> bool:
|
def can_run(self, forward_batch: ForwardBatch, shape_key: ShapeKey) -> bool:
|
||||||
return shape_key in self._graphs
|
return shape_key in self._graphs
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
@@ -105,7 +106,7 @@ class FullCudaGraphBackend(BaseCudaGraphBackend):
|
|||||||
|
|
||||||
def replay(
|
def replay(
|
||||||
self,
|
self,
|
||||||
shape_key: Any,
|
shape_key: ShapeKey,
|
||||||
static_forward_batch: ForwardBatch,
|
static_forward_batch: ForwardBatch,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> Any:
|
) -> Any:
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.layers.moe.utils import get_moe_a2a_backend
|
from sglang.srt.layers.moe.utils import get_moe_a2a_backend
|
||||||
from sglang.srt.layers.utils import MultiPlatformOp
|
from sglang.srt.layers.utils import MultiPlatformOp
|
||||||
|
from sglang.srt.model_executor.runner.shape_key import ShapeKey
|
||||||
from sglang.srt.model_executor.runner_backend.base_cuda_graph_backend import (
|
from sglang.srt.model_executor.runner_backend.base_cuda_graph_backend import (
|
||||||
BaseCudaGraphBackend,
|
BaseCudaGraphBackend,
|
||||||
)
|
)
|
||||||
@@ -193,7 +194,7 @@ class TcPiecewiseCudaGraphBackend(BaseCudaGraphBackend):
|
|||||||
|
|
||||||
def capture_one(
|
def capture_one(
|
||||||
self,
|
self,
|
||||||
shape_key: Any,
|
shape_key: ShapeKey,
|
||||||
forward_fn: Callable[[], Any],
|
forward_fn: Callable[[], Any],
|
||||||
dummies: Optional[Any] = None,
|
dummies: Optional[Any] = None,
|
||||||
post_warmup_hook: Optional[Callable[[], None]] = None,
|
post_warmup_hook: Optional[Callable[[], None]] = None,
|
||||||
@@ -207,7 +208,7 @@ class TcPiecewiseCudaGraphBackend(BaseCudaGraphBackend):
|
|||||||
if post_warmup_hook is not None:
|
if post_warmup_hook is not None:
|
||||||
post_warmup_hook()
|
post_warmup_hook()
|
||||||
|
|
||||||
def can_run(self, forward_batch: ForwardBatch, shape_key: Any) -> bool:
|
def can_run(self, forward_batch: ForwardBatch, shape_key: ShapeKey) -> bool:
|
||||||
# torch.compile manages its per-shape cache internally.
|
# torch.compile manages its per-shape cache internally.
|
||||||
# _run_compile_pass warms every shape in capture_num_tokens at __init__.
|
# _run_compile_pass warms every shape in capture_num_tokens at __init__.
|
||||||
return True
|
return True
|
||||||
@@ -219,7 +220,7 @@ class TcPiecewiseCudaGraphBackend(BaseCudaGraphBackend):
|
|||||||
|
|
||||||
def replay(
|
def replay(
|
||||||
self,
|
self,
|
||||||
shape_key: Any,
|
shape_key: ShapeKey,
|
||||||
static_forward_batch: ForwardBatch,
|
static_forward_batch: ForwardBatch,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> Any:
|
) -> Any:
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
|||||||
from sglang.srt.model_executor.runner import (
|
from sglang.srt.model_executor.runner import (
|
||||||
DecodeCudaGraphRunner,
|
DecodeCudaGraphRunner,
|
||||||
DeepEPCudaGraphRunnerAdapter,
|
DeepEPCudaGraphRunnerAdapter,
|
||||||
|
ShapeKey,
|
||||||
get_batch_sizes_to_capture,
|
get_batch_sizes_to_capture,
|
||||||
model_capture_mode,
|
model_capture_mode,
|
||||||
)
|
)
|
||||||
@@ -247,8 +248,8 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
return torch.int64
|
return torch.int64
|
||||||
|
|
||||||
def _make_graph_key(self, bs, stream_idx=None, variant_label=None):
|
def _make_graph_key(self, bs, stream_idx=None, variant_label=None):
|
||||||
# EAGLE doesn't use stream_idx / lora variants; key is just bs.
|
# EAGLE doesn't use stream_idx / lora variants.
|
||||||
return bs
|
return ShapeKey(size=bs)
|
||||||
|
|
||||||
# -----------------------------------------------------------------
|
# -----------------------------------------------------------------
|
||||||
# can_run
|
# can_run
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
|||||||
from sglang.srt.model_executor.runner import (
|
from sglang.srt.model_executor.runner import (
|
||||||
DecodeCudaGraphRunner,
|
DecodeCudaGraphRunner,
|
||||||
DeepEPCudaGraphRunnerAdapter,
|
DeepEPCudaGraphRunnerAdapter,
|
||||||
|
ShapeKey,
|
||||||
get_batch_sizes_to_capture,
|
get_batch_sizes_to_capture,
|
||||||
model_capture_mode,
|
model_capture_mode,
|
||||||
)
|
)
|
||||||
@@ -259,7 +260,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
return torch.int64
|
return torch.int64
|
||||||
|
|
||||||
def _make_graph_key(self, bs, stream_idx=None, variant_label=None):
|
def _make_graph_key(self, bs, stream_idx=None, variant_label=None):
|
||||||
return bs
|
return ShapeKey(size=bs)
|
||||||
|
|
||||||
def can_run(self, forward_batch: ForwardBatch):
|
def can_run(self, forward_batch: ForwardBatch):
|
||||||
if self.require_mlp_tp_gather:
|
if self.require_mlp_tp_gather:
|
||||||
|
|||||||
@@ -47,6 +47,7 @@ from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
|||||||
from sglang.srt.model_executor.runner import (
|
from sglang.srt.model_executor.runner import (
|
||||||
DecodeCudaGraphRunner,
|
DecodeCudaGraphRunner,
|
||||||
DeepEPCudaGraphRunnerAdapter,
|
DeepEPCudaGraphRunnerAdapter,
|
||||||
|
ShapeKey,
|
||||||
get_batch_sizes_to_capture,
|
get_batch_sizes_to_capture,
|
||||||
model_capture_mode,
|
model_capture_mode,
|
||||||
)
|
)
|
||||||
@@ -286,7 +287,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
return self.backend.replay(shape_key, forward_batch)
|
return self.backend.replay(shape_key, forward_batch)
|
||||||
|
|
||||||
def _make_graph_key(self, bs, stream_idx=None, variant_label=None):
|
def _make_graph_key(self, bs, stream_idx=None, variant_label=None):
|
||||||
return bs
|
return ShapeKey(size=bs)
|
||||||
|
|
||||||
def can_run(self, forward_batch: ForwardBatch):
|
def can_run(self, forward_batch: ForwardBatch):
|
||||||
if self.require_mlp_tp_gather:
|
if self.require_mlp_tp_gather:
|
||||||
|
|||||||
Reference in New Issue
Block a user