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
193 changes: 176 additions & 17 deletions tensorrt_llm/_torch/attention_backend/flashinfer.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,21 @@ def _slice_paged_kv_cache_heads(
return paged_kv_cache[tuple(index)]


def _get_page_table_num_blocks(kv_cache_manager, request_ids,
logical_num_blocks: list[int],
num_contexts: int) -> list[int]:
"""Keep context rows logical and expose every reserved generation page."""
generation_request_ids = request_ids[num_contexts:]
reserved_num_blocks = [
len(block_ids) for block_ids in
kv_cache_manager.get_batch_cache_indices(generation_request_ids)
]
return list(logical_num_blocks[:num_contexts]) + [
max(logical, reserved) for logical, reserved in zip(
logical_num_blocks[num_contexts:], reserved_num_blocks)
]


def _append_paged_kv_cache(
append_key: torch.Tensor,
append_value: torch.Tensor,
Expand Down Expand Up @@ -151,8 +166,10 @@ class FlashInferMultiItemParams:

@dataclass(kw_only=True, frozen=True)
class PlanParams:
"""
Parameters that affect the flashinfer execution plan
"""Parameters that affect FlashInfer wrapper planning.

Include values that change wrapper-owned CUDA graph state, even when a backend derives them
again at runtime.
"""

num_heads: int
Expand All @@ -167,6 +184,11 @@ class PlanParams:
sm_scale: Optional[float] = None
window_left: Optional[int] = None
kv_pool_id: Optional[int] = None
# Decode wrappers own persistent graph-visible buffers and counters. The speculative query width
# and generation batch size must distinguish cache entries; reusing a wrapper across either
# dimension can expose stale launch state.
q_len_per_req: int = 1
num_generations: int = 0


# NB: Some features (multi-item scoring) are only supported with the paged KV-cache wrapper.
Expand Down Expand Up @@ -242,6 +264,9 @@ class FlashInferAttentionMetadata(AttentionMetadata):
_ragged_qo_indptr_buf: torch.Tensor = field(init=False)
_ragged_kv_indptr_buf: torch.Tensor = field(init=False)
_cached_token_lens: torch.Tensor = field(init=False)
# Attention-visible KV lengths. These exclude pages reserved for future speculative tokens when
# the structural page table exposes extra capacity.
_logical_kv_lens: torch.Tensor = field(init=False, repr=False)
Comment thread
2ez4bz marked this conversation as resolved.
_plan_params_to_wrappers: Dict[PlanParams,
FlashInferWrappers] = field(init=False)

Expand Down Expand Up @@ -281,6 +306,11 @@ class FlashInferAttentionMetadata(AttentionMetadata):
_uses_full_draft_page_table: bool = field(init=False,
default=False,
repr=False)
# Whether target generation rows expose reserved speculative pages while keeping their
# attention-visible lengths in `_logical_kv_lens`.
_uses_full_generation_page_table: bool = field(init=False,
Comment thread
2ez4bz marked this conversation as resolved.
default=False,
repr=False)

def needs_plan(self, plan_params: PlanParams) -> bool:
if plan_params not in self._plan_params_to_wrappers:
Expand All @@ -307,9 +337,7 @@ def get_decode_wrapper(
raise ValueError(
"FlashInfer draft metadata views require the trtllm-gen "
"decode backend.")
num_seqs = self.num_seqs
result._kv_lens_buffer[:num_seqs].copy_(
self._draft_kv_runtime_lens[:num_seqs])
self._publish_decode_wrapper_kv_lens(result)
return result

def get_ragged_prefill_wrapper(
Expand Down Expand Up @@ -509,9 +537,12 @@ def _do_plan_mla_decode(self, plan_params: MLAPlanParams) -> None:
qo_indptr = self._qo_indptr[num_ctx:num_ctx + num_gen +
1] - self._qo_indptr[num_ctx]

num_pages_per_seq = kv_indptr[1:] - kv_indptr[:-1]
kv_len_arr = (num_pages_per_seq -
1) * plan_params.page_size + kv_last_page
if self._uses_full_generation_page_table:
kv_len_arr = self._logical_kv_lens[num_ctx:num_ctx + num_gen]
else:
num_pages_per_seq = kv_indptr[1:] - kv_indptr[:-1]
kv_len_arr = (num_pages_per_seq -
1) * plan_params.page_size + kv_last_page

self._mla_decode_wrapper.plan(
qo_indptr,
Expand Down Expand Up @@ -767,6 +798,69 @@ def update_shared_kv_draft_lengths(
num_accepted_tokens[num_contexts:num_seqs])
self._update_draft_kv_lengths()

def apply_spec_decode_kv_lens_offsets(
self,
offsets: torch.Tensor,
num_generations: int,
tokens_per_generation: int,
*,
num_chunked_contexts: int = 0,
restore: bool = False,
) -> None:
"""Apply overlap-scheduler corrections to FlashInfer's live KV state."""
if self._is_shared_kv_draft_view or self._is_separate_kv_draft_view:
raise RuntimeError(
"Speculative KV offsets must be applied to target metadata")
if not self._uses_full_generation_page_table:
return
if num_chunked_contexts == 0 and num_generations == 0:
return

direction = -1 if restore else 1
num_contexts = self.num_contexts
if num_chunked_contexts > 0:
row_slice = slice(num_contexts - num_chunked_contexts, num_contexts)
runtime_offsets = offsets[:num_chunked_contexts]
num_runtime_tokens = num_chunked_contexts * tokens_per_generation
token_slice = slice(self.num_ctx_tokens - num_runtime_tokens,
self.num_ctx_tokens)
else:
row_slice = slice(num_contexts, num_contexts + num_generations)
runtime_offsets = offsets[:num_generations]
num_runtime_tokens = num_generations * tokens_per_generation
token_slice = slice(self.num_ctx_tokens,
self.num_ctx_tokens + num_runtime_tokens)

self._cached_token_lens[row_slice].add_(runtime_offsets,
alpha=direction)
self._logical_kv_lens[row_slice].add_(runtime_offsets, alpha=direction)
token_offsets = runtime_offsets.repeat_interleave(tokens_per_generation)
self._positions[token_slice].add_(token_offsets, alpha=direction)

for wrappers in self._plan_params_to_wrappers.values():
self._publish_decode_wrapper_kv_lens(wrappers.decode_wrapper)

def _publish_decode_wrapper_kv_lens(self, decode_wrapper) -> None:
"""Publish device-logical lengths to a trtllm-gen decode wrapper."""
kv_lens_buffer = getattr(decode_wrapper, "_kv_lens_buffer", None)
if kv_lens_buffer is None or self.num_generations == 0:
return

if self._is_shared_kv_draft_view or self._is_separate_kv_draft_view:
kv_lens_buffer[:self.num_generations].copy_(
self._draft_kv_runtime_lens[:self.num_generations])
return
if not self._uses_full_generation_page_table:
return

start = self.num_contexts
end = start + self.num_generations
torch.add(
self._cached_token_lens[start:end],
self.seq_lens_kv_cuda[start:end],
out=kv_lens_buffer[:self.num_generations],
)

def _prepare_full_draft_page_table(self) -> None:
"""Expose every allocated draft page and use device KV lengths."""
if self._uses_full_draft_page_table:
Expand Down Expand Up @@ -886,6 +980,13 @@ def _post_init_with_buffers(self, buffers) -> None:
self._cached_token_lens = torch.empty((self.max_num_requests, ),
dtype=torch.int,
device='cuda')
self._logical_kv_lens = self.get_empty(
buffers,
(self.max_num_requests, ),
dtype=torch.int,
cache_name="_logical_kv_lens",
capture_graph=capture_graph,
)
self._draft_kv_runtime_lens = self.get_empty(
buffers,
(self.max_num_requests, ),
Expand All @@ -909,6 +1010,7 @@ def _post_init_with_buffers(self, buffers) -> None:
self._host_pool_indices: Dict[int, torch.Tensor] = {}
self._host_paged_kv_indices: Optional[torch.Tensor] = None
self._host_paged_kv_indptr_decode: Optional[torch.Tensor] = None
self._uses_full_generation_page_table = False
self._max_num_blocks_per_seq = 0

# VSWA (Variable Sliding Window Attention): models with per-layer
Expand Down Expand Up @@ -1391,12 +1493,33 @@ def _to_int32_tensor(arr: np.ndarray) -> torch.Tensor:
else:
self.num_ctx_cached_tokens = 0

# Number of tokens needed in the KV cache for each sequence after the
# next pass. Kept on the host: every consumer below needs host values,
# so a device-side computation would force a sync per step.
# Number of tokens needed in the KV cache for each sequence after the next pass.
# Kept on the host: every consumer below needs host values, so a device-side computation
# would force a sync per step. These logical block counts describe committed history, not
# all pages that the KV manager has reserved for fixed-width speculative appends.
kv_lens_host = np.asarray(num_cached_tokens_per_seq,
dtype=np.int64) + self.seq_lens_kv.numpy()
num_blocks = (kv_lens_host + self.page_size - 1) // self.page_size
logical_num_blocks = ((kv_lens_host + self.page_size - 1) //
Comment thread
2ez4bz marked this conversation as resolved.
self.page_size)
num_blocks = logical_num_blocks
use_full_generation_page_table = (
self.kv_cache_params.use_full_generation_page_table)
self._uses_full_generation_page_table = use_full_generation_page_table
if use_full_generation_page_table:
# Generation rows need every reserved page to keep a speculative append addressable
# across a page boundary. Treating that wider table as the logical sequence would make
# FlashInfer attend to uncommitted draft slots, so plans and positions use the separate
# device logical lengths populated below.
assert self.request_ids is not None
Comment thread
2ez4bz marked this conversation as resolved.
num_blocks = np.asarray(
_get_page_table_num_blocks(
self.kv_cache_manager,
self.request_ids,
logical_num_blocks.tolist(),
self.num_contexts,
),
dtype=np.int64,
)
self.num_blocks = num_blocks.tolist()
Comment thread
coderabbitai[bot] marked this conversation as resolved.

assert self.request_ids is not None
Expand Down Expand Up @@ -1464,9 +1587,9 @@ def _to_int32_tensor(arr: np.ndarray) -> torch.Tensor:
self._vswa_active_pool_id = primary_pool_id

# number of tokens in the last cache block used by each sequence,
# derived on the host so no GPU arithmetic or sync is needed.
# derived from the logical rather than reservation-width page count.
paged_kv_last_page_len = _to_int32_tensor(kv_lens_host -
(num_blocks - 1) *
(logical_num_blocks - 1) *
self.page_size)
self._paged_kv_last_page_len[:paged_kv_last_page_len.size(0)].copy_(
paged_kv_last_page_len, non_blocking=True)
Expand Down Expand Up @@ -1509,11 +1632,23 @@ def _to_int32_tensor(arr: np.ndarray) -> torch.Tensor:

# For cross attention, num_tokens is 0 during decode, and we don't need to update kv cache.
if self.num_tokens > 0:
if use_full_generation_page_table:
Comment thread
2ez4bz marked this conversation as resolved.
# The page-table indptr now describes addressable capacity. Deriving positions from
# it would count reserved pages as committed KV and potentially shift appends beyond
# the true sequence end.
logical_kv_lens = _to_int32_tensor(kv_lens_host)
self._logical_kv_lens[:logical_kv_lens.numel()].copy_(
logical_kv_lens, non_blocking=True)
position_kv_lens = self._logical_kv_lens[:self.num_seqs]
else:
position_kv_lens = flashinfer.get_seq_lens(
self.paged_kv_indptr,
self.paged_kv_last_page_len,
self.page_size,
)
batch_indices, positions = flashinfer.get_batch_indices_positions(
self.kv_indptr,
flashinfer.get_seq_lens(self.paged_kv_indptr,
self.paged_kv_last_page_len,
self.page_size),
position_kv_lens,
self.num_tokens,
)
self._batch_indices[:batch_indices.size(0)].copy_(batch_indices,
Expand Down Expand Up @@ -1687,6 +1822,20 @@ def plan(self,
if q_scaling is not None:
sm_scale = 1 / (math.sqrt(head_dim) * q_scaling)

# FlashInfer decode accepts one q_len_per_req. Current paths keep generation widths
# uniform (including padded speculative drafts); guard against malformed metadata or
# future per-request draft widths rather than launching the wrong shape.
q_len_per_req = 1
if self.num_generations > 0:
generation_seq_lens = self.seq_lens[self.
num_contexts:self.num_contexts +
self.num_generations]
q_len_per_req = int(generation_seq_lens[0])
if not torch.all(generation_seq_lens == q_len_per_req):
raise ValueError(
"FlashInfer decode requires a uniform query length per "
f"request, but got {generation_seq_lens.tolist()}")

plan_params = PlanParams(
num_heads=num_heads,
num_kv_heads=num_kv_heads,
Expand All @@ -1700,6 +1849,8 @@ def plan(self,
attention_mask_data=attention_mask_data,
multi_item_params=self._multi_item_params,
kv_pool_id=getattr(self, "_vswa_active_pool_id", None),
q_len_per_req=q_len_per_req,
num_generations=self.num_generations,
)
return self._plan_with_params(plan_params, flashinfer_backend)

Expand Down Expand Up @@ -1858,6 +2009,11 @@ def prefill_plan():

def decode_plan():
assert decode_wrapper is not None
if (self._uses_full_generation_page_table
and decode_wrapper._backend != "trtllm-gen"):
raise ValueError(
"Reservation-width FlashInfer page tables require the "
"trtllm-gen decode backend's independent KV lengths.")
# Host int32 indptr (retained by prepare, which always runs
# before plans): flashinfer moves it to the device itself, and
# its indptr.cpu()/get_seq_lens calls stay free of D2H syncs.
Expand All @@ -1884,7 +2040,10 @@ def decode_plan():
kv_data_type=plan_params.kv_dtype,
o_data_type=o_dtype,
block_tables=block_tables,
# Keep FlashInfer's recorded graph shape aligned with the wrapper cache key.
q_len_per_req=plan_params.q_len_per_req,
)
self._publish_decode_wrapper_kv_lens(decode_wrapper)

# Must sync after append_paged_kv_cache and before plan.
torch.cuda.current_stream().synchronize()
Expand Down
7 changes: 7 additions & 0 deletions tensorrt_llm/_torch/metadata.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,6 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

from dataclasses import dataclass
from enum import Enum
from typing import List, Optional
Expand Down Expand Up @@ -30,6 +33,10 @@ class KVCacheParams:
# The number of extra kv for draft tokens
num_extra_kv_tokens: Optional[int] = 0

# Whether generation page tables expose all reserved pages. Backends that
# use this must track the logical KV length separately.
use_full_generation_page_table: bool = False


class CacheType(Enum):
# Linear KV cache stores all the cached tokens of a sequence in a single page.
Expand Down
Loading
Loading