diff --git a/tests/test_config.py b/tests/test_config.py index 8171950cdd07..5024d47015e3 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -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, @@ -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() diff --git a/tests/v1/kv_connector/unit/test_nixl_connector_hma.py b/tests/v1/kv_connector/unit/test_nixl_connector_hma.py index 0124f6516ef9..8cbd5deccfa7 100644 --- a/tests/v1/kv_connector/unit/test_nixl_connector_hma.py +++ b/tests/v1/kv_connector/unit/test_nixl_connector_hma.py @@ -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 @@ -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. @@ -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) @@ -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 @@ -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," diff --git a/vllm/config/kv_transfer.py b/vllm/config/kv_transfer.py index 9f206ff5d4d9..d1344a5211cd 100644 --- a/vllm/config/kv_transfer.py +++ b/vllm/config/kv_transfer.py @@ -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", []) + ) diff --git a/vllm/config/vllm.py b/vllm/config/vllm.py index 82ca8ea26a38..39b9b20d4f5a 100644 --- a/vllm/config/vllm.py +++ b/vllm/config/vllm.py @@ -983,6 +983,30 @@ def __post_init__(self): "connectors (PD disaggregation, KV cache offload)." ) + # DCP+PD invariants:a side is either fully replicated or fully sharded; MLA only + if ( + 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), ( + f"decode_context_parallel_size={dcp_size} must be 1 or equal " + f"to tensor_parallel_size={tp_size} when using NixlConnector." + ) + if self.model_config is not None: + assert self.model_config.use_mla or dcp_size == 1, ( + "PD with decode_context_parallel_size > 1 is only " + "supported for MLA models." + ) + assert not (self.model_config.is_hybrid and dcp_size > 1), ( + "PD with decode_context_parallel_size > 1 is not " + "supported for hybrid Mamba/SSM models." + ) + if self.lora_config is not None: self.lora_config.verify_with_model_config(self.model_config) @@ -2244,6 +2268,52 @@ def _validate_v2_model_runner(self) -> None: "request parameter. Set VLLM_USE_V2_MODEL_RUNNER=0 if this is required." ) + def adjust_dcp_kv_cache_interleave_size( + self, kv_cache_config: "KVCacheConfig" + ) -> None: + """Normalize DCP interleave size against the resolved block_size for PD. + + Called by each worker (via ensure_kv_transfer_initialized), once it knows its + own final block_size via kv_cache_config. + """ + dcp_size = self.parallel_config.decode_context_parallel_size + if dcp_size <= 1: + return + # Get the kernel block_size, but don't use resolve_kv_cache_block_size to avoid + # scaling by dcp_size (we need the local block_size here). + local_block_size = min( + g.kv_cache_spec.block_size for g in kv_cache_config.kv_cache_groups + ) + if self.parallel_config.dcp_kv_cache_interleave_size > 1 and ( + self.parallel_config.cp_kv_cache_interleave_size + != self.parallel_config.dcp_kv_cache_interleave_size + ): + self.parallel_config.cp_kv_cache_interleave_size = ( + self.parallel_config.dcp_kv_cache_interleave_size + ) + logger.warning_once( + "cp_kv_cache_interleave_size is overridden by dcp_kv_cache" + "_interleave_size. And dcp-kv-cache-interleave-size will be " + "deprecated when PCP is fully supported." + ) + + if ( + self.kv_transfer_config is not None + and self.kv_transfer_config.kv_connector is not None + and self.parallel_config.cp_kv_cache_interleave_size != local_block_size + ): + interleave = self.parallel_config.cp_kv_cache_interleave_size + self.parallel_config.cp_kv_cache_interleave_size = local_block_size + logger.info_once( + "When using PD disaggregation with DCP " + "(decode_context_parallel_size=%d), " + "cp_kv_cache_interleave_size is automatically adjusted " + "from %d to block_size %d for block-level alignment.", + dcp_size, + interleave, + local_block_size, + ) + def validate_block_size(self) -> None: """Validate block_size against DCP and mamba constraints. @@ -2252,20 +2322,14 @@ def validate_block_size(self) -> None: """ block_size = self.cache_config.block_size - # DCP interleave-size compatibility - if self.parallel_config.decode_context_parallel_size > 1: - if self.parallel_config.dcp_kv_cache_interleave_size > 1 and ( - self.parallel_config.cp_kv_cache_interleave_size - != self.parallel_config.dcp_kv_cache_interleave_size - ): - self.parallel_config.cp_kv_cache_interleave_size = ( - self.parallel_config.dcp_kv_cache_interleave_size - ) - logger.warning_once( - "cp_kv_cache_interleave_size is overridden by dcp_kv_cache" - "_interleave_size. And dcp-kv-cache-interleave-size will be " - "deprecated when PCP is fully supported." - ) + # Skip DCP interleave-size compatibility when a KV connector is configured: + # cp_kv_cache_interleave_size is pinned to block_size for PD by each worker + pd_active = ( + self.kv_transfer_config is not None + and self.kv_transfer_config.kv_connector is not None + and self.kv_transfer_config.is_kv_transfer_instance + ) + if self.parallel_config.decode_context_parallel_size > 1 and not pd_active: assert ( self.parallel_config.cp_kv_cache_interleave_size <= block_size and block_size % self.parallel_config.cp_kv_cache_interleave_size == 0 @@ -2274,7 +2338,6 @@ def validate_block_size(self) -> None: "than or equal to and divisible by cp_kv_cache_interleave_size " f"({self.parallel_config.cp_kv_cache_interleave_size})." ) - # Mamba cache align-mode constraints if self.cache_config.mamba_cache_mode == "align": assert block_size <= self.scheduler_config.max_num_batched_tokens, ( diff --git a/vllm/distributed/kv_transfer/kv_connector/utils.py b/vllm/distributed/kv_transfer/kv_connector/utils.py index eff9bec8ee93..0e33b144ea75 100644 --- a/vllm/distributed/kv_transfer/kv_connector/utils.py +++ b/vllm/distributed/kv_transfer/kv_connector/utils.py @@ -399,6 +399,9 @@ class EngineTransferInfo: end_layer: int = 0 """Exclusive global index after the last layer owned by this PP rank.""" + remote_dcp_size: int = 1 + """Remote decode context parallel size.""" + # ---- Transfer topology ---- @@ -415,6 +418,7 @@ class TransferTopology: is_mamba: bool total_num_kv_heads: int attn_backends: list[type[AttentionBackend]] + dcp_size: int = 1 tensor_shape: torch.Size | None = None def __post_init__(self): @@ -516,6 +520,15 @@ def virtually_split_kv_in_blocks(self) -> bool: # interleaving means a simple half-split does not separate the parts). return self.is_mamba and not self._cross_layers_blocks + @property + def dcp_rank(self) -> int: + """This rank's DCP coverage rank. + + with ``dcp_size in (1, tp_size)`` enforced at the connector boundary, a + rank's DCP identity is always exactly ``tp_rank % dcp_size``. + """ + return self.tp_rank % self.dcp_size + # ============================================================ # Common methods # ============================================================ @@ -569,18 +582,70 @@ def local_replicates_kv_cache(self) -> bool: """Whether the local engine's KV cache is replicated.""" return self.is_mla or self.tp_size > self.total_num_kv_heads - def handshake_target_ranks(self, remote_tp_size: int) -> list[int]: + def dcp_source_ranks(self, remote_tp_size: int, remote_dcp_size: int) -> list[int]: + """Remote ranks whose DCP slice overlaps mine (MLA, ``remote_dcp_size > 1``). + + Shared by ``handshake_target_ranks`` (who to query metadata from) and + ``compute_tp_mapping`` (who to actually read from) — for MLA the two + questions have the identical answer, since DCP sharding is the only + thing keeping a remote rank from being interchangeable with any other. + """ + local_dcp_size = self.dcp_size + local_dcp_rank = self.dcp_rank + if local_dcp_size <= remote_dcp_size: + # Keep every remote rank whose slice sits inside mine. When + # local_dcp_size == 1 (replicated locally), local_dcp_rank == 0 reduces to + # every remote rank, since no single one holds the whole sequence. + return [ + r for r in range(remote_tp_size) if r % local_dcp_size == local_dcp_rank + ] + # Local finer-grained: exactly one remote rank covers my whole slice + return [local_dcp_rank % remote_dcp_size] + + def handshake_target_ranks( + self, remote_tp_size: int, remote_dcp_size: int = 1 + ) -> list[int]: """Pre-registration: compute which remote TP ranks to handshake with. - Pure math based on local/remote TP sizes — does not require - the remote engine to be registered yet. + Pure math based on local/remote TP (and DCP, when the remote shards + its KV cache) sizes — does not require the remote engine to be + registered yet. + + DCP support is scoped to ``dcp_size in (1, tp_size)`` on each side + and DCP sizes that divide one another: neither side ever has a + partially-duplicated, partially-sharded KV cache. When the + remote is not sharded (``remote_dcp_size == 1``) this reduces + exactly to the DTP>=PTP case, since a sharded local side + already has ``tp_size == dcp_size``. """ + if remote_dcp_size > 1: + return self.dcp_source_ranks(remote_tp_size, remote_dcp_size) + tp_ratio = self.tp_ratio(remote_tp_size) if tp_ratio > 0: return [self.tp_rank // tp_ratio] abs_ratio = -tp_ratio return [self.tp_rank * abs_ratio + i for i in range(abs_ratio)] + def dcp_consumer_count(self, remote_tp_size: int, remote_dcp_size: int) -> int: + """How many local ranks (in aggregate) read from a given remote rank. + + Used by the producer side to know how many reader notifications to + wait for before freeing a request's blocks. Reuses ``tp_ratio`` + whenever the remote isn't sharded — a sharded local side already + has ``tp_size == dcp_size``, so the existing TP-ratio formula is + already correct there unmodified. + """ + if remote_dcp_size > 1: + if self.dcp_size == 1: + # Replicated locally: every local rank reads every shard. + return self.tp_size + # Both sharded, different degrees. + return max(1, self.dcp_size // remote_dcp_size) + # Remote replicated: `tp_ratio` local ranks share each remote rank + # when local_tp >= remote_tp, else each remote rank has one reader. + return max(1, self.tp_ratio(remote_tp_size)) + def target_remote_ranks( self, remote_engine_id: EngineId, remote_pp_rank: int = 0 ) -> list[int]: @@ -606,6 +671,8 @@ def describe(self, remote_engine_id: EngineId, remote_pp_rank: int = 0) -> str: f"local_tp={self.tp_size}, " f"remote_tp={info.remote_tp_size}, " f"remote_pp={remote_pp_rank}, " + f"local_dcp={self.dcp_size}, " + f"remote_dcp={info.remote_dcp_size}, " f"local_rank={self.tp_rank}, " f"remote_block_len={info.remote_block_len})" ) diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_scheduler.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_scheduler.py index 7a3629c7229f..5cc435403ca9 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_scheduler.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_scheduler.py @@ -104,7 +104,9 @@ def __init__( # Requests that need to start recv/send. # New requests are added by update_state_after_alloc in # the scheduler. Used to make metadata passed to Worker. - self._reqs_need_recv: dict[ReqId, tuple[Request, BlockIds]] = {} + self._reqs_need_recv: dict[ + ReqId, tuple[Request, BlockIds, tuple[int, ...]] + ] = {} self._reqs_need_save: dict[ReqId, Request] = {} # Reqs to send and their expiration time self._reqs_need_send: dict[ReqId, float] = {} @@ -185,6 +187,7 @@ def on_new_request(self, request: "Request") -> None: host = params.get("remote_host") port = params.get("remote_port") tp_size = params.get("tp_size") + dcp_size = params.get("dcp_size", 1) pp_size = params.get("pp_size", 1) if ( remote_engine_id is None @@ -200,6 +203,7 @@ def on_new_request(self, request: "Request") -> None: host=host, port=port, tp_size=tp_size, + dcp_size=dcp_size, pp_size=pp_size, ) self._heartbeat_by_engine[remote_engine_id].req_ids.add(remote_request_id) @@ -406,12 +410,13 @@ def build_connector_meta( meta = NixlConnectorMetadata() # Loop through scheduled reqs and convert to ReqMeta. - for req_id, (req, block_ids) in self._reqs_need_recv.items(): + for req_id, (req, block_ids, cached) in self._reqs_need_recv.items(): assert req.kv_transfer_params is not None meta.add_new_req_to_recv( request_id=req_id, local_block_ids=block_ids, kv_transfer_params=req.kv_transfer_params, + local_num_computed_blocks=cached, ) if self.use_host_buffer: diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py index 12b9f0d937fa..f3f7e475241b 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py @@ -217,6 +217,17 @@ def _build_local_splits_from_plan( handle.append((addr + p_idx * chunk, chunk, dev)) yield handle + def _needs_split_local_xfer_handles(self, tp_ratio: int, plan: TPMapping) -> bool: + """Whether reads need per-source slices of the local KV region. + + Pure MLA attention is replicated across TP ranks and writes the whole + local region. Multiple physical remote workers may still participate + because DCP assigns them disjoint blocks, but that does not require + splitting the local region. Hybrid MLA+SSM is different: its mapping + contains multiple source ranks for the sharded SSM state. + """ + return tp_ratio < 0 and (not self.use_mla or len(plan.all_source_ranks) > 1) + def _fa_desc_replicated(self, num_fa_descs: int) -> list[bool]: """Per-FA-descriptor replicate flag, in _build_fa_local emission order (region-major; one desc per block, with K/V packed). Length ``num_fa_descs``. @@ -364,6 +375,11 @@ def __init__( self.tp_rank = get_tensor_model_parallel_rank() self.world_size = get_tensor_model_parallel_world_size() + # 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_size = vllm_config.parallel_config.decode_context_parallel_size + self.dcp_rank = self.tp_rank % self.dcp_size + self.num_blocks = kv_cache_config.num_blocks self.enable_permute_local_kv = False self.enable_heterogeneous_attn_post_process = False @@ -525,9 +541,10 @@ def __init__( self.compat_hash: str | None = None self.transfer_topo: TransferTopology | None = None - # With heterogeneous TP, P must wait for all assigned D TP workers to - # finish reading before safely freeing the blocks. + # With heterogeneous TP (or DCP), P must wait for all assigned D + # workers to finish reading before safely freeing the blocks. self.consumer_notification_counts_by_req = defaultdict[ReqId, int](int) + self.expected_consumer_notifications_by_req: dict[ReqId, int] = {} self.xfer_stats = NixlKVConnectorStats() self._physical_blocks_per_logical_kv_block = 1 @@ -538,7 +555,6 @@ def __init__( get_representative_spec_type(g.kv_cache_spec) for g in self.kv_cache_config.kv_cache_groups ) - # Per-region MLA flag, 1:1 with block_len_per_layer. True -> REPLICATE # (MLA), False -> SPLIT (head-sharded full-attn). Mixed only for models # combining both (e.g. GQA main + MLA Eagle-3 draft). @@ -580,6 +596,7 @@ def _nixl_handshake( port: int, remote_tp_size: int, expected_engine_id: str, + remote_dcp_size: int = 1, remote_pp_size: int = 1, notif_agents_only: bool = False, ) -> tuple[dict[tuple[int, int], str], float]: @@ -603,7 +620,9 @@ def _nixl_handshake( # local rank will read from. Note that With homogeneous TP, # this happens to be the same single rank_i. assert self.transfer_topo is not None - p_remote_ranks = self.transfer_topo.handshake_target_ranks(remote_tp_size) + p_remote_ranks = self.transfer_topo.handshake_target_ranks( + remote_tp_size, remote_dcp_size + ) remote_rank_to_agent_name: dict[tuple[int, int], str] = {} path = make_zmq_path("tcp", host, port) # Clock offset to the peer, estimated from the handshake round-trip. @@ -705,11 +724,11 @@ def _nixl_handshake( # Register Remote agent. if notif_agents_only: remote_agent_name = self._add_notif_only_remote_agent( - metadata, remote_tp_size + metadata, remote_tp_size, metadata.dcp_size ) else: remote_agent_name = self.add_remote_agent( - metadata, remote_rank, remote_tp_size + metadata, remote_rank, remote_tp_size, metadata.dcp_size ) setup_agent_time = time.perf_counter() logger.debug( @@ -724,7 +743,7 @@ def _nixl_handshake( return remote_rank_to_agent_name, best_offset def _add_notif_only_remote_agent( - self, metadata: NixlAgentMetadata, remote_tp_size: int + self, metadata: NixlAgentMetadata, remote_tp_size: int, remote_dcp_size: int = 1 ) -> str: """Load a remote agent for notifs only on the push-mode decode side. @@ -740,6 +759,7 @@ def _add_notif_only_remote_agent( remote_physical_blocks_per_logical=( metadata.physical_blocks_per_logical_kv_block ), + remote_dcp_size=remote_dcp_size, ), ) return self.nixl_wrapper.add_remote_agent(metadata.agent_metadata) @@ -856,6 +876,7 @@ def _ensure_handshake( host: str, port: int, tp_size: int, + dcp_size: int = 1, pp_size: int = 1, notif_agents_only: bool = False, ) -> Future[tuple[dict[tuple[int, int], str], float]] | None: @@ -881,6 +902,7 @@ def _ensure_handshake( port, tp_size, engine_id, + dcp_size, pp_size, notif_agents_only, ) @@ -918,6 +940,8 @@ def _background_nixl_handshake( meta.remote.host, meta.remote.port, meta.tp_size, + meta.dcp_size, + meta.pp_size, ) if fut is None: # Already handshaked — only happens if caller does not pre-check. @@ -967,8 +991,11 @@ def _register_packed_kv_cache( block_size=self.block_size, engine_id=self.engine_id, is_mla=self.use_mla, - total_num_kv_heads=self.model_config.get_total_num_kv_heads(), + total_num_kv_heads=1 + if self.use_mla + else self.model_config.get_total_num_kv_heads(), attn_backends=self.attn_backends, + dcp_size=self.dcp_size, tensor_shape=None, is_mamba=self._has_mamba, ) @@ -1026,6 +1053,7 @@ def _register_packed_kv_cache( physical_blocks_per_logical_kv_block=( self._physical_blocks_per_logical_kv_block ), + dcp_size=self.dcp_size, ) assert self.compat_hash is not None encoder = msgspec.msgpack.Encoder() @@ -1057,8 +1085,11 @@ def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): block_size=self.block_size, engine_id=self.engine_id, is_mla=self.use_mla, - total_num_kv_heads=self.model_config.get_total_num_kv_heads(), + total_num_kv_heads=1 + if self.use_mla + else self.model_config.get_total_num_kv_heads(), attn_backends=self.attn_backends, + dcp_size=self.dcp_size, # SSM States come in tuples (ssm, conv) tensor_shape=next(iter(kv_caches.values())).shape if not self._has_mamba @@ -1267,6 +1298,7 @@ def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): physical_blocks_per_logical_kv_block=( self._physical_blocks_per_logical_kv_block ), + dcp_size=self.dcp_size, ) # Wrap metadata in payload with hash for defensive decoding assert self.compat_hash is not None @@ -1492,6 +1524,7 @@ def add_remote_agent( nixl_agent_meta: NixlAgentMetadata, remote_tp_rank: int = 0, remote_tp_size: int = 1, + remote_dcp_size: int = 1, ) -> str: """ Add the remote NIXL agent and prepare the descriptors for reading cache @@ -1575,6 +1608,7 @@ def add_remote_agent( remote_block_size=nixl_agent_meta.block_size, remote_block_len=nixl_agent_meta.block_lens[0], remote_physical_blocks_per_logical=physical_blocks_per_logical, + remote_dcp_size=remote_dcp_size, ) transfer_topo.register_remote_engine(engine_id, transfer_info) logger.info("Transfer plan: %s", transfer_topo.describe(engine_id)) @@ -1583,6 +1617,7 @@ def add_remote_agent( transfer_topology=transfer_topo, remote_tp_size=remote_tp_size, group_spec_types=self._group_spec_types, + remote_dcp_size=remote_dcp_size, ) remote_agent_name = self.nixl_wrapper.add_remote_agent( @@ -1605,7 +1640,9 @@ def add_remote_agent( self.kv_caches_base_addr[engine_id][remote_tp_rank] = ( nixl_agent_meta.kv_caches_base_addr ) - self._validate_remote_agent_handshake(nixl_agent_meta, remote_tp_size) + self._validate_remote_agent_handshake( + nixl_agent_meta, remote_tp_size, remote_dcp_size + ) # This is 1 when P and D `--tensor-parallel-size` match. Otherwise, # this is the ratio between the two sizes. @@ -1635,10 +1672,8 @@ def add_remote_agent( ### (Optional) Register local agent memory regions. MLA is not split. split_key = (tp_ratio, remote_block_size) - if ( - tp_ratio < 0 - and (not self.use_mla or len(plan.all_source_ranks) > 1) - and split_key not in self.src_xfer_handles_by_tp_ratio + if self._needs_split_local_xfer_handles(tp_ratio, plan) and ( + split_key not in self.src_xfer_handles_by_tp_ratio ): # Remote tp_size > local tp_size: read from multiple remote ranks. # Logically "split" own regions into per-source chunks. Hybrid @@ -1696,7 +1731,10 @@ def add_remote_agent( return remote_agent_name def _validate_remote_agent_handshake( - self, nixl_agent_meta: NixlAgentMetadata, remote_tp_size: int + self, + nixl_agent_meta: NixlAgentMetadata, + remote_tp_size: int, + remote_dcp_size: int = 1, ): """ Validate the remote agent handshake metadata ensuring the @@ -1707,6 +1745,15 @@ def _validate_remote_agent_handshake( assert self.transfer_topo is not None remote_info = self.transfer_topo.get_engine_info(remote_engine_id) assert remote_info.remote_tp_size == remote_tp_size + assert remote_info.remote_dcp_size == remote_dcp_size + # DCP sizes must divide one another; this is what keeps the + # read-slicing math in pull_worker a closed form. + assert ( + self.dcp_size % remote_dcp_size == 0 or remote_dcp_size % self.dcp_size == 0 + ), ( + f"DCP sizes must divide one another: local={self.dcp_size}, " + f"remote={remote_dcp_size} (engine {remote_engine_id})." + ) tp_ratio = self.transfer_topo.tp_ratio(remote_tp_size) block_size_ratio = self.transfer_topo.block_size_ratio( @@ -2150,6 +2197,7 @@ def get_finished(self) -> tuple[set[str], set[str]]: if now < expires: break count = self.consumer_notification_counts_by_req.pop(req_id, 0) + self.expected_consumer_notifications_by_req.pop(req_id, None) self.xfer_stats.record_kv_expired_req() logger.warning( "Releasing expired KV blocks for request %s which were " @@ -2285,6 +2333,7 @@ def _send_heartbeats(self, metadata: NixlConnectorMetadata) -> None: hb_info.host, hb_info.port, hb_info.tp_size, + hb_info.dcp_size, hb_info.pp_size, self._hb_handshake_notif_only and hb_info.pp_size > 1, ) @@ -2370,7 +2419,7 @@ def _logical_to_kernel_block_ids(self, block_ids: BlockIds, ratio: int) -> Block This is required when the logical block size (the one set by the user) does not match the one required by the attn backend. `ratio` is the number of physical blocks per logical block. - We always receive logical blocks from the engine, so we expand them here eg: + We always receive logical blocks from the engine, so we expand them here eg: logical block ids: [(SW-clipped) [1], (FA) [2, 3]], ratio=2 physical block ids: [(SW-clipped) [2, 3], (FA) [4, 5, 6, 7]] """ diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/connector.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/connector.py index 7d9a34406b7d..67fee2f3d4ba 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/connector.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/connector.py @@ -357,6 +357,10 @@ def __init__( kv_cache_config: "KVCacheConfig", ): super().__init__(vllm_config, role, kv_cache_config) + if vllm_config.parallel_config.decode_context_parallel_size > 1: + raise ValueError( + "NixlPushConnector does not support decode_context_parallel_size > 1." + ) self.connector_scheduler: NixlPushConnectorScheduler | None = None self.connector_worker: NixlPushConnectorWorker | None = None if role == KVConnectorRole.SCHEDULER: diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/metadata.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/metadata.py index f6f3de0a1f29..aa0d270c6868 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/metadata.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/metadata.py @@ -41,8 +41,9 @@ # 4: Add KV block lease renewal through heartbeats # 5: Add remote_blocks_expiry_time to kv_transfer_params + handshake # clock-sync timestamp +# 6: Add dcp_size to NixlAgentMetadata/kv_transfer_params for PD+DCP (MLA) # -NIXL_CONNECTOR_VERSION: int = 5 +NIXL_CONNECTOR_VERSION: int = 6 @dataclass @@ -58,6 +59,7 @@ class NixlAgentMetadata: ssm_sizes: tuple[int, int] attn_backend_name: str physical_blocks_per_logical_kv_block: int + dcp_size: int = 1 @dataclass @@ -149,6 +151,7 @@ class HeartbeatInfo: host: str port: int tp_size: int + dcp_size: int = 1 pp_size: int = 1 @@ -168,6 +171,12 @@ class ReqMeta: # To be used when logical block size does not match the kernel block size local_physical_block_ids: BlockIds tp_size: int + dcp_size: int = 1 + # Per-KV-cache-group logical blocks this rank already holds, i.e. its + # prefix-cache hit. Fixes where this rank's DCP slice starts relative to + # the remote's; kept per-group since hybrid models (e.g. SWA+FA) can have + # different cache-hit counts per group. + local_num_computed_blocks: tuple[int, ...] = () remote: RemoteMeta | None = None # Remote block size, discovered during NIXL handshake (push mode). remote_block_size: int | None = None @@ -195,14 +204,17 @@ def _add_new_req( self, local_block_ids: BlockIds, kv_transfer_params: dict[str, Any], + local_num_computed_blocks: tuple[int, ...] = (), ) -> ReqMeta: return ReqMeta( local_block_ids=local_block_ids, local_physical_block_ids=local_block_ids, - # P workers don't need to receive tp_size from proxy here. + # P workers don't need to receive these from proxy here. tp_size=kv_transfer_params.get("tp_size", 1), + dcp_size=kv_transfer_params.get("dcp_size", 1), remote_block_size=kv_transfer_params.get("remote_block_size"), pp_size=kv_transfer_params.get("pp_size", 1), + local_num_computed_blocks=local_num_computed_blocks, ) def add_new_req_to_save( @@ -220,8 +232,11 @@ def add_new_req_to_recv( request_id: ReqId, local_block_ids: BlockIds, kv_transfer_params: dict[str, Any], + local_num_computed_blocks: tuple[int, ...] = (), ): - req = self._add_new_req(local_block_ids, kv_transfer_params) + req = self._add_new_req( + local_block_ids, kv_transfer_params, local_num_computed_blocks + ) req.remote = RemoteMeta( block_ids=kv_transfer_params["remote_block_ids"], engine_id=kv_transfer_params["remote_engine_id"], diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/pull_scheduler.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/pull_scheduler.py index 51a54d06f078..5043b6fd2b20 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/pull_scheduler.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/pull_scheduler.py @@ -158,12 +158,23 @@ def update_state_after_alloc( local_block_ids = self.get_sw_clipped_blocks( unhashed_local_block_ids ) + # Blocks covered by the local prefix cache, per KV cache group. + # Each count fixes where that group's DCP slice starts, which the + # worker needs to line up with the remote's slice. + local_num_computed_blocks = tuple( + sum( + block.block_hash is not None and not block.is_null + for block in group + ) + for group in blocks.blocks + ) # Get unhashed blocks to pull from remote. Mind that a full prefix # cache hit is indicated with an empty list. self._reqs_need_recv[request.request_id] = ( request, local_block_ids, + local_num_computed_blocks, ) else: @@ -216,7 +227,7 @@ def request_finished( # To avoid stranding the prefill blocks in the prefill instance, # we must add empty block_ids to _reqs_need_recv so that our # worker side will notify and free blocks in the prefill instance. - self._reqs_need_recv[request.request_id] = (request, []) + self._reqs_need_recv[request.request_id] = (request, [], ()) params["do_remote_prefill"] = False return False, None @@ -275,6 +286,8 @@ def request_finished( remote_host=self.side_channel_host, remote_port=self.side_channel_port, tp_size=self.vllm_config.parallel_config.tensor_parallel_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, remote_blocks_expiry_time=blocks_expiry_time, ) diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/pull_worker.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/pull_worker.py index 0c66596c4b7b..f50c7b69d688 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/pull_worker.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/pull_worker.py @@ -5,6 +5,7 @@ import time from typing import TYPE_CHECKING +from vllm.distributed.kv_transfer.kv_connector.utils import BlockIds from vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker import ( NixlBaseConnectorWorker, ) @@ -14,6 +15,7 @@ ) from vllm.distributed.kv_transfer.kv_connector.v1.nixl.tp_mapping import ( ReadSpec, + _is_ssm_spec, ) from vllm.logger import init_logger @@ -28,6 +30,68 @@ _KV_BLOCKS_EXPIRY_SAFETY_MARGIN = 5.0 +def _dcp_read_slice( + local_ids: list[int], + remote_ids: list[int], + remote_rank: int, + local_dcp_size: int, + local_dcp_rank: int, + remote_dcp_size: int, + local_num_computed_blocks: int, +) -> tuple[list[int], list[int]]: + """Slice a source rank's share of a DCP-sharded sequence. + + Scoped to MLA PD with ``dcp_size in (1, tp_size)``: a side is + either fully replicated (holds the whole sequence) or fully sharded + (holds exactly one out of every ``dcp_size`` blocks), and the ratio + between the two sizes is always a whole number. + + The closed-form slices below align the phase of the local uncached suffix + with each remote rank. They must operate on logical IDs, before HMA + expands either side into kernel physical IDs. The final truncation removes + non-overlapping allocation padding from an incomplete last DCP stripe. + + For example, the numbers below are global logical-block positions, not + block IDs allocated by either engine:: + + local_dcp_size=2, local_dcp_rank=0, remote_dcp_size=4 + local positions: [0, 2, 4, 6] + cached local positions: [0, 2, 4] + local_ids: [L6] # position 6 remains + + remote_rank=0 positions: [0, 4], remote_ids=[R0, R4] + start_local=(0-3)%2=1, start_remote=(3+1-0)//2=2 + matched slices: [], [] + + remote_rank=2 positions: [2, 6], remote_ids=[R2, R6] + start_local=(1-3)%2=0, start_remote=(3+0-1)//2=1 + matched slices: [L6], [R6] + """ + local_size, remote_size = local_dcp_size, remote_dcp_size + if local_size == remote_size: + # Same interleave on both sides. The remote list still starts at the + # beginning, so skip the logical blocks already cached locally. + local_slice = local_ids + remote_slice = remote_ids[local_num_computed_blocks:] + elif local_size < remote_size: + k = remote_size // local_size + p = (remote_rank - local_dcp_rank) // local_size + start_local = (p - local_num_computed_blocks) % k + start_remote = (local_num_computed_blocks + start_local - p) // k + local_slice = local_ids[start_local::k] + remote_slice = remote_ids[start_remote:] + else: + k = local_size // remote_size + remote_dcp_rank = remote_rank % remote_size + c = (local_dcp_rank - remote_dcp_rank) // remote_size + start_remote = c + local_num_computed_blocks * k + local_slice = local_ids + remote_slice = remote_ids[start_remote::k] + + matched_blocks = min(len(local_slice), len(remote_slice)) + return local_slice[:matched_blocks], remote_slice[:matched_blocks] + + class NixlPullConnectorWorker(NixlBaseConnectorWorker): """Pull-specific (READ) worker logic.""" @@ -132,36 +196,60 @@ def _read_blocks_for_req(self, req_id: str, meta: ReqMeta): remote_info = self.transfer_topo.get_engine_info(engine_id) tp_ratio = self.transfer_topo.tp_ratio(remote_info.remote_tp_size) + logical_remote_block_ids = meta.remote.block_ids meta.remote.block_ids = self._logical_to_kernel_block_ids( - meta.remote.block_ids, + logical_remote_block_ids, remote_info.remote_physical_blocks_per_logical, ) - remote_block_ids = meta.remote.block_ids - local_block_ids = meta.local_physical_block_ids - num_groups = len(local_block_ids) - read_specs = [ - ReadSpec( - remote_rank=rank, - local_block_ids=[ - list(local_block_ids[g]) - if rank in plan.source_ranks_per_group[g] - else [] - for g in range(num_groups) - ], - remote_block_ids=[ - list(remote_block_ids[g]) - if rank in plan.source_ranks_per_group[g] - else [] - for g in range(num_groups) - ], + logical_local_block_ids = meta.local_block_ids + num_groups = len(logical_local_block_ids) + dcp_active = self.dcp_size > 1 or remote_info.remote_dcp_size > 1 + + def group_ids(block_ids: BlockIds, rank: int) -> list[list[int]]: + return [ + list(block_ids[g]) if rank in plan.source_ranks_per_group[g] else [] + for g in range(num_groups) + ] + + read_specs = [] + for rank in plan.all_source_ranks: + local_ids = group_ids(logical_local_block_ids, rank) + remote_ids = group_ids(logical_remote_block_ids, rank) + if dcp_active: + # DCP interleaves at block granularity, so slicing happens + # here on logical blocks, before kernel-block expansion. + for g in range(num_groups): + if not local_ids[g] or _is_ssm_spec(self._group_spec_types[g]): + continue + # Prefix cache hit may lead to skip some of the remote reads + local_ids[g], remote_ids[g] = _dcp_read_slice( + local_ids[g], + remote_ids[g], + remote_rank=rank, + local_dcp_size=self.dcp_size, + local_dcp_rank=self.dcp_rank, + remote_dcp_size=remote_info.remote_dcp_size, + local_num_computed_blocks=meta.local_num_computed_blocks[g], + ) + local_ids = self._logical_to_kernel_block_ids( + local_ids, self._physical_blocks_per_logical_kv_block + ) + remote_ids = self._logical_to_kernel_block_ids( + remote_ids, remote_info.remote_physical_blocks_per_logical + ) + read_specs.append( + ReadSpec( + remote_rank=rank, + local_block_ids=local_ids, + remote_block_ids=remote_ids, + ) ) - for rank in plan.all_source_ranks - ] # D may have to perform multiple reads from different remote ranks. # Pure MLA reads once because its cache is replicated. Hybrid - # MLA+SSM still needs one read per SSM source rank. - if self.use_mla and tp_ratio < 0 and not self._has_mamba: + # MLA+SSM still needs one read per SSM source rank. With DCP, pure + # MLA may also read from multiple ranks (disjoint token slices). + if self.use_mla and tp_ratio < 0 and not self._has_mamba and not dcp_active: assert len(read_specs) == 1 for i, spec in enumerate(read_specs): @@ -199,12 +287,15 @@ def _read_blocks_for_req(self, req_id: str, meta: ReqMeta): remote_request_id=meta.remote.request_id, local_xfer_side_handle=local_xfer_side_handle, remote_xfer_side_handle=remote_xfer_side_handle, + expected_consumers=plan.local_consumers, ) if self.use_mla and tp_ratio < 0 and len(read_specs) == 1: # ..but we still need to notify the other remote ranks that we # have the blocks we need so they can update the request state. - notif_id = f"{meta.remote.request_id}:{self.world_size}".encode() + # Same thing for DCP (tp_size == dcp_size), so the raw tp_ratio already + # reflects whether any remote replica is left unchosen. + notif_id = f"{meta.remote.request_id}:{plan.local_consumers}".encode() remote_agents = self._remote_agents[meta.remote.engine_id] for rank_to_notify, agent in remote_agents.items(): if rank_to_notify != (0, read_specs[0].remote_rank): @@ -218,6 +309,7 @@ def _read_blocks( remote_request_id: str, local_xfer_side_handle: int, remote_xfer_side_handle: int, + expected_consumers: int, ): """ Post a READ point-to-point xfer request from a single local worker to @@ -247,13 +339,13 @@ def _read_blocks( # NOTE(rob): according to nvidia the staging blocks are used to # saturate IB with heterogeneous TP sizes. - # Number of D TP workers that will read from dst P. Propagate info - # on notification so that dst worker can wait before freeing blocks. - notif_id = f"{remote_request_id}:{self.world_size}".encode() + # Number of local workers that will notify this producer worker. + # Propagate on notification so dst worker can wait before freeing. + notif_id = f"{remote_request_id}:{expected_consumers}".encode() # Full prefix cache hit: do not need to read remote blocks, # just notify P worker that we have the blocks we need. - if len(local_block_ids) == 0: + if not any(len(group) > 0 for group in local_block_ids): # A full prefix cache hit is indicated with an empty list. agent_name = self._remote_agents[dst_engine_id][(0, remote_rank)] try: @@ -334,8 +426,8 @@ def _read_blocks( def _get_new_notifs(self) -> set[str]: """ Get req_ids which got a remote xfer message. When multiple consumers - are reading from the same producer (heterogeneous TP scenario), wait - for all consumers to be done pulling. + are reading from the same producer (heterogeneous TP or DCP + scenario), wait for all consumers to be done pulling. Also handles heartbeat notifications ("HB:req1,req2,...") by extending the lease on the referenced requests. @@ -351,7 +443,7 @@ def _get_new_notifs(self) -> set[str]: self._handle_heartbeat(msg[3:]) continue - req_id, tp_size = msg.rsplit(":", 1) + req_id, expected_consumers = msg.rsplit(":", 1) if ( req_id not in self._reqs_to_send and req_id not in self._reqs_to_process @@ -364,24 +456,22 @@ def _get_new_notifs(self) -> set[str]: ) continue - # NOTE: `tp_ratio` is the opposite when swapping local<>remote - n_consumers = int(tp_size) - tp_ratio = self.transfer_topo.tp_ratio(n_consumers) - - # Number of reads *per producer* to wait for. - # When remote D TP > local P TP we expect `tp_ratio` reads. - consumers_per_producer = ( - -tp_ratio if n_consumers > self.world_size else 1 + # Every reader of this req_id reports the same count (it's + # derived from aggregate topology, not the specific rank), + # so repeated notifications never disagree on it. + self.expected_consumer_notifications_by_req[req_id] = int( + expected_consumers ) self.consumer_notification_counts_by_req[req_id] += 1 # Wait all consumers (D) to be done reading before freeing. if ( self.consumer_notification_counts_by_req[req_id] - == consumers_per_producer + == self.expected_consumer_notifications_by_req[req_id] ): notified_req_ids.add(req_id) del self.consumer_notification_counts_by_req[req_id] + del self.expected_consumer_notifications_by_req[req_id] self._reqs_to_process.remove(req_id) self._reqs_to_send.pop(req_id, None) return notified_req_ids diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/push_scheduler.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/push_scheduler.py index 413ebdcdb09d..3710457a1bdc 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/push_scheduler.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/push_scheduler.py @@ -199,7 +199,7 @@ def update_state_after_alloc( # ReqMeta without a KeyError — the actual remote block IDs are # learned by P over the NIXL handshake at WRITE time. params["remote_block_ids"] = () - self._reqs_need_recv[request.request_id] = (request, local_block_ids) + self._reqs_need_recv[request.request_id] = (request, local_block_ids, ()) # Mark as processed so a re-entry (e.g. preemption + reschedule) # doesn't re-stage the registration. @@ -241,7 +241,7 @@ def request_finished( # recv so the worker emits a notif that lets P free them. # Seed remote_block_ids so add_new_req_to_recv won't KeyError. params["remote_block_ids"] = () - self._reqs_need_recv[request.request_id] = (request, []) + self._reqs_need_recv[request.request_id] = (request, [], ()) params["do_remote_prefill"] = False return False, None diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/tp_mapping.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/tp_mapping.py index b034b7605087..fa3ca6a3c3d7 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/tp_mapping.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/tp_mapping.py @@ -56,6 +56,10 @@ class TPMapping: # FA head offset factor for hetero-TP (D_TP > P_TP). rank_offset_factor: int + # Local ranks (in aggregate) that read from a given source rank. The producer frees + # a request's blocks only once that many notifications have come in. + local_consumers: int = 1 + # ====================================================================== # TP mapping computation @@ -66,22 +70,31 @@ def compute_tp_mapping( transfer_topology: TransferTopology, remote_tp_size: int, group_spec_types: tuple[type[KVCacheSpec], ...], + remote_dcp_size: int = 1, ) -> TPMapping: """Build the complete local-to-remote TP mapping. Computes source ranks, head slot assignments, and the rank offset factor in a single pass. + + DCP support is scoped to MLA only, with a side is either fully replicated or fully + sharded. DCP-branch reuses the same rank set used at handshake selection. """ tp_rank = transfer_topology.tp_rank tp_size = transfer_topology.tp_size total_num_kv_heads = transfer_topology.total_num_kv_heads # --- Attention source ranks --- if transfer_topology.is_mla or tp_size >= remote_tp_size: - # D (local TP) > P (remote TP): multiple local ranks read different chunks from - # *one* remote rank, corresponding to different kv heads. - # For MLA, we only need one remote since cache is duplicated. When P TP=k*TP k, - # this will spread mla ranks to read from remote k*tp_rank. - attn_ranks = [tp_rank * remote_tp_size // tp_size] + if transfer_topology.is_mla and remote_dcp_size > 1: + attn_ranks = transfer_topology.dcp_source_ranks( + remote_tp_size, remote_dcp_size + ) + else: + # D (local TP) > P (remote TP): multiple local ranks read different chunks + # from *one* remote rank, corresponding to different kv heads. + # For MLA, we only need one remote since cache is duplicated. When + # P TP=k*TP k, this will spread mla ranks to read from remote k*tp_rank. + attn_ranks = [tp_rank * remote_tp_size // tp_size] else: # P (remote TP) > D (local TP): one local rank # reads from multiple remote ranks. @@ -134,9 +147,14 @@ def compute_tp_mapping( # D TP > P TP: we index into remote to read different heads depending on rank. rank_offset_factor = tp_rank % (tp_size // remote_tp_size) + local_consumers = transfer_topology.dcp_consumer_count( + remote_tp_size, remote_dcp_size + ) + return TPMapping( source_ranks_per_group=source_ranks_per_group, all_source_ranks=tuple(all_ranks), rank_to_attention_slot=rank_to_attention_slot, rank_offset_factor=rank_offset_factor, + local_consumers=local_consumers, ) diff --git a/vllm/distributed/kv_transfer/kv_transfer_state.py b/vllm/distributed/kv_transfer/kv_transfer_state.py index f9209dc3e464..33193380f210 100644 --- a/vllm/distributed/kv_transfer/kv_transfer_state.py +++ b/vllm/distributed/kv_transfer/kv_transfer_state.py @@ -85,6 +85,8 @@ def ensure_kv_transfer_initialized( vllm_config.kv_transfer_config.is_kv_transfer_instance and _KV_CONNECTOR_AGENT is None ): + # PD only supports an interleave_size equal to block_size. + vllm_config.adjust_dcp_kv_cache_interleave_size(kv_cache_config) _sync_engine_id_across_tp(vllm_config) _KV_CONNECTOR_AGENT = KVConnectorFactory.create_connector( diff --git a/vllm/model_executor/layers/attention/mla_attention.py b/vllm/model_executor/layers/attention/mla_attention.py index 16fb961e563a..10fd1ac3b42a 100644 --- a/vllm/model_executor/layers/attention/mla_attention.py +++ b/vllm/model_executor/layers/attention/mla_attention.py @@ -2634,10 +2634,20 @@ def __init__( ) self.dcp_world_size: int = -1 + self._cp_kv_cache_interleave_size: int | None = None - self.cp_kv_cache_interleave_size: int = ( - get_current_vllm_config().parallel_config.cp_kv_cache_interleave_size - ) + @property + def cp_kv_cache_interleave_size(self) -> int: + """With PD+DCP, the real value isn't known until block_size is finalized, + which happens after this layer is built. Safe to cache after the first access, + as long as the adjustment always runs before any forward pass + (it's set up in Worker.initialize_from_config, ahead of warmup/serving). + """ + if self._cp_kv_cache_interleave_size is None: + self._cp_kv_cache_interleave_size = ( + get_current_vllm_config().parallel_config.cp_kv_cache_interleave_size + ) + return self._cp_kv_cache_interleave_size @abstractmethod def forward_mqa( diff --git a/vllm/model_executor/layers/sparse_attn_indexer.py b/vllm/model_executor/layers/sparse_attn_indexer.py index 9a671be55639..7ace1802cc59 100644 --- a/vllm/model_executor/layers/sparse_attn_indexer.py +++ b/vllm/model_executor/layers/sparse_attn_indexer.py @@ -768,14 +768,30 @@ def __init__( parallel_config = get_current_vllm_config().parallel_config self.dcp_world_size = parallel_config.decode_context_parallel_size self.dcp_rank = get_dcp_group().rank_in_group if self.dcp_world_size > 1 else 0 - self.cp_kv_cache_interleave_size = parallel_config.cp_kv_cache_interleave_size self.use_pcp = parallel_config.prefill_context_parallel_size > 1 + self._cp_kv_cache_interleave_size: int | None = None if current_platform.is_cuda() and not has_deep_gemm(): raise RuntimeError( "Sparse Attention Indexer CUDA op requires DeepGEMM support in " "the current vLLM environment." ) + @property + def cp_kv_cache_interleave_size(self) -> int: + """With PD+DCP, the real value isn't known until block_size is finalized, + which happens after this layer is built. Safe to cache after the first access, + as long as the adjustment always runs before any forward pass + (it's set up in Worker.initialize_from_config, ahead of warmup/serving). + """ + if self._cp_kv_cache_interleave_size is None: + value = ( + get_current_vllm_config().parallel_config.cp_kv_cache_interleave_size + ) + if isinstance(get_forward_context().attn_metadata, dict): + self._cp_kv_cache_interleave_size = value + return value + return self._cp_kv_cache_interleave_size + def forward_native( self, hidden_states: torch.Tensor,