Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
55 changes: 40 additions & 15 deletions tensorrt_llm/_torch/disaggregation/native/peer.py
Original file line number Diff line number Diff line change
Expand Up @@ -240,31 +240,56 @@ 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))
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}"
)

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
Expand Down
52 changes: 52 additions & 0 deletions tests/unittest/disaggregated/test_peer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Comment thread
coderabbitai[bot] marked this conversation as resolved.
"""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)
Expand Down
Loading