[KVConnector][1/N] PP-aware handshake aggregation and intermediate-PP output plumbing (#43720)

Signed-off-by: zixi-qi <[email protected]>
This commit is contained in:
qizixi
2026-06-04 22:04:19 -07:00
committed by GitHub
parent da1daf40bf
commit 96229fa99e
9 changed files with 459 additions and 26 deletions
@@ -0,0 +1,232 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from types import SimpleNamespace
from typing import Any
import pytest
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
KVConnectorBase_V1,
KVConnectorHandshakeMetadata,
)
from vllm.v1.engine import core as engine_core_module
pytestmark = pytest.mark.cpu_test
class _Metadata(KVConnectorHandshakeMetadata):
pass
class _FakeExecutor:
handshake_metadata_src: (
list[dict[tuple[int, int], KVConnectorHandshakeMetadata] | None] | None
)
last_instance: "_FakeExecutor | None" = None
def __init__(
self,
vllm_config: Any,
) -> None:
del vllm_config
self.handshake_metadata = self.handshake_metadata_src
self.handshake_calls = 0
_FakeExecutor.last_instance = self
def get_kv_connector_handshake_metadata(
self,
) -> list[dict[tuple[int, int], KVConnectorHandshakeMetadata] | None] | None:
self.handshake_calls += 1
return self.handshake_metadata
def init_kv_output_aggregator(self, connector: KVConnectorBase_V1) -> None:
pass
def _run_engine_core_handshake(
monkeypatch: pytest.MonkeyPatch,
connector: KVConnectorBase_V1,
*,
handshake_metadata: (
list[dict[tuple[int, int], KVConnectorHandshakeMetadata] | None] | None
),
) -> _FakeExecutor:
class _FakeScheduler:
def __init__(self, **kwargs: Any) -> None:
self.connector = connector
def get_kv_connector(self) -> KVConnectorBase_V1:
return connector
_FakeExecutor.handshake_metadata_src = handshake_metadata
_FakeExecutor.last_instance = None
monkeypatch.setattr("vllm.plugins.load_general_plugins", lambda: None)
monkeypatch.setattr(
engine_core_module.EngineCore,
"_initialize_kv_caches",
lambda self, vllm_config: SimpleNamespace(kv_cache_groups=[object()]),
)
monkeypatch.setattr(
engine_core_module,
"StructuredOutputManager",
lambda vllm_config: object(),
)
monkeypatch.setattr(
engine_core_module,
"resolve_kv_cache_block_sizes",
lambda kv_cache_config, vllm_config: (16, 16),
)
monkeypatch.setattr(
engine_core_module,
"MULTIMODAL_REGISTRY",
SimpleNamespace(engine_receiver_cache_from_config=lambda vllm_config: None),
)
monkeypatch.setattr(engine_core_module, "freeze_gc_heap", lambda: None)
monkeypatch.setattr(
engine_core_module, "maybe_attach_gc_debug_callback", lambda: None
)
monkeypatch.setattr(engine_core_module, "enable_envs_cache", lambda: None)
monkeypatch.setattr(engine_core_module, "get_hash_fn_by_name", lambda name: None)
monkeypatch.setattr(engine_core_module, "init_none_hash", lambda hash_fn: None)
monkeypatch.setattr(
engine_core_module, "get_request_block_hasher", lambda *args: None
)
vllm_config = SimpleNamespace(
parallel_config=SimpleNamespace(data_parallel_rank_local=0),
scheduler_config=SimpleNamespace(
get_scheduler_cls=lambda: _FakeScheduler,
enable_chunked_prefill=False,
async_scheduling=False,
),
speculative_config=None,
ec_transfer_config=None,
max_concurrent_batches=1,
model_config=SimpleNamespace(runner_type="generate"),
cache_config=SimpleNamespace(
enable_prefix_caching=False,
prefix_caching_hash_algo="builtin",
),
)
engine_core_module.EngineCore(vllm_config, _FakeExecutor, log_stats=False)
assert _FakeExecutor.last_instance is not None
return _FakeExecutor.last_instance
class _LegacyConnector(KVConnectorBase_V1):
def __init__(self) -> None:
self.legacy_metadata: dict[int, KVConnectorHandshakeMetadata] | None = None
def start_load_kv(self, forward_context: Any, **kwargs: Any) -> None:
pass
def wait_for_layer_load(self, layer_name: str) -> None:
pass
def save_kv_layer(
self,
layer_name: str,
kv_layer: Any,
attn_metadata: Any,
**kwargs: Any,
) -> None:
pass
def wait_for_save(self) -> None:
pass
def get_num_new_matched_tokens(
self, request: Any, num_computed_tokens: int
) -> tuple[int | None, bool]:
return 0, False
def update_state_after_alloc(
self, request: Any, blocks: Any, num_external_tokens: int
) -> None:
pass
def build_connector_meta(self, scheduler_output: Any) -> Any:
raise NotImplementedError
def set_xfer_handshake_metadata(
self, metadata: dict[int, KVConnectorHandshakeMetadata]
) -> None:
self.legacy_metadata = metadata
class _PPAwareConnector(_LegacyConnector):
def __init__(self) -> None:
super().__init__()
self.pp_aware_metadata: (
dict[tuple[int, int], KVConnectorHandshakeMetadata] | None
) = None
def set_xfer_handshake_metadata_pp_aware(
self, metadata: dict[tuple[int, int], KVConnectorHandshakeMetadata]
) -> None:
self.pp_aware_metadata = metadata
def test_engine_unwraps_handshake_metadata_for_legacy_connector(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Engine core always asks workers for `(pp_rank, tp_rank)`-keyed metadata,
then unwraps to `{tp_rank: metadata}` for a connector that has not opted
into PP-aware handshake (single-PP producer, all `pp_rank == 0`)."""
metadata_0 = _Metadata()
metadata_1 = _Metadata()
connector = _LegacyConnector()
executor = _run_engine_core_handshake(
monkeypatch,
connector,
handshake_metadata=[
{(0, 0): metadata_0},
None,
{(0, 1): metadata_1},
],
)
assert executor.handshake_calls == 1
assert connector.legacy_metadata == {0: metadata_0, 1: metadata_1}
def test_engine_rejects_pp_producer_for_legacy_connector(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A connector that has not opted into PP-aware handshake must not silently
drop metadata from `pp_rank > 0`; engine core init raises instead."""
connector = _LegacyConnector()
with pytest.raises(ValueError, match="does not support PP-disaggregated"):
_run_engine_core_handshake(
monkeypatch,
connector,
handshake_metadata=[{(0, 0): _Metadata()}, {(1, 0): _Metadata()}],
)
def test_engine_passes_handshake_metadata_through_for_pp_aware_connector(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A PP-aware connector receives the full `(pp_rank, tp_rank)`-keyed dict
unchanged."""
metadata_0 = _Metadata()
metadata_1 = _Metadata()
connector = _PPAwareConnector()
executor = _run_engine_core_handshake(
monkeypatch,
connector,
handshake_metadata=[{(0, 0): metadata_0}, {(1, 0): metadata_1}],
)
assert executor.handshake_calls == 1
assert connector.legacy_metadata is None
assert connector.pp_aware_metadata == {
(0, 0): metadata_0,
(1, 0): metadata_1,
}
@@ -261,11 +261,12 @@ def test_multi_example_connector_consistency():
storage1_scheduler_events = _ignore_event_collection(events["storage1-SCHEDULER"])
storage2_scheduler_events = _ignore_event_collection(events["storage2-SCHEDULER"])
# First event is bind_gpu_block_pool from initialization, then
# set_xfer_handshake_metadata, then on_new_request when the request is enqueued,
# then get_num_new_matched_tokens and update_state_after_alloc from generate().
# set_xfer_handshake_metadata_pp_aware, then on_new_request when the request is
# enqueued, then get_num_new_matched_tokens and update_state_after_alloc from
# generate().
assert storage1_scheduler_events[:6] == [
"bind_gpu_block_pool",
"set_xfer_handshake_metadata",
"set_xfer_handshake_metadata_pp_aware",
"on_new_request",
"get_num_new_matched_tokens 0",
"update_state_after_alloc num_blocks=[0] 0",
@@ -285,7 +286,7 @@ def test_multi_example_connector_consistency():
]
assert storage2_scheduler_events[:6] == [
"bind_gpu_block_pool",
"set_xfer_handshake_metadata",
"set_xfer_handshake_metadata_pp_aware",
"on_new_request",
"get_num_new_matched_tokens 0",
"update_state_after_alloc num_blocks=[0] 0",
@@ -0,0 +1,146 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import pytest
from vllm.distributed.kv_transfer.kv_connector.utils import (
EngineTransferInfo,
TransferTopology,
)
pytestmark = pytest.mark.cpu_test
class _FakeAttentionBackend:
@staticmethod
def get_kv_cache_shape(
num_blocks: int,
block_size: int,
num_kv_heads: int,
head_size: int,
) -> tuple[int, int, int, int, int]:
return (2, num_blocks, num_kv_heads, block_size, head_size)
def _make_topology(
*,
tp_rank: int = 1,
tp_size: int = 4,
total_num_kv_heads: int = 8,
) -> TransferTopology:
return TransferTopology(
tp_rank=tp_rank,
tp_size=tp_size,
block_size=16,
engine_id="local-engine",
is_mla=False,
is_mamba=False,
total_num_kv_heads=total_num_kv_heads,
attn_backends=[_FakeAttentionBackend],
)
def test_legacy_register_remote_engine_uses_pp_rank_zero() -> None:
topology = _make_topology()
info = EngineTransferInfo(
remote_tp_size=2,
remote_block_len=1024,
remote_block_size=16,
remote_physical_blocks_per_logical=1,
)
registered = topology.register_remote_engine("remote-engine", info)
assert registered == info
assert registered.remote_pp_rank == 0
assert topology.get_engine_info("remote-engine") == info
assert topology._engines[("remote-engine", 0)] == info
assert topology.target_remote_ranks("remote-engine") == [0]
def test_register_remote_engine_stores_pp_ranks_separately() -> None:
topology = _make_topology(tp_rank=0, tp_size=2)
info_0 = EngineTransferInfo(
remote_tp_size=2,
remote_block_len=1024,
remote_block_size=16,
remote_physical_blocks_per_logical=1,
remote_pp_rank=0,
start_layer=0,
end_layer=16,
)
info_1 = EngineTransferInfo(
remote_tp_size=1,
remote_block_len=512,
remote_block_size=8,
remote_physical_blocks_per_logical=2,
remote_pp_rank=1,
start_layer=16,
end_layer=32,
)
registered_0 = topology.register_remote_engine("remote-engine", info_0)
registered_1 = topology.register_remote_engine("remote-engine", info_1)
assert registered_0 == info_0
assert registered_1 == info_1
assert topology.get_engine_info("remote-engine") == info_0
assert topology.get_engine_info("remote-engine", 0) == info_0
assert topology.get_engine_info("remote-engine", 1) == info_1
assert set(topology._engines) == {
("remote-engine", 0),
("remote-engine", 1),
}
def test_helpers_use_requested_pp_rank() -> None:
topology = _make_topology(tp_rank=1, tp_size=2, total_num_kv_heads=2)
topology.register_remote_engine(
"remote-engine",
EngineTransferInfo(
remote_tp_size=1,
remote_block_len=1024,
remote_block_size=16,
remote_physical_blocks_per_logical=1,
remote_pp_rank=0,
start_layer=0,
end_layer=8,
),
)
topology.register_remote_engine(
"remote-engine",
EngineTransferInfo(
remote_tp_size=4,
remote_block_len=1024,
remote_block_size=16,
remote_physical_blocks_per_logical=1,
remote_pp_rank=1,
start_layer=8,
end_layer=16,
),
)
assert not topology.is_kv_replicated("remote-engine", 0)
assert topology.is_kv_replicated("remote-engine", 1)
assert topology.replicates_kv_cache("remote-engine", 1)
assert topology.target_remote_ranks("remote-engine", 0) == [0]
assert topology.target_remote_ranks("remote-engine", 1) == [2, 3]
assert "remote_pp=1" in topology.describe("remote-engine", 1)
def test_engine_info_fields_have_backward_compatible_defaults() -> None:
topology = _make_topology()
info = EngineTransferInfo(
remote_tp_size=2,
remote_block_len=1024,
remote_block_size=16,
remote_physical_blocks_per_logical=1,
)
registered = topology.register_remote_engine("remote-engine", info)
assert topology.get_engine_info("remote-engine") == registered
assert registered.remote_pp_rank == 0
assert registered.start_layer == 0
assert registered.end_layer == 0
@@ -368,7 +368,7 @@ def get_current_attn_backend(
class EngineTransferInfo:
"""Common per-remote-engine transfer state, computed at handshake.
Stored per ``engine_id`` inside ``TransferTopology._engines``.
Stored per ``(engine_id, pp_rank)`` inside ``TransferTopology._engines``.
"""
remote_tp_size: int
@@ -382,6 +382,15 @@ class EngineTransferInfo:
remote_physical_blocks_per_logical: int
"""Physical blocks per logical block."""
remote_pp_rank: int = 0
"""Remote producer PP rank for this engine."""
start_layer: int = 0
"""Global index of the first layer owned by this PP rank."""
end_layer: int = 0
"""Exclusive global index after the last layer owned by this PP rank."""
# ---- Transfer topology ----
@@ -403,7 +412,7 @@ class TransferTopology:
def __post_init__(self):
self.local_physical_heads = max(1, self.total_num_kv_heads // self.tp_size)
self._engines: dict[EngineId, EngineTransferInfo] = {}
self._engines: dict[tuple[EngineId, int], EngineTransferInfo] = {}
# Figure out whether the first dimension of the cache is K/V
# or num_blocks.
@@ -461,13 +470,16 @@ class TransferTopology:
f"Cannot register local engine {self.engine_id} as remote. "
f"Local identity is set via __init__ params."
)
if remote_engine_id in self._engines:
return self._engines[remote_engine_id]
self._engines[remote_engine_id] = info
engine_key = (remote_engine_id, info.remote_pp_rank)
if engine_key in self._engines:
return self._engines[engine_key]
self._engines[engine_key] = info
return info
def get_engine_info(self, remote_engine_id: EngineId) -> EngineTransferInfo:
return self._engines[remote_engine_id]
def get_engine_info(
self, remote_engine_id: EngineId, remote_pp_rank: int = 0
) -> EngineTransferInfo:
return self._engines[(remote_engine_id, remote_pp_rank)]
# ============================================================
# Layout properties
@@ -528,15 +540,22 @@ class TransferTopology:
)
return self.block_size // remote_block_size
def is_kv_replicated(self, remote_engine_id: EngineId) -> bool:
def is_kv_replicated(
self, remote_engine_id: EngineId, remote_pp_rank: int = 0
) -> bool:
"""Whether the KV cache is replicated across TP workers due to the
number of TP workers being greater than the number of KV heads.
"""
return self._engines[remote_engine_id].remote_tp_size > self.total_num_kv_heads
return (
self._engines[(remote_engine_id, remote_pp_rank)].remote_tp_size
> self.total_num_kv_heads
)
def replicates_kv_cache(self, remote_engine_id: EngineId) -> bool:
def replicates_kv_cache(
self, remote_engine_id: EngineId, remote_pp_rank: int = 0
) -> bool:
# MLA is always replicated as the hidden dim can't be split.
return self.is_mla or self.is_kv_replicated(remote_engine_id)
return self.is_mla or self.is_kv_replicated(remote_engine_id, remote_pp_rank)
@property
def local_replicates_kv_cache(self) -> bool:
@@ -555,12 +574,14 @@ class TransferTopology:
abs_ratio = -tp_ratio
return [self.tp_rank * abs_ratio + i for i in range(abs_ratio)]
def target_remote_ranks(self, remote_engine_id: EngineId) -> list[int]:
def target_remote_ranks(
self, remote_engine_id: EngineId, remote_pp_rank: int = 0
) -> list[int]:
"""Get the remote TP rank(s) that the current local TP rank will
read from. When remote tp_size > local tp_size, reads from
multiple remote ranks.
"""
info = self._engines[remote_engine_id]
info = self._engines[(remote_engine_id, remote_pp_rank)]
tp_ratio = self.tp_ratio(info.remote_tp_size)
if tp_ratio > 0:
return [self.tp_rank // tp_ratio]
@@ -593,15 +614,16 @@ class TransferTopology:
# Regular case: backends like FA register K/V in separate regions
return cache if self.split_k_and_v else [cache]
def describe(self, remote_engine_id: EngineId) -> str:
def describe(self, remote_engine_id: EngineId, remote_pp_rank: int = 0) -> str:
"""One-line summary of transfer config for logging."""
info = self._engines[remote_engine_id]
info = self._engines[(remote_engine_id, remote_pp_rank)]
return (
f"TransferTopology("
f"tp_ratio={self.tp_ratio(info.remote_tp_size)}, "
f"num_kv_heads={self.total_num_kv_heads if not self.is_mla else 1}, "
f"local_tp={self.tp_size}, "
f"remote_tp={info.remote_tp_size}, "
f"remote_pp={remote_pp_rank}, "
f"local_rank={self.tp_rank}, "
f"remote_block_len={info.remote_block_len})"
)
@@ -644,6 +644,23 @@ class KVConnectorBase_V1(ABC):
"""
return None
def set_xfer_handshake_metadata_pp_aware(
self, metadata: dict[tuple[int, int], KVConnectorHandshakeMetadata]
) -> None:
"""
Set handshake metadata keyed by (pp_rank, tp_rank).
- Default implementation assumes pp_rank is always 0
- PP-aware connectors override this to consume all PP producer shards.
"""
if any(pp_rank != 0 for pp_rank, _ in metadata):
raise ValueError(
f"{type(self).__name__} received pp_rank > 0 handshake metadata "
"but does not support PP-disaggregated KV transfer."
)
self.set_xfer_handshake_metadata(
{tp_rank: meta for (_, tp_rank), meta in metadata.items()}
)
@classmethod
def build_prom_metrics(
cls,
@@ -471,6 +471,12 @@ class MultiConnector(KVConnectorBase_V1, SupportsHMA):
for c in self._connectors:
c.set_xfer_handshake_metadata(metadata)
def set_xfer_handshake_metadata_pp_aware(
self, metadata: dict[tuple[int, int], KVConnectorHandshakeMetadata]
) -> None:
for c in self._connectors:
c.set_xfer_handshake_metadata_pp_aware(metadata)
def _aggregate_request_finished(
self,
request: "Request",
+3 -3
View File
@@ -177,13 +177,13 @@ class EngineCore:
if xfer_handshake_metadata:
# xfer_handshake_metadata is list of dicts from workers
# Each dict already has structure {tp_rank: metadata}
# Each dict already has structure {(pp_rank, tp_rank): metadata}
# Merge all worker dicts into a single dict
content: dict[int, Any] = {}
content: dict[tuple[int, int], Any] = {}
for worker_dict in xfer_handshake_metadata:
if worker_dict is not None:
content.update(worker_dict)
kv_connector.set_xfer_handshake_metadata(content)
kv_connector.set_xfer_handshake_metadata_pp_aware(content)
# Setup batch queue for pipeline parallelism.
# Batch queue for scheduled batches. This enables us to asynchronously
+1 -1
View File
@@ -203,7 +203,7 @@ class Executor(ABC):
def get_kv_connector_handshake_metadata(
self,
) -> list[dict[int, KVConnectorHandshakeMetadata]]:
) -> list[dict[tuple[int, int], KVConnectorHandshakeMetadata]]:
return self.collective_rpc("get_kv_connector_handshake_metadata")
@overload
+12 -3
View File
@@ -34,6 +34,9 @@ from vllm.distributed.kv_transfer import (
get_kv_transfer_group,
has_kv_transfer_group,
)
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
KVConnectorHandshakeMetadata,
)
from vllm.distributed.parallel_state import (
Handle,
get_pp_group,
@@ -513,8 +516,13 @@ class Worker(WorkerBase):
return int(self.available_kv_cache_memory_bytes)
def get_kv_connector_handshake_metadata(self) -> dict | None:
"""Get KV connector metadata from this worker if available."""
def get_kv_connector_handshake_metadata(
self,
) -> dict[tuple[int, int], KVConnectorHandshakeMetadata] | None:
"""Get KV connector metadata from this worker if available.
Returned dict is keyed by `(pp_rank, tp_rank)`.
"""
if not has_kv_transfer_group():
return None
@@ -525,8 +533,9 @@ class Worker(WorkerBase):
if (metadata := connector.get_handshake_metadata()) is None:
return None
pp_rank = get_pp_group().rank_in_group
tp_rank = get_tp_group().rank_in_group
return {tp_rank: metadata}
return {(pp_rank, tp_rank): metadata}
def get_kv_cache_spec(self) -> dict[str, KVCacheSpec]:
return self.model_runner.get_kv_cache_spec()