From 2b465b2c42e6f7d37fbbc67956dc9741e832dc29 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicol=C3=B2=20Lucchesi?= Date: Mon, 27 Jul 2026 18:51:37 +0200 Subject: [PATCH] [Misc][PD] Nixl cleanup `get_backend_aware_kv_block_len` and `virtually_split_kv_in_blocks` (#49988) Signed-off-by: NickLucche --- .../kv_connector/unit/test_nixl_connector.py | 12 -- .../kv_connector/v1/nixl/base_worker.py | 117 +++++++----------- 2 files changed, 44 insertions(+), 85 deletions(-) diff --git a/tests/v1/kv_connector/unit/test_nixl_connector.py b/tests/v1/kv_connector/unit/test_nixl_connector.py index 8bb554a0433..d58f117fa31 100644 --- a/tests/v1/kv_connector/unit/test_nixl_connector.py +++ b/tests/v1/kv_connector/unit/test_nixl_connector.py @@ -1115,18 +1115,6 @@ class TestNixlHandshake: block_lens=[remote_block_len], ) - assert worker.get_backend_aware_kv_block_len(0, mamba_view=False) == ( - local_block_len - ) - assert ( - worker.get_backend_aware_kv_block_len(0, first_split=True, mamba_view=True) - == worker._mamba_ssm_size[0] - ) - assert ( - worker.get_backend_aware_kv_block_len(0, first_split=False, mamba_view=True) - == worker._mamba_ssm_size[1] - ) - assert worker._build_fa_remote(plan, meta, block_size_ratio=1).tolist() == [ [0x1000 + local_block_len, local_block_len, 0] ] 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 671f5272dbf..d4e516d6df7 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 @@ -96,7 +96,6 @@ class NixlBaseConnectorWorker: physical_blocks_per_logical: int, ) -> np.ndarray: """Compute NIXL descriptor IDs for given block IDs.""" - num_fa_regions = self.num_regions num_ssm_regions = 0 if self._has_mamba: assert self._conv_decomp is not None @@ -108,7 +107,7 @@ class NixlBaseConnectorWorker: num_blocks = dst_num_blocks if block_size_ratio is not None: num_blocks = int(num_blocks * block_size_ratio) - num_fa_descs = num_fa_regions * num_blocks + num_fa_descs = self.num_regions * num_blocks # All-attention fast path: single vectorized broadcast. if num_ssm_regions == 0: @@ -119,7 +118,7 @@ class NixlBaseConnectorWorker: # 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] + region_ids = np.arange(self.num_regions)[:, None] return (region_ids * num_blocks + block_arr).flatten() # Compute desc ids per group using the right stride: FA descs have @@ -130,7 +129,7 @@ class NixlBaseConnectorWorker: for i, group in enumerate(block_ids): group_arr = np.asarray(group) if _is_attention_spec(self._group_spec_types[i]): - fa_region_ids = np.arange(num_fa_regions)[:, None] + fa_region_ids = np.arange(self.num_regions)[:, None] all_descs.append( (fa_region_ids * num_blocks + group_arr[None, :]).flatten() ) @@ -210,7 +209,7 @@ class NixlBaseConnectorWorker: def _fa_desc_replicated(self, num_fa_descs: int) -> list[bool]: """Per-FA-descriptor replicate flag, in _build_fa_local emission order - (region-major; K then optional V per region). Length ``num_fa_descs``. + (region-major; one desc per block, with K/V packed). Length ``num_fa_descs``. """ assert self.transfer_topo is not None n_regions = len(self.block_len_per_layer) @@ -1082,21 +1081,15 @@ class NixlBaseConnectorWorker: # With hybrid allocator, layers can share a kv cache tensor seen_base_addresses = [] - # Note(tms): I modified this from the original region setup code. - # K and V are now in different regions. Advantage is that we can - # elegantly support MLA and any cases where the K and V tensors - # are non-contiguous (it's not locally guaranteed that they will be) - # Disadvantage is that the encoded NixlAgentMetadata is now larger - # (roughly 8KB vs 5KB). - # Conversely for FlashInfer, K and V are registered in the same region - # to better exploit the memory layout (ie num_blocks is the first dim). + # K and V are packed into the content dim, so each attention layer is a + # single NIXL region whose block transfers as one unit. Mamba layers instead + # register separate conv/ssm sub-regions (see `_build_mamba_local`). tensor_size_bytes = None for layer_name, cache in xfer_buffers.items(): - # NOTE (NickLucche) Hybrid SSM models assume a layout that is similar to - # that of FI, with block laid out as in `get_backend_aware_kv_block_len`. - # However, physical page_size may differ when kernel requires a specific - # block size. This leads to SSM and FA layers having different num_blocks. + # NOTE (NickLucche) Hybrid SSM mamba/FA physical page_size may differ when + # kernel requires a specific block size. This leads to SSM and FA layers + # having different num_blocks. # `_physical_blocks_per_logical_kv_block` ratio is used to adjust for this. layer_spec = self._layer_specs.get(layer_name) if layer_spec is None: @@ -1107,7 +1100,7 @@ class NixlBaseConnectorWorker: ) continue if isinstance(layer_spec, UniformTypeKVCacheSpecs): - # MLA DSv32 Indexer case: UniformTypeKVCacheSpecs merges kv_cache_specs + # DSA Indexer case: UniformTypeKVCacheSpecs merges kv_cache_specs layer_spec = layer_spec.kv_cache_specs[layer_name] # `layer_spec.page_size_bytes` only accounts for logical page_size, that is # the page_size assuming constant `self._logical_num_blocks`. @@ -1270,7 +1263,32 @@ class NixlBaseConnectorWorker: block_size_ratio: int, ) -> np.ndarray: """Build desc regions (conv sub-projections + ssm) per layer for - local mamba blocks with DS conv layout, as an Nx3 uint64 array.""" + local mamba blocks with DS conv layout, as an Nx3 uint64 array. + + A Mamba block interleaves conv and SSM state, which crucially differ in + size, so the two are indexed as separate sub-regions. Attention blocks + instead pack K and V into the content dim and transfer as a single unit. + Reference diagram: + KVCacheTensor (Shared) + / \\ + / \\ + / \\ + Attention (FlashInfer) View Mamba View + | | + | | + +-------------------+ +-------------------+ + | KVCacheTensor | | KVCacheTensor | + | | | | + |<----- page ------>| |<----- page ------->| + | size | | size | + | Key 0 | Val 0 | |Conv 0 | SSM 0 | + | Key 1 | Val 1 | |Conv 1 | SSM 1 | + | ... | ... | | ... | ... | + | Key N-2 | Val N-2 | |Conv N-2| SSM N-2 | + | Key N-1 | Val N-1 | |Conv N-1| SSM N-1 | + +-------------------+ +--------------------+ + |1st_split-2nd_split| |1st_split-2nd_split | + """ assert block_size_ratio == 1, ( "Mamba 3-read transfer with block_size_ratio != 1 is not tested. " f"Got block_size_ratio={block_size_ratio}." @@ -1362,15 +1380,11 @@ class NixlBaseConnectorWorker: block_arange = np.arange(num_blocks, dtype=np.uint64) parts: list[np.ndarray] = [] for i, base_addr in enumerate(base_addresses): - kv_block_len = ( - self.get_backend_aware_kv_block_len( - layer_idx=i, first_split=True, mamba_view=False - ) - // block_size_ratio - ) - page_stride = self.block_len_per_layer[i] // block_size_ratio - addrs = base_addr + block_arange * page_stride - parts.append(self._stack_descs(addrs, kv_block_len, device_id)) + # K/V are packed into the content dim, so the whole block transfers + # as one unit: desc length equals the block stride. + block_len = self.block_len_per_layer[i] // block_size_ratio + addrs = base_addr + block_arange * block_len + parts.append(self._stack_descs(addrs, block_len, device_id)) return np.concatenate(parts) def _build_fa_remote( @@ -1397,9 +1411,7 @@ class NixlBaseConnectorWorker: for i, base_addr in enumerate(nixl_agent_meta.kv_caches_base_addr): replicated = self._is_region_replicated(i) # Read our whole local region size from remote.. - local_block_len = self.get_backend_aware_kv_block_len( - layer_idx=i, first_split=True, mamba_view=False - ) + local_block_len = self.block_len_per_layer[i] remote_kv_block_len = local_block_len // block_size_ratio if block_size_ratio > 1: # ..using remote kv_block_len as transfer unit @@ -1521,8 +1533,7 @@ class NixlBaseConnectorWorker: ) return self._remote_agents[engine_id][(0, remote_tp_rank)] - # Compare physical regions, not self.num_regions (doubled by - # FlashInfer's virtual K/V split). + # Number of physical regions registered locally (one per layer/tensor). num_local_regions = len(self.block_len_per_layer) if ( self.pp_size > 1 @@ -2301,46 +2312,6 @@ class NixlBaseConnectorWorker: remote_block_ids[i] = remote_group[:num_blocks] return local_block_ids, remote_block_ids - def get_backend_aware_kv_block_len( - self, layer_idx: int, first_split: bool = True, mamba_view: bool = False - ) -> int: - """ - Get the block length for one K/V element (K and V have the same size). - - For FA and other backends, this is equal to the length of the whole - block, as K and V are in separate regions. - For FlashInfer, this is half the length of the whole block, as K and V - share the same region. - Similarly, for SSM-based models, state and conv are interleaved, but crucially - the their size differs. - Reference diagram: - KVCacheTensor (Shared) - / \\ - / \\ - / \\ - Attention (FlashInfer) View Mamba View - | | - | | - +-------------------+ +-------------------+ - | KVCacheTensor | | KVCacheTensor | - | | | | - |<----- page ------>| |<----- page ------->| - | size | | size | - | Key 0 | Val 0 | |Conv 0 | SSM 0 | - | Key 1 | Val 1 | |Conv 1 | SSM 1 | - | ... | ... | | ... | ... | - | Key N-2 | Val N-2 | |Conv N-2| SSM N-2 | - | Key N-1 | Val N-1 | |Conv N-1| SSM N-1 | - +-------------------+ +--------------------+ - |1st_split-2nd_split| |1st_split-2nd_split | - """ - assert self.transfer_topo is not None - if self.transfer_topo.virtually_split_kv_in_blocks and mamba_view: - block_len = self._mamba_ssm_size[not first_split] - else: - block_len = self.block_len_per_layer[layer_idx] - return block_len - def get_kv_connector_stats(self) -> KVConnectorStats | None: """ Get the KV transfer stats for the connector.