From 5e60363960db96d56b519617d7a75be871f15a0e Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 7 Aug 2026 01:33:14 -0700 Subject: [PATCH] Fix prefill CP graph overflow with larger bucket search (#33906) --- python/sglang/srt/layers/cp/bcg.py | 66 ++++++- .../runner/prefill_cuda_graph_runner.py | 31 ++++ test/registered/cp/test_cp_strategy_unit.py | 167 ++++++++++++++++++ 3 files changed, 263 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/layers/cp/bcg.py b/python/sglang/srt/layers/cp/bcg.py index f657eca8e..6abbddcf0 100644 --- a/python/sglang/srt/layers/cp/bcg.py +++ b/python/sglang/srt/layers/cp/bcg.py @@ -17,16 +17,19 @@ from __future__ import annotations from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any, Dict +from typing import TYPE_CHECKING, Any, Dict, Optional import torch +from sglang.srt.layers.cp.base import get_cp_strategy +from sglang.srt.layers.cp.padding import get_cp_padding_align_size from sglang.srt.layers.cp.utils import ( cp_gather_after_forward, cp_split_before_forward, enable_cp_v2, prepare_cp_forward, ) +from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy from sglang.srt.model_executor.forward_batch_info import PPProxyTensors if TYPE_CHECKING: @@ -106,6 +109,67 @@ class PrefillCPBCGInput: ), ) + def required_local_tokens(self, extend_seq_lens: Any) -> Optional[int]: + """Return the aligned CP-local rows required by a live zigzag layout.""" + strategy = get_cp_strategy() + if not isinstance(strategy, ZigzagCPStrategy) or extend_seq_lens is None: + return None + + cp_segment_num = strategy.cp_size * 2 + per_rank_logical_tokens = [0] * strategy.cp_size + for raw_length in extend_seq_lens: + base, remainder = divmod(int(raw_length), cp_segment_num) + for rank in range(strategy.cp_size): + opposite_rank = cp_segment_num - 1 - rank + per_rank_logical_tokens[rank] += ( + base * 2 + int(rank < remainder) + int(opposite_rank < remainder) + ) + align_size = get_cp_padding_align_size() + return ( + (max(per_rank_logical_tokens) + align_size - 1) // align_size * align_size + ) + + def select_replay_bucket( + self, + *, + num_tokens: int, + required_local_tokens: int, + capture_num_tokens: list[int], + max_padding_factor: int, + ) -> Optional[int]: + """Return the smallest global capture whose CP-local rows fit.""" + max_num_tokens = num_tokens * max_padding_factor + for bucket in capture_num_tokens: + if bucket < num_tokens: + continue + if bucket > max_num_tokens: + break + captured_local_tokens = self.bucket_local_tokens.get(bucket) + if ( + captured_local_tokens is not None + and required_local_tokens <= captured_local_tokens + ): + return bucket + return None + + def select_replay_bucket_for_batch( + self, + *, + num_tokens: int, + extend_seq_lens: Any, + capture_num_tokens: list[int], + max_padding_factor: int, + ) -> Optional[int]: + required_local_tokens = self.required_local_tokens(extend_seq_lens) + if required_local_tokens is None: + return None + return self.select_replay_bucket( + num_tokens=num_tokens, + required_local_tokens=required_local_tokens, + capture_num_tokens=capture_num_tokens, + max_padding_factor=max_padding_factor, + ) + def prepare( self, runner: PrefillCudaGraphRunner, diff --git a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py index 2e7a0bf53..0aa5437e6 100644 --- a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py @@ -63,6 +63,7 @@ from sglang.srt.layers.cp.bcg import ( execute_prefill_cp_bcg, filter_prefill_cp_bcg_capture_num_tokens, ) +from sglang.srt.layers.cp.utils import is_cp_v2_active from sglang.srt.layers.dp_attention import ( DpPaddingMode, set_dp_buffer_len, @@ -1137,6 +1138,20 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): ), ): return False + if getattr(self, "enable_cp_v2_bcg_capture", False) and is_cp_v2_active( + forward_batch + ): + assert self.prefill_cp_bcg_input is not None + if ( + self.prefill_cp_bcg_input.select_replay_bucket_for_batch( + num_tokens=len(forward_batch.input_ids), + extend_seq_lens=forward_batch.extend_seq_lens_cpu, + capture_num_tokens=self.capture_num_tokens, + max_padding_factor=_MAX_PREFILL_CUDA_GRAPH_PADDING_FACTOR, + ) + is None + ): + return False # Multi-req replay is supported by body-capture backends via the # layer_model.forward monkey-patch in replay(): the captured graph runs # the transformer stack, then the outer model.forward runs @@ -1418,6 +1433,22 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): """ num_tokens = len(forward_batch.input_ids) static_num_tokens = self._pad_to_bucket(num_tokens, self.capture_num_tokens) + if getattr(self, "enable_cp_v2_bcg_capture", False) and is_cp_v2_active( + forward_batch + ): + assert self.prefill_cp_bcg_input is not None + static_num_tokens = ( + self.prefill_cp_bcg_input.select_replay_bucket_for_batch( + num_tokens=num_tokens, + extend_seq_lens=forward_batch.extend_seq_lens_cpu, + capture_num_tokens=self.capture_num_tokens, + max_padding_factor=_MAX_PREFILL_CUDA_GRAPH_PADDING_FACTOR, + ) + ) + if static_num_tokens is None: + raise RuntimeError( + "Prefill CUDA graph replay was admitted without a fitting bucket" + ) self.raw_num_tokens = num_tokens bs = forward_batch.batch_size diff --git a/test/registered/cp/test_cp_strategy_unit.py b/test/registered/cp/test_cp_strategy_unit.py index 392f7cde0..1e8956892 100644 --- a/test/registered/cp/test_cp_strategy_unit.py +++ b/test/registered/cp/test_cp_strategy_unit.py @@ -14,6 +14,7 @@ from sglang.srt.layers.cp.base import ( is_interleave, is_zigzag, ) +from sglang.srt.layers.cp.bcg import PrefillCPBCGInput from sglang.srt.layers.cp.interleave import InterleaveCPStrategy from sglang.srt.layers.cp.padding import ( get_cp_padding_align_size, @@ -28,6 +29,14 @@ from sglang.srt.layers.cp.utils import ( ) from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy from sglang.srt.mem_cache.memory_pool import KVWriteLoc +from sglang.srt.model_executor.cuda_graph_config import Backend +from sglang.srt.model_executor.forward_batch_info import ( + CaptureHiddenMode, + ForwardMode, +) +from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import ( + PrefillCudaGraphRunner, +) from sglang.srt.runtime_context import get_parallel from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -102,6 +111,164 @@ class TestCPStrategyUnit(CustomTestCase): self.assertIsNotNone(get_cp_strategy()) +class TestPrefillCPBCGReplay(CustomTestCase): + def tearDown(self): + init_cp_strategy(SimpleNamespace(enable_prefill_cp=False)) + + def _make_runner(self): + runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner) + runner._is_full_backend = False + runner.enable_lora = False + runner._capture_chunked_prefix = False + runner.prefill_backend_name = Backend.TC_PIECEWISE + runner.has_mha_companion_layers = False + runner.capture_hidden_mode = CaptureHiddenMode.NULL + runner.capture_num_tokens = [2048, 2304] + runner.max_num_tokens = 2304 + runner.enable_cp_v2_bcg_capture = True + return runner + + def _make_forward_batch(self): + return SimpleNamespace( + batch_size=3, + input_embeds=None, + replace_embeds=None, + mm_inputs=None, + forward_mode=ForwardMode.EXTEND, + capture_hidden_mode=CaptureHiddenMode.NULL, + global_num_tokens_cpu=None, + return_logprob=False, + input_ids=list(range(2048)), + seq_lens_cpu=[1534, 161, 353], + extend_seq_lens_cpu=[1534, 161, 353], + extend_prefix_lens_cpu=[0, 0, 0], + ) + + def _enable_zigzag(self): + init_cp_strategy( + SimpleNamespace( + enable_prefill_cp=True, + cp_strategy="zigzag", + attn_cp_size=4, + ) + ) + + def test_local_capacity_overflow_uses_next_capture_bucket(self): + runner = self._make_runner() + runner.capture_num_tokens.append(2560) + runner.max_num_tokens = 2560 + runner.prefill_cp_bcg_input = PrefillCPBCGInput( + input_embeds=torch.empty(0), + positions=torch.empty(0), + bucket_local_tokens={2048: 512, 2304: 576, 2560: 640}, + ) + forward_batch = self._make_forward_batch() + self._enable_zigzag() + + with ( + patch( + "sglang.srt.environ.envs.SGLANG_ENABLE_CP_V2.get", + return_value=True, + ), + patch( + "sglang.srt.layers.cp.bcg.get_cp_padding_align_size", + return_value=8, + ), + ): + selected_buckets = [] + for cp_rank in range(4): + with get_parallel().override(attn_cp_rank=cp_rank, attn_cp_size=4): + selected_buckets.append( + runner.prefill_cp_bcg_input.select_replay_bucket_for_batch( + num_tokens=2048, + extend_seq_lens=[1534, 161, 353], + capture_num_tokens=runner.capture_num_tokens, + max_padding_factor=2, + ) + ) + + self.assertEqual(selected_buckets, [2304, 2304, 2304, 2304]) + with get_parallel().override(attn_cp_rank=0, attn_cp_size=4): + self.assertTrue(runner.can_run_graph(forward_batch)) + + def test_bucket_search_preserves_two_x_padding_limit(self): + runner = self._make_runner() + runner.prefill_cp_bcg_input = PrefillCPBCGInput( + input_embeds=torch.empty(0), + positions=torch.empty(0), + bucket_local_tokens={2048: 512, 2304: 576}, + ) + + self.assertIsNone( + runner.prefill_cp_bcg_input.select_replay_bucket( + num_tokens=1024, + required_local_tokens=520, + capture_num_tokens=runner.capture_num_tokens, + max_padding_factor=2, + ) + ) + + def test_bucket_search_falls_back_when_no_capture_has_capacity(self): + runner = self._make_runner() + runner.prefill_cp_bcg_input = PrefillCPBCGInput( + input_embeds=torch.empty(0), + positions=torch.empty(0), + bucket_local_tokens={2048: 512, 2304: 516}, + ) + forward_batch = self._make_forward_batch() + self._enable_zigzag() + + with ( + get_parallel().override(attn_cp_rank=0, attn_cp_size=4), + patch( + "sglang.srt.environ.envs.SGLANG_ENABLE_CP_V2.get", + return_value=True, + ), + patch( + "sglang.srt.layers.cp.bcg.get_cp_padding_align_size", + return_value=8, + ), + ): + self.assertFalse(runner.can_run_graph(forward_batch)) + + def test_load_batch_uses_selected_larger_bucket(self): + class StopAfterRecordingFill(Exception): + pass + + class RecordingRegistry: + padded_num_tokens = None + + def fill_from(self, _source, **kwargs): + self.padded_num_tokens = kwargs["padded_num_tokens"] + raise StopAfterRecordingFill + + runner = self._make_runner() + runner.prefill_cp_bcg_input = PrefillCPBCGInput( + input_embeds=torch.empty(0), + positions=torch.empty(0), + bucket_local_tokens={2048: 512, 2304: 576}, + ) + runner.buffer_registry = RecordingRegistry() + forward_batch = self._make_forward_batch() + self._enable_zigzag() + + with ( + get_parallel().override(attn_cp_rank=0, attn_cp_size=4), + patch( + "sglang.srt.environ.envs.SGLANG_ENABLE_CP_V2.get", + return_value=True, + ), + patch( + "sglang.srt.layers.cp.bcg.get_cp_padding_align_size", + return_value=8, + ), + self.assertRaises(StopAfterRecordingFill), + ): + runner.load_batch(forward_batch) + + self.assertEqual(runner.buffer_registry.padded_num_tokens, 2304) + + class TestCPZigzagStrategy(CustomTestCase): def setUp(self): init_cp_strategy(