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
37 changes: 36 additions & 1 deletion tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,7 +51,7 @@
from tensorrt_llm.runtime.kv_cache_manager_v2._block_radix_tree import (
gen_multimodal_cache_key_tokens,
)
from tensorrt_llm.runtime.kv_cache_manager_v2._common import BAD_PAGE_INDEX, GPU_LEVEL
from tensorrt_llm.runtime.kv_cache_manager_v2._common import BAD_PAGE_INDEX, CACHE_LEVEL1, GPU_LEVEL
from tensorrt_llm.runtime.kv_cache_manager_v2._config import DataRole
from tensorrt_llm.runtime.kv_cache_manager_v2._utils import exact_div, typed_range
from tensorrt_llm.sampling_params import SamplingParams
Expand Down Expand Up @@ -695,6 +695,7 @@ def append_to_kv_heads_per_layer(

self.enable_block_reuse = kv_cache_config.enable_block_reuse
self.enable_partial_reuse = kv_cache_config.enable_partial_reuse
self.disk_prefetch_num_reqs = kv_cache_config.disk_prefetch_num_reqs

# With pipeline parallelism, multiple microbatches can be in-flight
# simultaneously, so we need slots for all concurrent sequences.
Expand Down Expand Up @@ -1829,5 +1830,39 @@ def _create_kv_cache(
kv_cache.set_base_page_index_buf(i, pool_idx, memoryview(buffer.numpy()))
return kv_cache

def prefetch_for_context_tokens(self, requests: list) -> bool:
"""Prefetch radix-tree blocks from disk→host for upcoming context requests.

Returns True if all prefetches succeeded, False if any failed.
"""
if not self.enable_block_reuse:
return False
# Prefetch via a transient KV cache that holds the reuse-matched blocks,
# prefetches disk->host, then closes. Holding blocks costs no GPU space
# (never resumed) and close() needs no stream sync. The transient cache
# is NOT registered in kv_cache_map / IndexMapper.
success = True
for req in requests:
all_tokens = req.get_tokens(DEFAULT_BEAM_INDEX)
tokens = self._augment_tokens_for_block_reuse(all_tokens, req, end=len(all_tokens) - 1)
# Match the ReuseScope salt derivation used in _create_kv_cache so
# the transient cache hits the same radix-tree blocks.
cache_salt = req.cache_salt
salt_int = (
int.from_bytes(hashlib.sha256(cache_salt.encode("utf-8")).digest()[:8], "little")
if cache_salt is not None
else None
)
kv_cache = self.impl.create_kv_cache(
ReuseScope(lora_id=req.lora_task_id, salt=salt_int), tokens
)
# Prefetch to the first tier below GPU (host if present, otherwise
# disk). prefetch() is a best-effort hint either way.
if not kv_cache.prefetch(CACHE_LEVEL1):
logger.warning("prefetch failed for request %s", req.py_request_id)
success = False
kv_cache.close()
return success

def reset_reuse_state(self):
self.impl.clear_reusable_blocks()
28 changes: 28 additions & 0 deletions tensorrt_llm/_torch/pyexecutor/py_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -441,6 +441,7 @@ def __init__(
# unnecessary). Several revert/skip paths gate on this flag.
self._is_kv_manager_v2 = isinstance(self.kv_cache_manager,
KVCacheManagerV2)
self._prefetched_request_ids: set[int] = set()
self.enable_kv_cache_events = self.kv_cache_manager is not None and self.kv_cache_manager.event_buffer_max_size > 0
self.enable_kv_cache_reuse = self.kv_cache_manager is not None and self.kv_cache_manager.enable_block_reuse
self.enable_partial_reuse_for_disagg = (
Expand Down Expand Up @@ -2549,6 +2550,30 @@ def _revert_ctx_alloc(self, dropped_context_requests):
for req in dropped_context_requests:
self.kv_cache_manager.revert_allocate_context(req)

@nvtx_range("_prefetch_for_context_requests")
def _prefetch_for_context_requests(self) -> None:
"""Pre-stage disk blocks to host for upcoming context requests with block reuse."""
if not isinstance(getattr(self, "kv_cache_manager", None),
KVCacheManagerV2):
return
if not self.kv_cache_manager.enable_block_reuse:
return
if self.kv_cache_manager.disk_prefetch_num_reqs <= 0:
return
max_prefetch = self.kv_cache_manager.disk_prefetch_num_reqs
candidates = []
for req in self.active_requests:
if len(candidates) >= max_prefetch:
break
if (req.is_first_context_chunk and req.py_request_id
not in self.kv_cache_manager.kv_cache_map
and req.py_request_id not in self._prefetched_request_ids):
candidates.append(req)
self._prefetched_request_ids.add(req.py_request_id)

if candidates:
Comment thread
reasonsolo marked this conversation as resolved.
self.kv_cache_manager.prefetch_for_context_tokens(candidates)

def _prepare_and_schedule_batch(self):
new_requests = self._fetch_and_activate_new_requests()
if self.should_stop_processing:
Expand All @@ -2567,6 +2592,8 @@ def _prepare_and_schedule_batch(self):

self._pad_attention_dp_dummy_request()

self._prefetch_for_context_requests()

if self.drafter is not None:
# Honor permanent disable flag based on rolling acceptance first
if self.drafter.draft_len_schedule is not None:
Expand Down Expand Up @@ -4816,6 +4843,7 @@ def _terminate_request(self, request: LlmRequest):

def _do_terminate_request(self, request: LlmRequest):
self.resource_manager.free_resources(request)
self._prefetched_request_ids.discard(request.py_request_id)

if self.gather_all_responses or self.dist.rank == 0:
self.result_wait_queues.pop(request.py_request_id, None)
Expand Down
8 changes: 8 additions & 0 deletions tensorrt_llm/llmapi/llm_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -2907,6 +2907,14 @@ class KvCacheConfig(StrictBaseModel, PybindMirror):
"your own risk. Only used when using KV cache manager v2 "
"(experimental).")

disk_prefetch_num_reqs: int = Field(
default=0,
ge=0,
description=
"Number of queued context requests to prefetch disk-tier KV cache blocks to host for. "
"Set to 0 to disable prefetch. Only effective with KV cache manager v2 and block reuse enabled."
)

def _to_pybind(self):
config = _KvCacheConfig(
enable_block_reuse=self.enable_block_reuse,
Expand Down
3 changes: 3 additions & 0 deletions tensorrt_llm/runtime/kv_cache_manager_v2/_common.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,9 @@ class PageIndexMode(enum.IntEnum):


GPU_LEVEL: Final[CacheLevel] = CacheLevel(0)
# First cache level below GPU. Its semantic tier depends on the configured
# cache_tiers: host when a host tier exists, otherwise disk.
CACHE_LEVEL1: Final[CacheLevel] = CacheLevel(1)

# Normal token id that falls in the tokenizer vocabulary.
TokenId = NewType("TokenId", int)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -105,6 +105,7 @@ class KvCacheConfigV2:
enable_partial_reuse: bool = False
copy_on_partial_reuse: bool = False
dtype: str = "auto"
disk_prefetch_num_reqs: int = 4
max_util_for_resume: float = 0.95


Expand Down
1 change: 1 addition & 0 deletions tests/unittest/disaggregated/test_kv_transfer.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,7 @@ class KvCacheConfigV2:
enable_partial_reuse: bool = False
copy_on_partial_reuse: bool = False
dtype: str = "auto"
disk_prefetch_num_reqs: int = 4
# V2 specific field
max_util_for_resume: float = 0.95

Expand Down
Loading