Skip to content
Draft
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
26 changes: 26 additions & 0 deletions tests/test_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,8 +15,10 @@
import vllm.envs as envs
from vllm.compilation.backends import VllmBackend
from vllm.config import (
CacheConfig,
CompilationConfig,
KernelConfig,
KVTransferConfig,
ModelConfig,
ParallelConfig,
PoolerConfig,
Expand All @@ -39,6 +41,30 @@
DEVICE_TYPE = current_platform.device_type


def test_pd_dcp_interleave_size_is_adjusted_to_block_size(caplog):
config = VllmConfig(
cache_config=CacheConfig(block_size=16),
parallel_config=ParallelConfig(
tensor_parallel_size=2,
decode_context_parallel_size=2,
cp_kv_cache_interleave_size=3,
),
kv_transfer_config=KVTransferConfig(
kv_connector="NixlConnector",
kv_role="kv_both",
),
)

kv_cache_config = SimpleNamespace(
kv_cache_groups=[SimpleNamespace(kv_cache_spec=SimpleNamespace(block_size=16))]
)
with caplog.at_level(logging.WARNING):
config.adjust_dcp_kv_cache_interleave_size(kv_cache_config)

assert config.parallel_config.cp_kv_cache_interleave_size == 16
assert "automatically adjusted from 3 to block_size 16" in caplog.text


def test_compile_config_repr_succeeds():
# setup: VllmBackend mutates the config object
config = VllmConfig()
Expand Down
237 changes: 236 additions & 1 deletion tests/v1/kv_connector/unit/test_nixl_connector_hma.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
"""Unit tests for NixlConnectorScheduler with HMA and Mamba N-1 prefill."""

import gc
from unittest.mock import patch
from unittest.mock import MagicMock, patch

import pytest
import torch
Expand Down Expand Up @@ -63,6 +63,163 @@ def test_sw_sizes(mock_platform, swa_enabled, expected_sw_sizes):
)


@pytest.mark.cpu_test
@pytest.mark.parametrize(
"use_mla,source_ranks,tp_ratio,expected",
[
pytest.param(True, (0,), -2, False, id="pure_mla_with_remote_dcp"),
pytest.param(True, (0, 1), -2, True, id="hybrid_mla_ssm"),
pytest.param(False, (0,), -2, True, id="full_attention"),
pytest.param(False, (0,), 2, False, id="local_tp_greater"),
],
)
def test_needs_split_local_xfer_handles(use_mla, source_ranks, tp_ratio, expected):
"""Handle creation and selection must use the same TP-mapping predicate.

In pure MLA, DCP can target several physical remote workers even though
the TP mapping has one replicated attention source. Those workers own
disjoint blocks, so every read uses the whole local-region handle.
"""
from vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker import (
NixlConnectorWorker,
)

worker = object.__new__(NixlConnectorWorker)
worker.use_mla = use_mla
plan = MagicMock()
plan.all_source_ranks = source_ranks

assert worker._needs_split_local_xfer_handles(tp_ratio, plan) is expected


@pytest.mark.cpu_test
def test_update_state_after_alloc_tracks_cached_blocks_per_group():
"""Hybrid SWA+FA requests can have different prefix-cache-hit counts per
KV cache group. Each group's count must be tracked independently rather
than collapsed to a single scalar, or DCP read-phase alignment
(local_num_computed_blocks) would misalign one of the groups."""
from vllm.distributed.kv_transfer.kv_connector.v1.nixl.pull_scheduler import (
NixlPullConnectorScheduler,
)
from vllm.v1.core.kv_cache_manager import KVCacheBlocks
from vllm.v1.core.kv_cache_utils import KVCacheBlock

scheduler = object.__new__(NixlPullConnectorScheduler)
scheduler._reqs_in_batch = set()
scheduler._reqs_need_save = {}
scheduler._reqs_need_recv = {}
scheduler.use_host_buffer = False
scheduler.is_bidirectional_kv_xfer_enabled = False
scheduler._is_hma_required = False

def cached(block_id):
return KVCacheBlock(block_id=block_id, _block_hash=object())

def uncached(block_id):
return KVCacheBlock(block_id=block_id)

# Group 0 (full-attention): 2 of 3 blocks cached.
# Group 1 (sliding-window): 1 of 3 blocks cached.
blocks = KVCacheBlocks(
blocks=(
[cached(0), cached(1), uncached(2)],
[cached(10), uncached(11), uncached(12)],
)
)

request = create_request(do_remote_prefill=True)

scheduler.update_state_after_alloc(request, blocks, num_external_tokens=2)

_, _, local_num_computed_blocks = scheduler._reqs_need_recv[request.request_id]
assert local_num_computed_blocks == (2, 1)


@pytest.mark.cpu_test
@pytest.mark.parametrize(
"local_ids,remote_ids,remote_rank,local_dcp_size,local_dcp_rank,"
"remote_dcp_size,cached,expected_local,expected_remote",
[
pytest.param(
[100, 101, 102],
[200, 201],
1,
1,
0,
2,
0,
[101],
[200],
id="remote_tail_padding",
),
pytest.param(
[100],
[200],
0,
2,
1,
1,
0,
[],
[],
id="local_tail_padding",
),
pytest.param(
[100],
[200, 201],
2,
2,
0,
4,
3,
[100],
[201],
id="prefix_phase",
),
pytest.param(
[100, 101],
[200, 201, 202, 203],
1,
2,
1,
2,
2,
[100, 101],
[202, 203],
id="equal_dcp_with_prefix",
),
],
)
def test_dcp_read_slice_matches_global_logical_positions(
local_ids,
remote_ids,
remote_rank,
local_dcp_size,
local_dcp_rank,
remote_dcp_size,
cached,
expected_local,
expected_remote,
):
from vllm.distributed.kv_transfer.kv_connector.v1.nixl.pull_worker import (
_dcp_read_slice,
)

actual_local, actual_remote = _dcp_read_slice(
local_ids,
remote_ids,
remote_rank=remote_rank,
local_dcp_size=local_dcp_size,
local_dcp_rank=local_dcp_rank,
remote_dcp_size=remote_dcp_size,
local_num_computed_blocks=cached,
)

assert actual_local == expected_local
assert actual_remote == expected_remote
assert len(actual_local) == len(actual_remote)


@pytest.mark.cpu_test
def test_logical_to_kernel_block_ids_with_hma():
"""Test _logical_to_kernel_block_ids expands blocks when HMA is enabled.
Expand Down Expand Up @@ -212,6 +369,10 @@ def test_read_blocks_for_req_expands_remote_ids(
worker._physical_blocks_per_logical_kv_block = local_physical_per_logical
worker._engine_last_active = {}
worker._bidirectional_kv_xfer_enabled = False
worker.dcp_size = 1
worker.dcp_rank = 0
worker._group_spec_types = resolved_types
worker._has_mamba = any(t is MambaSpec for t in resolved_types)

has_mamba = any(t is MambaSpec for t in resolved_types)
has_swa = any(t is SlidingWindowSpec for t in resolved_types)
Expand All @@ -227,6 +388,7 @@ def test_read_blocks_for_req_expands_remote_ids(
worker.transfer_topo.tp_ratio.return_value = tp_ratio
remote_info = MagicMock()
remote_info.remote_physical_blocks_per_logical = remote_physical_per_logical
remote_info.remote_dcp_size = 1
worker.transfer_topo.get_engine_info.return_value = remote_info
worker.use_mla = False

Expand Down Expand Up @@ -257,6 +419,79 @@ def test_read_blocks_for_req_expands_remote_ids(
)


@pytest.mark.cpu_test
def test_read_blocks_for_req_matches_dcp_before_hma_expansion():
"""DCP positions are logical-block coordinates, so HMA expansion must
happen only after local and remote logical IDs have been matched."""
from vllm.distributed.kv_transfer.kv_connector.v1.nixl.metadata import (
NixlConnectorMetadata,
)
from vllm.distributed.kv_transfer.kv_connector.v1.nixl.tp_mapping import (
TPMapping,
)
from vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker import (
NixlConnectorWorker,
)
from vllm.v1.kv_cache_interface import FullAttentionSpec

worker = object.__new__(NixlConnectorWorker)
worker._physical_blocks_per_logical_kv_block = 2
worker._engine_last_active = {}
worker._bidirectional_kv_xfer_enabled = False
worker.dcp_size = 1
worker.dcp_rank = 0
worker._group_spec_types = (FullAttentionSpec,)
worker._has_mamba = False
worker.use_mla = True
worker.kv_cache_config = make_kv_cache_config(block_size=16)

remote_engine_id = "remote-engine"
remote_info = MagicMock(
remote_physical_blocks_per_logical=2,
remote_dcp_size=2,
remote_tp_size=2,
remote_block_size=16,
)
worker.transfer_topo = MagicMock()
worker.transfer_topo.get_engine_info.return_value = remote_info
worker.transfer_topo.tp_ratio.return_value = 1

plan = MagicMock(spec=TPMapping)
plan.all_source_ranks = (1,)
plan.source_ranks_per_group = ((1,),)
plan.local_consumers = 1
worker.tp_mappings = {remote_engine_id: plan}
worker.src_xfer_handles_by_block_size = {16: 10}
worker.dst_xfer_side_handles = {remote_engine_id: {1: 20}}
worker._read_blocks = MagicMock()

metadata = NixlConnectorMetadata()
metadata.add_new_req_to_recv(
request_id="test-req",
local_block_ids=([10, 11],),
kv_transfer_params={
"remote_block_ids": ([20],),
"remote_engine_id": remote_engine_id,
"remote_request_id": "prefill-test-req",
"remote_host": "localhost",
"remote_port": 1234,
"tp_size": 2,
"dcp_size": 2,
},
local_num_computed_blocks=(0,),
)
meta = metadata.reqs_to_recv["test-req"]
meta.local_physical_block_ids = ([20, 21, 22, 23],)

worker._read_blocks_for_req("test-req", meta)

read_spec = worker._read_blocks.call_args.kwargs["read_spec"]
assert read_spec.local_block_ids == [[22, 23]]
assert read_spec.remote_block_ids == [[40, 41]]
# Keep the complete physical list for receive post-processing.
assert meta.remote.block_ids == [[40, 41]]


@pytest.mark.cpu_test
@pytest.mark.parametrize(
"local_physical_per_logical,remote_physical_per_logical,"
Expand Down
9 changes: 9 additions & 0 deletions vllm/config/kv_transfer.py
Original file line number Diff line number Diff line change
Expand Up @@ -119,3 +119,12 @@ def is_kv_consumer(self) -> bool:

def get_from_extra_config(self, key, default) -> Any:
return self.kv_connector_extra_config.get(key, default)

def has_connector(self, connector_name: str) -> bool:
"""Whether ``connector_name`` is configured, directly or in MultiConnector."""
if self.kv_connector == connector_name:
return True
return self.kv_connector == "MultiConnector" and any(
child.get("kv_connector") == connector_name
for child in self.kv_connector_extra_config.get("connectors", [])
)
Loading