diff --git a/tests/v1/kv_connector/unit/test_tp_mapping.py b/tests/v1/kv_connector/unit/test_tp_mapping.py index e57244a31f7..95d49faf042 100644 --- a/tests/v1/kv_connector/unit/test_tp_mapping.py +++ b/tests/v1/kv_connector/unit/test_tp_mapping.py @@ -9,6 +9,8 @@ No GPU or NIXL required. from __future__ import annotations +from types import SimpleNamespace + import pytest from vllm.distributed.kv_transfer.kv_connector.v1.nixl.tp_mapping import ( @@ -33,12 +35,15 @@ def _compute_mapping( num_kv_heads: int = 8, group_spec_types: tuple[type, ...] = (FullAttentionSpec,), ) -> TPMapping: - return compute_tp_mapping( + transfer_topology = SimpleNamespace( tp_rank=tp_rank, tp_size=tp_size, - remote_tp_size=remote_tp_size, is_mla=is_mla, total_num_kv_heads=num_kv_heads, + ) + return compute_tp_mapping( + transfer_topology=transfer_topology, + remote_tp_size=remote_tp_size, group_spec_types=group_spec_types, ) diff --git a/vllm/distributed/kv_transfer/kv_connector/utils.py b/vllm/distributed/kv_transfer/kv_connector/utils.py index b85416ab307..a86b25c75a4 100644 --- a/vllm/distributed/kv_transfer/kv_connector/utils.py +++ b/vllm/distributed/kv_transfer/kv_connector/utils.py @@ -593,7 +593,7 @@ class TransferTopology: return ( f"TransferTopology(" f"tp_ratio={self.tp_ratio(info.remote_tp_size)}, " - f"K={self.total_num_kv_heads}, " + 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"local_rank={self.tp_rank}, " 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 7115b8bed54..b034b760508 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 @@ -10,6 +10,7 @@ import numpy as np from vllm.distributed.kv_transfer.kv_connector.utils import ( BlockIds, + TransferTopology, ) from vllm.v1.kv_cache_interface import AttentionSpec, KVCacheSpec, MambaSpec @@ -62,11 +63,8 @@ class TPMapping: def compute_tp_mapping( - tp_rank: int, - tp_size: int, + transfer_topology: TransferTopology, remote_tp_size: int, - is_mla: bool, - total_num_kv_heads: int, group_spec_types: tuple[type[KVCacheSpec], ...], ) -> TPMapping: """Build the complete local-to-remote TP mapping. @@ -74,13 +72,15 @@ def compute_tp_mapping( Computes source ranks, head slot assignments, and the rank offset factor in a single pass. """ + 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 is_mla: - # All heads replicated across all ranks. - attn_ranks = [0] - elif tp_size >= remote_tp_size: + 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] else: # P (remote TP) > D (local TP): one local rank @@ -123,7 +123,7 @@ def compute_tp_mapping( } # --- Rank offset factor --- - if is_mla or tp_size <= remote_tp_size: + if transfer_topology.is_mla or tp_size <= remote_tp_size: # We don't index into remote for reading, no offset needed. rank_offset_factor = 0 elif tp_size > total_num_kv_heads: diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/utils.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/utils.py index d0b72464a27..2fa3829eaec 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/utils.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/utils.py @@ -10,6 +10,7 @@ import zmq from vllm.platforms import current_platform from vllm.utils.network_utils import make_zmq_socket +from vllm.v1.kv_cache_interface import KVCacheSpec, UniformTypeKVCacheSpecs # Supported platforms and types of kv transfer buffer. # {device: tuple of supported kv buffer types} @@ -46,3 +47,11 @@ def zmq_ctx(socket_type: Any, addr: str) -> Iterator[zmq.Socket]: finally: if ctx is not None: ctx.destroy(linger=0) + + +def get_representative_spec_type(spec: KVCacheSpec) -> type[KVCacheSpec]: + if isinstance(spec, UniformTypeKVCacheSpecs): + # All inner specs are the same type; pick any. + inner = next(iter(spec.kv_cache_specs.values())) + return type(inner) + return type(spec) diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/worker.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/worker.py index caa5f432c5e..eb90a4e1729 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/worker.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/worker.py @@ -53,6 +53,7 @@ from vllm.distributed.kv_transfer.kv_connector.v1.nixl.tp_mapping import ( ) from vllm.distributed.kv_transfer.kv_connector.v1.nixl.utils import ( _NIXL_SUPPORTED_DEVICE, + get_representative_spec_type, zmq_ctx, ) from vllm.distributed.kv_transfer.kv_connector.v1.ssm_conv_transfer_utils import ( @@ -100,24 +101,24 @@ class NixlConnectorWorker: num_blocks = dst_num_blocks if block_size_ratio is not None: num_blocks = int(num_blocks * block_size_ratio) - ratio = physical_blocks_per_logical - logical_blocks = num_blocks // ratio - num_fa_descs = num_fa_regions * num_blocks # All-attention fast path: single vectorized broadcast. if num_ssm_regions == 0: + # NOTE (NickLucche) With HMA, every kv group has the same number of layers + # and layers from different groups share the same kv tensor. + # eg block_ids=[[1, 2], [3]]->blocks [1, 2] need to be + # read across all regions, same for [3], but group0-group1 blocks will + # always differ (different areas). Therefore we can just flatten the + # block_ids and compute the descs ids for all groups at once. block_arr = np.concatenate(block_ids)[None, :] region_ids = np.arange(num_fa_regions)[:, None] return (region_ids * num_blocks + block_arr).flatten() - # NOTE (NickLucche) With HMA, every kv group has the same number - # of layers and layers from different groups share the same kv - # tensor. Therefore we compute desc IDs per group using the - # right stride: - # FA descs have num_blocks entries per region (kernel granularity), - # SSM descs have logical_blocks entries per region (no kernel - # splitting). + # Compute desc ids per group using the right stride: FA descs have + # num_blocks entries per region (kernel granularity), SSM descs have + # logical_blocks entries per region (no kernel splitting). + logical_blocks = num_blocks // physical_blocks_per_logical all_descs: list[np.ndarray] = [] for i, group in enumerate(block_ids): group_arr = np.asarray(group) @@ -426,8 +427,10 @@ class NixlConnectorWorker: self._physical_blocks_per_logical_kv_block = 1 self._sync_block_size_with_kernel() + # Unwrap UniformTypeKVCacheSpecs to get the representative spec type self._group_spec_types = tuple( - type(g.kv_cache_spec) for g in self.kv_cache_config.kv_cache_groups + get_representative_spec_type(g.kv_cache_spec) + for g in self.kv_cache_config.kv_cache_groups ) # Per-engine TP mappings. Generated during handshake. @@ -1259,12 +1262,9 @@ class NixlConnectorWorker: logger.info("Transfer plan: %s", transfer_topo.describe(engine_id)) self.tp_mappings[engine_id] = compute_tp_mapping( - transfer_topo.tp_rank, - transfer_topo.tp_size, - transfer_info.remote_tp_size, - transfer_topo.is_mla, - transfer_topo.total_num_kv_heads, - self._group_spec_types, + transfer_topology=transfer_topo, + remote_tp_size=remote_tp_size, + group_spec_types=self._group_spec_types, ) remote_agent_name = self.nixl_wrapper.add_remote_agent( @@ -1391,7 +1391,8 @@ class NixlConnectorWorker: ) # num_kv_heads > tp_size with P_TP > D_TP not supported for non-mamba. # Mamba models can have replicated FA KV with tp_ratio < 0. - if not self._has_mamba: + # MLA models do not need to handle kv replication. + if not self.use_mla and not self._has_mamba: assert not ( tp_ratio < 0 and self.transfer_topo.is_kv_replicated(remote_engine_id) ) @@ -1915,9 +1916,9 @@ class NixlConnectorWorker: # D may have to perform multiple reads from different remote ranks. # MLA opt: when P TP > D TP, only a single read is executed for - # the first remote rank (cache is duplicated). + # the first remote rank (cache is duplicated).. if self.use_mla and tp_ratio < 0: - read_specs = read_specs[:1] + assert len(read_specs) == 1 for i, spec in enumerate(read_specs): remote_block_size = remote_info.remote_block_size @@ -1959,11 +1960,10 @@ class NixlConnectorWorker: if self.use_mla and tp_ratio < 0 and read_specs: # ..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"{req_id}:{self.world_size}".encode() + notif_id = f"{meta.remote.request_id}:{self.world_size}".encode() remote_agents = self._remote_agents[meta.remote.engine_id] - read_ranks = {s.remote_rank for s in read_specs} for rank_to_notify, agent in remote_agents.items(): - if rank_to_notify not in read_ranks: + if rank_to_notify != read_specs[0].remote_rank: self.nixl_wrapper.send_notif(agent, notif_msg=notif_id) def _read_blocks(