diff --git a/.buildkite/test_areas/misc.yaml b/.buildkite/test_areas/misc.yaml index 6709dbb56585..ad1c58f08717 100644 --- a/.buildkite/test_areas/misc.yaml +++ b/.buildkite/test_areas/misc.yaml @@ -518,6 +518,7 @@ steps: - vllm/v1/ - tests/test_envs.py - tests/test_outputs.py + - tests/test_pcp_dp.py - tests/test_pooling_params.py - tests/test_ray_env.py - tests/test_sampling_params.py @@ -538,6 +539,7 @@ steps: - python3 standalone_tests/lazy_imports.py - pytest -v -s test_envs.py - pytest -v -s test_outputs.py + - pytest -v -s test_pcp_dp.py - pytest -v -s test_pooling_params.py - pytest -v -s test_ray_env.py - pytest -v -s test_sampling_params.py diff --git a/tests/test_pcp_dp.py b/tests/test_pcp_dp.py new file mode 100644 index 000000000000..f6eed57a6f70 --- /dev/null +++ b/tests/test_pcp_dp.py @@ -0,0 +1,161 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from types import SimpleNamespace + +import pytest +import torch + +from vllm.config import ParallelConfig +from vllm.distributed.device_communicators.all2all import AgRsAll2AllManager +from vllm.forward_context import DPMetadata +from vllm.v1.attention.ops.pcp import maybe_gather_mla_latent_cache_inputs + + +@pytest.mark.parametrize( + "pcp_size,sp_size,enable_ep,expected", + [ + (1, 1, True, [5, 7]), + (1, 2, True, [3, 3, 4, 4]), + (2, 1, False, [10, 14]), + (2, 1, True, [5, 5, 7, 7]), + (2, 2, True, [3, 3, 3, 3, 4, 4, 4, 4]), + ], +) +def test_dispatch_sizes_expand_pcp_before_tp(pcp_size, sp_size, enable_ep, expected): + config = ParallelConfig( + distributed_executor_backend="mp", + data_parallel_size=2, + prefill_context_parallel_size=pcp_size, + tensor_parallel_size=sp_size, + enable_expert_parallel=enable_ep, + ) + metadata = DPMetadata.make(config, 5, torch.tensor([5, 7])) + with metadata.sp_local_sizes(sp_size, pcp_size=pcp_size, use_ep=enable_ep) as sizes: + assert sizes == expected + assert metadata.local_sizes is None + assert metadata.num_tokens_across_dp_cpu.tolist() == [5, 7] + + +@pytest.mark.parametrize( + "dp_size,pcp_size,tp_size,use_ep,is_sp,expected", + [ + (2, 2, 1, True, False, "ep"), + (2, 2, 2, True, True, "ep"), + (2, 2, 2, False, False, "dp"), + (2, 1, 2, True, False, "dp"), + (1, 2, 2, True, False, "pcp"), + (2, 2, 2, True, False, None), + ], +) +def test_dispatch_reuses_existing_groups( + dp_size, pcp_size, tp_size, use_ep, is_sp, expected, monkeypatch +): + groups = { + "dp": SimpleNamespace(world_size=dp_size), + "pcp": SimpleNamespace(world_size=pcp_size), + "ep": object(), + } + for name, group in groups.items(): + monkeypatch.setattr( + f"vllm.distributed.device_communicators.all2all.get_{name}_group", + lambda group=group: group, + ) + manager = AgRsAll2AllManager.__new__(AgRsAll2AllManager) + manager.dp_world_size = dp_size + manager.tp_group = SimpleNamespace(world_size=tp_size) + manager.use_ep = use_ep + if expected is None: + with pytest.raises(AssertionError, match="requires sequence-parallel MoE"): + manager._get_comm_group(is_sp) + else: + assert manager._get_comm_group(is_sp) is groups[expected] + + +@pytest.mark.parametrize("enable_ep", [False, True]) +def test_ag_rs_dispatch_and_combine_use_dp_pcp_sizes(monkeypatch, enable_ep): + calls = [] + local_tokens = [20, 21] if enable_ep else [20, 21, 22, 23] + + class FakeGroup: + world_size = 4 if enable_ep else 2 + rank_in_group = 2 if enable_ep else 1 + + def all_gatherv(self, tensors, dim, sizes): + calls.append(("gather", sizes)) + return [torch.tensor([10, 11, 20, 21, 22, 23]) for _ in tensors] + + def reduce_scatterv(self, tensor, dim, sizes): + calls.append(("scatter", sizes)) + return tensor[2 : 2 + len(local_tokens)] + + config = ParallelConfig( + distributed_executor_backend="mp", + data_parallel_size=2, + data_parallel_rank=1, + prefill_context_parallel_size=2, + enable_expert_parallel=enable_ep, + ) + metadata = DPMetadata.make(config, 2, torch.tensor([1, 2])) + manager = AgRsAll2AllManager.__new__(AgRsAll2AllManager) + manager.dp_world_size = 2 + manager.use_ep = enable_ep + + manager.tp_group = SimpleNamespace(world_size=1) + for name in ("dp", "ep"): + monkeypatch.setattr( + f"vllm.distributed.device_communicators.all2all.get_{name}_group", + FakeGroup, + ) + monkeypatch.setattr( + "vllm.distributed.device_communicators.all2all.get_pcp_group", + lambda: SimpleNamespace(world_size=2), + ) + monkeypatch.setattr( + "vllm.distributed.device_communicators.all2all.get_forward_context", + lambda: SimpleNamespace(dp_metadata=metadata), + ) + with metadata.sp_local_sizes(1, pcp_size=2, use_ep=enable_ep): + hidden_states, _, _ = manager.dispatch( + torch.tensor(local_tokens), + torch.ones(len(local_tokens)), + torch.zeros(len(local_tokens)), + ) + combined = manager.combine(hidden_states) + assert combined.tolist() == local_tokens + sizes = [1, 1, 2, 2] if enable_ep else [2, 4] + assert calls == [("gather", sizes), ("scatter", sizes)] + + +@pytest.mark.parametrize("slots", [[3, 4], [3, 4, -1, -1, -1, -1]]) +def test_decode_cache_write_ignores_dp_padding(slots): + kv = torch.arange(12).reshape(6, 2) + pe = torch.arange(6).reshape(6, 1, 1) + slots = torch.tensor(slots) + cache_kv, cache_pe, cache_slots = maybe_gather_mla_latent_cache_inputs( + kv, pe, slots, num_decode_tokens=2, use_pcp=True + ) + torch.testing.assert_close(cache_kv, kv[:2]) + torch.testing.assert_close(cache_pe, pe[:2]) + torch.testing.assert_close(cache_slots, slots[:2]) + + +def test_expanded_slot_mapping_keeps_pcp_prefill_padding(monkeypatch): + calls = [] + + def all_gather(tensor, dim): + calls.append(tensor.shape[0]) + return torch.cat((tensor, tensor), dim=dim) + + monkeypatch.setattr( + "vllm.v1.attention.ops.pcp.get_pcp_group", + lambda: SimpleNamespace(world_size=2, all_gather=all_gather), + ) + kv = torch.zeros(3, 2) # Two decodes followed by one PCP padding row. + pe = torch.zeros(3, 1, 1) + slots = torch.tensor([3, 4, 8, 3, 4, -1]) + cache_kv, _, cache_slots = maybe_gather_mla_latent_cache_inputs( + kv, pe, slots, num_decode_tokens=2, use_pcp=True + ) + assert calls == [1, 1] + assert cache_kv.shape == (4, 2) + assert cache_slots.tolist() == [3, 4, 8, -1] diff --git a/tests/v1/worker/test_gpu_autoregressive_speculator.py b/tests/v1/worker/test_gpu_autoregressive_speculator.py index f73023cee359..4fc7d4d2e5d5 100644 --- a/tests/v1/worker/test_gpu_autoregressive_speculator.py +++ b/tests/v1/worker/test_gpu_autoregressive_speculator.py @@ -19,6 +19,7 @@ ) from vllm.v1.attention.backends import flash_attn as flash_attn_module from vllm.v1.attention.backends.flash_attn import FlashAttentionMetadata +from vllm.v1.attention.backends.utils import split_decodes_and_prefills from vllm.v1.worker.gpu.cudagraph_utils import BatchExecutionDescriptor from vllm.v1.worker.gpu.spec_decode import speculator as base_spec_module from vllm.v1.worker.gpu.spec_decode.autoregressive import speculator as spec_module @@ -69,6 +70,49 @@ def embed_input_ids( raise AssertionError("embed_input_ids should not be called during loading") +@pytest.mark.parametrize("cg_mode", [CUDAGraphMode.NONE, CUDAGraphMode.FULL]) +def test_pcp_draft_metadata_keeps_graph_padding_in_decode(cg_mode): + def build(common_prefix_len, common_attn_metadata): + return split_decodes_and_prefills( + common_attn_metadata, + decode_threshold=1, + require_uniform=True, + treat_short_extends_as_decodes=False, + ) + + speculator = object.__new__(_TestSpeculator) + speculator.arange_np = torch.arange(5, dtype=torch.int32).numpy() + speculator.max_model_len = speculator.draft_max_seq_len = 32 + speculator.draft_is_prefilling = torch.zeros(4, dtype=torch.bool) + speculator.input_buffers = SimpleNamespace( + query_start_loc=torch.tensor([0, 1, 2, 2, 2], dtype=torch.int32), + seq_lens=torch.tensor([11, 21, 0, 0], dtype=torch.int32), + ) + speculator.block_tables = SimpleNamespace( + cp_size=1, + input_block_tables=[torch.zeros(4, 1, dtype=torch.int32)], + slot_mappings=torch.tensor([[10, 20, -1, -1]]), + ) + speculator.kv_cache_config = SimpleNamespace(kv_cache_groups=[object()]) + speculator.attn_groups = [ + [ + SimpleNamespace( + get_metadata_builder=lambda _: SimpleNamespace(build=build), + layer_names=["draft"], + ) + ] + ] + num_reqs_padded = 4 if cg_mode == CUDAGraphMode.FULL else 2 + metadata = speculator._build_uniform_attn_metadata( + batch_desc=BatchExecutionDescriptor(cg_mode, 4, num_reqs_padded), + num_reqs=2, + num_query_per_req=1, + seq_lens_cpu_upper_bound=torch.tensor([10, 20], dtype=torch.int32), + step=1, + ) + assert metadata["draft"] == (num_reqs_padded, 0, num_reqs_padded, 0) + + def _mock_base_model_load(monkeypatch): monkeypatch.setattr( base_spec_module, diff --git a/tests/v1/worker/test_gpu_pcp_manager.py b/tests/v1/worker/test_gpu_pcp_manager.py index 547e985c179c..0c4b5aff68c2 100644 --- a/tests/v1/worker/test_gpu_pcp_manager.py +++ b/tests/v1/worker/test_gpu_pcp_manager.py @@ -209,6 +209,18 @@ def test_partition_padding_is_derived_from_batch_descriptor( assert local_batch.num_reqs_after_padding == expected_reqs +def test_dummy_draft_does_not_reuse_previous_graph_batch(): + manager, _ = _make_capture_manager(torch.ones((4, 2), dtype=torch.int32)) + dummy_batch = InputBatch.make_dummy(1, 4, manager.input_buffers) + manager.draft_prefill_batch = replace(dummy_batch) + + manager.prepare_draft_prefill(dummy_batch, dummy_batch.input_ids) + + assert ( + manager.get_draft_input_buffers(manager.input_buffers) is manager.input_buffers + ) + + def test_capture_uses_pcp_persistent_inputs(): manager, _ = _make_capture_manager(torch.ones((4, 2), dtype=torch.int32)) diff --git a/vllm/distributed/device_communicators/all2all.py b/vllm/distributed/device_communicators/all2all.py index 9e53c12e24ed..21a1f00da4a6 100644 --- a/vllm/distributed/device_communicators/all2all.py +++ b/vllm/distributed/device_communicators/all2all.py @@ -56,11 +56,17 @@ class AgRsAll2AllManager(All2AllManagerBase): def __init__(self, cpu_group, tcp_store_group=None): super().__init__(cpu_group, tcp_store_group) + self.use_ep = get_current_vllm_config().parallel_config.enable_expert_parallel def _get_comm_group(self, is_sequence_parallel: bool) -> Any: if is_sequence_parallel: return get_ep_group() if self.dp_world_size > 1: + if self.use_ep and get_pcp_group().world_size > 1: + assert self.tp_group.world_size == 1, ( + "DP+PCP with TP>1 requires sequence-parallel MoE inputs" + ) + return get_ep_group() return get_dp_group() return get_pcp_group() @@ -72,6 +78,7 @@ def _get_sizes(self, num_local_tokens: int, comm_group: Any) -> list[int]: assert dp_metadata is not None sizes = dp_metadata.get_chunk_sizes_across_dp_rank() assert sizes is not None + assert len(sizes) == comm_group.world_size return sizes def dispatch_router_logits( diff --git a/vllm/distributed/elastic_ep/standby_state.py b/vllm/distributed/elastic_ep/standby_state.py index c9b12f447d9a..084ed51f8c52 100644 --- a/vllm/distributed/elastic_ep/standby_state.py +++ b/vllm/distributed/elastic_ep/standby_state.py @@ -6,6 +6,7 @@ from vllm.distributed.parallel_state import ( _init_stateless_group, _node_count, + get_pcp_group, get_pp_group, get_tp_group, get_world_group, @@ -72,12 +73,13 @@ def create_standby_groups( _STANDBY_WORLD_NODE_COUNT = _node_count(_STANDBY_WORLD.tcp_store_group) tp_size = get_tp_group().world_size + pcp_size = get_pcp_group().world_size pp_size = get_pp_group().world_size all_ranks = torch.arange(new_world_size_across_dp).reshape( - -1, new_dp_size, pp_size, tp_size + -1, new_dp_size, pp_size, pcp_size, tp_size ) - standby_dp_ranks = all_ranks.transpose(1, 3).reshape(-1, new_dp_size).unbind(0) + standby_dp_ranks = all_ranks.transpose(1, 4).reshape(-1, new_dp_size).unbind(0) standby_dp_ranks = [x.tolist() for x in standby_dp_ranks] # Deferred to commit so the warm-up runs while the engine is paused. @@ -87,7 +89,9 @@ def create_standby_groups( ) standby_ep_ranks = ( - all_ranks.transpose(1, 2).reshape(-1, new_dp_size * tp_size).unbind(0) + all_ranks.transpose(1, 2) + .reshape(-1, new_dp_size * pcp_size * tp_size) + .unbind(0) ) standby_ep_ranks = [x.tolist() for x in standby_ep_ranks] _STANDBY_EP = _init_stateless_group( diff --git a/vllm/forward_context.py b/vllm/forward_context.py index 6a31b3faff1d..9ddc8e9255ed 100644 --- a/vllm/forward_context.py +++ b/vllm/forward_context.py @@ -58,8 +58,17 @@ class BatchDescriptor: def _compute_sp_num_tokens( - num_tokens_across_dp_cpu: torch.Tensor, sequence_parallel_size: int + num_tokens_across_dp_cpu: torch.Tensor, + sequence_parallel_size: int, + pcp_size: int = 1, + use_ep: bool = True, ) -> list[int]: + if pcp_size > 1: + num_tokens_across_dp_cpu = ( + num_tokens_across_dp_cpu.repeat_interleave(pcp_size) + if use_ep + else num_tokens_across_dp_cpu * pcp_size + ) sp_tokens = ( num_tokens_across_dp_cpu + sequence_parallel_size - 1 ) // sequence_parallel_size @@ -98,12 +107,14 @@ def make( return DPMetadata(num_tokens_across_dp_cpu) @contextmanager - def sp_local_sizes(self, sequence_parallel_size: int): + def sp_local_sizes( + self, sequence_parallel_size: int, pcp_size: int = 1, use_ep: bool = False + ): """Context manager for setting self.local_sizes. Same as self.chunked_sizes but without any chunking. """ self.local_sizes = _compute_sp_num_tokens( - self.num_tokens_across_dp_cpu, sequence_parallel_size + self.num_tokens_across_dp_cpu, sequence_parallel_size, pcp_size, use_ep ) try: yield self.local_sizes diff --git a/vllm/model_executor/layers/fused_moe/runner/moe_runner.py b/vllm/model_executor/layers/fused_moe/runner/moe_runner.py index 870d04c0c856..08279d1c8938 100644 --- a/vllm/model_executor/layers/fused_moe/runner/moe_runner.py +++ b/vllm/model_executor/layers/fused_moe/runner/moe_runner.py @@ -657,7 +657,11 @@ def _sequence_parallel_context(self): """ ctx = get_forward_context() return ( - ctx.dp_metadata.sp_local_sizes(self.moe_config.sp_size) + ctx.dp_metadata.sp_local_sizes( + self.moe_config.sp_size, + pcp_size=self.moe_config.pcp_size, + use_ep=self.moe_config.use_ep, + ) if ctx.dp_metadata else nullcontext() ) diff --git a/vllm/platforms/cuda.py b/vllm/platforms/cuda.py index a46b607c8e47..fb1540b60066 100644 --- a/vllm/platforms/cuda.py +++ b/vllm/platforms/cuda.py @@ -329,12 +329,6 @@ def check_and_update_config(cls, vllm_config: VllmConfig) -> None: parallel_config = vllm_config.parallel_config model_config = vllm_config.model_config - if ( - parallel_config.prefill_context_parallel_size > 1 - and parallel_config.data_parallel_size > 1 - ): - raise ValueError("PCP does not support data parallelism on CUDA yet.") - if parallel_config.worker_cls == "auto": parallel_config.worker_cls = "vllm.v1.worker.gpu_worker.Worker" diff --git a/vllm/platforms/rocm.py b/vllm/platforms/rocm.py index 844fe4884b0e..e724ee2cdfe5 100644 --- a/vllm/platforms/rocm.py +++ b/vllm/platforms/rocm.py @@ -934,12 +934,6 @@ def check_and_update_config(cls, vllm_config: "VllmConfig") -> None: compilation_config = vllm_config.compilation_config parallel_config = vllm_config.parallel_config - if ( - parallel_config.prefill_context_parallel_size > 1 - and parallel_config.data_parallel_size > 1 - ): - raise ValueError("PCP does not support data parallelism on ROCm yet.") - if ( compilation_config.cudagraph_mode.has_full_cudagraphs() and parallel_config.prefill_context_parallel_size > 1 diff --git a/vllm/v1/attention/backends/utils.py b/vllm/v1/attention/backends/utils.py index 28366951ac7d..e918774088f4 100644 --- a/vllm/v1/attention/backends/utils.py +++ b/vllm/v1/attention/backends/utils.py @@ -828,7 +828,7 @@ def split_decodes_and_prefills( (query_lens == query_lens[0]) | (query_lens == 0) ): return num_reqs, 0, num_tokens, 0 # all decodes - is_prefill = query_lens != query_lens[0] + is_prefill = (query_lens != query_lens[0]) & (query_lens != 0) else: is_prefill = query_lens > decode_threshold diff --git a/vllm/v1/attention/ops/pcp.py b/vllm/v1/attention/ops/pcp.py index 75ab1c9e8e13..7e7dbe224575 100644 --- a/vllm/v1/attention/ops/pcp.py +++ b/vllm/v1/attention/ops/pcp.py @@ -18,8 +18,15 @@ def _gather_prefill_cache_inputs( assert all(tensor.shape[0] == local_num_tokens for tensor in tensors) assert 0 <= num_decode_tokens <= local_num_tokens - if num_decode_tokens == local_num_tokens: - return tensors, slot_mapping[:num_decode_tokens] + # Replicated draft decodes use unexpanded slot mappings, even with DP padding. + if ( + num_decode_tokens == local_num_tokens + or slot_mapping.shape[0] <= local_num_tokens + ): + return ( + tuple(tensor[:num_decode_tokens] for tensor in tensors), + slot_mapping[:num_decode_tokens], + ) pcp_group = get_pcp_group() gathered_prefills = tuple( diff --git a/vllm/v1/worker/gpu/pcp_manager.py b/vllm/v1/worker/gpu/pcp_manager.py index cf3941813dbc..aff8f04cf300 100644 --- a/vllm/v1/worker/gpu/pcp_manager.py +++ b/vllm/v1/worker/gpu/pcp_manager.py @@ -754,6 +754,7 @@ def get_draft_input_buffers( def prepare_draft_prefill( self, input_batch: InputBatch, input_ids: torch.Tensor ) -> None: + self.draft_prefill_batch = None if input_batch is not self._global_batch or self._local_batch is None: return local_batch = self._local_batch diff --git a/vllm/v1/worker/gpu/spec_decode/speculator.py b/vllm/v1/worker/gpu/spec_decode/speculator.py index 90e02395e076..d1d6d888aabe 100644 --- a/vllm/v1/worker/gpu/spec_decode/speculator.py +++ b/vllm/v1/worker/gpu/spec_decode/speculator.py @@ -364,7 +364,7 @@ def _build_attn_metadata( kv_cache_config=self.kv_cache_config, causal=causal, seq_lens_cpu_upper_bound=draft_seq_lens_cpu_upper_bound, - is_prefilling=self.draft_is_prefilling[:num_reqs], + is_prefilling=self.draft_is_prefilling[:num_reqs_padded], ) return attn_metadata diff --git a/vllm/v1/worker/gpu_worker.py b/vllm/v1/worker/gpu_worker.py index 58af8211278e..856e70e6bd8e 100644 --- a/vllm/v1/worker/gpu_worker.py +++ b/vllm/v1/worker/gpu_worker.py @@ -406,13 +406,14 @@ def init_device(self): if dp_local_rank is None: dp_local_rank = self.parallel_config.data_parallel_index - tp_pp_world_size = ( + tp_pcp_pp_world_size = ( self.parallel_config.pipeline_parallel_size + * self.parallel_config.prefill_context_parallel_size * self.parallel_config.tensor_parallel_size ) - # DP_LOCAL_RANK * TP_PP_WORLD_SIZE + TP_LOCAL_RANK - self.local_rank += dp_local_rank * tp_pp_world_size + # DP_LOCAL_RANK * TP_PCP_PP_WORLD_SIZE + TP_LOCAL_RANK + self.local_rank += dp_local_rank * tp_pcp_pp_world_size # Publish the logical-to-physical mapping for topology queries # such as NIC affinity and P2P checks.