[Tiny] Cuda Graph Refactor Code Style Follow up (#28107)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
29ac249be1
commit
6c3e429ba1
@@ -2205,27 +2205,15 @@ Please consult the documentation below and [server_args.py](https://github.com/s
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>False</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>bool flag (set to enable)</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>--prefill-cuda-graph-backend</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><strong>Deprecated alias</strong> for <code>--cuda-graph-backend-prefill</code>.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`None`</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>breakable</code>, <code>tc_piecewise</code>, <code>disabled</code></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>--decode-cuda-graph-backend</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><strong>Deprecated alias</strong> for <code>--cuda-graph-backend-decode</code>.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`None`</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>full</code>, <code>breakable</code>, <code>tc_piecewise</code>, <code>disabled</code></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>--disable-prefill-cuda-graph</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><strong>Deprecated.</strong> Use <code>--cuda-graph-backend-prefill=disabled</code>.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Disable the prefill-phase CUDA graph. Convenience for <code>--cuda-graph-backend-prefill=disabled</code>.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>False</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>bool flag (set to enable)</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>--disable-decode-cuda-graph</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><strong>Deprecated.</strong> Use <code>--cuda-graph-backend-decode=disabled</code>.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Disable the decode-phase CUDA graph. Convenience for <code>--cuda-graph-backend-decode=disabled</code>.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>False</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>bool flag (set to enable)</td>
|
||||
</tr>
|
||||
|
||||
@@ -39,9 +39,8 @@ def install_canary(
|
||||
|
||||
assert not check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE), (
|
||||
"kv-canary: piecewise cuda graph is not supported by the current "
|
||||
"SingleForwardManager design; set "
|
||||
"--cuda-graph-backend-prefill=disabled (or =breakable) when canary "
|
||||
"is enabled"
|
||||
"SingleForwardManager design; set --cuda-graph-backend-prefill=disabled "
|
||||
"(or =breakable) when canary is enabled"
|
||||
)
|
||||
|
||||
perturb_config = PerturbConfig.from_env()
|
||||
|
||||
@@ -1,3 +1,16 @@
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
"""FB-shared slot registry for the CUDA graph forward paths.
|
||||
|
||||
``CudaGraphBufferRegistry`` is the ForwardBatch → graph-resident buffer mirror
|
||||
|
||||
@@ -1,3 +1,16 @@
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
"""Phase / backend identifiers, the canonical default for
|
||||
cuda_graph_config, and the --cuda-graph-config JSON CLI parser.
|
||||
|
||||
|
||||
@@ -1,3 +1,16 @@
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
"""Shared scaffolding for the prefill and decode CUDA graph runners."""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -73,10 +86,14 @@ def get_batch_sizes_to_capture(
|
||||
if mul_base % get_attention_cp_size() != 0:
|
||||
mul_base *= get_attention_cp_size()
|
||||
|
||||
# pad `num_max_requests` to avoid being filtered out
|
||||
num_max_requests = (num_max_requests + mul_base - 1) // mul_base * mul_base
|
||||
if max(capture_bs) > num_max_requests:
|
||||
# In some cases (e.g., with a small GPU or --max-running-requests), the #max-running-requests
|
||||
# is very small. We add more values here to make sure we capture the maximum bs.
|
||||
capture_bs += [num_max_requests]
|
||||
|
||||
# Model input token count = bs * num_tokens_per_bs; must be a multiple of attn_tp_size.
|
||||
capture_bs = [bs for bs in capture_bs if bs * num_tokens_per_bs % mul_base == 0]
|
||||
capture_bs = [bs for bs in capture_bs if bs <= num_max_requests]
|
||||
capture_bs = list(sorted(set(capture_bs)))
|
||||
|
||||
@@ -359,6 +359,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
self.num_tokens_per_bs = 1
|
||||
if model_runner.spec_algorithm.is_speculative():
|
||||
if self.model_runner.is_draft_worker:
|
||||
# Draft workers can use TARGET_VERIFY mode.
|
||||
if not self.model_runner.spec_algorithm.is_dflash():
|
||||
raise RuntimeError("This should not happen")
|
||||
self.capture_forward_mode = ForwardMode.TARGET_VERIFY
|
||||
@@ -379,13 +380,16 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
if KTRANSFORMERS_AVAILABLE:
|
||||
KTMoEWrapper.set_capture_batch_sizes(self.capture_bs)
|
||||
|
||||
# If returning hidden states is enabled, set initial capture hidden mode to full to avoid double-capture on startup
|
||||
if model_runner.server_args.enable_return_hidden_states:
|
||||
self.capture_hidden_mode = CaptureHiddenMode.FULL
|
||||
|
||||
# Attention backend
|
||||
self.max_bs = max(self.capture_bs)
|
||||
self.max_num_token = self.max_bs * self.num_tokens_per_bs
|
||||
self.attn_backend.init_cuda_graph_state(self.max_bs, self.max_num_token)
|
||||
|
||||
# Init PDMux if needed
|
||||
self.maybe_init_pdmux()
|
||||
self.seq_len_fill_value = (
|
||||
self.attn_backend.get_cuda_graph_seq_len_fill_value()
|
||||
@@ -404,6 +408,9 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
set_torch_compile_config()
|
||||
|
||||
if self.model_runner.server_args.enable_lora:
|
||||
# Phase 2 of LoRA CUDA graph init: dense LoRA batch metadata.
|
||||
# Phase 1 (MoE buffers) was handled earlier in ModelRunner via
|
||||
# lora_manager.init_cuda_graph_moe_buffers().
|
||||
self.model_runner.lora_manager.init_cuda_graph_batch_info(
|
||||
max_bs_in_cuda_graph=self.max_bs,
|
||||
num_tokens_per_bs=self.num_tokens_per_bs,
|
||||
@@ -475,9 +482,6 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
f"Capture cuda graph failed: {e}\n" f"{CUDA_GRAPH_CAPTURE_FAILED_MSG}"
|
||||
)
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# Helpers
|
||||
# -----------------------------------------------------------------
|
||||
def maybe_init_pdmux(self):
|
||||
if self.enable_pdmux:
|
||||
self.stream_groups = get_stream_groups()
|
||||
@@ -503,10 +507,8 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
return "lora"
|
||||
return "nolora"
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# can_run
|
||||
# -----------------------------------------------------------------
|
||||
def can_run(self, forward_batch: ForwardBatch):
|
||||
# Disable for token embedding overrides (dynamic per-request)
|
||||
if forward_batch.replace_embeds is not None:
|
||||
return False
|
||||
if self.require_mlp_tp_gather:
|
||||
@@ -533,6 +535,9 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
if self.require_mlp_sync:
|
||||
is_bs_supported = is_bs_supported and forward_batch.can_run_dp_cuda_graph
|
||||
|
||||
# NOTE: cuda graph cannot handle mixed batch (encoder_len = 0)
|
||||
# If mixed batch cannot be supported, then encoder_lens can be removed in cuda graph
|
||||
# because the full_text_row_masked_out_mask tensor will always be ones
|
||||
is_encoder_lens_supported = (
|
||||
torch.all(forward_batch.encoder_lens > 0)
|
||||
if self.is_encoder_decoder
|
||||
@@ -573,9 +578,6 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
and is_ngram_supported
|
||||
)
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# Profiling helpers
|
||||
# -----------------------------------------------------------------
|
||||
def _init_profile_context_and_memory_record(self):
|
||||
profile_context = profile(
|
||||
activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA],
|
||||
@@ -600,9 +602,6 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
)
|
||||
logger.info(log_message)
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# capture_prepare
|
||||
# -----------------------------------------------------------------
|
||||
def capture_prepare(
|
||||
self,
|
||||
size: int,
|
||||
@@ -645,6 +644,8 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
else None
|
||||
)
|
||||
|
||||
# Adjust for attention TP if needed (matching replay path in
|
||||
# populate_from_forward_batch).
|
||||
buffers.num_token_non_padded[...] = num_tokens
|
||||
if (
|
||||
enable_num_token_non_padded()
|
||||
@@ -658,6 +659,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
buffers.num_token_non_padded.copy_(local)
|
||||
|
||||
pp_proxy_tensors = None
|
||||
# pipeline parallelism
|
||||
if self.pp_size > 1:
|
||||
pp_proxy_tensors = PPProxyTensors(
|
||||
{k: v[:num_tokens] for k, v in buffers.pp_proxy_tensors.items()}
|
||||
@@ -687,10 +689,13 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
)
|
||||
|
||||
if self.model_runner.server_args.enable_lora:
|
||||
# It is safe to capture CUDA graph using empty LoRA id, as the LoRA kernels will always be launched whenever
|
||||
# `--enable-lora` is set to True (and return immediately if the LoRA id is empty for perf optimization).
|
||||
lora_ids = [None] * bs
|
||||
else:
|
||||
lora_ids = None
|
||||
|
||||
# mamba state tracking (registry-owned when enabled)
|
||||
mamba_track_indices = (
|
||||
_slot("mamba_track_indices")
|
||||
if registry.has_slot("mamba_track_indices")
|
||||
@@ -739,6 +744,8 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
bootstrap_room_ids_int=bootstrap_room_ids_int,
|
||||
)
|
||||
|
||||
# Trip the coordinator so the hisparse code path is captured into the
|
||||
# graph; backends read it from self.model_runner.hisparse_coordinator.
|
||||
forward_batch.hisparse_coordinator = self.model_runner.hisparse_coordinator
|
||||
if forward_batch.hisparse_coordinator is not None:
|
||||
forward_batch.hisparse_coordinator.num_real_reqs.fill_(bs)
|
||||
@@ -748,14 +755,14 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
|
||||
return forward_batch, attn_backend, pp_proxy_tensors
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# capture
|
||||
# -----------------------------------------------------------------
|
||||
def capture(self) -> None:
|
||||
profile_context = empty_context()
|
||||
if self.enable_profile_cuda_graph:
|
||||
profile_context = self._init_profile_context_and_memory_record()
|
||||
|
||||
# Trigger CUDA graph capture for specific shapes.
|
||||
# Capture the large shapes first so that the smaller shapes
|
||||
# can reuse the memory pool allocated for the large shapes.
|
||||
with freeze_gc(self.model_runner.server_args.enable_cudagraph_gc):
|
||||
if not self.enable_pdmux:
|
||||
with graph_capture() as graph_capture_context, profile_context as prof:
|
||||
@@ -814,9 +821,6 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
) as forward:
|
||||
self.capture_one_shape(bs, forward, stream_idx, variant_label)
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# capture_one_shape
|
||||
# -----------------------------------------------------------------
|
||||
def capture_one_shape(
|
||||
self,
|
||||
size: int,
|
||||
@@ -905,10 +909,11 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
),
|
||||
)
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# recapture
|
||||
# -----------------------------------------------------------------
|
||||
def recapture_if_needed(self, forward_batch: ForwardBatch):
|
||||
|
||||
# If the required capture_hidden_mode changes, we need to recapture the graph
|
||||
|
||||
# These are the different factors that can influence the capture_hidden_mode
|
||||
capture_hidden_mode_required_by_forward_batch = (
|
||||
forward_batch.capture_hidden_mode
|
||||
)
|
||||
@@ -922,20 +927,21 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
else CaptureHiddenMode.NULL
|
||||
)
|
||||
|
||||
# Determine the highest capture_hidden_mode required
|
||||
# (If we have FULL, we can emulate LAST or NULL)
|
||||
# (If we have LAST, we can emulate NULL)
|
||||
required_capture_hidden_mode = max(
|
||||
capture_hidden_mode_required_by_forward_batch,
|
||||
capture_hidden_mode_required_by_spec_info,
|
||||
capture_hidden_mode_required_for_returning_hidden_states,
|
||||
)
|
||||
|
||||
# If the current hidden mode is no longer aligned with the required hidden mode, we need to set it to what is required and re-capture
|
||||
if self.capture_hidden_mode != required_capture_hidden_mode:
|
||||
self.capture_hidden_mode = required_capture_hidden_mode
|
||||
self.backend.cleanup()
|
||||
self.capture()
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# replay_prepare
|
||||
# -----------------------------------------------------------------
|
||||
def replay_prepare(
|
||||
self,
|
||||
forward_batch: ForwardBatch,
|
||||
@@ -944,6 +950,8 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
self.deepep_adapter.replay()
|
||||
|
||||
if not forward_batch.needs_forward_metadata_init():
|
||||
# Pre-planned (plan-stream replay_prepare already ran).
|
||||
# In speculative decoding, these two fields are still needed.
|
||||
self.buffers.input_ids[: self.raw_num_token].copy_(forward_batch.input_ids)
|
||||
self.buffers.positions[: self.raw_num_token].copy_(forward_batch.positions)
|
||||
if (
|
||||
@@ -995,6 +1003,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
and forward_batch.input_embeds is not None
|
||||
):
|
||||
buffers.input_embeds[:raw_num_token].copy_(forward_batch.input_embeds)
|
||||
# Padded tokens aren't read, so skip zeroing them.
|
||||
if self.enable_two_batch_overlap:
|
||||
self.tbo_plugin.replay_prepare(
|
||||
forward_mode=self.capture_forward_mode,
|
||||
@@ -1021,6 +1030,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
)
|
||||
attn_backend.init_forward_metadata_out_graph(fb_view)
|
||||
|
||||
# Store fields
|
||||
self.raw_bs = raw_bs
|
||||
self.raw_num_token = raw_num_token
|
||||
self.bs = bs
|
||||
@@ -1034,9 +1044,6 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
self.bs, stream_idx, variant_label
|
||||
)
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# replay
|
||||
# -----------------------------------------------------------------
|
||||
def replay(
|
||||
self,
|
||||
forward_batch: ForwardBatch,
|
||||
@@ -1083,9 +1090,6 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
assert isinstance(output, PPProxyTensors)
|
||||
return PPProxyTensors({k: v[: self.bs] for k, v in output.tensors.items()})
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# spec info
|
||||
# -----------------------------------------------------------------
|
||||
def get_spec_info(self, num_tokens: int):
|
||||
spec_info = None
|
||||
if (
|
||||
@@ -1130,6 +1134,8 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
resolve_dflash_verify_mask_policy,
|
||||
)
|
||||
|
||||
# Avoid enabling custom-mask modes during graph capture for backends that
|
||||
# can express DFLASH verify via their built-in causal path.
|
||||
_, build_custom_mask = resolve_dflash_verify_mask_policy(
|
||||
self.model_runner.attn_backend
|
||||
)
|
||||
|
||||
@@ -263,9 +263,6 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
||||
|
||||
self.raw_num_tokens = 0
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# Helpers
|
||||
# -----------------------------------------------------------------
|
||||
def _is_mamba_track_enabled(self) -> bool:
|
||||
return (
|
||||
self.model_runner.server_args.enable_mamba_extra_buffer()
|
||||
@@ -423,9 +420,6 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
||||
static_forward_batch=static_forward_batch,
|
||||
)
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# can_run
|
||||
# -----------------------------------------------------------------
|
||||
def can_run(self, forward_batch: ForwardBatch) -> bool:
|
||||
if forward_batch.input_embeds is not None:
|
||||
return False
|
||||
@@ -468,9 +462,6 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
||||
# logits_processor eagerly on top with live multi-req metadata.
|
||||
return True
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# capture_prepare
|
||||
# -----------------------------------------------------------------
|
||||
def capture_prepare(self, num_tokens: int) -> tuple[ForwardBatch, AttentionBackend]:
|
||||
"""Build a dummy prefill ForwardBatch for capture/warmup at this shape.
|
||||
|
||||
@@ -595,9 +586,6 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
||||
self.tbo_plugin.capture_one_batch_size(forward_batch, num_tokens=num_tokens)
|
||||
return forward_batch, self.model_runner.attn_backend
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# capture
|
||||
# -----------------------------------------------------------------
|
||||
def capture(self) -> None:
|
||||
with freeze_gc(self.model_runner.server_args.enable_cudagraph_gc):
|
||||
with graph_capture() as graph_capture_context:
|
||||
@@ -628,9 +616,6 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
||||
)
|
||||
self.capture_one_shape(num_tokens)
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# capture_one_shape
|
||||
# -----------------------------------------------------------------
|
||||
def capture_one_shape(self, size: int) -> None:
|
||||
"""Per-shape capture: build dummy ForwardBatch + run_once,
|
||||
delegate to backend. size is the prefill token count.
|
||||
@@ -663,9 +648,6 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
||||
post_warmup_hook=post_warmup_hook,
|
||||
)
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# replay_prepare
|
||||
# -----------------------------------------------------------------
|
||||
def replay_prepare(self, forward_batch: ForwardBatch, **kwargs) -> ForwardBatch:
|
||||
"""Pad, populate static buffers, and build the static_forward_batch
|
||||
the model code reads during replay.
|
||||
@@ -800,9 +782,6 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
||||
self._static_num_tokens = static_num_tokens
|
||||
return static_forward_batch
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# replay
|
||||
# -----------------------------------------------------------------
|
||||
def replay(
|
||||
self, forward_batch: ForwardBatch, **kwargs
|
||||
) -> Union[LogitsProcessorOutput, PPProxyTensors, EmbeddingPoolerOutput]:
|
||||
|
||||
@@ -1,3 +1,16 @@
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
"""ShapeKey — typed identifier for one captured CUDA-graph shape."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -1,3 +1,16 @@
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
"""Backend interface for CUDA graph capture/replay."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -1,3 +1,16 @@
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
"""BreakableCudaGraphBackend — segment-captured graphs with eager break
|
||||
markers (eager_on_graph decorators on attention / mamba layers).
|
||||
No torch.compile.
|
||||
|
||||
@@ -1,3 +1,16 @@
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
"""FullCudaGraphBackend — captures the entire model forward as one
|
||||
torch.cuda.CUDAGraph per shape.
|
||||
"""
|
||||
|
||||
@@ -1,3 +1,16 @@
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
"""TcPiecewiseCudaGraphBackend — torch.compile-driven piecewise CUDA graph.
|
||||
|
||||
FX-splits the model forward at attention layers; per-shape compiled
|
||||
|
||||
@@ -1,3 +1,16 @@
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
"""runner_backend utilities — phase → BaseCudaGraphBackend resolution.
|
||||
|
||||
Centralizes per-phase backend resolution so platform overrides (NPU,
|
||||
|
||||
@@ -15,8 +15,10 @@ Backends in cuda_graph_backend/ import from here. Runners do not.
|
||||
# piecewise_cuda_graph.context_manager and points users at
|
||||
# --disable-piecewise-cuda-graph, which doesn't apply here.
|
||||
CUDA_GRAPH_CAPTURE_FAILED_MSG = (
|
||||
"CUDA graph capture failed.\n"
|
||||
"To work around this error, add --disable-cuda-graph to your launch command\n"
|
||||
"(or use --disable-decode-cuda-graph to disable only the decode phase).\n"
|
||||
"Please report this issue at https://github.com/sgl-project/sglang/issues/new/choose"
|
||||
"Possible solutions:\n"
|
||||
"1. set --mem-fraction-static to a smaller value (e.g., 0.8 or 0.7)\n"
|
||||
"2. set --cuda-graph-max-bs-decode to a smaller value (e.g., 16)\n"
|
||||
"3. disable torch compile by not using --enable-torch-compile\n"
|
||||
"4. disable CUDA graph by --cuda-graph-backend-decode=disabled. (Not recommended. Huge performance loss)\n"
|
||||
"Open an issue on GitHub https://github.com/sgl-project/sglang/issues/new/choose \n"
|
||||
)
|
||||
|
||||
+2
-2
@@ -112,7 +112,7 @@ def set_tc_piecewise_forward_context(
|
||||
|
||||
|
||||
TC_PIECEWISE_CUDA_GRAPH_CAPTURE_FAILED_MSG = (
|
||||
"Piecewise CUDA Graph is enabled by default as an experimental feature.\n"
|
||||
"To work around this error, add --disable-piecewise-cuda-graph to your launch command.\n"
|
||||
"Piecewise CUDA Graph capture failed.\n"
|
||||
"To work around this error, add --cuda-graph-backend-prefill=disabled to your launch command.\n"
|
||||
"Please report this issue at https://github.com/sgl-project/sglang/issues/new/choose"
|
||||
)
|
||||
|
||||
@@ -1,3 +1,16 @@
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
"""Static-buffer dataclasses used by the CUDA graph runners.
|
||||
|
||||
DecodeInputBuffers backs the decode-phase capture/replay path.
|
||||
|
||||
@@ -1,3 +1,16 @@
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
"""Process-global capture-mode flags shared by the decode runner and the
|
||||
speculative-draft runners. Read by model code that needs to take a
|
||||
capture-time branch (e.g. lora dual-graph capture decides per-batch
|
||||
|
||||
@@ -1,3 +1,16 @@
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
"""DeepEP capture/replay adapter — records the dispatch mode used during
|
||||
capture and re-applies it during replay so DeepEP all-to-all has
|
||||
consistent expert routing across the captured graph.
|
||||
|
||||
@@ -1,3 +1,16 @@
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
"""Shared graph memory pool used by the speculative-draft cuda graph
|
||||
runners. The new DecodeCudaGraphRunner and PrefillCudaGraphRunner
|
||||
backends each own their pool internally; this global is retained for the
|
||||
|
||||
@@ -763,18 +763,14 @@ class ServerArgs:
|
||||
cuda_graph_bs_decode: Optional[List[int]] = None
|
||||
cuda_graph_bs_prefill: Optional[List[int]] = None
|
||||
cuda_graph_tc_compiler: Optional[Literal["eager", "inductor"]] = None
|
||||
# Boolean per-phase off-switches; convenience for
|
||||
# --cuda-graph-backend-{prefill,decode}=disabled.
|
||||
disable_prefill_cuda_graph: bool = False
|
||||
disable_decode_cuda_graph: bool = False
|
||||
|
||||
# Legacy CLI inputs that fold into cuda_graph_config (with a CLI
|
||||
# deprecation warning). Internal-only after parsing.
|
||||
disable_cuda_graph: bool = False
|
||||
disable_prefill_cuda_graph: bool = False
|
||||
disable_decode_cuda_graph: bool = False
|
||||
prefill_cuda_graph_backend: Optional[
|
||||
Literal["breakable", "tc_piecewise", "disabled"]
|
||||
] = None
|
||||
decode_cuda_graph_backend: Optional[
|
||||
Literal["full", "breakable", "tc_piecewise", "disabled"]
|
||||
] = None
|
||||
enable_layerwise_nvtx_marker: bool = False
|
||||
enable_nccl_nvls: bool = False
|
||||
enable_symm_mem: bool = False
|
||||
@@ -1414,15 +1410,13 @@ class ServerArgs:
|
||||
_set(Phase.DECODE, "backend", Backend.DISABLED)
|
||||
_set(Phase.PREFILL, "backend", Backend.DISABLED)
|
||||
|
||||
# ---- Legacy convenience flags ----
|
||||
# ---- Boolean per-phase off-switches ----
|
||||
# Below the explicit backend selectors so --cuda-graph-backend-*
|
||||
# wins if both are given.
|
||||
if self.disable_prefill_cuda_graph:
|
||||
_set(Phase.PREFILL, "backend", Backend.DISABLED)
|
||||
if self.disable_decode_cuda_graph:
|
||||
_set(Phase.DECODE, "backend", Backend.DISABLED)
|
||||
if self.prefill_cuda_graph_backend is not None:
|
||||
_set(Phase.PREFILL, "backend", self.prefill_cuda_graph_backend)
|
||||
if self.decode_cuda_graph_backend is not None:
|
||||
_set(Phase.DECODE, "backend", self.decode_cuda_graph_backend)
|
||||
|
||||
# ---- Per-phase convenience flags ----
|
||||
if self.cuda_graph_backend_decode is not None:
|
||||
@@ -6788,6 +6782,18 @@ class ServerArgs:
|
||||
default=ServerArgs.cuda_graph_tc_compiler,
|
||||
help="Compiler used by the tc_piecewise backend (currently only the prefill phase consumes it).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--disable-prefill-cuda-graph",
|
||||
action="store_true",
|
||||
help="Disable the prefill-phase CUDA graph. Convenience for "
|
||||
"--cuda-graph-backend-prefill=disabled.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--disable-decode-cuda-graph",
|
||||
action="store_true",
|
||||
help="Disable the decode-phase CUDA graph. Convenience for "
|
||||
"--cuda-graph-backend-decode=disabled.",
|
||||
)
|
||||
|
||||
# --- CUDA graph: debug / profiling flags -------------------------
|
||||
parser.add_argument(
|
||||
@@ -6846,34 +6852,6 @@ class ServerArgs:
|
||||
new_flag="--cuda-graph-backend-prefill=breakable",
|
||||
help="Deprecated alias for --cuda-graph-backend-prefill=breakable.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prefill-cuda-graph-backend",
|
||||
type=str,
|
||||
choices=Backend.ALL,
|
||||
action=DeprecatedAliasStoreAction,
|
||||
new_flag="--cuda-graph-backend-prefill",
|
||||
help="Deprecated alias for --cuda-graph-backend-prefill.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--decode-cuda-graph-backend",
|
||||
type=str,
|
||||
choices=Backend.ALL,
|
||||
action=DeprecatedAliasStoreAction,
|
||||
new_flag="--cuda-graph-backend-decode",
|
||||
help="Deprecated alias for --cuda-graph-backend-decode.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--disable-prefill-cuda-graph",
|
||||
action=DeprecatedStoreTrueAction,
|
||||
new_flag="--cuda-graph-backend-prefill=disabled",
|
||||
help="Deprecated. Use --cuda-graph-backend-prefill=disabled instead.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--disable-decode-cuda-graph",
|
||||
action=DeprecatedStoreTrueAction,
|
||||
new_flag="--cuda-graph-backend-decode=disabled",
|
||||
help="Deprecated. Use --cuda-graph-backend-decode=disabled instead.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--disable-piecewise-cuda-graph",
|
||||
action=DeprecatedStoreConstAction,
|
||||
|
||||
@@ -401,8 +401,6 @@ class TestDSV4BreakableCudaGraphMetadataContract(CustomTestCase):
|
||||
sa = ServerArgs(model_path="dummy")
|
||||
self.assertNotEqual(sa.cuda_graph_backend_decode, "breakable")
|
||||
self.assertNotEqual(sa.cuda_graph_backend_prefill, "breakable")
|
||||
self.assertNotEqual(sa.decode_cuda_graph_backend, "breakable")
|
||||
self.assertNotEqual(sa.prefill_cuda_graph_backend, "breakable")
|
||||
self.assertFalse(
|
||||
AttentionBackend.use_captured_forward_metadata_for_breakable_cuda_graph
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user