diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py index e4749ee7cb9c..86ad02deecb7 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py @@ -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 @@ -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. @@ -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() diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 1b624f72d88b..cb7e3241532b 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -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 = ( @@ -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: + 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: @@ -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: @@ -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) diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 6425cb7d2f2d..731aec9d4b39 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -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, diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_common.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_common.py index dca283ef87df..6e8029acfa33 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_common.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_common.py @@ -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) diff --git a/tests/unittest/disaggregated/test_cache_transceiver_single_process.py b/tests/unittest/disaggregated/test_cache_transceiver_single_process.py index eecf1a2d819c..6bb0bf2c1b8a 100644 --- a/tests/unittest/disaggregated/test_cache_transceiver_single_process.py +++ b/tests/unittest/disaggregated/test_cache_transceiver_single_process.py @@ -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 diff --git a/tests/unittest/disaggregated/test_kv_transfer.py b/tests/unittest/disaggregated/test_kv_transfer.py index 3b8d59e42edd..e87ce6fe7bcd 100644 --- a/tests/unittest/disaggregated/test_kv_transfer.py +++ b/tests/unittest/disaggregated/test_kv_transfer.py @@ -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