[CP] Support breakable CUDA graphs for zigzag strategy (#33136)
This commit is contained in:
@@ -0,0 +1,240 @@
|
|||||||
|
# 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 helpers for context-parallel prefill."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import TYPE_CHECKING, Any, Dict
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.layers.cp.utils import (
|
||||||
|
cp_gather_after_forward,
|
||||||
|
cp_split_before_forward,
|
||||||
|
enable_cp_v2,
|
||||||
|
prepare_cp_forward,
|
||||||
|
)
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import PPProxyTensors
|
||||||
|
from sglang.srt.model_executor.runner.shape_key import ShapeKey
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import (
|
||||||
|
PrefillCudaGraphRunner,
|
||||||
|
)
|
||||||
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
|
||||||
|
|
||||||
|
def supports_prefill_cp_bcg(server_args: ServerArgs) -> bool:
|
||||||
|
"""Return whether the selected prefill-CP configuration supports BCG."""
|
||||||
|
resolved = server_args._resolved()
|
||||||
|
prefill_attention_backend, _ = server_args._resolved_attention_backends()
|
||||||
|
return (
|
||||||
|
server_args.enable_prefill_cp
|
||||||
|
and resolved.attn_cp_size == server_args.tp_size
|
||||||
|
and server_args.cp_strategy == "zigzag"
|
||||||
|
and prefill_attention_backend == "trtllm_mha"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def enable_cp_v2_bcg_capture(server_args: ServerArgs) -> bool:
|
||||||
|
"""Return whether CP-v2 breakable prefill capture is enabled."""
|
||||||
|
return enable_cp_v2() and supports_prefill_cp_bcg(server_args)
|
||||||
|
|
||||||
|
|
||||||
|
def filter_prefill_cp_bcg_capture_num_tokens(
|
||||||
|
capture_num_tokens: list[int], server_args: ServerArgs
|
||||||
|
) -> list[int]:
|
||||||
|
"""Keep only token buckets where the zigzag CP strategy can run."""
|
||||||
|
min_num_tokens = server_args._resolved().attn_cp_size * 2
|
||||||
|
filtered = [size for size in capture_num_tokens if size >= min_num_tokens]
|
||||||
|
if not filtered:
|
||||||
|
raise ValueError(
|
||||||
|
"Prefill CP breakable CUDA graph requires at least one token bucket "
|
||||||
|
f">= {min_num_tokens}, but got {capture_num_tokens}."
|
||||||
|
)
|
||||||
|
return filtered
|
||||||
|
|
||||||
|
|
||||||
|
def _slice_output_rows(output: Any, num_tokens: int) -> Any:
|
||||||
|
if output is None:
|
||||||
|
return None
|
||||||
|
if torch.is_tensor(output) or isinstance(output, PPProxyTensors):
|
||||||
|
return output[:num_tokens]
|
||||||
|
if isinstance(output, tuple):
|
||||||
|
return tuple(_slice_output_rows(item, num_tokens) for item in output)
|
||||||
|
if isinstance(output, list):
|
||||||
|
return [_slice_output_rows(item, num_tokens) for item in output]
|
||||||
|
raise TypeError(f"Unsupported prefill CP BCG output: {type(output)}")
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class PrefillCPBCGInput:
|
||||||
|
"""Fixed-address CP-local inputs and per-bucket replay state."""
|
||||||
|
|
||||||
|
input_embeds: torch.Tensor
|
||||||
|
positions: torch.Tensor
|
||||||
|
bucket_local_tokens: Dict[int, int] = field(default_factory=dict)
|
||||||
|
live_local_tokens: int = 0
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(cls, runner: PrefillCudaGraphRunner) -> PrefillCPBCGInput:
|
||||||
|
with torch.device(runner.device):
|
||||||
|
return cls(
|
||||||
|
input_embeds=torch.zeros(
|
||||||
|
(
|
||||||
|
runner.max_num_tokens,
|
||||||
|
runner.model_runner.model_config.hidden_size,
|
||||||
|
),
|
||||||
|
dtype=runner.model_runner.dtype,
|
||||||
|
),
|
||||||
|
positions=torch.zeros(
|
||||||
|
(runner.max_num_tokens,),
|
||||||
|
dtype=torch.int64,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
def prepare(
|
||||||
|
self,
|
||||||
|
runner: PrefillCudaGraphRunner,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
*,
|
||||||
|
static_num_tokens: int,
|
||||||
|
capture: bool,
|
||||||
|
) -> None:
|
||||||
|
"""Shard global prefill inputs into fixed-address CP-local buffers."""
|
||||||
|
# Replay batches may reuse a ForwardBatch object whose metadata was
|
||||||
|
# built for a different request layout. Always rebuild before sharding.
|
||||||
|
forward_batch.attn_cp_metadata = None
|
||||||
|
prepare_cp_forward(forward_batch)
|
||||||
|
|
||||||
|
captured_local_tokens = None
|
||||||
|
if not capture:
|
||||||
|
try:
|
||||||
|
captured_local_tokens = self.bucket_local_tokens[static_num_tokens]
|
||||||
|
except KeyError as exc:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Missing CP-local capture capacity for global prefill bucket "
|
||||||
|
f"{static_num_tokens}"
|
||||||
|
) from exc
|
||||||
|
|
||||||
|
# Breakable graph segments retain their captured row geometry when
|
||||||
|
# a smaller live batch reuses the same global token bucket.
|
||||||
|
metadata = forward_batch.attn_cp_metadata
|
||||||
|
if hasattr(metadata, "per_rank_actual_token"):
|
||||||
|
live_physical_tokens = max(metadata.per_rank_actual_token)
|
||||||
|
if live_physical_tokens > captured_local_tokens:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Live batch needs {live_physical_tokens} local CP rows, "
|
||||||
|
f"but global prefill bucket {static_num_tokens} has "
|
||||||
|
f"captured capacity {captured_local_tokens}"
|
||||||
|
)
|
||||||
|
cp_size = len(metadata.per_rank_actual_token)
|
||||||
|
metadata.per_rank_actual_token = [captured_local_tokens] * cp_size
|
||||||
|
metadata.max_rank_len = [captured_local_tokens] * cp_size
|
||||||
|
|
||||||
|
raw_tokens = int(forward_batch.extend_num_tokens)
|
||||||
|
global_input_ids = forward_batch.input_ids[:raw_tokens]
|
||||||
|
global_positions = forward_batch.positions[:raw_tokens]
|
||||||
|
global_input_embeds = runner.model_runner.model.get_input_embeddings()(
|
||||||
|
global_input_ids
|
||||||
|
)
|
||||||
|
local_input_embeds, local_positions = cp_split_before_forward(
|
||||||
|
global_input_embeds,
|
||||||
|
global_positions,
|
||||||
|
forward_batch,
|
||||||
|
)
|
||||||
|
live_local_tokens = int(local_input_embeds.shape[0])
|
||||||
|
|
||||||
|
if capture:
|
||||||
|
captured_local_tokens = live_local_tokens
|
||||||
|
self.bucket_local_tokens[static_num_tokens] = captured_local_tokens
|
||||||
|
else:
|
||||||
|
assert captured_local_tokens is not None
|
||||||
|
if live_local_tokens > captured_local_tokens:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Live batch needs {live_local_tokens} local CP rows, but global "
|
||||||
|
f"prefill bucket {static_num_tokens} has captured capacity "
|
||||||
|
f"{captured_local_tokens}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if captured_local_tokens > self.input_embeds.shape[0]:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"CP-local capture needs {captured_local_tokens} rows, but the "
|
||||||
|
f"fixed input buffer has capacity {self.input_embeds.shape[0]}"
|
||||||
|
)
|
||||||
|
|
||||||
|
input_embeds = self.input_embeds[:captured_local_tokens]
|
||||||
|
positions = self.positions[:captured_local_tokens]
|
||||||
|
input_embeds.zero_()
|
||||||
|
positions.zero_()
|
||||||
|
input_embeds[:live_local_tokens].copy_(local_input_embeds)
|
||||||
|
positions[:live_local_tokens].copy_(local_positions)
|
||||||
|
forward_batch.input_embeds = input_embeds
|
||||||
|
forward_batch.positions = positions
|
||||||
|
self.live_local_tokens = live_local_tokens
|
||||||
|
|
||||||
|
|
||||||
|
def execute_prefill_cp_bcg(
|
||||||
|
runner: PrefillCudaGraphRunner,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
static_forward_batch: ForwardBatch,
|
||||||
|
static_num_tokens: int,
|
||||||
|
raw_num_tokens: int,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
"""Replay a CP-local body and run the global gather/logits tail eagerly."""
|
||||||
|
cp_input = runner.prefill_cp_bcg_input
|
||||||
|
assert cp_input is not None
|
||||||
|
model = runner.model_runner.model
|
||||||
|
with runner._prefill_forward_context(
|
||||||
|
static_forward_batch,
|
||||||
|
num_tokens=static_num_tokens,
|
||||||
|
raw_num_tokens=raw_num_tokens,
|
||||||
|
):
|
||||||
|
local_output = runner.backend.replay(
|
||||||
|
ShapeKey(size=static_num_tokens),
|
||||||
|
static_forward_batch,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
local_output = _slice_output_rows(local_output, cp_input.live_local_tokens)
|
||||||
|
|
||||||
|
capture_aux_hidden_states = getattr(model, "capture_aux_hidden_states", False)
|
||||||
|
aux_hidden_states = None
|
||||||
|
if capture_aux_hidden_states:
|
||||||
|
hidden_states, aux_hidden_states = local_output
|
||||||
|
else:
|
||||||
|
hidden_states = local_output
|
||||||
|
|
||||||
|
if not model.pp_group.is_last_rank:
|
||||||
|
return (
|
||||||
|
(hidden_states, aux_hidden_states)
|
||||||
|
if capture_aux_hidden_states
|
||||||
|
else hidden_states
|
||||||
|
)
|
||||||
|
|
||||||
|
hidden_states = cp_gather_after_forward(
|
||||||
|
hidden_states,
|
||||||
|
static_forward_batch,
|
||||||
|
torch.cuda.current_stream(),
|
||||||
|
)
|
||||||
|
return model.logits_processor(
|
||||||
|
forward_batch.input_ids,
|
||||||
|
hidden_states,
|
||||||
|
model.lm_head,
|
||||||
|
forward_batch,
|
||||||
|
aux_hidden_states,
|
||||||
|
)
|
||||||
@@ -76,6 +76,7 @@ from sglang.srt.layers import deep_gemm_wrapper, model_parallel
|
|||||||
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
|
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
|
||||||
from sglang.srt.layers.cp.utils import (
|
from sglang.srt.layers.cp.utils import (
|
||||||
get_cp_strategy,
|
get_cp_strategy,
|
||||||
|
is_cp_v2_active,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||||
from sglang.srt.layers.sampler import create_sampler
|
from sglang.srt.layers.sampler import create_sampler
|
||||||
@@ -234,6 +235,16 @@ elif current_platform.is_out_of_tree():
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _prefill_cuda_graph_allows_context_parallel(
|
||||||
|
prefill_runner, forward_batch: ForwardBatch
|
||||||
|
) -> bool:
|
||||||
|
"""Allow CP only through a runner that captured the validated CP-v2 body."""
|
||||||
|
return get_cp_strategy() is None or (
|
||||||
|
bool(getattr(prefill_runner, "enable_cp_v2_bcg_capture", False))
|
||||||
|
and is_cp_v2_active(forward_batch)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ModelRunnerOutput:
|
class ModelRunnerOutput:
|
||||||
logits_output: Union[LogitsProcessorOutput, PPProxyTensors]
|
logits_output: Union[LogitsProcessorOutput, PPProxyTensors]
|
||||||
@@ -1520,7 +1531,9 @@ class ModelRunner:
|
|||||||
and not isinstance(self.prefill_cuda_graph_runner, EagerRunner)
|
and not isinstance(self.prefill_cuda_graph_runner, EagerRunner)
|
||||||
and self.prefill_cuda_graph_runner is not None
|
and self.prefill_cuda_graph_runner is not None
|
||||||
and self.prefill_cuda_graph_runner.can_run_graph(forward_batch)
|
and self.prefill_cuda_graph_runner.can_run_graph(forward_batch)
|
||||||
and get_cp_strategy() is None
|
and _prefill_cuda_graph_allows_context_parallel(
|
||||||
|
self.prefill_cuda_graph_runner, forward_batch
|
||||||
|
)
|
||||||
):
|
):
|
||||||
# Prefill cuda graph (piecewise).
|
# Prefill cuda graph (piecewise).
|
||||||
kwargs = self._extend_forward_kwargs(forward_batch, pp_proxy_tensors)
|
kwargs = self._extend_forward_kwargs(forward_batch, pp_proxy_tensors)
|
||||||
|
|||||||
@@ -53,6 +53,16 @@ from sglang.kernels.ops.kvcache.kv_indices import (
|
|||||||
from sglang.srt.configs.model_config import is_deepseek_dsa
|
from sglang.srt.configs.model_config import is_deepseek_dsa
|
||||||
from sglang.srt.distributed.parallel_state import graph_capture
|
from sglang.srt.distributed.parallel_state import graph_capture
|
||||||
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
|
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
|
||||||
|
from sglang.srt.layers.cp.bcg import (
|
||||||
|
PrefillCPBCGInput,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.cp.bcg import (
|
||||||
|
enable_cp_v2_bcg_capture as should_enable_cp_v2_bcg_capture,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.cp.bcg import (
|
||||||
|
execute_prefill_cp_bcg,
|
||||||
|
filter_prefill_cp_bcg_capture_num_tokens,
|
||||||
|
)
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
DpPaddingMode,
|
DpPaddingMode,
|
||||||
set_dp_buffer_len,
|
set_dp_buffer_len,
|
||||||
@@ -350,6 +360,8 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
self._is_full_backend = False
|
self._is_full_backend = False
|
||||||
# Same ordering requirement: capture_prepare reads this.
|
# Same ordering requirement: capture_prepare reads this.
|
||||||
self._capture_lora = False
|
self._capture_lora = False
|
||||||
|
self.enable_cp_v2_bcg_capture = False
|
||||||
|
self.prefill_cp_bcg_input: Optional[PrefillCPBCGInput] = None
|
||||||
# TcPiecewise does its compile pass during backend construction.
|
# TcPiecewise does its compile pass during backend construction.
|
||||||
# Wrap only that path with the prefill CUDA graph failure hint.
|
# Wrap only that path with the prefill CUDA graph failure hint.
|
||||||
try:
|
try:
|
||||||
@@ -447,6 +459,16 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
for name in _PREFILL_STATIC_FIELDS
|
for name in _PREFILL_STATIC_FIELDS
|
||||||
}
|
}
|
||||||
|
|
||||||
|
server_args = model_runner.server_args
|
||||||
|
self.enable_cp_v2_bcg_capture = isinstance(
|
||||||
|
self.backend, BreakableCudaGraphBackend
|
||||||
|
) and should_enable_cp_v2_bcg_capture(server_args)
|
||||||
|
if self.enable_cp_v2_bcg_capture:
|
||||||
|
self.capture_num_tokens = filter_prefill_cp_bcg_capture_num_tokens(
|
||||||
|
self.capture_num_tokens, server_args
|
||||||
|
)
|
||||||
|
self.prefill_cp_bcg_input = PrefillCPBCGInput.create(self)
|
||||||
|
|
||||||
# Static hidden_states buffer giving the captured graph a stable
|
# Static hidden_states buffer giving the captured graph a stable
|
||||||
# address; load_batch refreshes it from live spec_info at replay.
|
# address; load_batch refreshes it from live spec_info at replay.
|
||||||
# Draft consumes aux-concatenated hidden states from the target
|
# Draft consumes aux-concatenated hidden states from the target
|
||||||
@@ -1290,6 +1312,14 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
"""
|
"""
|
||||||
num_tokens = size
|
num_tokens = size
|
||||||
forward_batch, attn_backend = self.capture_prepare(num_tokens)
|
forward_batch, attn_backend = self.capture_prepare(num_tokens)
|
||||||
|
if self.enable_cp_v2_bcg_capture:
|
||||||
|
assert self.prefill_cp_bcg_input is not None
|
||||||
|
self.prefill_cp_bcg_input.prepare(
|
||||||
|
self,
|
||||||
|
forward_batch,
|
||||||
|
static_num_tokens=num_tokens,
|
||||||
|
capture=True,
|
||||||
|
)
|
||||||
if forward_batch.lora_ids is not None:
|
if forward_batch.lora_ids is not None:
|
||||||
# Fill the static prefill LoRA batch info the captured kernels
|
# Fill the static prefill LoRA batch info the captured kernels
|
||||||
# will read (all-None ids: ranks stay 0, kernels no-op).
|
# will read (all-None ids: ranks stay 0, kernels no-op).
|
||||||
@@ -1526,8 +1556,19 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
forward_batch.spec_info.hidden_states
|
forward_batch.spec_info.hidden_states
|
||||||
)
|
)
|
||||||
|
|
||||||
|
metadata_forward_batch = forward_batch
|
||||||
|
if self.enable_cp_v2_bcg_capture:
|
||||||
|
assert self.prefill_cp_bcg_input is not None
|
||||||
|
self.prefill_cp_bcg_input.prepare(
|
||||||
|
self,
|
||||||
|
static_forward_batch,
|
||||||
|
static_num_tokens=static_num_tokens,
|
||||||
|
capture=False,
|
||||||
|
)
|
||||||
|
metadata_forward_batch = static_forward_batch
|
||||||
|
|
||||||
self._prepare_forward_metadata_for_replay(
|
self._prepare_forward_metadata_for_replay(
|
||||||
forward_batch, static_forward_batch, static_num_tokens
|
metadata_forward_batch, static_forward_batch, static_num_tokens
|
||||||
)
|
)
|
||||||
|
|
||||||
return static_forward_batch
|
return static_forward_batch
|
||||||
@@ -1667,7 +1708,16 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
if shape_key.variant_label is not None:
|
if shape_key.variant_label is not None:
|
||||||
self._prepare_chunked_prefix_replay(shape_key, forward_batch)
|
self._prepare_chunked_prefix_replay(shape_key, forward_batch)
|
||||||
|
|
||||||
if self._uses_eager_prefill_tail():
|
if self.enable_cp_v2_bcg_capture:
|
||||||
|
output = execute_prefill_cp_bcg(
|
||||||
|
self,
|
||||||
|
forward_batch,
|
||||||
|
static_forward_batch,
|
||||||
|
static_num_tokens,
|
||||||
|
raw_num_tokens,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
elif self._uses_eager_prefill_tail():
|
||||||
output = self._execute_body_capture(
|
output = self._execute_body_capture(
|
||||||
forward_batch,
|
forward_batch,
|
||||||
static_forward_batch,
|
static_forward_batch,
|
||||||
|
|||||||
@@ -4456,6 +4456,7 @@ class ServerArgs:
|
|||||||
is_deepseek_v4,
|
is_deepseek_v4,
|
||||||
is_nemotron_h,
|
is_nemotron_h,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.layers.cp.bcg import supports_prefill_cp_bcg
|
||||||
|
|
||||||
rules = [
|
rules = [
|
||||||
# MLA prefill under BCG takes forward_mha, which has no eager
|
# MLA prefill under BCG takes forward_mha, which has no eager
|
||||||
@@ -4482,7 +4483,8 @@ class ServerArgs:
|
|||||||
# CP all_gather replay size mismatch under BCG.
|
# CP all_gather replay size mismatch under BCG.
|
||||||
(
|
(
|
||||||
"context parallel (attn_cp_size > 1)",
|
"context parallel (attn_cp_size > 1)",
|
||||||
lambda: self._resolved().attn_cp_size > 1,
|
lambda: self._resolved().attn_cp_size > 1
|
||||||
|
and not supports_prefill_cp_bcg(self),
|
||||||
),
|
),
|
||||||
# Capture builds a dummy extend forward with attn_dcp_metadata=None.
|
# Capture builds a dummy extend forward with attn_dcp_metadata=None.
|
||||||
(
|
(
|
||||||
|
|||||||
@@ -22,6 +22,8 @@ class TestGptOss4GpuMxfp4CP(BaseTestGptOss):
|
|||||||
"4",
|
"4",
|
||||||
"--cp-strategy",
|
"--cp-strategy",
|
||||||
"zigzag",
|
"zigzag",
|
||||||
|
"--cuda-graph-backend-prefill",
|
||||||
|
"breakable",
|
||||||
"--cuda-graph-max-bs-decode",
|
"--cuda-graph-max-bs-decode",
|
||||||
"200",
|
"200",
|
||||||
],
|
],
|
||||||
|
|||||||
Reference in New Issue
Block a user