diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py index ffe7b6f5f83e..6715e86047d3 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py @@ -17,6 +17,7 @@ import os import sys from collections import OrderedDict, defaultdict +from dataclasses import dataclass from typing import TYPE_CHECKING, Dict, Iterable, List, NamedTuple, Optional, Sequence, Tuple, Union import numpy as np @@ -58,6 +59,7 @@ LayerId, LifeCycleId, PageIndexMode, + PlannedDropHandle, PoolGroupPeakBlockStats, ReuseScope, SwaScratchReuseConfig, @@ -130,6 +132,90 @@ class Role: class BlockReusePolicy(StrEnum): ALL_REUSABLE = "all_reusable" PER_REQUEST = "per_request" + PER_CONVERSATION = "per_conversation" + + +def _request_conversation_id(request: LlmRequest) -> Optional[str]: + if request.is_dummy_request: + return None + conversation_params = request.py_conversation_params + if conversation_params is None: + return None + conversation_id = conversation_params.conversation_id.strip() + return conversation_id or None + + +@dataclass(slots=True) +class _ConversationState: + current_request_id: Optional[int] = None + planned_drop_handle: Optional[PlannedDropHandle] = None + + +class ConversationManager: + """Track the current request and drop plan for each conversation.""" + + def __init__(self) -> None: + self._conversation_states: Dict[str, _ConversationState] = {} + + def save_drop_plan(self, request: LlmRequest, kv_cache: _KVCache) -> None: + """Save a completed context's drop plan and apply the preceding plan on success.""" + request_id = request.py_request_id + conversation_id = _request_conversation_id(request) + if conversation_id is None: + return + + state = self._conversation_states[conversation_id] + if state.current_request_id != request_id: + return + + drop_handle = kv_cache.plan_committed_block_drop() + if drop_handle is None: + logger.warning( + f"Committed blocks for request {request_id} in conversation " + f"{conversation_id} have been dropped." + ) + else: + previous_handle = state.planned_drop_handle + state.planned_drop_handle = drop_handle + if previous_handle is not None: + previous_handle.drop() + + self.finish_request(request) + + def prepare_request(self, request: LlmRequest) -> None: + """Register a context request unless its conversation has another active one.""" + conversation_id = _request_conversation_id(request) + if conversation_id is None: + return + request_id = request.py_request_id + state = self._conversation_states.setdefault(conversation_id, _ConversationState()) + current_request_id = state.current_request_id + if current_request_id is not None and current_request_id != request_id: + logger.warning( + f"Conversation {conversation_id} already has current request " + f"{current_request_id}. Request {request_id} will ignore " + "conversation params." + ) + return + + state.current_request_id = request_id + + def finish_request(self, request: LlmRequest) -> None: + """Clear a request as active while preserving any saved drop plan.""" + conversation_id = _request_conversation_id(request) + if conversation_id is None: + return + state = self._conversation_states.get(conversation_id) + if state is None or state.current_request_id != request.py_request_id: + return + + state.current_request_id = None + if state.planned_drop_handle is None: + self._conversation_states.pop(conversation_id) + + def clear(self) -> None: + """Clear state after reusable KV-cache blocks have been cleared.""" + self._conversation_states.clear() def _estimate_full_attn_size_per_token( @@ -990,6 +1076,12 @@ 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 + enable_conversation_manager = ( + self.enable_block_reuse + and self.block_reuse_policy == BlockReusePolicy.PER_CONVERSATION + and not self.is_draft + ) + self.conversation_manager = ConversationManager() if enable_conversation_manager else None # With pipeline parallelism, multiple microbatches can be in-flight # simultaneously, so we need slots for all concurrent sequences. @@ -1954,6 +2046,8 @@ def prepare_context(self, req: LlmRequest) -> bool: def _prepare_context_impl(self, req: LlmRequest) -> bool: if req.is_first_context_chunk: + if self.conversation_manager is not None: + self.conversation_manager.prepare_request(req) kv_cache = self.kv_cache_map.get(req.py_request_id) if kv_cache is None: all_tokens = req.get_tokens(DEFAULT_BEAM_INDEX) @@ -2825,6 +2919,8 @@ def release_index_slot(self, request_id: int) -> None: self._early_freed_index_requests.add(request_id) def free_resources(self, request: LlmRequest, pin_on_release: bool = False): + if self.conversation_manager is not None: + self.conversation_manager.finish_request(request) self._allocated_draft_lens.pop(request.py_request_id, None) kv_cache = self.kv_cache_map.pop(request.py_request_id, None) if kv_cache is None: @@ -3065,6 +3161,8 @@ def shutdown(self): kv_cache.close() self.kv_cache_map.clear() self.impl.shutdown() + if self.conversation_manager is not None: + self.conversation_manager.clear() def get_max_resource_count(self) -> int: # TODO: implement this @@ -3146,6 +3244,8 @@ def update_context_resources(self, scheduled_batch: ScheduledRequests): if should_commit: self.try_commit_blocks(req) if req.context_remaining_length == 0: + if self.conversation_manager is not None: + self.conversation_manager.save_drop_plan(req, kv_cache) # Scratch blocks are only for prefill chunks. Disable them at # the context/generation boundary so generation uses normal KV # pages before the first generation allocation. @@ -3328,3 +3428,5 @@ def prefetch_for_context_tokens(self, requests: list) -> bool: def reset_reuse_state(self): self.impl.clear_reusable_blocks() + if self.conversation_manager is not None: + self.conversation_manager.clear() diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index ca2ae9cc0062..f416a1d2f7a7 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -3666,12 +3666,19 @@ class KvCacheConfig(StrictBaseModel, PybindMirror): "pool_ratio is set.") # This is a pure python field, not a pybind field. It is only for the Pytorch backend. - block_reuse_policy: Literal["all_reusable", "per_request"] = Field( - default="all_reusable", - status="prototype", - description="KV cache manager v2 block reuse policy. " - "With SWA scratch reuse and 'all_reusable', only non-scratch " - "blocks are saved for reuse.") + block_reuse_policy: Literal[ + "all_reusable", "per_request", "per_conversation"] = Field( + default="all_reusable", + status="prototype", + description="KV cache manager v2 block reuse policy. " + "'all_reusable' commits reusable blocks after every context chunk; " + "'per_request' commits them only after the final context chunk; " + "'per_conversation' uses 'per_request' commits and drops the previous " + "turn's committed SWA-window blocks after the current turn's final context " + "chunk. All reusable blocks remain subject to normal cache eviction. " + "Requests without conversation params use 'per_request' behavior. When " + "'all_reusable' and SWA scratch reuse are both enabled, only non-scratch " + "blocks are committed for reuse.") def _to_pybind(self): config = _KvCacheConfig( diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py index 34399289b9ae..5557a1d2f32e 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py @@ -53,6 +53,7 @@ ExpandedBuffer, KVCacheManager, PageIndexConverter, + PlannedDropHandle, PoolDesc, PoolGroupDesc, PoolGroupPeakBlockStats, @@ -118,6 +119,7 @@ "KVCacheUpdatedData", "KvCacheStatus", "LayerGroupId", + "PlannedDropHandle", "LayerId", "LifeCycleId", "MemAddress", diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi index 61ddc4bdbb4f..0945db4d2394 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi @@ -51,6 +51,9 @@ CacheLevel = NewType("CacheLevel", int) TokenId = NewType("TokenId", int) TokenIdExt = Union[TokenId, bytes] +class PlannedDropHandle: + def drop(self) -> None: ... + class ReuseScope(NamedTuple): lora_id: int | None = None salt: int | None = None @@ -349,6 +352,7 @@ class _KVCache: def committed_tokens(self) -> list[TokenIdExt]: ... @property def reuse_scope(self) -> ReuseScope: ... + def plan_committed_block_drop(self) -> PlannedDropHandle | None: ... def stop_committing(self) -> None: ... def suspend(self) -> None: ... def resume(self, cuda_stream: CudaStream | None = None) -> bool: ... diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_core/__init__.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_core/__init__.py index ea7360d69c83..4e52e154c47b 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_core/__init__.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_core/__init__.py @@ -14,7 +14,7 @@ # limitations under the License. from .._common import DEFAULT_BEAM_INDEX, BeamIndex -from ._kv_cache import _KVCache +from ._kv_cache import PlannedDropHandle, _KVCache from ._kv_cache_manager import ( AggregatedPageDesc, ExpandedBuffer, @@ -29,6 +29,7 @@ __all__ = [ "KVCacheManager", "_KVCache", + "PlannedDropHandle", "BeamIndex", "DEFAULT_BEAM_INDEX", "AggregatedPageDesc", diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_core/_kv_cache.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_core/_kv_cache.py index 63e2244d9b98..5494449c0bd6 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_core/_kv_cache.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_core/_kv_cache.py @@ -16,7 +16,7 @@ import array import enum import math -from collections.abc import Sequence +from collections.abc import Iterable, Sequence from contextlib import contextmanager from dataclasses import dataclass from itertools import chain @@ -37,6 +37,7 @@ CudaStream, PageIndex, PageIndexMode, + PageStatus, Priority, TokenIdExt, ) @@ -129,6 +130,58 @@ def __del__(self) -> None: self.pages.clear() +class PlannedDropHandle: + """Track committed pages planned for dropping without owning them. + + The handle stores weak references and does not keep pages alive. Dropping it + decrements each live page's planned-drop count and removes an already-droppable + page from eviction tracking when no plans remain. + """ + + __slots__ = ("_page_refs",) + + _page_refs: tuple[rawref.ref[CommittedPage], ...] | None + + def __init__(self, pages: Iterable[CommittedPage]) -> None: + planned_pages = tuple({id(page): page for page in pages}.values()) + self._page_refs = tuple(rawref.ref(page) for page in planned_pages) + for page in planned_pages: + page.planned_drop_count += 1 + + def drop(self) -> None: + """Apply this drop plan and invalidate the handle. + + A live page is removed from eviction tracking only when this is its final + plan and it is already droppable and queued for eviction. Calling this + method twice is invalid. + """ + page_refs = self._page_refs + if page_refs is None: + raise ValueError("Planned drop handle has already been dropped") + + pages = list[CommittedPage]() + for page_ref in page_refs: + page = page_ref() + if page is not None: + if page.planned_drop_count <= 0: + raise ValueError("Committed page has no planned drop") + pages.append(page) + + self._page_refs = None + for page in pages: + page.planned_drop_count -= 1 + if ( + page.planned_drop_count == 0 + and page.status == PageStatus.DROPPABLE + and page.scheduled_for_eviction + ): + page.manager.exclude_from_eviction(page) + + def __del__(self) -> None: + if self._page_refs is not None: + self.drop() + + class _Status(enum.Enum): ACTIVE = enum.auto() SUSPENDED = enum.auto() @@ -993,6 +1046,43 @@ def committed_tokens(self) -> list[TokenIdExt]: def reuse_scope(self) -> ReuseScope: return self._reuse_scope + def plan_committed_block_drop(self) -> PlannedDropHandle | None: + """Plan dropping SWA blocks needed only by the next conversation turn. + + The plan covers committed pages in each SWA life cycle's current + attention window. Full-attention and attention-sink blocks are excluded + because later turns may still need them. SSM state is not yet supported. + This must be called after stop_committing(). Returns None without + creating a plan if any required SWA page is unavailable. + """ + if self._commit_state != self.CommitState.USER_STOP: + raise LogicError("plan_committed_block_drop() requires stop_committing()") + + end = self._num_committed_blocks + pages_to_drop: list[CommittedPage] = [] + for lc_idx, lc in self.manager._life_cycles.items(): + if isinstance(lc, SsmLifeCycle): + # TODO: Support recording reusable SSM state pages. + continue + if lc.window_size is None: + continue + stale_range = _KVCache._get_stale_range( + self.tokens_per_block, self.num_committed_tokens, lc + ) + window_start = min(stale_range.end, end) + for ordinal in typed_range(window_start, end): + tree_block = self._blocks[ordinal].tree_block + if tree_block is None: + return None + page_ref = tree_block.storage[lc_idx] + if page_ref is None: + return None + page = page_ref() + if page is None: + return None + pages_to_drop.append(page) + return PlannedDropHandle(pages_to_drop) + # Users promise to not commit any more tokens. For cases where we shouldn't reuse generated tokens # (eg. CoT), this helps us drop (instead of evict) out-of-window blocks for SWA layers. # If there is a uncommitted block containing committed tokens, we will commit the block immediately. diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_page.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_page.py index 43bac9243622..70a82c090c01 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_page.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_page.py @@ -241,6 +241,7 @@ class CommittedPage(Page): """ block: rawref.ref["Block"] + planned_drop_count: int __rawref__: rawref.ref["CommittedPage"] def is_committed(self) -> bool: @@ -256,6 +257,7 @@ def __init__( priority: Priority, ): self.block = rawref.ref(block) + self.planned_drop_count = 0 self.__rawref__ = rawref.NULL Page.__init__( self, diff --git a/tensorrt_llm/usage/llm_args_golden_manifest.json b/tensorrt_llm/usage/llm_args_golden_manifest.json index 28760a844f9d..739db09bfb23 100644 --- a/tensorrt_llm/usage/llm_args_golden_manifest.json +++ b/tensorrt_llm/usage/llm_args_golden_manifest.json @@ -561,9 +561,10 @@ { "allowed_values": [ "all_reusable", - "per_request" + "per_request", + "per_conversation" ], - "annotation": "Literal['all_reusable', 'per_request']", + "annotation": "Literal['all_reusable', 'per_request', 'per_conversation']", "converter": "", "kind": "categorical", "path": "kv_cache_config.block_reuse_policy" diff --git a/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py b/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py index be1f2936abf8..307588e0d2a0 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py +++ b/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py @@ -1,14 +1,29 @@ # Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +from dataclasses import dataclass, field from types import SimpleNamespace +from unittest.mock import patch import pytest +import torch from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import BlockReusePolicy, KVCacheManagerV2 -from tensorrt_llm.bindings.internal.batch_manager import CacheType as CacheTypeCpp +from tensorrt_llm._torch.pyexecutor.scheduler import ScheduledRequests +from tensorrt_llm.bindings import DataType +from tensorrt_llm.bindings.internal.batch_manager import CacheType +from tensorrt_llm.conversation_params import ConversationParams from tensorrt_llm.llmapi.llm_args import KvCacheConfig -from tensorrt_llm.runtime.kv_cache_manager_v2 import GpuCacheTierConfig, KVCacheManagerConfig +from tensorrt_llm.mapping import Mapping +from tensorrt_llm.runtime.kv_cache_manager_v2 import ( + DEFAULT_BEAM_INDEX, + GpuCacheTierConfig, + KVCacheManagerConfig, +) +from tensorrt_llm.runtime.kv_cache_manager_v2._utils import init_cuda_once + +TOKENS_PER_BLOCK = 4 +MAX_SEQ_LEN = 16 class _FakeKVCache: @@ -29,7 +44,7 @@ def _build_cache_config_for_test( kv_cache_config: KvCacheConfig, *, is_draft: bool = False ) -> KVCacheManagerConfig: cache_manager = object.__new__(KVCacheManagerV2) - cache_manager.kv_cache_type = CacheTypeCpp.SELFKONLY + cache_manager.kv_cache_type = CacheType.SELFKONLY cache_manager.head_dim_per_layer = [128] cache_manager.enable_swa_scratch_reuse = False cache_manager.num_extra_kv_tokens = 0 @@ -104,3 +119,292 @@ def test_try_commit_blocks_commits_partial_block_at_context_end() -> None: assert kv_cache.committed_tokens == [4, 5, 6, 7, 8, 9] assert kv_cache.num_committed_tokens == 10 assert kv_cache.stopped_committing + + +@dataclass +class _ContextRequest: + request_id: int + tokens: list[int] + context_remaining_length: int + conversation_id: str + py_request_id: int = field(init=False) + py_conversation_params: ConversationParams | None = field(init=False) + use_conversation_params: bool = True + lora_task_id: int | None = None + cache_salt: str | None = None + is_first_context_chunk: bool = True + is_last_context_chunk: bool = True + is_disagg_generation_init_state: bool = False + is_dummy_request: bool = False + context_current_position: int = 0 + prepopulated_prompt: tuple[int, int] | None = None + multimodal_hashes: None = None + multimodal_positions: None = None + multimodal_lengths: None = None + + def __post_init__(self) -> None: + self.py_request_id = self.request_id + if not self.use_conversation_params: + self.py_conversation_params = None + return + self.py_conversation_params = ConversationParams(conversation_id=self.conversation_id) + + @property + def prompt_len(self) -> int: + return len(self.tokens) + + @property + def is_dummy(self) -> bool: + return self.is_dummy_request + + @property + def prepopulated_prompt_len(self) -> int: + if self.prepopulated_prompt is None: + return 0 + return self.prepopulated_prompt[0] + + def get_tokens(self, beam_id: int = DEFAULT_BEAM_INDEX) -> list[int]: + assert beam_id == DEFAULT_BEAM_INDEX + return self.tokens + + def set_prepopulated_prompt_len(self, length: int, tokens_per_block: int) -> None: + self.prepopulated_prompt = (length, tokens_per_block) + + +@pytest.fixture +def manager() -> KVCacheManagerV2: + if not torch.cuda.is_available(): + pytest.skip("requires CUDA") + init_cuda_once() + manager = KVCacheManagerV2( + KvCacheConfig( + enable_block_reuse=True, + enable_partial_reuse=True, + max_gpu_total_bytes=16 << 20, + max_attention_window=[MAX_SEQ_LEN, TOKENS_PER_BLOCK], + max_util_for_resume=1.0, + block_reuse_policy="per_conversation", + ), + CacheType.SELF, + num_layers=2, + num_kv_heads=128, + head_dim=1024, + tokens_per_block=TOKENS_PER_BLOCK, + max_seq_len=MAX_SEQ_LEN, + max_batch_size=2, + mapping=Mapping(world_size=1, rank=0, tp_size=1, pp_size=1), + dtype=DataType.HALF, + vocab_size=4096, + enable_stats=False, + ) + try: + yield manager + finally: + manager.shutdown() + + +def _context_batch(*requests: _ContextRequest) -> ScheduledRequests: + batch = ScheduledRequests() + for request in requests: + batch.append_context_request(request) + return batch + + +def _prepare_context_resources( + manager: KVCacheManagerV2, + *requests: _ContextRequest, +) -> ScheduledRequests: + batch = _context_batch(*requests) + manager.prepare_resources(batch) + return batch + + +def _update_context_resources( + manager: KVCacheManagerV2, + batch: ScheduledRequests, +) -> None: + manager.update_context_resources(batch) + + +def _free_if_active( + manager: KVCacheManagerV2, + request: _ContextRequest, +) -> None: + manager.free_resources(request) + + +def _run_context( + manager: KVCacheManagerV2, + request: _ContextRequest, +) -> None: + batch = _prepare_context_resources(manager, request) + assert manager.prepare_context(request) + request.context_remaining_length = request.prompt_len - request.context_current_position + assert manager.resize_context(request, num_tokens=request.context_remaining_length) + request.context_current_position = request.prompt_len + request.context_remaining_length = 0 + _update_context_resources(manager, batch) + + +def test_per_conversation_policy_delays_commit_until_last_context_chunk( + manager: KVCacheManagerV2, +) -> None: + request = _ContextRequest(1, list(range(8)), 8, "conv-1") + + try: + batch = _prepare_context_resources(manager, request) + assert manager.prepare_context(request) + assert manager.resize_context(request, num_tokens=4) + request.context_current_position = 4 + request.context_remaining_length = 4 + _update_context_resources(manager, batch) + + kv_cache = manager.kv_cache_map[request.py_request_id] + assert kv_cache.num_committed_tokens == 0 + assert kv_cache.history_length == 4 + + request.is_first_context_chunk = False + batch = _prepare_context_resources(manager, request) + assert manager.prepare_context(request) + assert manager.resize_context(request, num_tokens=4) + request.context_current_position = 8 + request.context_remaining_length = 0 + _update_context_resources(manager, batch) + + assert kv_cache.num_committed_tokens == 8 + assert kv_cache.history_length == 8 + finally: + _free_if_active(manager, request) + + +def test_per_conversation_policy_without_params_uses_per_request_commit( + manager: KVCacheManagerV2, +) -> None: + request = _ContextRequest( + 1, + list(range(8)), + 8, + "conv-1", + use_conversation_params=False, + ) + batch = _context_batch(request) + + try: + assert manager.prepare_context(request) + assert manager.resize_context(request, num_tokens=4) + request.context_current_position = 4 + request.context_remaining_length = 4 + _update_context_resources(manager, batch) + + kv_cache = manager.kv_cache_map[request.py_request_id] + assert kv_cache.num_committed_tokens == 0 + assert kv_cache.history_length == 4 + finally: + if request.py_request_id in manager.kv_cache_map: + manager.free_resources(request) + + +def test_per_conversation_policy_releases_cancelled_request( + manager: KVCacheManagerV2, +) -> None: + request_a = _ContextRequest(1, list(range(8)), 8, "conv-1") + request_b = _ContextRequest(2, list(range(8)), 8, "conv-1") + + try: + batch_a = _prepare_context_resources(manager, request_a) + assert manager.prepare_context(request_a) + assert manager.resize_context(request_a, num_tokens=4) + request_a.context_current_position = 4 + request_a.context_remaining_length = 4 + _update_context_resources(manager, batch_a) + _free_if_active(manager, request_a) + + batch_b = _prepare_context_resources(manager, request_b) + with patch( + "tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2.logger.warning" + ) as mock_warning: + assert manager.prepare_context(request_b) + mock_warning.assert_not_called() + assert manager.resize_context(request_b, num_tokens=request_b.prompt_len) + request_b.context_current_position = request_b.prompt_len + request_b.context_remaining_length = 0 + _update_context_resources(manager, batch_b) + finally: + _free_if_active(manager, request_b) + _free_if_active(manager, request_a) + + +def test_per_conversation_policy_drops_previous_divergent_blocks( + manager: KVCacheManagerV2, +) -> None: + request_a = _ContextRequest(1, list(range(8)), 8, "conv-1") + request_b = _ContextRequest( + 2, + [*range(8), 100, 101, 102, 103], + 12, + "conv-1", + ) + request_old_prompt = _ContextRequest(3, list(range(8)), 8, "conv-2") + try: + _run_context(manager, request_a) + _free_if_active(manager, request_a) + + _run_context(manager, request_b) + assert request_b.prepopulated_prompt_len == 8 + _free_if_active(manager, request_b) + + assert manager.prepare_context(request_old_prompt) + assert request_old_prompt.prepopulated_prompt_len == 0 + finally: + _free_if_active(manager, request_old_prompt) + _free_if_active(manager, request_b) + _free_if_active(manager, request_a) + + +def test_per_conversation_policy_ignores_overlapping_request( + manager: KVCacheManagerV2, +) -> None: + request_a = _ContextRequest(1, list(range(8)), 8, "conv-1") + request_b = _ContextRequest(2, [0, 1, 2, 3, 100, 101, 102, 103], 8, "conv-1") + request_old_prompt = _ContextRequest(3, list(range(8)), 8, "conv-2") + conversation_params = request_b.py_conversation_params + + try: + batch_a = _prepare_context_resources(manager, request_a) + assert manager.prepare_context(request_a) + assert manager.resize_context(request_a, num_tokens=4) + request_a.context_current_position = 4 + request_a.context_remaining_length = 4 + _update_context_resources(manager, batch_a) + + batch_b = _prepare_context_resources(manager, request_b) + with patch( + "tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2.logger.warning" + ) as mock_warning: + assert manager.prepare_context(request_b) + mock_warning.assert_called_once_with( + "Conversation conv-1 already has current request 1. " + "Request 2 will ignore conversation params." + ) + assert request_b.py_conversation_params is conversation_params + assert manager.resize_context(request_b, num_tokens=request_b.prompt_len) + request_b.context_current_position = request_b.prompt_len + request_b.context_remaining_length = 0 + _update_context_resources(manager, batch_b) + _free_if_active(manager, request_b) + + request_a.is_first_context_chunk = False + batch_a = _prepare_context_resources(manager, request_a) + assert manager.prepare_context(request_a) + assert manager.resize_context(request_a, num_tokens=4) + request_a.context_current_position = 8 + request_a.context_remaining_length = 0 + _update_context_resources(manager, batch_a) + _free_if_active(manager, request_a) + + assert manager.prepare_context(request_old_prompt) + assert request_old_prompt.prepopulated_prompt_len == request_old_prompt.prompt_len - 1 + finally: + _free_if_active(manager, request_old_prompt) + _free_if_active(manager, request_b) + _free_if_active(manager, request_a) diff --git a/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_manager_v2.py b/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_manager_v2.py index 9b8bda967ca3..0922de2a5b44 100755 --- a/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_manager_v2.py +++ b/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_manager_v2.py @@ -46,6 +46,7 @@ KVCacheManagerConfig, LayerGroupId, LayerId, + PlannedDropHandle, ReuseScope, SsmLayerConfig, SwaScratchReuseConfig, @@ -99,6 +100,7 @@ KVCacheManagerConfig, LayerGroupId, LayerId, + PlannedDropHandle, ReuseScope, SsmLayerConfig, SwaScratchReuseConfig, @@ -695,6 +697,46 @@ def test_commit_min_snapshot_reuses_swa_post_commit_prefix(self) -> None: self.assertEqual(kv2.num_committed_tokens, len(prompt)) kv2.close() + def test_planned_drop_handle(self) -> None: + window_size = 8 + self.prepare(16 << 20, 0, 0, 2, window_size, 0, tokens_per_block=8) + long_tokens = [self.next_token() for _ in range(24)] + short_tokens = long_tokens[:8] + + def plan_drop(tokens: list[TokenIdExt]) -> PlannedDropHandle: + kv_cache = self.manager.create_kv_cache(None, tokens) + with TemporaryCudaStream([]) as stream_holder: + stream = cast(CudaStream, stream_holder.handle) + self.assertTrue(kv_cache.resume(stream)) + self.assertTrue(kv_cache.resize(len(tokens))) + uncommitted = tokens[kv_cache.num_committed_tokens :] + if uncommitted: + kv_cache.commit(uncommitted) + kv_cache.stop_committing() + drop_handle = kv_cache.plan_committed_block_drop() + self.assertIsNotNone(drop_handle) + self.assertIsInstance(drop_handle, PlannedDropHandle) + _ = stream_holder.take_finish_event() + kv_cache.close() + assert drop_handle is not None + return drop_handle + + long_handle = plan_drop(long_tokens) + short_handle = plan_drop(short_tokens) + self.assertEqual(self.manager.probe_reuse(None, short_tokens), len(short_tokens)) + + short_handle.drop() + self.assertEqual(self.manager.probe_reuse(None, short_tokens), 0) + self.assertEqual(self.manager.probe_reuse(None, long_tokens), len(long_tokens)) + + long_handle.drop() + # The SWA window is dropped, while older full-attention blocks remain reusable. + self.assertEqual( + self.manager.probe_reuse(None, long_tokens), len(long_tokens) - window_size + ) + with self.assertRaisesRegex(ValueError, "already been dropped"): + long_handle.drop() + def test_reuse_scope_isolates_reuse(self) -> None: self.prepare(16 << 20, 0, 0, 2, None, 0, tokens_per_block=8) tokens = [TokenId(i) for i in range(64)] diff --git a/tests/unittest/llmapi/test_llm_args.py b/tests/unittest/llmapi/test_llm_args.py index 3365ce6685ab..4f17a87123c3 100644 --- a/tests/unittest/llmapi/test_llm_args.py +++ b/tests/unittest/llmapi/test_llm_args.py @@ -601,6 +601,8 @@ def test_KvCacheConfig_declaration(): assert pybind_config.enable_partial_reuse == True assert pybind_config.copy_on_partial_reuse == True assert pybind_config.attention_dp_events_gather_period_ms == 10 + assert (KvCacheConfig(block_reuse_policy="per_conversation"). + block_reuse_policy == "per_conversation") with pytest.raises(ValidationError): KvCacheConfig(block_reuse_policy="invalid")