diff --git a/vllm/config/parallel.py b/vllm/config/parallel.py index 5769d5345cdd..bb1d2095de7b 100644 --- a/vllm/config/parallel.py +++ b/vllm/config/parallel.py @@ -547,8 +547,6 @@ def _validate_parallel_config(self) -> Self: tp = self.tensor_parallel_size pcp = self.prefill_context_parallel_size dcp = self.decode_context_parallel_size - if pcp > 1 and self.data_parallel_size > 1: - raise ValueError("PCP does not support data parallelism yet.") if pcp == 1: # DCP reuses the TP ranks when PCP is disabled. if tp % dcp != 0: diff --git a/vllm/platforms/cuda.py b/vllm/platforms/cuda.py index d1df152c63de..1e5a5a3ad02b 100644 --- a/vllm/platforms/cuda.py +++ b/vllm/platforms/cuda.py @@ -322,6 +322,12 @@ 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 280dec2d3181..fe4a19369e88 100644 --- a/vllm/platforms/rocm.py +++ b/vllm/platforms/rocm.py @@ -900,6 +900,12 @@ 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