[Bugfix][PD] Fix multi-node TP (TP>8) (#39907)

Signed-off-by: NickLucche <nlucches@redhat.com>
This commit is contained in:
Nicolò Lucchesi
2026-05-12 22:20:57 -07:00
committed by GitHub
parent dcacdf9a88
commit 71bcd02ef3
2 changed files with 29 additions and 2 deletions
@@ -1,6 +1,8 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from unittest.mock import MagicMock, patch
from vllm.distributed.kv_transfer.kv_connector.v1.example_connector import ( # noqa: E501
ExampleConnectorMetadata,
)
@@ -38,11 +40,20 @@ def test_kv_connector_mixin_clears_metadata():
vllm_config.kv_transfer_config.kv_role = "kv_both"
vllm_config.kv_transfer_config.kv_connector_extra_config["name"] = "unit"
# Initialize the global connector instance
kv_cache_config = KVCacheConfig(
num_blocks=0, kv_cache_tensors=[], kv_cache_groups=[]
)
ensure_kv_transfer_initialized(vllm_config, kv_cache_config)
# Initialize the global connector instance.
# kv_transfer init now syncs engine_id across TP, so unit tests need
# a minimal mocked TP group.
mock_tp_group = MagicMock()
mock_tp_group.broadcast_object.side_effect = lambda value, src=0: value
with patch(
"vllm.distributed.parallel_state.get_tp_group",
return_value=mock_tp_group,
):
ensure_kv_transfer_initialized(vllm_config, kv_cache_config)
try:
# Minimal scheduler output with empty metadata; mixin should still
@@ -48,6 +48,20 @@ def is_v1_kv_transfer_group(connector: KVConnectorBaseType | None = None) -> boo
return isinstance(connector, KVConnectorBase_V1)
def _sync_engine_id_across_tp(vllm_config: "VllmConfig") -> None:
"""Broadcast engine_id from TP rank 0 so all workers in a
multi-node TP group share the same value."""
from vllm.distributed.parallel_state import (
get_tp_group,
)
assert vllm_config.kv_transfer_config is not None
synced_id = get_tp_group().broadcast_object(
vllm_config.kv_transfer_config.engine_id, src=0
)
vllm_config.kv_transfer_config.engine_id = synced_id
def ensure_kv_transfer_initialized(
vllm_config: "VllmConfig", kv_cache_config: "KVCacheConfig"
) -> None:
@@ -64,6 +78,8 @@ def ensure_kv_transfer_initialized(
vllm_config.kv_transfer_config.is_kv_transfer_instance
and _KV_CONNECTOR_AGENT is None
):
_sync_engine_id_across_tp(vllm_config)
_KV_CONNECTOR_AGENT = KVConnectorFactory.create_connector(
config=vllm_config,
role=KVConnectorRole.WORKER,