Fix prefill CP graph overflow with larger bucket search (#33906)
This commit is contained in:
@@ -17,16 +17,19 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import TYPE_CHECKING, Any, Dict
|
from typing import TYPE_CHECKING, Any, Dict, Optional
|
||||||
|
|
||||||
import torch
|
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 (
|
from sglang.srt.layers.cp.utils import (
|
||||||
cp_gather_after_forward,
|
cp_gather_after_forward,
|
||||||
cp_split_before_forward,
|
cp_split_before_forward,
|
||||||
enable_cp_v2,
|
enable_cp_v2,
|
||||||
prepare_cp_forward,
|
prepare_cp_forward,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy
|
||||||
from sglang.srt.model_executor.forward_batch_info import PPProxyTensors
|
from sglang.srt.model_executor.forward_batch_info import PPProxyTensors
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
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(
|
def prepare(
|
||||||
self,
|
self,
|
||||||
runner: PrefillCudaGraphRunner,
|
runner: PrefillCudaGraphRunner,
|
||||||
|
|||||||
@@ -63,6 +63,7 @@ from sglang.srt.layers.cp.bcg import (
|
|||||||
execute_prefill_cp_bcg,
|
execute_prefill_cp_bcg,
|
||||||
filter_prefill_cp_bcg_capture_num_tokens,
|
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 (
|
from sglang.srt.layers.dp_attention import (
|
||||||
DpPaddingMode,
|
DpPaddingMode,
|
||||||
set_dp_buffer_len,
|
set_dp_buffer_len,
|
||||||
@@ -1137,6 +1138,20 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
),
|
),
|
||||||
):
|
):
|
||||||
return False
|
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
|
# Multi-req replay is supported by body-capture backends via the
|
||||||
# layer_model.forward monkey-patch in replay(): the captured graph runs
|
# layer_model.forward monkey-patch in replay(): the captured graph runs
|
||||||
# the transformer stack, then the outer model.forward runs
|
# the transformer stack, then the outer model.forward runs
|
||||||
@@ -1418,6 +1433,22 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
"""
|
"""
|
||||||
num_tokens = len(forward_batch.input_ids)
|
num_tokens = len(forward_batch.input_ids)
|
||||||
static_num_tokens = self._pad_to_bucket(num_tokens, self.capture_num_tokens)
|
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
|
self.raw_num_tokens = num_tokens
|
||||||
|
|
||||||
bs = forward_batch.batch_size
|
bs = forward_batch.batch_size
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from sglang.srt.layers.cp.base import (
|
|||||||
is_interleave,
|
is_interleave,
|
||||||
is_zigzag,
|
is_zigzag,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.layers.cp.bcg import PrefillCPBCGInput
|
||||||
from sglang.srt.layers.cp.interleave import InterleaveCPStrategy
|
from sglang.srt.layers.cp.interleave import InterleaveCPStrategy
|
||||||
from sglang.srt.layers.cp.padding import (
|
from sglang.srt.layers.cp.padding import (
|
||||||
get_cp_padding_align_size,
|
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.layers.cp.zigzag import ZigzagCPStrategy
|
||||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
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.srt.runtime_context import get_parallel
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
@@ -102,6 +111,164 @@ class TestCPStrategyUnit(CustomTestCase):
|
|||||||
self.assertIsNotNone(get_cp_strategy())
|
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):
|
class TestCPZigzagStrategy(CustomTestCase):
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
init_cp_strategy(
|
init_cp_strategy(
|
||||||
|
|||||||
Reference in New Issue
Block a user