forked from Karylab-cklius/vllm
[Bugfix][PD] Fix multi-node TP (TP>8) (#39907)
Signed-off-by: NickLucche <nlucches@redhat.com>
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user