diff --git a/tensorrt_llm/_torch/disaggregation/native/bounce/impl.py b/tensorrt_llm/_torch/disaggregation/native/bounce/impl.py index 5d88485ec8b1..931f33794e79 100644 --- a/tensorrt_llm/_torch/disaggregation/native/bounce/impl.py +++ b/tensorrt_llm/_torch/disaggregation/native/bounce/impl.py @@ -16,9 +16,11 @@ the contract in core.py. Holds the buffers, the gather and scatter kernels, and the scatter worker, and runs the side effects that drive each region's state machine. Never imports transfer.py.""" +from __future__ import annotations + import queue import threading -from typing import Callable, Dict, List, Optional +from typing import TYPE_CHECKING, Callable, Dict, List, Optional import numpy as np @@ -41,6 +43,9 @@ from .core import BounceTransport, Disposition, Settlement, TransferContext from .gather_scatter import Plan, gather_contiguous, scatter_contiguous +if TYPE_CHECKING: + from tensorrt_llm._torch.disaggregation.resource.page import KVCachePageTable + RidSlice = tuple # the request id and slice id a region serves _MIB = 1024 * 1024 _SCATTER_POLL_S = 0.5 # how often the scatter worker wakes to re-check the stop flag and reclaim @@ -57,8 +62,8 @@ class VmmBounceTransport(BounceTransport): @classmethod def from_config( - cls, agent, cfg, *, device_id: int, block_bytes_per_group: List[int] - ) -> Optional["VmmBounceTransport"]: + cls, agent, cfg, *, device_id: int, block_bytes_per_group: list[int | None] + ) -> VmmBounceTransport | None: """Build a transport sized from the config and clamped to free memory, or None if not even one chunk fits.""" chunk = cfg.chunk_mb * _MIB @@ -94,7 +99,7 @@ def __init__( device_id: int, capacity_bytes: int, phys_chunk_size: int, - block_bytes_per_group: List[int], + block_bytes_per_group: list[int | None], min_bytes: int = DEFAULT_MIN_BYTES, min_blocks: int = 96, quarantine_grace_s: float = _QUARANTINE_GRACE_S, @@ -587,21 +592,28 @@ def decode_result_tail(message): return None, None, None -def block_bytes_per_group(page_table) -> list: - """Byte size of one cache block for each layer group, aligned with the layer-group indices a - recv request uses. Non-attention groups (mamba/KDA recurrent state) hold ``None``: they carry - no paged blocks (their KVSlice entry is always empty — see ``_create_kv_slice``) and their - payload is sized separately via ``MambaPolicy.payload_bytes``. Keeping them as placeholders - instead of truncating means a trailing (or hypothetically interleaved) mamba group can never - shift an attention group off the end of this list and poison the bounce gate.""" +def block_bytes_per_group(page_table: KVCachePageTable) -> list[int | None]: + """Return transferred bytes per cache block for each layer group. + + All distinct physical pools exposed by an attention group contribute to its + transfer size. Multiple logical views of the same physical pool contribute + only once. Non-attention groups retain a ``None`` placeholder so the result + remains aligned with receive-request layer-group indices. + """ from tensorrt_llm._torch.disaggregation.resource.page import AttentionLayerGroup from tensorrt_llm._torch.disaggregation.resource.utils import get_physical_pool assert page_table is not None - out: list = [] + out: list[int | None] = [] for lg_idx, lg in enumerate(page_table.layer_groups): if not isinstance(lg, AttentionLayerGroup): out.append(None) continue - out.append(int(get_physical_pool(page_table, lg_idx, 0).slot_bytes)) + pool_indices = {pool_view.pool_idx for pool_view in lg.pool_views} + out.append( + sum( + int(get_physical_pool(page_table, lg_idx, pool_idx).slot_bytes) + for pool_idx in pool_indices + ) + ) return out diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index 673d088e8443..a578cf3b179a 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -61,13 +61,13 @@ ) from tensorrt_llm._torch.disaggregation.native.messenger import ZMQMessenger, decode_message from tensorrt_llm._torch.disaggregation.native.mixers.ssm.peer import MambaPolicy -from tensorrt_llm._torch.disaggregation.native.peer import PeerRegistrar +from tensorrt_llm._torch.disaggregation.native.peer import PeerOverlap, PeerRegistrar from tensorrt_llm._torch.disaggregation.native.perf_logger import PerfTimer, perf_log_manager from tensorrt_llm._torch.disaggregation.native.rank_info import RankInfo from tensorrt_llm._torch.disaggregation.native.utils import get_local_ip from tensorrt_llm._torch.disaggregation.nixl.agent import NixlTransferAgent from tensorrt_llm._torch.disaggregation.resource.kv_extractor import KVRegionExtractorV1 -from tensorrt_llm._torch.disaggregation.resource.page import MapperKind +from tensorrt_llm._torch.disaggregation.resource.page import KVCachePageTable, MapperKind from tensorrt_llm._torch.disaggregation.resource.utils import get_unique_pool_memory_descs from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager @@ -1587,7 +1587,11 @@ def _build_recv_req_info(self, task: KVRecvTask) -> RecvReqInfo: ) @staticmethod - def _fanin_bounce_safe(overlap, peer_ri) -> bool: + def _fanin_bounce_safe( + overlap: PeerOverlap, + peer_ri: RankInfo, + receiver_page_table: Optional[KVCachePageTable], + ) -> bool: """Whether multi-writer bounce's equal total//num_writers split is valid for this overlap. The split assumes every writer contributes the same size, which holds when: * duplicate_head_factor == 1 -- else some ranks don't send KV (should_send_kv) yet still @@ -1607,15 +1611,20 @@ def _fanin_bounce_safe(overlap, peer_ri) -> bool: return False # Replicated pools (e.g. MiniMax M3 index-key) are sent by one elected # fan-in owner only, so with multiple writers their contributions - # differ in size and the equal split is invalid. - if len(overlap.ranks) > 1 and peer_ri.page_table is not None: - for layer_group in peer_ri.page_table.layer_groups: - for pool_view in getattr(layer_group, "pool_views", ()): - if pool_view.mapper_kind == MapperKind.REPLICATED: - return False + # differ in size and the equal split is invalid. Inspect both endpoints: + # a masked PP stage may advertise no replicated view even though another + # stage owns one that is visible in the receiver's page table. + if len(overlap.ranks) > 1: + for page_table in (peer_ri.page_table, receiver_page_table): + if page_table is None: + continue + for layer_group in page_table.layer_groups: + for pool_view in getattr(layer_group, "pool_views", ()): + if pool_view.mapper_kind == MapperKind.REPLICATED: + return False return True - def dispatch_task(self, task: KVRecvTask): + def dispatch_task(self, task: KVRecvTask) -> None: params = task._params logger.debug( f"Receiver.dispatch_task: unique_rid={task._unique_rid}, ctx_dp_rank={params.ctx_dp_rank}" @@ -1657,7 +1666,12 @@ def dispatch_task(self, task: KVRecvTask): # None), where the real writer count exceeds expected_transfers and would overflow the slot. topo_overlap = peer_overlap if sender_dp_rank is not None else dp0_overlap allow_bounce = task.expected_transfers == 1 or ( - sender_dp_rank is not None and self._fanin_bounce_safe(topo_overlap, peer_infos) + sender_dp_rank is not None + and self._fanin_bounce_safe( + topo_overlap, + peer_infos, + self._registrar.self_extractor.page_table, + ) ) # Recurrent (mamba/KDA) state rides the SAME coalesced write as the KV blocks (the sender # appends its MambaPolicy fragments in _build_kv_write_meta), so the bounce region must be diff --git a/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py b/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py index bf8e91006b7d..8d772c0de6e9 100644 --- a/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py +++ b/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py @@ -326,60 +326,65 @@ def build_page_table(kv_cache_manager: KVCacheManager) -> KVCachePageTable: pool_views = [kv_view] # Indexer K cache support. The DSA indexer K cache is identical on - # every TP rank (single index head), so its view is REPLICATED with - # one synthesized buffer entry per local layer: the slot packs the - # layers equal-sized in local-layer order. - if getattr(kv_cache_manager, "enable_indexer_k_cache", False): - local_indexer_mask = getattr(kv_cache_manager, "indexer_k_cache_local_layer_mask", None) - if local_indexer_mask is not None and not all( - local_indexer_mask[lid] for lid in local_layer_ids - ): - raise NotImplementedError( - "The Python KV transceiver runtime does not support a " - "per-layer masked indexer k-cache pool yet: " - f"{sum(local_indexer_mask[lid] for lid in local_layer_ids)}" - f" of {len(local_layer_ids)} layers in this layer group " - "own an indexer k-cache. Use the C++ cache transceiver " - "for models with cross-layer indexer sharing (e.g. " - "GLM 5.2)." + # every TP rank (single index head), so its view is REPLICATED. With a + # per-layer indexer mask (cross-layer indexer sharing, e.g. GLM 5.2) + # only the "full" indexer-owning layers get a pool row, so the view + # covers that subset: one buffer entry per owning layer, each mapped to + # its packed row in the (possibly masked) pool. When the mask is absent + # every layer owns a row (dense/legacy layout) and this reduces to the + # equal-sized packing in local-layer order. + if kv_cache_manager.enable_indexer_k_cache: + local_indexer_mask = kv_cache_manager.indexer_k_cache_local_layer_mask + owning_layer_ids = [ + lid + for lid in local_layer_ids + if local_indexer_mask is None or local_indexer_mask[lid] + ] + # A layer group whose layers are all masked out owns no indexer pool + # row on this rank (the pool getter would raise); skip it so the peer + # simply transfers nothing for this rank's indexer. + if owning_layer_ids: + indexer_pool = kv_cache_manager.impl.get_indexer_k_cache_pool() + if indexer_pool.shape[1] != len(owning_layer_ids): + raise RuntimeError( + "The DSA indexer K-cache pool row count does not match " + "the number of indexer-owning layers in its layer group: " + f"{indexer_pool.shape[1]} rows for {len(owning_layer_ids)} layers" + ) + # indexer_pool shape: (numBlocks, numIndexerLayers, kvFactor, + # blockSize), dtype=UINT8. numIndexerLayers is the number of + # owning layers on this rank (== the attention layer count when + # unmasked). slot_bytes packs every owning-layer row. + per_block_elems = 1 + for d in indexer_pool.shape[1:]: # skip numBlocks dim + per_block_elems *= d + indexer_slot_bytes = per_block_elems * indexer_pool.element_size() + indexer_bytes_per_layer = indexer_slot_bytes // indexer_pool.shape[1] + indexer_physical = PhysicalPool( + base_address=int(indexer_pool.data_ptr()), + slot_bytes=indexer_slot_bytes, + num_slots=num_blocks, ) - indexer_pool = kv_cache_manager.impl.get_indexer_k_cache_pool() - # indexer_pool shape: (numBlocks, numLayers, kvFactor, blockSize), dtype=UINT8 - # slot_bytes = numLayers * kvFactor * blockSize * element_size - if indexer_pool.shape[1] != len(local_layer_ids): - raise NotImplementedError( - "Disaggregated KV transfer does not support a per-layer " - "masked indexer k-cache pool yet: the indexer " - f"pool holds {indexer_pool.shape[1]} layer rows but the " - f"layer group has {len(local_layer_ids)} layers. Disable " - "disaggregated serving for models with cross-layer " - "indexer sharing." + indexer_view = PoolView( + pool_idx=len(physical_pools), + buffer_entries=np.array( + [ + ( + lid, + kv_cache_manager.impl.get_indexer_k_cache_pool_layer_idx(lid) + * indexer_bytes_per_layer, + indexer_bytes_per_layer, + ) + for lid in owning_layer_ids + ], + dtype=BUFFER_ENTRY_DTYPE, + ), + pool_role=frozenset({"indexer_k"}), + mapper_kind=MapperKind.REPLICATED, + bytes_per_layer=indexer_bytes_per_layer, ) - per_block_elems = 1 - for d in indexer_pool.shape[1:]: # skip numBlocks dim - per_block_elems *= d - indexer_slot_bytes = per_block_elems * indexer_pool.element_size() - indexer_physical = PhysicalPool( - base_address=int(indexer_pool.data_ptr()), - slot_bytes=indexer_slot_bytes, - num_slots=num_blocks, - ) - indexer_bytes_per_layer = indexer_slot_bytes // len(local_layer_ids) - indexer_view = PoolView( - pool_idx=1, - buffer_entries=np.array( - [ - (lid, i * indexer_bytes_per_layer, indexer_bytes_per_layer) - for i, lid in enumerate(local_layer_ids) - ], - dtype=BUFFER_ENTRY_DTYPE, - ), - pool_role=frozenset({"indexer_k"}), - mapper_kind=MapperKind.REPLICATED, - bytes_per_layer=indexer_bytes_per_layer, - ) - physical_pools.append(indexer_physical) - pool_views.append(indexer_view) + physical_pools.append(indexer_physical) + pool_views.append(indexer_view) pool_groups.append(PhysicalPoolGroup(pools=physical_pools)) local_layers = [ diff --git a/tensorrt_llm/_torch/models/modeling_deepseekv3.py b/tensorrt_llm/_torch/models/modeling_deepseekv3.py index 9755f0fc2806..26998721c843 100755 --- a/tensorrt_llm/_torch/models/modeling_deepseekv3.py +++ b/tensorrt_llm/_torch/models/modeling_deepseekv3.py @@ -1901,7 +1901,6 @@ def forward( return hidden_states -@register_auto_model("GlmMoeDsaForCausalLM") @register_auto_model("DeepseekV32ForCausalLM") @register_auto_model("DeepseekV3ForCausalLM") class DeepseekV3ForCausalLM(SpecDecOneEngineForCausalLM[DeepseekV3Model, @@ -1919,25 +1918,14 @@ def get_preferred_transceiver_runtime( cls, pretrained_config: Any = None ) -> Optional[Literal["CPP", "PYTHON"]]: - """Preferred KV-cache transceiver runtime, differentiated per checkpoint. - - ``DeepseekV3ForCausalLM`` / ``DeepseekV32ForCausalLM`` use MLA attention, which transfers - a large latent KV that the Python (v2) transceiver handles better in disaggregated - serving, so they prefer the Python transceiver. GLM 5.2 (``GlmMoeDsaForCausalLM`` / - ``glm_moe_dsa``) uses a per-layer masked DSA indexer k-cache pool (cross-layer indexer - sharing) that the Python transceiver does not support, so GLM checkpoints must use the - C++ transceiver, which handles both the masked pool and dense indexer layouts. Applied - only when ``cache_transceiver_config.transceiver_runtime`` is 'auto'; an explicit runtime + """Preferred KV-cache transceiver runtime. + + ``DeepseekV3ForCausalLM`` and ``DeepseekV32ForCausalLM`` use MLA + attention, which transfers a large latent KV that the Python (v2) + transceiver handles better in disaggregated serving. Applied only when + ``cache_transceiver_config.transceiver_runtime`` is 'auto'; an explicit runtime is always respected. """ - if pretrained_config is not None: - architectures = getattr(pretrained_config, 'architectures', - None) or [] - # model_type is checked as a fallback: it is 'glm_moe_dsa' on GLM - # checkpoints until __init__ rewrites it to 'deepseek_v32'. - if ("GlmMoeDsaForCausalLM" in architectures or getattr( - pretrained_config, 'model_type', None) == 'glm_moe_dsa'): - return "CPP" return "PYTHON" def __init__(self, model_config: ModelConfig[PretrainedConfig]): @@ -2103,3 +2091,16 @@ def setup_aliases(self) -> None: layer.mlp.experts.fuse_shared_expert( layer.mlp.shared_experts) layer.mlp.shared_experts = None + + +@register_auto_model("GlmMoeDsaForCausalLM") +class GlmMoeDsaForCausalLM(DeepseekV3ForCausalLM): + """GLM 5.2 model flavor with an independent transceiver preference.""" + + @classmethod + def get_preferred_transceiver_runtime( + cls, + pretrained_config: Any = None, + ) -> Optional[Literal["CPP", "PYTHON"]]: + """Prefer Python for GLM 5.2's masked DSA indexer K-cache transfer.""" + return "PYTHON" diff --git a/tests/integration/test_lists/qa/llm_function_core.txt b/tests/integration/test_lists/qa/llm_function_core.txt index 6f222f7bab1f..3efab398ec36 100644 --- a/tests/integration/test_lists/qa/llm_function_core.txt +++ b/tests/integration/test_lists/qa/llm_function_core.txt @@ -815,6 +815,7 @@ accuracy/test_llm_api_pytorch_multimodal.py::TestQwen3_5_35B_A3B_VL::test_fp8_pr accuracy/test_llm_api_pytorch_multimodal.py::TestQwen3_5_27B_VL::test_auto_dtype accuracy/test_llm_api_pytorch_multimodal.py::TestVILA1_5_3B::test_auto_dtype accuracy/test_llm_api_pytorch_ray.py::TestLlama3_1_8BInstruct::test_pp2_ray +unittest/disaggregated/test_cache_transceiver_single_process.py::test_cache_transceiver_v1_masked_dsa_indexer_across_asymmetric_pp unittest/disaggregated/test_openai_disagg_server.py disaggregated/test_ad_disagg.py::test_async_eagle3_full_model_handoff disaggregated/test_ad_disagg.py::test_async_generation_matches_aggregate diff --git a/tests/integration/test_lists/test-db/l0_h100.yml b/tests/integration/test_lists/test-db/l0_h100.yml index 88130605c153..13f4bb62d60e 100644 --- a/tests/integration/test_lists/test-db/l0_h100.yml +++ b/tests/integration/test_lists/test-db/l0_h100.yml @@ -97,6 +97,7 @@ l0_h100: - unittest/disaggregated/test_cache_transceiver_single_process.py::test_cache_transceiver_boundary_lengths -k "v2" # DSA indexer-K side cache (V1, REPLICATED). - unittest/disaggregated/test_cache_transceiver_single_process.py::test_cache_transceiver_v1_dsa_indexer + - unittest/disaggregated/test_cache_transceiver_single_process.py::test_cache_transceiver_v1_masked_dsa_indexer_across_asymmetric_pp - unittest/disaggregated/test_cache_transceiver_harness_report.py - unittest/disaggregated/test_cache_transceiver_harness.py - unittest/disaggregated/test_cache_transceiver_precheck_e2e.py diff --git a/tests/unittest/disaggregated/test_bounce.py b/tests/unittest/disaggregated/test_bounce.py index 84a7eda4fe1a..6789c6fefab6 100644 --- a/tests/unittest/disaggregated/test_bounce.py +++ b/tests/unittest/disaggregated/test_bounce.py @@ -221,48 +221,65 @@ def test_make_kv_result_msg_uses_binary_frame(result_name): # --------------------------------------------------------------------------- # -# fan-in safety gate — equal total//num_writers split only for uniform TP-by-head +# fan-in safety gate — equal total//num_writers split only for uniform writers # --------------------------------------------------------------------------- # -def test_fanin_bounce_safe_gate(): +def test_fanin_bounce_safe_gate() -> None: """Restrict multi-writer equal-split bounce to uniform TP-by-head. - PP (overlap_pp_size>1 -> unequal per-writer sizes) and duplicate_head_factor>1 - (MLA / duplicate TP heads -> some ranks don't send KV yet count in - expected_transfers) must fall back to the per-fragment path. + Uneven PP splits and duplicate_head_factor>1 (MLA / duplicate TP heads -> + some ranks don't send KV yet count in expected_transfers) must fall back to + the per-fragment path. """ tfr = pytest.importorskip("tensorrt_llm._torch.disaggregation.native.transfer") from tensorrt_llm._torch.disaggregation.resource.page import MapperKind safe = tfr.Receiver._fanin_bounce_safe - def ov(dup, pp, ranks=(0,)): + def ov(dup: int, pp: int, ranks: tuple[int, ...] = (0,)) -> SimpleNamespace: return SimpleNamespace(duplicate_head_factor=dup, overlap_pp_size=pp, ranks=list(ranks)) - def ri(lpp, page_table=None): + def ri(lpp: list[int], page_table: SimpleNamespace | None = None) -> SimpleNamespace: return SimpleNamespace(layer_num_per_pp=lpp, page_table=page_table) - def pt(mapper_kind): + def pt(mapper_kind: MapperKind) -> SimpleNamespace: view = SimpleNamespace(mapper_kind=mapper_kind) return SimpleNamespace(layer_groups=[SimpleNamespace(pool_views=[view])]) # single PP stage (overlap_pp_size <= 1): only duplicate_head_factor matters - assert safe(ov(1, 1), ri([24])) is True - assert safe(ov(1, 0), ri([24])) is True - assert safe(ov(2, 1), ri([24])) is False # duplicate heads / MLA -> some don't send + assert safe(ov(1, 1), ri([24]), None) is True + assert safe(ov(1, 0), ri([24]), None) is True + assert safe(ov(2, 1), ri([24]), None) is False # duplicate heads / MLA -> some don't send # EVEN PP fan-in (equal layers per overlapping stage) -> allowed - assert safe(ov(1, 4), ri([20, 20, 20, 20])) is True + assert safe(ov(1, 4), ri([20, 20, 20, 20]), None) is True # UNEVEN PP fan-in -> per-writer sizes differ -> fall back - assert safe(ov(1, 4), ri([20, 20, 20, 19])) is False + assert safe(ov(1, 4), ri([20, 20, 20, 19]), None) is False # incomplete per-stage info (single element for a multi-stage fan-in) -> conservative fall back - assert safe(ov(1, 4), ri([20])) is False + assert safe(ov(1, 4), ri([20]), None) is False # duplicate heads blocks even an otherwise-even PP split - assert safe(ov(2, 4), ri([20, 20, 20, 20])) is False + assert safe(ov(2, 4), ri([20, 20, 20, 20]), None) is False # replicated views (one elected sender per destination) make multi-writer # contributions unequal -> fall back; single-writer overlap stays safe, - # and sharded-only view schemes are unaffected - assert safe(ov(1, 1, ranks=(0, 1)), ri([24], pt(MapperKind.REPLICATED))) is False - assert safe(ov(1, 1, ranks=(0,)), ri([24], pt(MapperKind.REPLICATED))) is True - assert safe(ov(1, 1, ranks=(0, 1)), ri([24], pt(MapperKind.NHD))) is True + # and sharded-only view schemes are unaffected. Both page tables must be + # checked because the representative peer rank may be a fully masked PP + # stage while another sender stage owns the replicated rows. + assert safe(ov(1, 1, ranks=(0, 1)), ri([24], pt(MapperKind.REPLICATED)), None) is False + assert safe(ov(1, 1, ranks=(0,)), ri([24], pt(MapperKind.REPLICATED)), None) is True + assert ( + safe( + ov(1, 1, ranks=(0, 1)), + ri([24], pt(MapperKind.NHD)), + pt(MapperKind.REPLICATED), + ) + is False + ) + assert ( + safe( + ov(1, 1, ranks=(0, 1)), + ri([24], pt(MapperKind.NHD)), + pt(MapperKind.NHD), + ) + is True + ) # --------------------------------------------------------------------------- # @@ -369,6 +386,33 @@ def _recv_req(block_counts, rid=1, slice_id=0): @pytest.mark.skipif(not _HAVE_TRANSPORT, reason="bounce.transport import needs CUDA bindings") class TestFanInReserve: + def test_block_bytes_include_distinct_physical_pools_once(self) -> None: + entries = np.array([], dtype=BUFFER_ENTRY_DTYPE) + page_table = KVCachePageTable( + tokens_per_block=32, + layer_groups=[ + AttentionLayerGroup( + pool_group_idx=0, + local_layers=[LocalLayer(local_layer_id=0, global_layer_id=0)], + pool_views=[ + PoolView(pool_idx=0, buffer_entries=entries), + PoolView(pool_idx=1, buffer_entries=entries), + PoolView(pool_idx=1, buffer_entries=entries), + ], + ) + ], + pool_groups=[ + PhysicalPoolGroup( + pools=[ + PhysicalPool(base_address=0x200000, slot_bytes=100, num_slots=8), + PhysicalPool(base_address=0x300000, slot_bytes=25, num_slots=8), + ] + ) + ], + ) + + assert btr.block_bytes_per_group(page_table) == [125] + def test_reserve_stamps_base_and_per_writer(self, monkeypatch): t = _make_transport(monkeypatch, block_bytes_per_group=[100]) req = _recv_req([2]) # total = 2 * 100 = 200 diff --git a/tests/unittest/disaggregated/test_cache_transceiver_single_process.py b/tests/unittest/disaggregated/test_cache_transceiver_single_process.py index 00eaf14a10da..418ea4352634 100644 --- a/tests/unittest/disaggregated/test_cache_transceiver_single_process.py +++ b/tests/unittest/disaggregated/test_cache_transceiver_single_process.py @@ -25,6 +25,9 @@ import uuid from types import SimpleNamespace +# Do not inherit a NIC pin from the host: the selected interface may not exist +# in the test container and would prevent the NIXL agent from initializing. +os.environ.pop("UCX_NET_DEVICES", None) # Exclude UCX IB transport (avoid NIXL setup hangs without IB) and gdr_copy # (avoid SIGSEGV at process exit from UCX rcache cleanup; gdr_copy disabled # falls back to cuda_ipc / cuda_copy without affecting correctness). @@ -326,6 +329,7 @@ def _create_cache_manager( num_layers: int = NUM_LAYERS, max_batch_size: int = MAX_BATCH_SIZE, enable_indexer_k_cache: bool = False, + indexer_k_cache_layer_mask: list[bool] | None = None, ) -> "KVCacheManager | KVCacheManagerV2": """Create a KVCacheManager (V1) or KVCacheManagerV2 for the given mapping.""" assert not (enable_indexer_k_cache and use_v2), "DSA indexer K cache is V1-only" @@ -398,6 +402,7 @@ def _create_cache_manager( enable_indexer_k_cache=enable_indexer_k_cache, indexer_k_cache_quant_block_size=128, indexer_k_cache_index_head_dim=INDEXER_HEAD_DIM if enable_indexer_k_cache else 0, + indexer_k_cache_layer_mask=indexer_k_cache_layer_mask, ) @@ -411,7 +416,8 @@ def _create_managers_for_instance( num_layers: int = NUM_LAYERS, max_batch_size: int = MAX_BATCH_SIZE, enable_indexer_k_cache: bool = False, -) -> List: + indexer_k_cache_layer_mask: list[bool] | None = None, +) -> list[KVCacheManager | KVCacheManagerV2]: """Create cache managers for all ranks in an instance.""" managers = [] for pp_rank in range(pp): @@ -433,6 +439,7 @@ def _create_managers_for_instance( num_layers, max_batch_size, enable_indexer_k_cache, + indexer_k_cache_layer_mask, ) ) return managers @@ -447,7 +454,7 @@ def _init_pool_data_v1( is_mla: bool, fill_random: bool = True, seed_base: int = 1000, -): +) -> None: """Initialize pool data for V1 managers.""" num_kv_heads = 1 if is_mla else NUM_KV_HEADS for rank, mgr in enumerate(managers): @@ -468,6 +475,9 @@ def _init_pool_data_v1( pool_tensor.zero_() if getattr(mgr, "enable_indexer_k_cache", False): + local_mask = mgr.indexer_k_cache_local_layer_mask + if local_mask is not None and not any(local_mask): + continue # DSA indexer K is TP-replicated: seed by PP stage only so every # TP rank of a stage holds identical bytes. indexer_pool = mgr.impl.get_indexer_k_cache_pool().view(torch.uint8) @@ -865,7 +875,16 @@ def verify_all_requests( ) -def _get_indexer_block_data(mgr, request_id, layer_idx, num_layers, pp, tp, enable_dp, req_idx): +def _get_indexer_block_data( + managers: list[KVCacheManager], + request_id: int, + layer_idx: int, + num_layers: int, + pp: int, + tp: int, + enable_dp: bool, + req_idx: int, +) -> torch.Tensor | None: """Per-request indexer-K bytes for one global layer on the owning rank. Indexer K is TP-replicated, so any TP rank of the layer's PP stage works; @@ -873,21 +892,26 @@ def _get_indexer_block_data(mgr, request_id, layer_idx, num_layers, pp, tp, enab """ pp_rank = _pp_rank_of_layer(layer_idx, num_layers, pp) tp_rank = req_idx % tp if enable_dp else 0 - owner = mgr[pp_rank * tp + tp_rank] + owner = managers[pp_rank * tp + tp_rank] block_indices = owner.get_batch_cache_indices([request_id], layer_idx)[0] valid = [idx for idx in block_indices if idx >= 0] if not valid: return None local_layer = layer_idx - _pp_layer_start(pp_rank, num_layers, pp) + local_mask = owner.indexer_k_cache_local_layer_mask + if local_mask is not None and not local_mask[local_layer]: + return None # Pool shape: (numBlocks, numLayers, kvFactor, blockSize), dtype uint8. pool = owner.impl.get_indexer_k_cache_pool().view(torch.uint8) - return pool[valid, local_layer] + pool_layer = owner.impl.get_indexer_k_cache_pool_layer_idx(local_layer) + assert pool_layer >= 0 + return pool[valid, pool_layer] def _verify_indexer_k_all_requests( request_lengths: List[int], - ctx_managers: List, - gen_managers: List, + ctx_managers: list[KVCacheManager], + gen_managers: list[KVCacheManager], ctx_tp: int, ctx_pp: int, gen_tp: int, @@ -897,7 +921,7 @@ def _verify_indexer_k_all_requests( ctx_request_ids: List[int], gen_request_ids: List[int], num_layers: int, -): +) -> None: """Compare the transferred DSA indexer K bytes for every request/layer.""" for req_idx, _req_len in enumerate(request_lengths): for layer_idx in range(num_layers): @@ -921,7 +945,11 @@ def _verify_indexer_k_all_requests( gen_enable_dp, req_idx, ) - if ctx_data is None or gen_data is None: + assert (ctx_data is None) == (gen_data is None), ( + f"Indexer ownership mismatch at req={req_idx} layer={layer_idx}: " + f"ctx_present={ctx_data is not None} gen_present={gen_data is not None}" + ) + if ctx_data is None: continue assert ctx_data.shape == gen_data.shape, ( f"Indexer shape mismatch at req={req_idx} layer={layer_idx}: " @@ -952,7 +980,8 @@ def run_transfer_test( num_layers: int = NUM_LAYERS, request_lengths: Optional[List[int]] = None, enable_indexer_k_cache: bool = False, -): + indexer_k_cache_layer_mask: list[bool] | None = None, +) -> None: """Run a full KV transfer test using KvCacheTransceiverV2.""" if request_lengths is None: request_lengths = REQUEST_LENGTHS @@ -971,6 +1000,7 @@ def run_transfer_test( num_layers, max_batch_size, enable_indexer_k_cache, + indexer_k_cache_layer_mask, ) gen_managers = _create_managers_for_instance( gen_tp, @@ -982,6 +1012,7 @@ def run_transfer_test( num_layers, max_batch_size, enable_indexer_k_cache, + indexer_k_cache_layer_mask, ) # 2. Initialize data: random for ctx, zeros for gen @@ -1493,6 +1524,29 @@ def test_cache_transceiver_v1_dsa_indexer( ) +@pytest.mark.timeout(180) +def test_cache_transceiver_v1_masked_dsa_indexer_across_asymmetric_pp() -> None: + """Transfer a masked DSA indexer cache from CTX PP2 to GEN PP1. + + CTX rank 0 is fully masked while rank 1 owns both indexer rows. This + exercises the real ``KvCacheTransceiverV2`` path where the representative + sender page table has no REPLICATED view but the receiver does. + """ + run_transfer_test( + ctx_tp=1, + ctx_pp=2, + gen_tp=1, + gen_pp=1, + ctx_enable_dp=False, + gen_enable_dp=False, + is_mla=True, + use_v2=False, + request_lengths=[30], + enable_indexer_k_cache=True, + indexer_k_cache_layer_mask=[False, False, True, True], + ) + + if __name__ == "__main__": # Quick smoke test run_transfer_test(1, 1, 1, 1, False, False, False, False) diff --git a/tests/unittest/disaggregated/test_extractor.py b/tests/unittest/disaggregated/test_extractor.py index ce136cf5bd9e..1c96ee972df7 100644 --- a/tests/unittest/disaggregated/test_extractor.py +++ b/tests/unittest/disaggregated/test_extractor.py @@ -208,8 +208,17 @@ def test_build_page_table(): manager.shutdown() -def _make_v1_dsa_manager(pp_size: int = 1, pp_rank: int = 0) -> KVCacheManager: - """V1 KVCacheManager with the DSA indexer K cache enabled (MLA-style).""" +def _make_v1_dsa_manager( + pp_size: int = 1, + pp_rank: int = 0, + indexer_k_cache_layer_mask: list[bool] | None = None, +) -> KVCacheManager: + """V1 KVCacheManager with the DSA indexer K cache enabled (MLA-style). + + ``indexer_k_cache_layer_mask`` is a global per-model ``list[bool]`` marking + the "full" indexer-owning layers (cross-layer indexer sharing, e.g. GLM + 5.2); ``None`` keeps the dense layout where every layer owns an indexer row. + """ return KVCacheManager( kv_cache_config=KvCacheConfig( max_tokens=512, @@ -228,6 +237,7 @@ def _make_v1_dsa_manager(pp_size: int = 1, pp_rank: int = 0) -> KVCacheManager: enable_indexer_k_cache=True, indexer_k_cache_quant_block_size=128, indexer_k_cache_index_head_dim=128, + indexer_k_cache_layer_mask=indexer_k_cache_layer_mask, ) @@ -264,8 +274,67 @@ def test_v1_dsa_indexer_page_table_is_replicated_with_per_layer_entries(): @pytest.mark.cuda -def test_v1_dsa_indexer_replicated_transfer_across_pp(): - """PP1 ctx sends the DSA indexer K cache into two PP2 gen ranks. +def test_v1_dsa_masked_indexer_page_table_covers_owning_layers() -> None: + """The masked indexer view covers only the owning layers. + + A per-layer indexer mask (cross-layer indexer sharing, e.g. GLM 5.2) gives + only the owning layers a pool row, so the REPLICATED indexer view covers + exactly that subset -- one entry per owning layer mapped to its packed row + -- instead of one entry per LG layer. + """ + # Of the 4 layers, only local layers 0 and 2 own an indexer K cache row. + manager = _make_v1_dsa_manager(indexer_k_cache_layer_mask=[True, False, True, False]) + try: + page_table = build_page_table(manager) + lg = page_table.layer_groups[0] + assert len(lg.pool_views) == 2 + _, idx_view = lg.pool_views + assert idx_view.mapper_kind == MapperKind.REPLICATED + assert idx_view.pool_role == frozenset({"indexer_k"}) + + # Only the two owning layers appear -- a strict subset of the LG. + owning = sorted(int(e["local_layer_id"]) for e in idx_view.buffer_entries) + assert owning == [0, 2] + + # The pool holds one row per owning layer; entries pack contiguously in + # owning (local-layer) order, so layer 0 -> row 0, layer 2 -> row 1. + idx_pool = get_physical_pool(page_table, 0, idx_view.pool_idx) + sizes = {int(e["size"]) for e in idx_view.buffer_entries} + assert len(sizes) == 1 + per_layer = sizes.pop() + assert per_layer * len(idx_view.buffer_entries) == idx_pool.slot_bytes + offset_by_layer = { + int(e["local_layer_id"]): int(e["offset"]) for e in idx_view.buffer_entries + } + assert offset_by_layer == {0: 0, 2: per_layer} + finally: + manager.shutdown() + + +@pytest.mark.cuda +def test_v1_dsa_fully_masked_indexer_group_omits_pool() -> None: + """A PP stage without an indexer-owning layer advertises no indexer pool.""" + manager = _make_v1_dsa_manager(indexer_k_cache_layer_mask=[False] * 4) + try: + page_table = build_page_table(manager) + assert len(page_table.layer_groups[0].pool_views) == 1 + assert len(page_table.pool_groups[0].pools) == 1 + finally: + manager.shutdown() + + +@pytest.mark.cuda +@pytest.mark.parametrize( + "indexer_k_cache_layer_mask", + [ + pytest.param(None, id="dense"), + pytest.param([True, False, True, False], id="masked"), + ], +) +def test_v1_dsa_indexer_replicated_transfer_across_pp( + indexer_k_cache_layer_mask: list[bool] | None, +) -> None: + """PP1 ctx sends a dense or masked DSA indexer K cache into two PP2 gen ranks. Exercises the full python path on real V1 managers: page-table build, role-set matching, ReplicatedMapper layer-strided offsets, and a @@ -278,8 +347,15 @@ def test_v1_dsa_indexer_replicated_transfer_across_pp(): from tensorrt_llm._torch.disaggregation.native.peer import PeerRegistrar from tensorrt_llm._torch.disaggregation.native.rank_info import RankInfo - ctx = _make_v1_dsa_manager() - gens = [_make_v1_dsa_manager(pp_size=2, pp_rank=r) for r in range(2)] + ctx = _make_v1_dsa_manager(indexer_k_cache_layer_mask=indexer_k_cache_layer_mask) + gens = [ + _make_v1_dsa_manager( + pp_size=2, + pp_rank=r, + indexer_k_cache_layer_mask=indexer_k_cache_layer_mask, + ) + for r in range(2) + ] try: ctx_extractor = KVRegionExtractorV1(ctx) ctx_ri = RankInfo.from_kv_cache_manager("ctx", ctx, device_id=0) @@ -295,8 +371,10 @@ def test_v1_dsa_indexer_replicated_transfer_across_pp(): block_ids = np.array([0, 2, 3], dtype=np.int64) ctx_pt = ctx_extractor.page_table idx_pool = get_physical_pool(ctx_pt, 0, 1) - layers_per_gen = 2 - per_layer = idx_pool.slot_bytes // 4 + num_indexer_layers = int(ctx_pool_tensor.shape[1]) + assert num_indexer_layers % len(gens) == 0 + layers_per_gen = num_indexer_layers // len(gens) + per_layer = idx_pool.slot_bytes // num_indexer_layers for gen_pp_rank, gen in enumerate(gens): gen_ri = RankInfo.from_kv_cache_manager("gen", gen, device_id=0) @@ -313,9 +391,8 @@ def test_v1_dsa_indexer_replicated_transfer_across_pp(): ctx_extractor.extract(block_ids, layer_group_id=0, pool_idx=1), gen_extractor.extract(block_ids, layer_group_id=0, pool_idx=1), ) - # This gen rank holds 2 of the 4 layers; the fragment is the - # contiguous 2-layer range at this PP stage's offset within - # the ctx slot. + # Each gen rank holds a contiguous subset of the owning layers; + # masked-out layers do not consume bytes in either pool. assert pair.src.memory.bytes_per_region == layers_per_gen * per_layer expected_src_off = gen_pp_rank * layers_per_gen * per_layer base_ptrs = idx_pool.base_address + block_ids * idx_pool.slot_bytes diff --git a/tests/unittest/llmapi/test_llm_args.py b/tests/unittest/llmapi/test_llm_args.py index ceef1c8de786..3c2e9fb76e2b 100644 --- a/tests/unittest/llmapi/test_llm_args.py +++ b/tests/unittest/llmapi/test_llm_args.py @@ -4042,10 +4042,8 @@ def test_resolve_default_backend_env_priority(self, monkeypatch): class TestDeepseekRuntimePreferences: """DeepSeek KV-cache manager and transceiver preferences. - DeepseekV3ForCausalLM and DeepseekV32ForCausalLM prefer the Python KV-cache - transceiver, while GlmMoeDsaForCausalLM (GLM 5.2) requires the C++ transceiver - because its per-layer masked DSA indexer k-cache pool is not supported by the - Python (v2) transceiver. + DeepseekV3ForCausalLM, DeepseekV32ForCausalLM, and GlmMoeDsaForCausalLM + prefer the Python KV-cache transceiver. """ @staticmethod @@ -4055,19 +4053,32 @@ def _pretrained_config(architectures, model_type): cfg.model_type = model_type return cfg - @pytest.mark.parametrize("architectures,model_type,expected", [ - (["GlmMoeDsaForCausalLM"], "glm_moe_dsa", "CPP"), - (["DeepseekV3ForCausalLM"], "deepseek_v3", "PYTHON"), - (["DeepseekV32ForCausalLM"], "deepseek_v32", "PYTHON"), + @pytest.mark.parametrize("architectures,model_type", [ + (["GlmMoeDsaForCausalLM"], "glm_moe_dsa"), + (["DeepseekV3ForCausalLM"], "deepseek_v3"), + (["DeepseekV32ForCausalLM"], "deepseek_v32"), ]) def test_preference_per_architecture(self, architectures: list[str], - model_type: str, - expected: str) -> None: - from tensorrt_llm._torch.models.modeling_deepseekv3 import \ - DeepseekV3ForCausalLM + model_type: str) -> None: + from tensorrt_llm._torch.models.modeling_utils import \ + get_registered_model_class cfg = self._pretrained_config(architectures, model_type) - assert DeepseekV3ForCausalLM.get_preferred_transceiver_runtime( - cfg) == expected + model_cls = get_registered_model_class(architectures[0]) + assert model_cls is not None + assert model_cls.get_preferred_transceiver_runtime(cfg) == "PYTHON" + + def test_glm_uses_independent_model_flavor(self) -> None: + """GLM can evolve independently from its DeepSeek implementation base.""" + from tensorrt_llm._torch.models.modeling_deepseekv3 import ( + DeepseekV3ForCausalLM, GlmMoeDsaForCausalLM) + from tensorrt_llm._torch.models.modeling_utils import \ + get_registered_model_class + + assert get_registered_model_class( + "DeepseekV3ForCausalLM") is DeepseekV3ForCausalLM + assert get_registered_model_class( + "GlmMoeDsaForCausalLM") is GlmMoeDsaForCausalLM + assert issubclass(GlmMoeDsaForCausalLM, DeepseekV3ForCausalLM) def test_prefers_python_without_config(self) -> None: """Preference is unconditional without a pretrained config."""