Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
34 commits
Select commit Hold shift + click to select a range
348748a
Support prefill context parallelism with data parallelism
pisceskkk Aug 28, 2026
10f43e9
Use a combined MoE DP-PCP communication group
pisceskkk Aug 28, 2026
4ed8abf
Use the MoE DP-PCP group for static non-SP dispatch
pisceskkk Aug 28, 2026
0c3b063
Route non-SP MoE through the DP-PCP group
pisceskkk Aug 28, 2026
3cc259b
Coordinate MoE token counts across DP and PCP
pisceskkk Aug 31, 2026
0696b03
[MoE] Rename non-SP communication metadata
pisceskkk Aug 31, 2026
168e2b9
[MoE] Preserve AgRs dispatch statement order
pisceskkk Aug 31, 2026
a1f56ce
Adapt #54131 to current main
LucasWilkinson Sep 16, 2026
c2d5a74
Drop the platform-level PCP-with-DP guard
LucasWilkinson Sep 16, 2026
c64295e
Ignore DP padding rows in the PCP decode cache write
LucasWilkinson Sep 16, 2026
fefc688
Name the condition for MoE dispatch spanning DP x PCP
LucasWilkinson Sep 16, 2026
e3c0ca7
Carry MoE non-SP token counts through DBO microbatching
LucasWilkinson Sep 16, 2026
de81d52
Thread the DP sync state through the draft speculators
LucasWilkinson Sep 16, 2026
e86a59b
Simplify PCP and DP coordination and fix draft graph padding
LucasWilkinson Sep 16, 2026
a91da6a
Unify PCP dispatch sizes and group selection
LucasWilkinson Sep 16, 2026
cd83b66
Reuse EP group for PCP and DP dispatch
LucasWilkinson Sep 16, 2026
4b2b65b
Simplify PCP dispatch metadata
LucasWilkinson Sep 16, 2026
cedaa81
Merge main into PCP and DP support
LucasWilkinson Sep 16, 2026
9dfda40
Drop GLM-4 MoE changes from PCP and DP support
LucasWilkinson Sep 16, 2026
bf0dca8
Merge branch 'main' into pcp-dp
robertgshaw2-redhat Sep 17, 2026
3fb1a32
Update vllm/forward_context.py
LucasWilkinson Sep 17, 2026
3af3a66
Merge branch 'main' into pcp-dp
LucasWilkinson Sep 17, 2026
84720ec
Merge branch 'main' into pcp-dp
robertgshaw2-redhat Sep 17, 2026
b714395
Merge branch 'main' into pcp-dp
mergify[bot] Sep 18, 2026
854c055
Merge branch 'main' into pcp-dp
LucasWilkinson Sep 18, 2026
03afdaf
Merge branch 'main' into pcp-dp
LucasWilkinson Sep 21, 2026
22947aa
Merge branch 'main' into pcp-dp
LucasWilkinson Sep 21, 2026
cff6374
Merge branch 'main' into pcp-dp
mergify[bot] Sep 22, 2026
cbdaef5
[CI] Run tests/test_pcp_dp.py in CPU job
LucasWilkinson Sep 23, 2026
2765c63
Merge branch 'main' into pcp-dp
mergify[bot] Sep 23, 2026
ebbcc08
Merge branch 'main' into pcp-dp
LucasWilkinson Sep 23, 2026
cc6cf84
Merge branch 'main' into pcp-dp
robertgshaw2-redhat Sep 23, 2026
57a7d05
Merge branch 'main' into pcp-dp
robertgshaw2-redhat Sep 23, 2026
cb35957
Merge branch 'main' into pcp-dp
robertgshaw2-redhat Sep 24, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions .buildkite/test_areas/misc.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down
161 changes: 161 additions & 0 deletions tests/test_pcp_dp.py
Original file line number Diff line number Diff line change
@@ -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]
44 changes: 44 additions & 0 deletions tests/v1/worker/test_gpu_autoregressive_speculator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down
12 changes: 12 additions & 0 deletions tests/v1/worker/test_gpu_pcp_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))

Expand Down
7 changes: 7 additions & 0 deletions vllm/distributed/device_communicators/all2all.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand All @@ -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(
Expand Down
10 changes: 7 additions & 3 deletions vllm/distributed/elastic_ep/standby_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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.
Expand All @@ -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(
Expand Down
17 changes: 14 additions & 3 deletions vllm/forward_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
6 changes: 5 additions & 1 deletion vllm/model_executor/layers/fused_moe/runner/moe_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
)
Expand Down
6 changes: 0 additions & 6 deletions vllm/platforms/cuda.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"

Expand Down
6 changes: 0 additions & 6 deletions vllm/platforms/rocm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion vllm/v1/attention/backends/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Loading
Loading