From f37b456449297d2fe44c1463f4cb9f64b35f680d Mon Sep 17 00:00:00 2001 From: Yao Yao Date: Wed, 15 Jul 2026 13:30:55 +0000 Subject: [PATCH 1/2] [None][fix] disagg: derive KV transfer layer offset from physical slot order PeerRegistrar.get_kv_map() sorted the pool's global layer ids before computing each layer's slot offset, relying on the convention that managers assign global_layer_id monotonically with the layer's byte offset. That discards the physical-offset order recovered by get_pool_view_global_layer_ids() and, for a non-monotonic layout, maps a layer to the wrong slot and transfers the wrong KV bytes. Use the physical slot order directly (buffer_entries offsets for INDEXED pools, local_layers packing order for FLAT pools) so a layer's index is its slot offset, and anchor the transfer on the physically-first overlapping layer. Replace the implicit monotonicity invariant with an explicit check that the shared layers form an aligned contiguous slot run on both peers, raising instead of silently corrupting the transfer. Add unit tests covering a non-monotonic layout and the non-contiguous-overlap rejection. Signed-off-by: Yao Yao --- .../_torch/disaggregation/native/peer.py | 49 ++++++++++++----- tests/unittest/disaggregated/test_peer.py | 52 +++++++++++++++++++ 2 files changed, 87 insertions(+), 14 deletions(-) diff --git a/tensorrt_llm/_torch/disaggregation/native/peer.py b/tensorrt_llm/_torch/disaggregation/native/peer.py index 1902c9d14f4f..98ef88a4eeb4 100644 --- a/tensorrt_llm/_torch/disaggregation/native/peer.py +++ b/tensorrt_llm/_torch/disaggregation/native/peer.py @@ -240,31 +240,52 @@ def get_kv_map( f"(local={self_pv.mapper_kind.name}, peer={peer_pv.mapper_kind.name})" ) - # FLAT pools carry no per-buffer layer info, so layer ids and - # layer count come from the layer_group itself. - # - # Sort by global_layer_id so that ``.index(first_overlap_layer)`` - # below returns the layer's slot position. This relies on the - # convention that managers (V1 / V2 / DSv4) assign global_layer_id - # monotonically with the layer's byte offset in the slot. + # Order both layer-id lists by physical slot position so that a layer's + # index in the list *is* its slot offset. The KV transfer maps layers + # positionally (byte offset = index * per-layer stride) and the mappers + # copy one contiguous ``[offset, offset + transfer_layers)`` fragment, so + # the order must reflect the actual physical layout. We derive it from + # the KV-cache manager's layout rather than assuming ``global_layer_id`` + # is monotonic with the layer's byte offset in the slot: + # INDEXED: ``get_pool_view_global_layer_ids`` orders layers by their + # ``buffer_entries`` offsets (the V2 pool-view layout). + # FLAT: the pool has no per-buffer layer info; it packs the whole + # layer group equal-sized in ``local_layers`` order, so that + # order already *is* the physical order. if self_pv.mapper_kind == MapperKind.FLAT: - self_global_ids = sorted(get_global_layer_ids(self_lg)) - peer_global_ids = sorted(get_global_layer_ids(peer_lg)) + self_global_ids = get_global_layer_ids(self_lg) + peer_global_ids = get_global_layer_ids(peer_lg) self_num_layers = get_layer_group_num_layers(self_lg) peer_num_layers = get_layer_group_num_layers(peer_lg) else: - self_global_ids = sorted(get_pool_view_global_layer_ids(self_pv, self_lg)) - peer_global_ids = sorted(get_pool_view_global_layer_ids(peer_pv, peer_lg)) + self_global_ids = get_pool_view_global_layer_ids(self_pv, self_lg) + peer_global_ids = get_pool_view_global_layer_ids(peer_pv, peer_lg) self_num_layers = get_pool_view_num_layers(self_pv) peer_num_layers = get_pool_view_num_layers(peer_pv) - overlapping_layers = sorted(set(self_global_ids) & set(peer_global_ids)) - transfer_layers = len(overlapping_layers) + overlap = set(self_global_ids) & set(peer_global_ids) + transfer_layers = len(overlap) if transfer_layers > 0: - first_overlap_layer = overlapping_layers[0] + # Anchor on the overlap layer that comes first in self's physical + # order and locate the *same* global layer in peer's physical order. + # Since the mapper copies a single contiguous fragment, the shared + # layers must occupy an aligned, contiguous run of slots on both + # peers. Validate that here instead of relying on a global-layer-id + # ordering convention and silently transferring the wrong bytes. + first_overlap_layer = next(g for g in self_global_ids if g in overlap) self_layer_offset = self_global_ids.index(first_overlap_layer) peer_layer_offset = peer_global_ids.index(first_overlap_layer) + self_run = self_global_ids[self_layer_offset : self_layer_offset + transfer_layers] + peer_run = peer_global_ids[peer_layer_offset : peer_layer_offset + transfer_layers] + if set(self_run) != overlap or self_run != peer_run: + raise ValueError( + "PeerRegistrar.get_kv_map: shared layers do not form an " + "aligned contiguous run of physical slots on both peers " + f"(self={self_global_ids}, peer={peer_global_ids}, " + f"overlap={sorted(overlap)}); the KV transfer requires shared " + "layers to occupy matching contiguous slot ranges." + ) else: self_layer_offset = 0 peer_layer_offset = 0 diff --git a/tests/unittest/disaggregated/test_peer.py b/tests/unittest/disaggregated/test_peer.py index d2329d74e1aa..8ddfdcec3d42 100644 --- a/tests/unittest/disaggregated/test_peer.py +++ b/tests/unittest/disaggregated/test_peer.py @@ -404,6 +404,58 @@ def test_peer_registrar_get_kv_map_head_mismatch(): assert isinstance(mapper, HeadMismatchMapper) +def test_peer_registrar_get_kv_map_uses_physical_offset_order(): + """The layer slot offset must follow physical buffer order, not sorted global id. + + ``self`` lays out its two global layers in physical order ``[5, 3]`` (layer 3 + sits at physical slot 1); ``peer`` holds only layer 3. The transfer must copy + ``self``'s slot at offset 1 -- a sort-by-global-id would wrongly pick offset 0 + (layer 5's bytes). + """ + self_pt = make_page_table(global_layer_ids=[5, 3]) + peer_pt = make_page_table(global_layer_ids=[3], block_bytes=[512]) + + self_rankinfo = make_rankinfo(instance_name="local", page_table=self_pt) + reg = _make_peer_registrar(self_rankinfo) + peer_ri = make_rankinfo( + instance_name="peer", + instance_rank=2, + layer_num_per_pp=[1], + page_table=peer_pt, + ) + reg.register(peer_ri.instance_name, peer_ri.instance_rank, peer_ri) + mapper = reg.get_kv_map(peer_ri, (0, 0), (0, 0)) + assert isinstance(mapper, HeadMatchMapper) + # slot_size_per_layer = 1024 // 2 = 512; layer 3 is at physical slot 1 on + # self and slot 0 on peer. + assert mapper._src_block_off == 512 + assert mapper._dst_block_off == 0 + + +def test_peer_registrar_get_kv_map_rejects_non_contiguous_overlap(): + """Shared layers that are not an aligned contiguous slot run must be rejected. + + ``self`` holds layers ``[0, 1, 2]`` and ``peer`` holds ``[0, 2]``: the overlap + ``{0, 2}`` is not contiguous within ``self``'s slot, so a single contiguous + fragment transfer would corrupt layer 1's bytes. ``get_kv_map`` must raise + rather than emit a wrong mapping. + """ + self_pt = make_page_table(global_layer_ids=[0, 1, 2]) + peer_pt = make_page_table(global_layer_ids=[0, 2]) + + self_rankinfo = make_rankinfo(instance_name="local", page_table=self_pt) + reg = _make_peer_registrar(self_rankinfo) + peer_ri = make_rankinfo( + instance_name="peer", + instance_rank=2, + layer_num_per_pp=[2], + page_table=peer_pt, + ) + reg.register(peer_ri.instance_name, peer_ri.instance_rank, peer_ri) + with pytest.raises(ValueError, match="aligned contiguous run"): + reg.get_kv_map(peer_ri, (0, 0), (0, 0)) + + def test_peer_registrar_tpb_divisible_warns_but_compatible(): # local=16, peer=32: 32 % 16 == 0 → compatible with warning, register succeeds self_rankinfo = make_rankinfo(instance_name="local", tokens_per_block=16) From d7c7030943b2a9fbdb3e49d8b894dbea41f45d16 Mon Sep 17 00:00:00 2001 From: Yao Yao Date: Fri, 17 Jul 2026 10:34:42 +0000 Subject: [PATCH 2/2] [None][fix] explicitly handle MapperKind.INDEXED in get_kv_map Address review nit: branch explicitly on MapperKind.INDEXED instead of relying on else, and raise ValueError for unexpected mapper kinds. Signed-off-by: Yao Yao --- tensorrt_llm/_torch/disaggregation/native/peer.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/disaggregation/native/peer.py b/tensorrt_llm/_torch/disaggregation/native/peer.py index 98ef88a4eeb4..0762782ebab9 100644 --- a/tensorrt_llm/_torch/disaggregation/native/peer.py +++ b/tensorrt_llm/_torch/disaggregation/native/peer.py @@ -257,11 +257,15 @@ def get_kv_map( peer_global_ids = get_global_layer_ids(peer_lg) self_num_layers = get_layer_group_num_layers(self_lg) peer_num_layers = get_layer_group_num_layers(peer_lg) - else: + elif self_pv.mapper_kind == MapperKind.INDEXED: self_global_ids = get_pool_view_global_layer_ids(self_pv, self_lg) peer_global_ids = get_pool_view_global_layer_ids(peer_pv, peer_lg) self_num_layers = get_pool_view_num_layers(self_pv) peer_num_layers = get_pool_view_num_layers(peer_pv) + else: + raise ValueError( + f"PeerRegistrar.get_kv_map: unexpected mapper kind {self_pv.mapper_kind!r}" + ) overlap = set(self_global_ids) & set(peer_global_ids) transfer_layers = len(overlap)