Skip to content
Closed
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
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
# limitations under the License.

from collections.abc import Sequence
from typing import Optional

import numpy as np

Expand Down Expand Up @@ -177,9 +178,15 @@ def __init__(
peer_bytes_per_layer: int,
self_buffers_per_layer: int,
peer_buffers_per_layer: int,
self_kv_heads: Optional[int] = None,
peer_kv_heads: Optional[int] = None,
):
self._ri = self_ri
self._peer_ri = peer_ri
# Head counts are per layer group (a merged draft group differs from
# the target's); None falls back to the rank-level value.
self_kv_heads = self_kv_heads or self_ri.attention.kv_heads_per_rank
peer_kv_heads = peer_kv_heads or peer_ri.attention.kv_heads_per_rank

self_tp_per_dp = self_ri.tp_size_per_dp_group
peer_tp_per_dp = peer_ri.tp_size_per_dp_group
Expand Down Expand Up @@ -212,28 +219,21 @@ def __init__(
buffers_per_layer=peer_buffers_per_layer,
side="peer",
)
bytes_per_head = self._bytes_per_head(
src_buffer_bytes, self._ri.attention.kv_heads_per_rank, side="local"
)
peer_bytes_per_head = self._bytes_per_head(
dst_buffer_bytes, peer_ri.attention.kv_heads_per_rank, side="peer"
)
bytes_per_head = self._bytes_per_head(src_buffer_bytes, self_kv_heads, side="local")
peer_bytes_per_head = self._bytes_per_head(dst_buffer_bytes, peer_kv_heads, side="peer")
if bytes_per_head != peer_bytes_per_head:
raise ValueError(
f"HND bytes per head mismatch: local={bytes_per_head}, peer={peer_bytes_per_head}"
)
self._bytes_cont_heads = (
min(self._ri.attention.kv_heads_per_rank, peer_ri.attention.kv_heads_per_rank)
* bytes_per_head
)
self._bytes_cont_heads = min(self_kv_heads, peer_kv_heads) * bytes_per_head

self._src_head_off, self._dst_head_off = self._compute_head_offsets(
self_tp_per_dp,
peer_tp_per_dp,
self_tp_rank,
peer_tp_rank,
self_kv_heads=self._ri.attention.kv_heads_per_rank,
peer_kv_heads=peer_ri.attention.kv_heads_per_rank,
self_kv_heads=self_kv_heads,
peer_kv_heads=peer_kv_heads,
bytes_per_head=bytes_per_head,
)

Expand Down Expand Up @@ -363,6 +363,8 @@ def __init__(
peer_bytes_per_layer: int,
self_buffers_per_layer: int,
peer_buffers_per_layer: int,
self_kv_heads: Optional[int] = None,
peer_kv_heads: Optional[int] = None,
) -> None:
# Deliberately do not call HNDHeadMismatchMapper.__init__: its offsets
# assume HND-contiguous heads. Initialize the three attributes consumed
Expand All @@ -375,8 +377,8 @@ def __init__(
f"local={self_tpb}, peer={peer_tpb}"
)

self_heads = self_ri.attention.kv_heads_per_rank
peer_heads = peer_ri.attention.kv_heads_per_rank
self_heads = self_kv_heads or self_ri.attention.kv_heads_per_rank
peer_heads = peer_kv_heads or peer_ri.attention.kv_heads_per_rank
if self_buffers_per_layer != peer_buffers_per_layer:
raise ValueError(
"NHD buffer count per layer mismatch: "
Expand Down Expand Up @@ -608,6 +610,8 @@ def build_kv_mapper(
peer_bytes_per_layer: int,
self_buffers_per_layer: int = 1,
peer_buffers_per_layer: int = 1,
self_kv_heads: Optional[int] = None,
peer_kv_heads: Optional[int] = None,
) -> RegionMapperBase:
"""Pick the mapper for one view pair.

Expand All @@ -633,7 +637,11 @@ def build_kv_mapper(
peer_bytes_per_layer,
)

head_match, _ = self.head_match(peer_ri)
# Per-group head counts (merged draft groups differ from the target's
# rank-level value); None falls back to rank-level.
group_self_heads = self_kv_heads or self._ri.attention.kv_heads_per_rank
group_peer_heads = peer_kv_heads or peer_ri.attention.kv_heads_per_rank
head_match = self._ri.attention.is_mla or group_self_heads == group_peer_heads
if head_match:
return IntactMapper(
self_layer_offsets,
Expand All @@ -653,6 +661,8 @@ def build_kv_mapper(
peer_bytes_per_layer=peer_bytes_per_layer,
self_buffers_per_layer=self_buffers_per_layer,
peer_buffers_per_layer=peer_buffers_per_layer,
self_kv_heads=group_self_heads,
peer_kv_heads=group_peer_heads,
)

return HNDHeadMismatchMapper(
Expand All @@ -664,4 +674,6 @@ def build_kv_mapper(
peer_bytes_per_layer=peer_bytes_per_layer,
self_buffers_per_layer=self_buffers_per_layer,
peer_buffers_per_layer=peer_buffers_per_layer,
self_kv_heads=group_self_heads,
peer_kv_heads=group_peer_heads,
)
11 changes: 11 additions & 0 deletions tensorrt_llm/_torch/disaggregation/native/peer.py
Original file line number Diff line number Diff line change
Expand Up @@ -342,6 +342,15 @@ def get_kv_map(
pool_idx=peer_pi,
)

# Head counts are a per-layer-group property (a merged draft group's
# head count differs from the target's); fall back to the rank-level
# value for tables that predate per-group head metadata.
self_kv_heads = (
getattr(self_lg, "kv_head_num_per_rank", 0) or self._ri.attention.kv_heads_per_rank
)
peer_kv_heads = (
getattr(peer_lg, "kv_head_num_per_rank", 0) or peer_ri.attention.kv_heads_per_rank
)
mapper = self._attention_policy.build_kv_mapper(
peer_ri=peer_ri,
mapper_kind=self_pv.mapper_kind,
Expand All @@ -351,6 +360,8 @@ def get_kv_map(
peer_bytes_per_layer=peer_bytes_per_layer,
self_buffers_per_layer=self_buffers_per_layer,
peer_buffers_per_layer=peer_buffers_per_layer,
self_kv_heads=self_kv_heads,
peer_kv_heads=peer_kv_heads,
)

self._kv_map_cache[cache_key] = mapper
Expand Down
81 changes: 72 additions & 9 deletions tensorrt_llm/_torch/disaggregation/native/transfer.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,8 +62,12 @@
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.kv_extractor import (
KVRegionExtractorV1,
build_page_table_from_manager,
merge_draft_page_table,
)
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
Expand Down Expand Up @@ -809,9 +813,13 @@ def _build_kv_write_meta(self, task: KVSendTask, req_info: RecvReqInfo) -> Write
f"src/dst block count mismatch: {src_block_ids.size} vs "
f"{dst_block_ids.size} (dst must not exceed src)"
)
tpb = extractor.page_table.tokens_per_block
token_range = task._slice.token_range
lg_info = extractor.page_table.layer_groups[self_lg]
# Merged draft groups carry their own page size; fall back to the
# table-wide value for regular target groups.
tpb = (
getattr(lg_info, "tokens_per_block", None) or extractor.page_table.tokens_per_block
)
window_size = getattr(lg_info, "sliding_window_size", None)

# Block lists are the suffix of [..., slice_end); cached prefix
Expand Down Expand Up @@ -2242,6 +2250,10 @@ class TransferWorkerConfig:
tx_timeout_s: Optional[float] = None
rx_timeout_s: Optional[float] = None
bounce: Optional["Config"] = None
# Separate one-model draft KV cache manager whose (small) prompt KV is
# transferred alongside the target's. Its layer groups are merged into
# the page table with offset global ids and a per-group tokens_per_block.
draft_kv_cache_manager: Optional[KVCacheManager] = None


class TransferWorker:
Expand All @@ -2257,9 +2269,35 @@ def __init__(self, config: TransferWorkerConfig):
config.device_id,
self._aux_buffer.meta if self._aux_buffer is not None else None,
)
assert self._rank_info.page_table is not None
self._num_target_layer_groups = len(self._rank_info.page_table.layer_groups)
self._num_target_pool_groups = len(self._rank_info.page_table.pool_groups)
if config.draft_kv_cache_manager is not None:
if kvm.mapping.pp_size != 1:
raise ValueError(
"Draft KV cache transfer is not supported with pipeline "
"parallelism yet (layer_num_per_pp bookkeeping covers "
"target layers only)."
)
draft_pt = build_page_table_from_manager(config.draft_kv_cache_manager)
self._rank_info.page_table = merge_draft_page_table(
self._rank_info.page_table, draft_pt
)
logger.info(
"Registered separate draft KV cache for disaggregated transfer: "
f"{len(draft_pt.layer_groups)} layer group(s), "
f"tokens_per_block={draft_pt.tokens_per_block} "
f"(target uses {self._rank_info.attention.tokens_per_block})"
)
self._setup_peer_infrastructure(kvm)
self._setup_transfer_engine()

@property
def num_target_layer_groups(self) -> int:
"""Layer groups in the page table that belong to the target manager;
indices >= this are merged draft-manager groups."""
return self._num_target_layer_groups

def populate_instance_and_rank_info(self, endpoints: list[str], layer_num_per_pp: list[int]):
assert self._rank_info is not None
self._rank_info.sender_endpoints = endpoints
Expand Down Expand Up @@ -2300,7 +2338,9 @@ def sweep_stale_req_infos(self):

def _setup_peer_infrastructure(self, kvm: KVCacheManager):
self._rank_info_server = RankInfoServer(self._rank_info) if kvm.mapping.rank == 0 else None
self._kv_extractor = KVRegionExtractorV1(kvm)
# Build the extractor from the (possibly draft-merged) page table so
# extraction covers every registered pool, not just the target's.
self._kv_extractor = KVRegionExtractorV1(self._rank_info.page_table)
self._peer_registrar = PeerRegistrar(self._rank_info, self._kv_extractor)

def _setup_transfer_engine(self):
Expand Down Expand Up @@ -2339,13 +2379,36 @@ def _setup_transfer_engine(self):

def _register_kv_cache(self):
assert self._rank_info.page_table is not None
memory_descs = get_unique_pool_memory_descs(
self._rank_info.page_table, self._rank_info.device_id
)
if memory_descs:
pt = self._rank_info.page_table
n_target = self._num_target_layer_groups
# Pools from different managers can use different VMM chunk sizes
# (the V2 manager derives its cuMemCreate granularity from the pool
# quota, and a merged draft manager's quota is much smaller). The
# NIXL agent requires a uniform chunk size within one registration
# batch, while its region bookkeeping is per-region — so register
# the target's and the draft's pools as separate batches.
sub_tables = [
KVCachePageTable(pt.tokens_per_block, pt.layer_groups[:n_target], pt.pool_groups)
]
if n_target < len(pt.layer_groups):
sub_tables.append(
KVCachePageTable(pt.tokens_per_block, pt.layer_groups[n_target:], pt.pool_groups)
)
for batch_idx, sub_table in enumerate(sub_tables):
memory_descs = get_unique_pool_memory_descs(sub_table, self._rank_info.device_id)
if not memory_descs:
continue
if batch_idx > 0:
# Keep descriptor names unique across registration batches.
memory_descs = [
(ptr, size, dev, f"{name}_draft{batch_idx}")
for ptr, size, dev, name in memory_descs
]
reg_memory_desc = RegMemoryDescs("VRAM", memory_descs)
self._agent.register_memory(reg_memory_desc)
logger.debug(f"Registered KV cache memory with transfer agent: {memory_descs}")
logger.debug(
f"Registered KV cache memory batch {batch_idx} with transfer agent: {memory_descs}"
)
self._registered_mem.append(reg_memory_desc)

def _register_aux_buffer(self):
Expand Down
52 changes: 52 additions & 0 deletions tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py
Original file line number Diff line number Diff line change
Expand Up @@ -648,3 +648,55 @@ def build_page_table_from_manager(manager) -> KVCachePageTable:
return _build_page_table_v2(manager)
else:
return build_page_table(manager)


# Offset applied to the draft manager's synthetic global layer ids when its
# layer groups are merged into the target's page table. Keeps them
# collision-free against every target id scheme (plain pp_layers ids and the
# virtual-layer encodings alike), and both peers apply the same constant, so
# peer matching by global-id overlap pairs draft groups with draft groups.
DRAFT_GLOBAL_LAYER_ID_OFFSET = 1 << 30


def merge_draft_page_table(
target_pt: KVCachePageTable,
draft_pt: KVCachePageTable,
) -> KVCachePageTable:
"""Append a separate draft manager's page table to the target's.

The merged table lets the disaggregated transfer machinery treat draft
layers as additional attention layer groups: NIXL registration, peer
matching (by offset global ids + pool_role) and the per-pool mappers all
operate on the merged table. Draft groups carry their own
``tokens_per_block`` (the draft manager's page size may differ from the
target's, e.g. MiniMax-M3's 32 vs 128).
"""
pool_base = len(target_pt.pool_groups)
merged_groups: List[LayerGroup] = list(target_pt.layer_groups)
for lg in draft_pt.layer_groups:
if not isinstance(lg, AttentionLayerGroup):
raise ValueError(
"merge_draft_page_table supports attention draft layer groups "
f"only, got {type(lg).__name__}"
)
merged_groups.append(
AttentionLayerGroup(
pool_group_idx=lg.pool_group_idx + pool_base,
kv_head_num_per_rank=lg.kv_head_num_per_rank,
sliding_window_size=lg.sliding_window_size,
local_layers=[
LocalLayer(
local_layer_id=ll.local_layer_id,
global_layer_id=ll.global_layer_id + DRAFT_GLOBAL_LAYER_ID_OFFSET,
)
for ll in lg.local_layers
],
pool_views=lg.pool_views,
tokens_per_block=draft_pt.tokens_per_block,
)
)
return KVCachePageTable(
tokens_per_block=target_pt.tokens_per_block,
layer_groups=merged_groups,
pool_groups=list(target_pt.pool_groups) + list(draft_pt.pool_groups),
)
8 changes: 8 additions & 0 deletions tensorrt_llm/_torch/disaggregation/resource/page.py
Original file line number Diff line number Diff line change
Expand Up @@ -254,6 +254,11 @@ class AttentionLayerGroup(LayerGroup):
sliding_window_size: Optional[int] = None
local_layers: List[LocalLayer] = field(default_factory=list)
pool_views: List[PoolView] = field(default_factory=list)
# Page size override for groups whose manager uses a different
# tokens_per_block than the table default (e.g. a separate one-model
# draft KV cache manager merged into the target's table). None means
# "use KVCachePageTable.tokens_per_block".
tokens_per_block: Optional[int] = None

def to_dict(self) -> dict:
return {
Expand All @@ -262,16 +267,19 @@ def to_dict(self) -> dict:
"sliding_window_size": self.sliding_window_size,
"local_layers": [ll.to_dict() for ll in self.local_layers],
"pool_views": [pv.to_dict() for pv in self.pool_views],
"tokens_per_block": self.tokens_per_block,
}

@classmethod
def from_dict(cls, data: dict) -> "AttentionLayerGroup":
tokens_per_block = data.get("tokens_per_block")
return cls(
pool_group_idx=int(data["pool_group_idx"]),
kv_head_num_per_rank=int(data["kv_head_num_per_rank"]),
sliding_window_size=data.get("sliding_window_size"),
local_layers=[LocalLayer.from_dict(x) for x in data.get("local_layers", [])],
pool_views=[PoolView.from_dict(pv) for pv in data.get("pool_views", [])],
tokens_per_block=int(tokens_per_block) if tokens_per_block is not None else None,
)


Expand Down
Loading
Loading