Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
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
71 changes: 56 additions & 15 deletions tests/v1/kv_connector/unit/test_nixl_connector.py
Original file line number Diff line number Diff line change
Expand Up @@ -576,19 +576,38 @@ def _nixl_handshake(


class TestNixlHandshake:
@pytest.mark.parametrize("pcp_rank", [0, 1])
@pytest.mark.parametrize(
("pcp_rank", "pcp_size", "dcp_size", "expected_tracked"),
[
(0, 2, 1, True),
(1, 2, 1, False),
(0, 2, 2, True),
(1, 2, 2, True),
(0, 4, 4, True),
(1, 4, 4, True),
(2, 4, 4, True),
(3, 4, 4, True),
],
)
@patch(
"vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
FakeNixlWrapper,
)
def test_pcp_producer_uses_canonical_replica(
self, default_vllm_config, dist_init, pcp_rank
def test_pcp_producer_exposes_dcp_shards_or_canonical_replica(
self,
default_vllm_config,
dist_init,
pcp_rank,
pcp_size,
dcp_size,
expected_tracked,
):
"""Only PCP rank zero publishes, but every rank reports completion."""
"""Replicated PCP is canonicalized; PCP-DCP publishes every shard."""
from vllm.v1.attention.backends.flash_attn import FlashAttentionBackend

vllm_config = create_vllm_config(kv_role="kv_producer")
vllm_config.parallel_config.prefill_context_parallel_size = 2
vllm_config.parallel_config.prefill_context_parallel_size = pcp_size
vllm_config.parallel_config.decode_context_parallel_size = dcp_size
with (
patch(
"vllm.distributed.kv_transfer.kv_connector.v1.nixl."
Expand All @@ -610,29 +629,31 @@ def test_pcp_producer_uses_canonical_replica(
worker = connector.connector_worker
assert worker is not None
assert worker.pcp_rank == pcp_rank
assert worker.pcp_dcp_sharded is (dcp_size > 1)
assert worker.transfer_tp_rank == (pcp_rank if dcp_size > 1 else 0)
assert worker.transfer_tp_size == (pcp_size if dcp_size > 1 else 1)

req_id = "req"
metadata = NixlConnectorMetadata()
metadata.reqs_in_batch.add(req_id)
metadata.reqs_to_send[req_id] = time.perf_counter() + 10
worker.start_load_kv(metadata)
expected_tracked = pcp_rank == 0
assert (req_id in worker._reqs_to_process) == expected_tracked
assert (req_id in worker._reqs_to_send) == expected_tracked

payload = MagicMock(spec=NixlHandshakePayload)
worker.xfer_handshake_metadata = payload
worker.transfer_topo = MagicMock()
worker._get_new_notifs = MagicMock(
side_effect=lambda: {"sent"} if pcp_rank == 0 else set()
side_effect=lambda: {"sent"} if expected_tracked else set()
)

expected_payload = payload if pcp_rank == 0 else None
expected_payload = payload if expected_tracked else None
assert connector.get_handshake_metadata() is expected_payload
done_sending, done_recving = connector.get_finished(set())
assert done_sending == ({"sent"} if pcp_rank == 0 else {req_id})
assert done_sending == ({"sent"} if expected_tracked else {req_id})
assert done_recving == set()
if pcp_rank > 0:
if not expected_tracked:
assert connector.get_finished(set()) == (set(), set())

@patch(
Expand Down Expand Up @@ -3210,18 +3231,38 @@ def test_transfer_mode_changes_compatibility_hash():


@pytest.mark.skip_global_cleanup
def test_scheduler_advertises_transfer_mode():
# Each scheduler advertises its transfer mode in kv_transfer_params so an
# external router can route pull (READ) vs push (WRITE) producers.
@pytest.mark.parametrize("mode", ["pull", "push"])
@pytest.mark.parametrize(
"tp_size,pcp_size,dcp_size,transfer_tp_size",
[(1, 1, 1, 1), (4, 1, 1, 4), (4, 1, 4, 4), (1, 4, 1, 1), (1, 4, 4, 4)],
)
def test_scheduler_advertises_transfer_topology(
mode, tp_size, pcp_size, dcp_size, transfer_tp_size
):
"""The consumer must address every distinct producer KV shard."""
from vllm.distributed.kv_transfer.kv_connector.v1.nixl.pull_scheduler import (
NixlPullConnectorScheduler,
)
from vllm.distributed.kv_transfer.kv_connector.v1.nixl.push_scheduler import (
NixlPushConnectorScheduler,
)

assert NixlPullConnectorScheduler._TRANSFER_MODE == "pull"
assert NixlPushConnectorScheduler._TRANSFER_MODE == "push"
config = create_vllm_config(kv_role="kv_producer")
config.parallel_config.tensor_parallel_size = tp_size
config.parallel_config.prefill_context_parallel_size = pcp_size
config.parallel_config.decode_context_parallel_size = dcp_size
cls = NixlPullConnectorScheduler if mode == "pull" else NixlPushConnectorScheduler
scheduler = cls(config, "prefiller", make_kv_cache_config(block_size=16))
request = create_request(request_id=1, num_tokens=32, do_remote_decode=True)
request.status = RequestStatus.FINISHED_LENGTH_CAPPED
try:
delay, params = scheduler.request_finished(request, ([0, 1],))
assert delay
assert params["transfer_mode"] == mode
assert params["tp_size"] == transfer_tp_size
assert params["dcp_size"] == dcp_size
finally:
scheduler.shutdown()


@pytest.mark.parametrize(
Expand Down
15 changes: 9 additions & 6 deletions tests/v1/kv_connector/unit/test_nixl_push_connector.py
Original file line number Diff line number Diff line change
Expand Up @@ -353,6 +353,7 @@ def fresh(cls) -> _StubWriterWorker:
w.consumer_notification_counts_by_req = defaultdict(int)
w.tp_rank = 0
w.pcp_rank = 0
w.pcp_dcp_sharded = False
w.world_size = 1
w.engine_id = "test-decode-engine"
w._remote_agents = {}
Expand Down Expand Up @@ -540,9 +541,11 @@ def test_start_load_kv_enqueues_to_writer(self):
assert w._push_writer_wake.is_set()
assert w.start_push_calls == []

def test_noncanonical_pcp_rank_skips_producer_work(self):
@pytest.mark.parametrize("sharded", [False, True])
def test_noncanonical_pcp_rank_only_pushes_distinct_shards(self, sharded):
w = _StubWriterWorker.fresh()
w.pcp_rank = 1
w.pcp_dcp_sharded = sharded
w._send_heartbeats = MagicMock()

meta = NixlConnectorMetadata()
Expand All @@ -552,11 +555,11 @@ def test_noncanonical_pcp_rank_skips_producer_work(self):

w.start_load_kv(meta)

assert w._finished_blocks_inbox.empty()
assert "req" not in w._reqs_to_process
assert "req" not in w._reqs_to_send
assert not w._push_writer_wake.is_set()
w._send_heartbeats.assert_not_called()
assert w._finished_blocks_inbox.empty() is (not sharded)
assert ("req" in w._reqs_to_process) is sharded
assert ("req" in w._reqs_to_send) is sharded
assert w._push_writer_wake.is_set() is sharded
assert w._send_heartbeats.called is sharded


# The P→D handshake must run on the base worker's background executor, never
Expand Down
1 change: 1 addition & 0 deletions tests/v1/kv_connector/unit/test_tp_mapping.py
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,7 @@ def test_source_ranks_p_gt_d(self):
@pytest.mark.parametrize(
"tp_rank,tp_size,remote_tp_size,dcp_size,remote_dcp_size,expected_ranks",
[
(0, 1, 8, 1, 8, tuple(range(8))),
(0, 4, 4, 1, 4, (0, 1, 2, 3)),
(2, 4, 4, 4, 4, (2,)),
(0, 2, 4, 2, 4, (0, 2)),
Expand Down
2 changes: 2 additions & 0 deletions tests/v1/kv_connector/unit/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -545,6 +545,7 @@ def make_nixl_scheduler(
sched._reqs_need_save = {}
sched.use_host_buffer = False
sched.engine_id = "test-engine"
sched.transfer_tp_size = 1
sched.side_channel_host = "localhost"
sched.side_channel_port = 5555
sched.blocks_per_sw = []
Expand Down Expand Up @@ -584,6 +585,7 @@ def make_nixl_push_scheduler(
sched.decoder_kv_blocks_ttl = decoder_kv_blocks_ttl
sched.use_host_buffer = False
sched.engine_id = "decode-engine"
sched.transfer_tp_size = 1
sched.side_channel_host = "127.0.0.1"
sched.side_channel_port = 5600
sched.is_bidirectional_kv_xfer_enabled = is_bidirectional_kv_xfer_enabled
Expand Down
12 changes: 6 additions & 6 deletions vllm/config/vllm.py
Original file line number Diff line number Diff line change
Expand Up @@ -1246,14 +1246,14 @@ def __post_init__(self):
self.kv_transfer_config is not None
and self.kv_transfer_config.has_connector("NixlConnector")
):
assert self.parallel_config.prefill_context_parallel_size == 1, (
"NIXL does not support prefill context parallelism."
)
dcp_size = self.parallel_config.decode_context_parallel_size
tp_size = self.parallel_config.tensor_parallel_size
assert dcp_size in (1, tp_size), (
transfer_tp_size = max(
self.parallel_config.tensor_parallel_size,
self.parallel_config.prefill_context_parallel_size,
)
assert dcp_size in (1, transfer_tp_size), (
f"decode_context_parallel_size={dcp_size} must be 1 or equal "
f"to tensor_parallel_size={tp_size} when using NixlConnector."
f"to the NIXL transfer parallel size={transfer_tp_size}."
)
if self.model_config is not None:
assert self.model_config.use_mla or dcp_size == 1, (
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,12 @@ def __init__(
kv_cache_config: "KVCacheConfig",
):
self.vllm_config = vllm_config
parallel_config = vllm_config.parallel_config
# TP1 PCP+DCP exposes its DCP shards as transfer ranks.
self.transfer_tp_size = max(
parallel_config.tensor_parallel_size,
parallel_config.decode_context_parallel_size,
)
self.block_size = vllm_config.cache_config.block_size
self.engine_id: EngineId = engine_id
self.kv_cache_config = kv_cache_config
Expand Down
28 changes: 19 additions & 9 deletions vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -529,9 +529,15 @@ def __init__(
self.dcp_size = vllm_config.parallel_config.decode_context_parallel_size
self.pcp_rank = get_pcp_group().rank_in_group if self.pcp_size > 1 else 0

# DCP support is scoped to MLA, with dcp_size in (1, tp_size): either fully
# replicated or fully sharded. A DCP rank is always derivable this way.
self.dcp_rank = self.tp_rank % self.dcp_size
# TP1 PCP+DCP owns distinct KV shards; replicated PCP uses rank zero.
self.pcp_dcp_sharded = self.pcp_size > 1 and self.dcp_size > 1
self.transfer_tp_size = (
self.pcp_size if self.pcp_dcp_sharded else self.world_size
)
self.transfer_tp_rank = self.pcp_rank if self.pcp_dcp_sharded else self.tp_rank

# MLA is either fully replicated or fully sharded across transfer ranks.
self.dcp_rank = self.transfer_tp_rank % self.dcp_size

self.num_blocks = kv_cache_config.num_blocks
self.enable_permute_local_kv = False
Expand Down Expand Up @@ -762,12 +768,16 @@ def _validate_remote_parallel_config(
local_dcp_size = self.dcp_size
remote_pcp_size = agent_metadata.pcp_size
remote_dcp_size = agent_metadata.dcp_size
if (local_pcp_size > 1 and remote_dcp_size > 1) or (
remote_pcp_size > 1 and local_dcp_size > 1
if remote_pcp_size > 1 and remote_dcp_size not in (1, remote_pcp_size):
raise NotImplementedError(
"Remote NixlConnector PCP+DCP does not span the full PCP group. "
f"Remote PCP/DCP={remote_pcp_size}/{remote_dcp_size}."
)
if (local_pcp_size > 1 and local_dcp_size == 1 and remote_dcp_size > 1) or (
remote_pcp_size > 1 and remote_dcp_size == 1 and local_dcp_size > 1
):
raise NotImplementedError(
"NixlConnector PCP requires decode_context_parallel_size=1 "
"on both instances. "
"Replicated PCP cannot be paired with a DCP-sharded NIXL peer. "
f"Local PCP/DCP={local_pcp_size}/{local_dcp_size}; "
f"remote PCP/DCP={remote_pcp_size}/{remote_dcp_size}."
)
Expand Down Expand Up @@ -1191,8 +1201,8 @@ def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]):
"""Register the KV Cache data in nixl."""

self.transfer_topo = TransferTopology(
tp_rank=self.tp_rank,
tp_size=self.world_size,
tp_rank=self.transfer_tp_rank,
tp_size=self.transfer_tp_size,
block_size=self.block_size,
engine_id=self.engine_id,
is_mla=self.use_mla,
Expand Down
12 changes: 9 additions & 3 deletions vllm/distributed/kv_transfer/kv_connector/v1/nixl/connector.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,10 +107,15 @@ def __init__(
"Consumers and kv_both require "
"prefill_context_parallel_size=1."
)
if pcp_size > 1 and parallel_config.decode_context_parallel_size > 1:
dcp_size = parallel_config.decode_context_parallel_size
if (
pcp_size > 1
and dcp_size > 1
and (parallel_config.tensor_parallel_size != 1 or dcp_size != pcp_size)
):
raise NotImplementedError(
"NixlConnector PCP producers currently require "
"decode_context_parallel_size=1."
"NixlConnector PCP+DCP requires TP1 with DCP spanning "
"the full PCP group."
)
# TODO: Support PCP with bidirectional KV transfer by tracking separate
# send and receive completion counts.
Expand Down Expand Up @@ -311,6 +316,7 @@ def get_handshake_metadata(self) -> KVConnectorHandshakeMetadata | None:
if (
self.kv_transfer_config.kv_role == "kv_producer"
and self.connector_worker.pcp_rank > 0
and not self.connector_worker.pcp_dcp_sharded
):
return None
return self.connector_worker.xfer_handshake_metadata
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -282,7 +282,7 @@ def request_finished(
remote_request_id=request.request_id,
remote_host=self.side_channel_host,
remote_port=self.side_channel_port,
tp_size=self.vllm_config.parallel_config.tensor_parallel_size,
tp_size=self.transfer_tp_size,
dcp_size=self.vllm_config.parallel_config.decode_context_parallel_size,
pp_size=self.vllm_config.parallel_config.pipeline_parallel_size,
remote_num_tokens=remote_num_tokens,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -80,7 +80,7 @@ def start_load_kv(self, metadata: NixlConnectorMetadata):
while not self._ready_requests.empty():
self._read_blocks_for_req(*self._ready_requests.get_nowait())

if self.pcp_rank > 0:
if self.pcp_rank > 0 and not self.pcp_dcp_sharded:
# Replicated-KV PCP: only PCP rank 0 serves the KV, so this rank
# has nothing to send. Report the requests as sent right away so
# the scheduler-side aggregation (world_size workers, and any
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -286,7 +286,8 @@ def request_finished(
remote_request_id=request.request_id,
remote_host=self.side_channel_host,
remote_port=self.side_channel_port,
tp_size=self.vllm_config.parallel_config.tensor_parallel_size,
tp_size=self.transfer_tp_size,
dcp_size=self.vllm_config.parallel_config.decode_context_parallel_size,
pp_size=self.vllm_config.parallel_config.pipeline_parallel_size,
remote_num_tokens=remote_num_tokens,
transfer_mode=self._TRANSFER_MODE,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -163,7 +163,7 @@ def shutdown(self):

def start_load_kv(self, metadata: NixlConnectorMetadata):
"""Pre-process metadata; defer NIXL ops to the writer thread."""
if self.pcp_rank > 0:
if self.pcp_rank > 0 and not self.pcp_dcp_sharded:
return

# D-side: track reqs waiting for P to push.
Expand Down
9 changes: 9 additions & 0 deletions vllm/v1/worker/gpu_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,7 @@
Handle,
checkpoint_prepare_distributed_state,
checkpoint_restore_distributed_state,
get_pcp_group,
get_pp_group,
get_tp_group,
resume_device_comms,
Expand Down Expand Up @@ -703,6 +704,14 @@ def get_kv_connector_handshake_metadata(

pp_rank = get_pp_group().rank_in_group
tp_rank = get_tp_group().rank_in_group
parallel_config = self.vllm_config.parallel_config
if (
parallel_config.prefill_context_parallel_size > 1
and parallel_config.decode_context_parallel_size > 1
):
tp_rank += (
get_pcp_group().rank_in_group * parallel_config.tensor_parallel_size
)
return {(pp_rank, tp_rank): metadata}

def get_kv_cache_spec(self) -> dict[str, KVCacheSpec]:
Expand Down