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
Original file line number Diff line number Diff line change
Expand Up @@ -335,7 +335,8 @@ def get_draft_subpage_view(self) -> Optional["MiniMaxM3DraftSubpageView"]:
)
logger.info(
f"[unified-kv] draft sub-page view active "
f"(tokens_per_block={self._draft_subpage_view_obj.tokens_per_block})"
f"(tokens_per_block={self._draft_subpage_view_obj.tokens_per_block}, "
f"flat_page_bound={self._draft_subpage_view_obj.blocks_in_primary_pool})"
)
return self._draft_subpage_view_obj

Expand Down Expand Up @@ -698,6 +699,7 @@ def __init__(self, manager, draft_layer_ids: Sequence[int], subpage_tokens: int)
f"{manager.kv_cache_pool_mapping[int(local)].tolist()}"
)
addr_key, _dt, num_slots, scale, _shape = manager._kv_slot_geometry(layer_id, None)
self._num_slots = num_slots
self._slot_units = scale * self._subdiv # slot stride in 32-tok units
self.num_pools = 1
self.num_attention_op_pools = 1
Expand All @@ -719,6 +721,20 @@ def __init__(self, manager, draft_layer_ids: Sequence[int], subpage_tokens: int)
self._slots_host: Optional[np.ndarray] = None
self._arange: Optional[torch.Tensor] = None

@property
def blocks_in_primary_pool(self) -> int:
"""Flattened sub-page index bound relative to the draft K pointer.

``FlashInferTrtllmGenFmha`` uses this value to size the flat paged-KV
tensor passed to FlashInfer. The wrapped V2 manager reports its bound
in 128-token page units relative to a different pool base, so
delegating that property through ``__getattr__`` under-describes this
32-token, draft-K-rooted view. The final slot contributes only this
layer's K and V pages; inter-layer padding after V is not addressable
from the view and need not be included.
"""
return (self._num_slots - 1) * self._slot_units + 2 * self._subdiv

def __getattr__(self, name):
manager = self.__dict__.get("_manager")
if manager is None:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,9 @@ class _FakeManager:
tokens_per_block = 128
max_blocks_per_seq = 16
num_pools = 1
# V2's flattened bound uses the target pool's 128-token page units and
# base pointer. The draft view must not delegate this value.
blocks_in_primary_pool = 1024 * SCALE

def __init__(self):
self.layer_offsets = {DRAFT_LAYER: DRAFT_LAYER}
Expand Down Expand Up @@ -70,6 +73,11 @@ def test_view_geometry():
assert view.kv_cache_pool_pointers.tolist() == [[ADDR, 0]]
# The draft layer's mapping row is rewritten to the view's single pool.
assert view.kv_cache_pool_mapping[DRAFT_LAYER].tolist() == [0, 0]
# FlashInfer wraps the pool as a flat tensor rooted at the draft K
# pointer. Its upper bound must use 32-token sub-page units and stop after
# the last slot's V pages, not delegate the target manager's 128-token
# page bound.
assert view.blocks_in_primary_pool == (1024 - 1) * SCALE * 4 + 8


def test_block_table_expansion():
Expand Down Expand Up @@ -127,6 +135,7 @@ def test_subdiv_one_degenerates_to_identity():
assert view._subdiv == 1
assert view.tokens_per_block == 128
assert view.max_blocks_per_seq == 16
assert view.blocks_in_primary_pool == (1024 - 1) * SCALE + 2
dst = torch.zeros((1, 1, 2, view.max_blocks_per_seq), dtype=torch.int32)
view.copy_batch_block_offsets(dst, request_ids=[1], beam_width=1, num_contexts=1, num_seqs=1)
assert dst[0, 0, 0, :2].tolist() == [5 * SCALE, 7 * SCALE]
Expand Down
Loading