[Fix] Coordinate FullCG prefix variants across DP ranks (#37888)
Co-authored-by: Lianmin Zheng <lianminzheng@gmail.com>
This commit is contained in:
co-authored by
Lianmin Zheng
parent
c8ba8996c4
commit
fbf8f1dbf6
@@ -756,6 +756,7 @@ class TboForwardBatchPreparer:
|
|||||||
"return_logprob",
|
"return_logprob",
|
||||||
"can_run_decode_cuda_graph",
|
"can_run_decode_cuda_graph",
|
||||||
"can_run_dp_prefill_cuda_graph",
|
"can_run_dp_prefill_cuda_graph",
|
||||||
|
"dp_prefill_cuda_graph_max_prefix_len",
|
||||||
"dp_padding_mode",
|
"dp_padding_mode",
|
||||||
"global_forward_mode",
|
"global_forward_mode",
|
||||||
"is_prefill_only",
|
"is_prefill_only",
|
||||||
|
|||||||
@@ -2301,6 +2301,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
is_extend_in_batch: bool = False
|
is_extend_in_batch: bool = False
|
||||||
can_run_decode_cuda_graph: bool = False
|
can_run_decode_cuda_graph: bool = False
|
||||||
can_run_dp_prefill_cuda_graph: bool = False
|
can_run_dp_prefill_cuda_graph: bool = False
|
||||||
|
dp_prefill_cuda_graph_max_prefix_len: int = 0
|
||||||
tbo_split_seq_index: Optional[int] = None
|
tbo_split_seq_index: Optional[int] = None
|
||||||
# Rank-consistent forward mode for the recv skipper, derived from the MLP
|
# Rank-consistent forward mode for the recv skipper, derived from the MLP
|
||||||
# sync all-gather (the TBO-only `global_forward_mode` is None without TBO).
|
# sync all-gather (the TBO-only `global_forward_mode` is None without TBO).
|
||||||
@@ -3600,6 +3601,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
global_num_tokens_for_logprob=self.global_num_tokens_for_logprob,
|
global_num_tokens_for_logprob=self.global_num_tokens_for_logprob,
|
||||||
can_run_decode_cuda_graph=self.can_run_decode_cuda_graph,
|
can_run_decode_cuda_graph=self.can_run_decode_cuda_graph,
|
||||||
can_run_dp_prefill_cuda_graph=self.can_run_dp_prefill_cuda_graph,
|
can_run_dp_prefill_cuda_graph=self.can_run_dp_prefill_cuda_graph,
|
||||||
|
dp_prefill_cuda_graph_max_prefix_len=self.dp_prefill_cuda_graph_max_prefix_len,
|
||||||
is_extend_in_batch=self.is_extend_in_batch,
|
is_extend_in_batch=self.is_extend_in_batch,
|
||||||
is_prefill_only=self.is_prefill_only,
|
is_prefill_only=self.is_prefill_only,
|
||||||
seq_lens_cpu=self.seq_lens_cpu,
|
seq_lens_cpu=self.seq_lens_cpu,
|
||||||
|
|||||||
@@ -97,6 +97,7 @@ class MLPSyncBatchInfo:
|
|||||||
is_extend_in_batch: bool
|
is_extend_in_batch: bool
|
||||||
local_can_run_tbo: bool
|
local_can_run_tbo: bool
|
||||||
local_forward_mode: int
|
local_forward_mode: int
|
||||||
|
prefill_cuda_graph_max_prefix_len: int = 0
|
||||||
|
|
||||||
# some gathered elements
|
# some gathered elements
|
||||||
tp0_info_cpu: torch.Tensor = None
|
tp0_info_cpu: torch.Tensor = None
|
||||||
@@ -116,6 +117,7 @@ class MLPSyncBatchInfo:
|
|||||||
int(self.local_can_run_tbo),
|
int(self.local_can_run_tbo),
|
||||||
self.local_forward_mode,
|
self.local_forward_mode,
|
||||||
int(self.can_run_prefill_cuda_graph),
|
int(self.can_run_prefill_cuda_graph),
|
||||||
|
self.prefill_cuda_graph_max_prefix_len,
|
||||||
],
|
],
|
||||||
device=device,
|
device=device,
|
||||||
dtype=dtype,
|
dtype=dtype,
|
||||||
@@ -131,6 +133,7 @@ class MLPSyncBatchInfo:
|
|||||||
1, # local_can_run_tbo
|
1, # local_can_run_tbo
|
||||||
ForwardMode.IDLE.value, # local_forward_mode
|
ForwardMode.IDLE.value, # local_forward_mode
|
||||||
0, # can_run_prefill_cuda_graph
|
0, # can_run_prefill_cuda_graph
|
||||||
|
0, # prefill_cuda_graph_max_prefix_len
|
||||||
],
|
],
|
||||||
device=device,
|
device=device,
|
||||||
dtype=dtype,
|
dtype=dtype,
|
||||||
@@ -212,6 +215,7 @@ class MLPSyncBatchInfo:
|
|||||||
self.can_run_decode_cuda_graph = bool(tp0_info_cpu[:, 2].min())
|
self.can_run_decode_cuda_graph = bool(tp0_info_cpu[:, 2].min())
|
||||||
self.is_extend_in_batch = bool(tp0_info_cpu[:, 3].max())
|
self.is_extend_in_batch = bool(tp0_info_cpu[:, 3].max())
|
||||||
self.can_run_prefill_cuda_graph = bool(tp0_info_cpu[:, 6].min())
|
self.can_run_prefill_cuda_graph = bool(tp0_info_cpu[:, 6].min())
|
||||||
|
self.prefill_cuda_graph_max_prefix_len = int(tp0_info_cpu[:, 7].max())
|
||||||
if _ENABLE_METRICS_DP_ATTENTION:
|
if _ENABLE_METRICS_DP_ATTENTION:
|
||||||
self.dp_cooperation_info = DPCooperationInfo.create(
|
self.dp_cooperation_info = DPCooperationInfo.create(
|
||||||
tp0_info_cpu[:, 5].tolist()
|
tp0_info_cpu[:, 5].tolist()
|
||||||
@@ -241,6 +245,9 @@ def _update_gather_batch(
|
|||||||
# Check forward mode for cuda graph
|
# Check forward mode for cuda graph
|
||||||
batch.can_run_decode_cuda_graph = mlp_sync_info.can_run_decode_cuda_graph
|
batch.can_run_decode_cuda_graph = mlp_sync_info.can_run_decode_cuda_graph
|
||||||
batch.can_run_dp_prefill_cuda_graph = mlp_sync_info.can_run_prefill_cuda_graph
|
batch.can_run_dp_prefill_cuda_graph = mlp_sync_info.can_run_prefill_cuda_graph
|
||||||
|
batch.dp_prefill_cuda_graph_max_prefix_len = (
|
||||||
|
mlp_sync_info.prefill_cuda_graph_max_prefix_len
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def should_skip_scheduler_all_gather(dp_size: int) -> bool:
|
def should_skip_scheduler_all_gather(dp_size: int) -> bool:
|
||||||
@@ -390,12 +397,18 @@ def prepare_mlp_sync_batch_raw(
|
|||||||
local_batch=local_batch, disable_cuda_graph=disable_cuda_graph
|
local_batch=local_batch, disable_cuda_graph=disable_cuda_graph
|
||||||
)
|
)
|
||||||
breakable_prefill = check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE)
|
breakable_prefill = check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE)
|
||||||
coordinated_prefill = breakable_prefill or check_cuda_graph_backend(
|
full_prefill = check_cuda_graph_backend(Phase.PREFILL, Backend.FULL)
|
||||||
Phase.PREFILL, Backend.FULL
|
coordinated_prefill = breakable_prefill or full_prefill
|
||||||
)
|
|
||||||
prefill_graph_runner = (
|
prefill_graph_runner = (
|
||||||
model_runner.prefill_cuda_graph_runner if coordinated_prefill else None
|
model_runner.prefill_cuda_graph_runner if coordinated_prefill else None
|
||||||
)
|
)
|
||||||
|
prefill_cuda_graph_max_prefix_len = (
|
||||||
|
max(local_batch.prefix_lens, default=0)
|
||||||
|
if full_prefill
|
||||||
|
and local_batch is not None
|
||||||
|
and local_batch.forward_mode in (ForwardMode.EXTEND, ForwardMode.MIXED)
|
||||||
|
else 0
|
||||||
|
)
|
||||||
can_run_prefill_cuda_graph = _local_prefill_cuda_graph_vote(
|
can_run_prefill_cuda_graph = _local_prefill_cuda_graph_vote(
|
||||||
local_batch=local_batch,
|
local_batch=local_batch,
|
||||||
prefill_graph_runner=prefill_graph_runner,
|
prefill_graph_runner=prefill_graph_runner,
|
||||||
@@ -448,6 +461,7 @@ def prepare_mlp_sync_batch_raw(
|
|||||||
is_extend_in_batch=is_extend_in_batch,
|
is_extend_in_batch=is_extend_in_batch,
|
||||||
local_can_run_tbo=local_can_run_tbo,
|
local_can_run_tbo=local_can_run_tbo,
|
||||||
local_forward_mode=local_forward_mode,
|
local_forward_mode=local_forward_mode,
|
||||||
|
prefill_cuda_graph_max_prefix_len=prefill_cuda_graph_max_prefix_len,
|
||||||
)
|
)
|
||||||
|
|
||||||
if dp_size == 1:
|
if dp_size == 1:
|
||||||
|
|||||||
@@ -455,6 +455,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
is_extend_in_batch: bool = False
|
is_extend_in_batch: bool = False
|
||||||
can_run_decode_cuda_graph: bool = False
|
can_run_decode_cuda_graph: bool = False
|
||||||
can_run_dp_prefill_cuda_graph: bool = False
|
can_run_dp_prefill_cuda_graph: bool = False
|
||||||
|
dp_prefill_cuda_graph_max_prefix_len: int = 0
|
||||||
global_forward_mode: Optional[ForwardMode] = None
|
global_forward_mode: Optional[ForwardMode] = None
|
||||||
|
|
||||||
# For two-batch overlap
|
# For two-batch overlap
|
||||||
@@ -827,6 +828,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
is_extend_in_batch=batch.is_extend_in_batch,
|
is_extend_in_batch=batch.is_extend_in_batch,
|
||||||
can_run_decode_cuda_graph=batch.can_run_decode_cuda_graph,
|
can_run_decode_cuda_graph=batch.can_run_decode_cuda_graph,
|
||||||
can_run_dp_prefill_cuda_graph=batch.can_run_dp_prefill_cuda_graph,
|
can_run_dp_prefill_cuda_graph=batch.can_run_dp_prefill_cuda_graph,
|
||||||
|
dp_prefill_cuda_graph_max_prefix_len=batch.dp_prefill_cuda_graph_max_prefix_len,
|
||||||
global_forward_mode=batch.global_forward_mode,
|
global_forward_mode=batch.global_forward_mode,
|
||||||
is_prefill_only=batch.is_prefill_only,
|
is_prefill_only=batch.is_prefill_only,
|
||||||
spec_algorithm=batch.spec_algorithm,
|
spec_algorithm=batch.spec_algorithm,
|
||||||
|
|||||||
@@ -857,13 +857,6 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
int(num_tokens) == 0 for num_tokens in global_num_tokens
|
int(num_tokens) == 0 for num_tokens in global_num_tokens
|
||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _has_prefix_hit(forward_batch: ForwardBatch) -> bool:
|
|
||||||
prefix_lens = forward_batch.extend_prefix_lens_cpu
|
|
||||||
return prefix_lens is not None and any(
|
|
||||||
int(length) > 0 for length in prefix_lens
|
|
||||||
)
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _max_addressable_prefix_len(model_runner) -> int:
|
def _max_addressable_prefix_len(model_runner) -> int:
|
||||||
table_width = model_runner.req_to_token_pool.req_to_token.shape[1]
|
table_width = model_runner.req_to_token_pool.req_to_token.shape[1]
|
||||||
@@ -899,32 +892,39 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
prefix_chunk_len = max(requested_capacity // capture_req_slots, 1)
|
prefix_chunk_len = max(requested_capacity // capture_req_slots, 1)
|
||||||
return prefix_chunk_len, prefix_chunk_len * capture_req_slots
|
return prefix_chunk_len, prefix_chunk_len * capture_req_slots
|
||||||
|
|
||||||
def _select_prefix_capture_chunks(
|
def _select_prefix_capture_chunks(self, max_prefix_len: int) -> Optional[int]:
|
||||||
self, prefix_lens: Sequence[int]
|
"""Smallest captured variant covering max_prefix_len, or None."""
|
||||||
) -> Optional[int]:
|
|
||||||
"""Smallest captured variant covering the batch's max prefix, or None."""
|
|
||||||
max_prefix_len = max(int(length) for length in prefix_lens)
|
|
||||||
real_n = _ceil_div(max_prefix_len, self._prefix_chunk_len)
|
real_n = _ceil_div(max_prefix_len, self._prefix_chunk_len)
|
||||||
return next((n for n in self._prefix_capture_variants if n >= real_n), None)
|
return next((n for n in self._prefix_capture_variants if n >= real_n), None)
|
||||||
|
|
||||||
def _has_uncapturable_chunked_prefix(
|
def _has_uncapturable_chunked_prefix(
|
||||||
self, prefix_lens: Sequence[int] | None
|
self, prefix_lens: Sequence[int] | None
|
||||||
) -> bool:
|
) -> bool:
|
||||||
|
if not self._capture_chunked_prefix or prefix_lens is None:
|
||||||
|
return False
|
||||||
|
max_prefix_len = max(prefix_lens, default=0)
|
||||||
return (
|
return (
|
||||||
self._capture_chunked_prefix
|
max_prefix_len > 0
|
||||||
and prefix_lens is not None
|
and self._select_prefix_capture_chunks(max_prefix_len) is None
|
||||||
and any(int(length) > 0 for length in prefix_lens)
|
|
||||||
and self._select_prefix_capture_chunks(prefix_lens) is None
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def _shape_key(self, num_tokens: int, forward_batch: ForwardBatch) -> ShapeKey:
|
def _shape_key(self, num_tokens: int, forward_batch: ForwardBatch) -> ShapeKey:
|
||||||
variant = None
|
variant = None
|
||||||
if self._capture_chunked_prefix and self._has_prefix_hit(forward_batch):
|
if self._capture_chunked_prefix:
|
||||||
captured_n = self._select_prefix_capture_chunks(
|
prefix_lens = forward_batch.extend_prefix_lens_cpu
|
||||||
forward_batch.extend_prefix_lens_cpu
|
local_max_prefix_len = (
|
||||||
|
max(prefix_lens, default=0) if prefix_lens is not None else 0
|
||||||
)
|
)
|
||||||
assert captured_n is not None, "prefix batch has no captured FullCG variant"
|
max_prefix_len = max(
|
||||||
variant = _chunked_prefix_variant(captured_n)
|
local_max_prefix_len,
|
||||||
|
forward_batch.dp_prefill_cuda_graph_max_prefix_len,
|
||||||
|
)
|
||||||
|
if max_prefix_len > 0:
|
||||||
|
captured_n = self._select_prefix_capture_chunks(max_prefix_len)
|
||||||
|
assert captured_n is not None, (
|
||||||
|
"prefix batch has no captured FullCG variant"
|
||||||
|
)
|
||||||
|
variant = _chunked_prefix_variant(captured_n)
|
||||||
return ShapeKey(size=num_tokens, variant_label=variant)
|
return ShapeKey(size=num_tokens, variant_label=variant)
|
||||||
|
|
||||||
def _create_chunked_prefix_buffers(self) -> _ChunkedPrefixCaptureBuffers:
|
def _create_chunked_prefix_buffers(self) -> _ChunkedPrefixCaptureBuffers:
|
||||||
|
|||||||
@@ -458,6 +458,7 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
|
|||||||
forward_mode=SimpleNamespace(is_target_verify=lambda: False),
|
forward_mode=SimpleNamespace(is_target_verify=lambda: False),
|
||||||
capture_hidden_mode=CaptureHiddenMode.NULL,
|
capture_hidden_mode=CaptureHiddenMode.NULL,
|
||||||
global_num_tokens_cpu=None,
|
global_num_tokens_cpu=None,
|
||||||
|
dp_prefill_cuda_graph_max_prefix_len=0,
|
||||||
return_logprob=False,
|
return_logprob=False,
|
||||||
extend_prefix_lens_cpu=[8],
|
extend_prefix_lens_cpu=[8],
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user