diff --git a/cpp/include/tensorrt_llm/batch_manager/llmRequest.h b/cpp/include/tensorrt_llm/batch_manager/llmRequest.h index 644813e381b7..32558cdc78b3 100644 --- a/cpp/include/tensorrt_llm/batch_manager/llmRequest.h +++ b/cpp/include/tensorrt_llm/batch_manager/llmRequest.h @@ -1908,6 +1908,15 @@ class GenericLlmRequest return mPerfMetrics.kvCacheMetrics.numNewAllocatedBlocks; } + void updateKvCachePerfMetrics( + SizeType32 allocTotalBlocks, SizeType32 allocNewBlocks, SizeType32 reusedBlocks, SizeType32 missedBlocks) + { + updateAllocTotalBlocksPerRequest(allocTotalBlocks); + updateAllocNewBlocksPerRequest(allocNewBlocks); + updateReusedBlocksPerRequest(reusedBlocks); + updateMissedBlocksPerRequest(missedBlocks); + } + void updateReusedBlocksPerRequest(SizeType32 reusedBlocksPerRequest) { mPerfMetrics.kvCacheMetrics.numReusedBlocks += reusedBlocksPerRequest; diff --git a/cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp b/cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp index ee257713b85d..0846663dafad 100644 --- a/cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp @@ -191,11 +191,19 @@ void initBindings(nb::module_& m) .def_prop_ro("is_disagg_context_complete_state", &GenLlmReq::isDisaggContextCompleteState) .def_prop_ro("stage", &GenLlmReq::getRequestStage) .def_prop_ro("kv_cache_transfer_time_ms", &GenLlmReq::getKvCacheTransferTimeMS) + .def_prop_ro("kv_cache_transfer_start", &GenLlmReq::getKvCacheTransferStart) + .def_prop_ro("kv_cache_transfer_end", &GenLlmReq::getKvCacheTransferEnd) .def_prop_ro("kv_cache_size", &GenLlmReq::getKvCacheSize) + .def("set_kv_cache_transfer_start", &GenLlmReq::setKvCacheTransferStart, nb::arg("time")) + .def("set_kv_cache_transfer_end", &GenLlmReq::setKvCacheTransferEnd, nb::arg("time")) + .def("set_kv_cache_size", &GenLlmReq::setKvCacheSize, nb::arg("target_buffer_size")) + .def("update_kv_cache_size", &GenLlmReq::updateKvCacheSize, nb::arg("target_buffer_size")) .def_prop_ro("avg_decoded_tokens_per_iter", &GenLlmReq::getAvgDecodedTokensPerIter) .def_prop_ro("alloc_total_blocks", &GenLlmReq::getAllocTotalBlocksPerRequest) .def_prop_ro("alloc_new_blocks", &GenLlmReq::getAllocNewBlocksPerRequest) .def("alloc_context_logits", &GenLlmReq::allocContextLogitsHost, nb::arg("vocab_size"), nb::arg("logit_dtype")) + .def("update_kv_cache_perf_metrics", &GenLlmReq::updateKvCachePerfMetrics, nb::arg("alloc_total_blocks"), + nb::arg("alloc_new_blocks"), nb::arg("reused_blocks"), nb::arg("missed_blocks")) .def_prop_ro("reused_blocks", &GenLlmReq::getReusedBlocksPerRequest) .def_prop_ro("missed_blocks", &GenLlmReq::getMissedBlocksPerRequest) .def_prop_ro("kv_cache_hit_rate", &GenLlmReq::getKVCacheHitRatePerRequest) diff --git a/tensorrt_llm/_torch/auto_deploy/_compat.py b/tensorrt_llm/_torch/auto_deploy/_compat.py index 4d0c8e387862..6e79568df948 100644 --- a/tensorrt_llm/_torch/auto_deploy/_compat.py +++ b/tensorrt_llm/_torch/auto_deploy/_compat.py @@ -207,6 +207,7 @@ class KvCacheConfig(BaseModel): cross_kv_cache_fraction: Optional[float] = None secondary_offload_min_priority: Optional[int] = None event_buffer_max_size: int = 0 + kv_cache_event_hash_algo: str = "auto" enable_partial_reuse: bool = True copy_on_partial_reuse: bool = True use_uvm: bool = False diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index f3a6568efd87..d07cfa732fa6 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -236,8 +236,6 @@ def __init__( self._max_beam_width = max_beam_width self._kv_connector_manager = kv_connector_manager self._llm_args = llm_args - # For V2 fallback use only, will be removed after V2 is stable - self._cache_transceiver_config = llm_args.cache_transceiver_config self._speculative_config = speculative_config self._sparse_attention_config = sparse_attention_config self._tokens_per_block = tokens_per_block @@ -247,6 +245,7 @@ def __init__( self._dummy_reqs = None self._profiling_stage_data = profiling_stage_data self._is_disagg = is_disagg + self._cache_transceiver_config = llm_args.cache_transceiver_config self._execution_stream = execution_stream self._kv_cache_manager_cls = self._get_model_kv_cache_manager_cls( model_engine) @@ -261,30 +260,46 @@ def _get_model_kv_cache_manager_cls( kv_cache_config = (kv_cache_config_override if kv_cache_config_override is not None else self._kv_cache_config) model_config = model_engine.model.model_config - config = model_config.pretrained_config cls = get_kv_cache_manager_cls( model_config, kv_cache_config, is_disagg=self._is_disagg, cache_transceiver_config=self._cache_transceiver_config) + cls = self._fallback_if_unsupported_kv_cache_manager_v2( + cls, model_config, kv_cache_config) + # The V1-route hybrid mamba managers (disagg, TRTLLM_USE_CPP_MAMBA, + # TRTLLM_USE_PY_MAMBA, or one-model speculative decoding) keep mamba + # state in a separate cache that doesn't honor block reuse. Warn at + # the routing site so users see the warning where the decision is + # actually made. + if is_hybrid_linear(model_engine.model.model_config.pretrained_config) \ + and kv_cache_config.enable_block_reuse: + uses_v1_mamba_route = self._is_disagg \ + or os.environ.get('TRTLLM_USE_CPP_MAMBA', '0') == '1' \ + or os.environ.get('TRTLLM_USE_PY_MAMBA', '0') == '1' \ + or self._speculative_config is not None + if uses_v1_mamba_route: + logger.warning( + "Block reuse does not work with MTP for hybrid linear models " + "when using the legacy MambaCacheManager (TRTLLM_USE_CPP_MAMBA=1)" + ) + return cls + + def _fallback_if_unsupported_kv_cache_manager_v2( + self, + kv_cache_manager_cls, + model_config: ModelConfig, + kv_cache_config: Optional[KvCacheConfig] = None): + config = model_config.pretrained_config # Use ``issubclass`` rather than identity equality so V2 subclasses # (e.g. ``MiniMaxM3KVCacheManagerV2`` from the sparse-attention path) - # also go through the V2-incompatible-feature gate below. The earlier - # ``cls == KVCacheManagerV2`` check silently bypassed all V2 subclasses - # and let the bare ``assert event_buffer_max_size == 0`` in - # ``KVCacheManagerV2.__init__`` trip at executor construction with - # a useless error message. - if issubclass(cls, KVCacheManagerV2): + # also go through the V2-incompatible-feature gate below. + if issubclass(kv_cache_manager_cls, KVCacheManagerV2): incompat: List[str] = [] if self._kv_connector_manager is not None: incompat.append("kv_connector_manager") if self._max_beam_width is not None and self._max_beam_width > 1: incompat.append("beam_width > 1") - if kv_cache_config.event_buffer_max_size > 0: - incompat.append("event_buffer_max_size > 0") - if (self._cache_transceiver_config is not None - and self._cache_transceiver_config.backend is not None): - incompat.append("cache_transceiver") if incompat: incompat_str = ", ".join(incompat) # Some models are structurally bound to V2 and cannot fall @@ -315,24 +330,8 @@ def _get_model_kv_cache_manager_cls( logger.warning( "KVCacheManagerV2 is not supported with %s. " "Falling back to KVCacheManager.", incompat_str) - cls = KVCacheManager - # The V1-route hybrid mamba managers (disagg, TRTLLM_USE_CPP_MAMBA, - # TRTLLM_USE_PY_MAMBA, or one-model speculative decoding) keep mamba - # state in a separate cache that doesn't honor block reuse. Warn at - # the routing site so users see the warning where the decision is - # actually made. - if is_hybrid_linear(model_engine.model.model_config.pretrained_config) \ - and kv_cache_config.enable_block_reuse: - uses_v1_mamba_route = self._is_disagg \ - or os.environ.get('TRTLLM_USE_CPP_MAMBA', '0') == '1' \ - or os.environ.get('TRTLLM_USE_PY_MAMBA', '0') == '1' \ - or self._speculative_config is not None - if uses_v1_mamba_route: - logger.warning( - "Block reuse does not work with MTP for hybrid linear models " - "when using the legacy MambaCacheManager (TRTLLM_USE_CPP_MAMBA=1)" - ) - return cls + return KVCacheManager + return kv_cache_manager_cls def _per_manager_cache_cost(self, manager_cls, @@ -943,18 +942,8 @@ def _create_one_model_draft_kv_cache_manager( # Get the appropriate KV cache manager class for the draft model draft_kv_cache_manager_cls = get_kv_cache_manager_cls( effective_draft_config, draft_kv_config, is_disagg=self._is_disagg) - - # Use V2 if enabled and the base class is KVCacheManager - if draft_kv_cache_manager_cls == KVCacheManagerV2: - if self._kv_connector_manager is not None or ( - self._max_beam_width is not None and self._max_beam_width - > 1) or draft_kv_config.event_buffer_max_size > 0 or ( - self._cache_transceiver_config is not None - and self._cache_transceiver_config.backend is not None): - logger.warning( - "KVCacheManagerV2 is not supported with disaggregated serving or beam width > 1 or event buffer max size > 0 or disagg config. " - "Falling back to KVCacheManager for draft model.") - draft_kv_cache_manager_cls = KVCacheManager + draft_kv_cache_manager_cls = self._fallback_if_unsupported_kv_cache_manager_v2( + draft_kv_cache_manager_cls, effective_draft_config, draft_kv_config) estimating_kv_cache = estimating_kv_cache and not self._skip_est # For MTP with models using sparse attention (e.g., DeepSeek V3 with DSA), @@ -2086,6 +2075,18 @@ def create_py_executor_instance( if cross_kv_cache_manager is not None else LlmRequestState.CONTEXT_INIT) + # V2 scheduler uses scheduler_capacity as the per-iteration request + # budget (BudgetTracker.max_num_requests). Unlike V1 which has a + # separate CapacityScheduler (needs pp_size * max_batch_size to hold + # requests across PP stages) and MicroBatchScheduler (uses + # max_batch_size for per-forward batch limit), V2 merges both into + # one loop. PP on-the-fly is handled by inflight_request_ids + # filtering, so its budget should be based on max_batch_size, not + # max_num_sequences (which includes the pp_size multiplier). + v2_scheduler_capacity = max_batch_size + if v2_scheduler_capacity == 1 and mapping.enable_attention_dp and kv_cache_manager: + v2_scheduler_capacity += 1 + if isinstance(kv_cache_manager, KVCacheManagerV2): # V2: interleaved scheduler handles both capacity and budget draft_kv_cache_manager = resources.get( @@ -2101,7 +2102,7 @@ def create_py_executor_instance( ctx_chunk_config=ctx_chunk_config, peft_cache_manager=peft_cache_manager.impl if peft_cache_manager is not None else None, - scheduler_capacity=scheduler_capacity, + scheduler_capacity=v2_scheduler_capacity, draft_kv_cache_manager=draft_kv_cache_manager, cross_kv_cache_manager=cross_kv_cache_manager, no_schedule_until_state=no_schedule_until_state, diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py index 02e48743465c..caf80ecba7b5 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py @@ -12,11 +12,12 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. - import hashlib import math import os -from typing import TYPE_CHECKING, Iterable, List, NamedTuple, Optional, Sequence, Tuple, Union +from collections import OrderedDict, defaultdict +from dataclasses import fields +from typing import TYPE_CHECKING, Dict, Iterable, List, NamedTuple, Optional, Sequence, Tuple, Union import torch @@ -27,12 +28,13 @@ get_size_in_bytes, prefer_pinned, ) -from tensorrt_llm.bindings.internal.batch_manager import KvCacheStats +from tensorrt_llm.bindings.internal.batch_manager import KvCacheIterationStats, KvCacheStats from tensorrt_llm.bindings.internal.batch_manager.kv_cache_manager_v2_utils import ( IndexMapper, copy_batch_block_offsets_to_device, ) from tensorrt_llm.llmapi.llm_args import KvCacheConfig +from tensorrt_llm.runtime.kv_cache_hash import get_effective_kv_cache_event_hash_algo from tensorrt_llm.runtime.kv_cache_manager_v2 import ( DEFAULT_BEAM_INDEX, AttentionLayerConfig, @@ -41,6 +43,7 @@ DiskCacheTierConfig, GpuCacheTierConfig, HostCacheTierConfig, + KVCacheIterationStatsDelta, LayerId, ReuseScope, TokenIdExt, @@ -51,15 +54,31 @@ 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, CACHE_LEVEL1, GPU_LEVEL +from tensorrt_llm.runtime.kv_cache_manager_v2._common import ( + BAD_PAGE_INDEX, + CACHE_LEVEL1, + GPU_LEVEL, + CacheLevel, +) from tensorrt_llm.runtime.kv_cache_manager_v2._config import DataRole +from tensorrt_llm.runtime.kv_cache_manager_v2._event_manager import KVCacheEventManager +from tensorrt_llm.runtime.kv_cache_manager_v2._exceptions import CuError +from tensorrt_llm.runtime.kv_cache_manager_v2._exceptions import ( + OutOfMemoryError as KVCacheOutOfMemoryError, +) +from tensorrt_llm.runtime.kv_cache_manager_v2._life_cycle_registry import AttnLifeCycle, LifeCycleId from tensorrt_llm.runtime.kv_cache_manager_v2._utils import exact_div, typed_range from tensorrt_llm.sampling_params import SamplingParams -from ..._utils import binding_to_torch_dtype, nvtx_range, str_dtype_to_torch +from ..._utils import binding_to_torch_dtype, mpi_rank, nvtx_range, str_dtype_to_torch from ...logger import logger from ...mapping import CpType, Mapping from .connectors.kv_cache_connector import KvCacheConnectorManager +from .kv_cache_stats import ( + KVCacheV2IterationStatsReport, + KVCacheV2LifeCycleIterationStats, + KVCacheV2PoolGroupIterationStats, +) from .llm_request import LlmRequest, LlmRequestState, SamplingConfig, get_draft_token_length from .resource_manager import ( BaseResourceManager, @@ -77,6 +96,21 @@ if TYPE_CHECKING: from tensorrt_llm._torch.attention_backend.interface import AttentionMetadata +KV_CACHE_ITERATION_STATS_DELTA_FIELDS = tuple( + field.name for field in fields(KVCacheIterationStatsDelta) +) +KV_CACHE_ITERATION_STATS_REUSE_FIELDS = ( + "iter_reused_blocks", + "iter_full_reused_blocks", + "iter_partial_reused_blocks", + "iter_missed_blocks", +) +KV_CACHE_ITERATION_STATS_POOL_GROUP_FIELDS = tuple( + field_name + for field_name in KV_CACHE_ITERATION_STATS_DELTA_FIELDS + if field_name not in KV_CACHE_ITERATION_STATS_REUSE_FIELDS +) + class Role: KEY = DataRole("key") @@ -431,6 +465,7 @@ def __init__( kv_connector_manager: Optional[KvCacheConnectorManager] = None, execution_stream: Optional[torch.cuda.Stream] = None, is_disagg: bool = False, + enable_stats: bool = False, **kwargs, ) -> None: self.mapping = mapping @@ -486,8 +521,11 @@ def __init__( self._kv_reserve_draft_tokens = max(self.max_total_draft_tokens, draft_loop_tokens) self.event_buffer_max_size = kv_cache_config.event_buffer_max_size - - assert self.event_buffer_max_size == 0, "event_buffer_max_size must be 0" + self.enable_stats = enable_stats + kv_cache_event_hash_algo = get_effective_kv_cache_event_hash_algo( + kv_cache_config.kv_cache_event_hash_algo, + use_kv_cache_manager_v2=True, + ) self._stream = ( execution_stream if execution_stream is not None else torch.cuda.current_stream() @@ -513,6 +551,27 @@ def __init__( else: self.max_attention_window_vec = [None] + event_window_size = max( + self.max_seq_len if window_size is None else int(window_size) + for window_size in self.max_attention_window_vec + ) + self.event_manager: Optional[KVCacheEventManager] = None + if self.event_buffer_max_size > 0: + if mapping.enable_attention_dp: + self.event_manager = KVCacheEventManager( + self.event_buffer_max_size, + window_size=event_window_size, + attention_dp_rank=mapping.rank, + attention_dp_gather=Distributed.get(mapping).allgather, + hash_algo=kv_cache_event_hash_algo, + ) + elif mpi_rank() == 0: + self.event_manager = KVCacheEventManager( + self.event_buffer_max_size, + window_size=event_window_size, + hash_algo=kv_cache_event_hash_algo, + ) + if isinstance(num_kv_heads, int): self.num_kv_heads_per_layer = [ (num_kv_heads + tp_size - 1) // tp_size for _ in range(self.num_local_layers) @@ -577,9 +636,9 @@ def append_to_kv_heads_per_layer( ) quota = min(quota, quota_from_max_tokens) logger.info( - f"max_tokens {kv_cache_config.max_tokens} is provided. Allowed quota from " - f"max_tokens is {quota_from_max_tokens / (1 << 30)}GiB. New quota is " - f"{quota / (1 << 30)}GiB" + f"max_tokens {kv_cache_config.max_tokens} is provided. " + f"Allowed quota from max_tokens is {quota_from_max_tokens / (1 << 30)}GiB. " + f"New quota is {quota / (1 << 30)}GiB" ) assert quota != float("inf"), ( @@ -600,7 +659,7 @@ def append_to_kv_heads_per_layer( logger.info(f"KV cache manager v2 device quota set to {quota / (1 << 30)}GiB") cache_tiers: List[CacheTierConfig] = [GpuCacheTierConfig(quota=quota)] - if kv_cache_config.host_cache_size is not None and kv_cache_config.host_cache_size > 0: + if kv_cache_config.host_cache_size is not None and kv_cache_config.host_cache_size >= 0: host_quota = kv_cache_config.host_cache_size else: # The V2 MAX_UTILIZATION scheduler relies on suspend/resume to @@ -611,12 +670,24 @@ def append_to_kv_heads_per_layer( # # Automatically provision a host tier matching the GPU quota so # suspend/resume works out of the box. Cap at available host - # memory to avoid allocation failures. + # memory and pinnable memory limit to avoid allocation failures. + import resource + try: mem_available = os.sysconf("SC_PAGE_SIZE") * os.sysconf("SC_AVPHYS_PAGES") except (ValueError, OSError): mem_available = float("inf") - host_quota = min(quota, int(mem_available * 0.5)) + try: + _soft, _hard = resource.getrlimit(resource.RLIMIT_MEMLOCK) + memlock_limit = _soft if _soft != resource.RLIM_INFINITY else float("inf") + except (ValueError, OSError): + memlock_limit = float("inf") + candidates = [quota] + if mem_available != float("inf"): + candidates.append(int(mem_available * 0.5)) + if memlock_limit != float("inf"): + candidates.append(int(memlock_limit * 0.8)) + host_quota = min(candidates) if host_quota <= 0: host_quota = quota if host_quota > 0: @@ -644,7 +715,35 @@ def append_to_kv_heads_per_layer( self.kv_cache_manager_py_config = config - self.impl = KVCacheManagerPy(config) + try: + self.impl = KVCacheManagerPy(config, event_manager=self.event_manager) + except (CuError, KVCacheOutOfMemoryError): + if len(cache_tiers) > 1: + logger.warning( + "Failed to initialize KV cache manager with host cache " + "tier (cuMemHostRegister may have failed). " + "Retrying without host cache tier." + ) + cache_tiers_gpu_only = [t for t in cache_tiers if isinstance(t, GpuCacheTierConfig)] + config = self._build_cache_config( + kv_cache_config, + tokens_per_block=tokens_per_block, + vocab_size=vocab_size, + cache_tiers=cache_tiers_gpu_only, + ) + cache_tiers = cache_tiers_gpu_only + self.kv_cache_manager_py_config = config + self.impl = KVCacheManagerPy(config, event_manager=self.event_manager) + else: + raise + if self.event_manager is not None: + self.event_manager.set_layer_group_window_sizes( + self._get_event_window_sizes_by_layer_group() + ) + self.event_manager.add_created_event( + self._get_event_num_blocks_per_cache_level(cache_tiers, tokens_per_block), + self._get_event_layer_group_ids(), + ) self.num_pools = len(self.impl.layer_grouping) @@ -680,8 +779,8 @@ def append_to_kv_heads_per_layer( if max_seq_len > max_num_tokens: logger.warning( - f"max_seq_len {max_seq_len} is greater than max_num_tokens " - f"{max_num_tokens} that can be allocated in kv cache manager, setting " + f"max_seq_len {max_seq_len} is greater than max_num_tokens {max_num_tokens} " + "that can be allocated in kv cache manager, setting " f"max_seq_len to {max_num_tokens}" ) # max_num_tokens is a float from clamp_max_seq_len_for_mem; cast @@ -749,9 +848,64 @@ def append_to_kv_heads_per_layer( device="cpu", ) + self._log_kv_cache_pool_lifecycle_mapping() + def _get_quota_from_max_tokens(self, max_tokens: int) -> int: return int(max_tokens * self.get_cache_bytes_per_token()) + def _get_event_num_blocks_per_cache_level( + self, + cache_tiers: List[CacheTierConfig], + tokens_per_block: int, + ) -> List[int]: + bytes_per_block = self.get_cache_bytes_per_token() * tokens_per_block + if bytes_per_block <= 0: + return [] + return [int(tier.quota // bytes_per_block) for tier in cache_tiers] + + def _get_event_layer_group_ids(self) -> List[int]: + return [int(layer_group_id) for layer_group_id in range(len(self.impl.layer_grouping))] + + def _get_event_window_sizes_by_layer_group(self) -> Dict[int, int]: + # Assumes every layer in a group shares the same sliding_window_size, + # which is how `impl.layer_grouping` partitions layers today. Only the + # first layer's window is read; if the grouping policy ever permits + # mixed windows in one group, this needs to fan out per-layer. + + def get_event_window_size(layer_id: int) -> int: + window_size = self.kv_cache_manager_py_config.layers[layer_id].sliding_window_size + return self.max_seq_len if window_size is None else int(window_size) + + return { + int(layer_group_id): get_event_window_size(int(layer_ids[0])) + for layer_group_id, layer_ids in enumerate(self.impl.layer_grouping) + } + + def _format_kv_cache_pool_lifecycle_entry(self, layer_id: LayerId, role: DataRole) -> str: + attr = self.impl._storage.get_buffer_attr(layer_id, role) + pool_group_id = self.impl._storage.get_pool_group_index(attr.life_cycle_id) + lifecycle = self.impl._life_cycles.get_life_cycle(attr.life_cycle_id) + return ( + f"role={str(role)}, pool_group_id={int(pool_group_id)}, " + f"lifecycle_id={int(attr.life_cycle_id)}, " + f"lifecycle={lifecycle}" + ) + + def _log_kv_cache_pool_lifecycle_mapping(self) -> None: + entries = OrderedDict() + for layer in self.kv_cache_manager_py_config.layers: + for buffer in layer.buffers: + entries.setdefault( + self._format_kv_cache_pool_lifecycle_entry(layer.layer_id, buffer.role), None + ) + + if not entries: + return + + logger.info(f"{type(self).__name__} role-to-pool/lifecycle mapping:") + for entry in entries: + logger.info(entry) + def _build_pool_mapping_tensors(self) -> Tuple[torch.Tensor, torch.Tensor]: kv_cache_pool_pointers = torch.tensor( [ @@ -892,6 +1046,7 @@ def _build_cache_config( vocab_size=vocab_size, cache_tiers=cache_tiers, max_util_for_resume=kv_cache_config.max_util_for_resume, + enable_stats=self.enable_stats, layers=layer_configs, ) @@ -1103,6 +1258,24 @@ def get_num_free_blocks(self) -> int: ) return max_num_pages // self.kv_factor + def commit_scheduled_kv_cache_stats(self, scheduled_batch: ScheduledRequests) -> None: + if self.is_draft or not self.enable_stats: + return + dirty_req_ids = self.impl.get_dirty_stats_kv_cache_ids() + for req in scheduled_batch.all_requests(): + if req.py_request_id in dirty_req_ids: + kv_cache = self.kv_cache_map.get(req.py_request_id) + if kv_cache is None: + continue + request_stats = kv_cache.commit_pending_stats() + if not req.is_dummy and not request_stats.empty: + req.update_kv_cache_perf_metrics( + request_stats.alloc_total_blocks, + request_stats.alloc_new_blocks, + request_stats.reused_blocks, + request_stats.missed_blocks, + ) + # ---- Scheduling API (called by KVCacheV2Scheduler) ---- def is_request_active(self, request_id: int) -> bool: @@ -1110,13 +1283,36 @@ def is_request_active(self, request_id: int) -> bool: kv_cache = self.kv_cache_map.get(request_id) return kv_cache is not None and kv_cache.is_active + def _effective_draft_len(self, req: LlmRequest) -> int: + """Draft token length to use for next-step KV capacity calculation. + + For a disagg gen request whose KV transmission just completed + (state == DISAGG_GENERATION_TRANS_COMPLETE), py_draft_tokens is + still [] when the scheduler asks for capacity, because it gets + mirrored from context_phase_params.draft_tokens later in + _prepare_disagg_gen_transmission_complete (which runs AFTER the + scheduler in the executor loop). Without compensating here, the + first gen forward writes 1 + len(ctx_draft_tokens) tokens into + KV cache but only +1 was reserved, OOB-ing the KV block table at + the next tokens_per_block-aligned boundary. + """ + draft_len = get_draft_token_length(req) + if ( + draft_len == 0 + and req.is_disagg_generation_transmission_complete + and req.context_phase_params is not None + ): + ctx_draft_tokens = req.context_phase_params.draft_tokens + if ctx_draft_tokens is not None: + draft_len = len(ctx_draft_tokens) + return draft_len + def _required_gen_capacity(self, req: LlmRequest, current_capacity: int) -> int: """Compute generation KV cache capacity for a request. Grows *current_capacity* by 1 + draft tokens. """ - draft_len = get_draft_token_length(req) - return current_capacity + 1 + draft_len + return current_capacity + 1 + self._effective_draft_len(req) def try_allocate_generation(self, req: LlmRequest) -> bool: """Try to allocate one additional KV cache slot for a generation request. @@ -1133,10 +1329,50 @@ def try_allocate_generation(self, req: LlmRequest) -> bool: return False self._restore_page_index_bufs(req.py_request_id, kv_cache) - draft_len = get_draft_token_length(req) + draft_len = self._effective_draft_len(req) self._allocated_draft_lens[req.py_request_id] = draft_len return kv_cache.resize(self._required_gen_capacity(req, kv_cache.capacity)) + def trim_to_history(self, req: LlmRequest, history_length: int) -> bool: + """Mark *history_length* tokens of this request's KV as historic. + + For sliding-window-style life cycles (AttnLifeCycle with non-None + window_size), this triggers ``_unlock_stale_blocks`` inside V2's + ``resize()`` so blocks before ``(history_length + 1 - window) // + tokens_per_block`` get released back to their pool group. For + full-context life cycles (``window_size=None``) and SSM cycles, it + is a no-op — the stale range stays empty. + + Used by the disagg-gen transceiver right after KV transfer + completes: at that moment the cache has the entire prompt KV + written, so ``history_length=prompt_len`` correctly classifies + every prompt token as historic. Without this call, ``history_length`` + stays 0 until ``update_resources`` runs after the first forward + pass — and in benchmark fill-phase the first forward never fires + until every disagg-gen request is ready, so SWA / sparse-attn + pool groups would otherwise stay 100% occupied with pre-window + prompt blocks and the V2 scheduler would deadlock on the next + ``resize(+1)``. + + Returns True on success (or no-op), False if the underlying + ``kv_cache.resize`` rejected the call. + """ + kv_cache = self.kv_cache_map.get(req.py_request_id) + if kv_cache is None or not kv_cache.is_active: + return True + if history_length <= kv_cache.history_length: + return True + # resize() requires capacity >= history_length; clamp for safety. + target_capacity = max(kv_cache.capacity, history_length) + try: + return kv_cache.resize(target_capacity, history_length=history_length) + except Exception as e: + logger.warning( + f"trim_to_history failed for req {req.py_request_id} " + f"(capacity={kv_cache.capacity}, target_history={history_length}): {e}" + ) + return False + def revert_allocate_generation(self, req: LlmRequest) -> None: """Undo the capacity growth from try_allocate_generation. @@ -1146,11 +1382,16 @@ def revert_allocate_generation(self, req: LlmRequest) -> None: This method shrinks capacity back to undo that spurious growth so it does not accumulate across iterations and overflow the host page-index buffer. + + Mirror the effective draft length used in _required_gen_capacity + so disagg-gen-trans-complete revert stays symmetric. """ kv_cache = self.kv_cache_map.get(req.py_request_id) if kv_cache is None or not kv_cache.is_active: return - draft_len = get_draft_token_length(req) + draft_len = self._allocated_draft_lens.pop( + req.py_request_id, self._effective_draft_len(req) + ) reverted_cap = kv_cache.capacity - 1 - draft_len if reverted_cap < 0: return @@ -1234,17 +1475,22 @@ def prepare_context(self, req: LlmRequest) -> bool: if req.is_first_context_chunk: kv_cache = self.kv_cache_map.get(req.py_request_id) if kv_cache is None: + all_tokens = req.get_tokens(DEFAULT_BEAM_INDEX) # Last token cannot be recovered, so we don't include it in # the input tokens to look up for the block that can be reused. if self.enable_block_reuse: - all_tokens = req.get_tokens(DEFAULT_BEAM_INDEX) tokens = self._augment_tokens_for_block_reuse( all_tokens, req, end=len(all_tokens) - 1 ) else: tokens = None kv_cache = self._create_kv_cache( - req.py_request_id, req.lora_task_id, tokens, cache_salt=req.cache_salt + req.py_request_id, + req.lora_task_id, + tokens, + cache_salt=req.cache_salt, + is_dummy=req.is_dummy, + expected_prompt_length=req.prompt_len - 1, ) if kv_cache is None: return False @@ -1269,13 +1515,13 @@ def prepare_context(self, req: LlmRequest) -> bool: ) return self._resume_and_restore(req.py_request_id, kv_cache) - def resize_context(self, req: LlmRequest, num_tokens: int) -> bool: + def resize_context( + self, req: LlmRequest, num_tokens: int, history_length: int | None = None + ) -> bool: """Resize KV cache to cover context_current_position + num_tokens. - num_tokens is the number of tokens to be processed (i.e., - context_remaining_length or a chunk thereof). The target capacity is - computed as context_current_position + num_tokens so that block reuse - overlaps with existing capacity are handled correctly. + history_length, when set, lets SWA life cycles compute their stale + range at allocation time so pre-window blocks are never allocated. Returns True on success, False if resize failed (first chunk is suspended on failure). @@ -1291,7 +1537,8 @@ def resize_context(self, req: LlmRequest, num_tokens: int) -> bool: capacity = max(kv_cache.capacity, target) pre_cap = kv_cache.capacity - if not kv_cache.resize(capacity): + success = kv_cache.resize(capacity, history_length) + if not success: if req.is_first_context_chunk: kv_cache.suspend() return False @@ -1371,7 +1618,11 @@ def _prepare_draft_resources(self, scheduled_batch: ScheduledRequests): kv_cache = self.kv_cache_map.get(req.py_request_id) if kv_cache is None: kv_cache = self._create_kv_cache( - req.py_request_id, req.lora_task_id, None, cache_salt=req.cache_salt + req.py_request_id, + req.lora_task_id, + None, + cache_salt=req.cache_salt, + is_dummy=req.is_dummy, ) kv_cache.stop_committing() if not self._resume_and_restore(req.py_request_id, kv_cache): @@ -1461,15 +1712,315 @@ def _augment_tokens_for_block_reuse( chunk_end, ) + def _stats_window_size(self, window_size: Optional[int]) -> int: + return self.max_seq_len if window_size is None else int(window_size) + + def _stats_life_cycle_window_size(self, life_cycle) -> Optional[int]: + if not isinstance(life_cycle, AttnLifeCycle): + return None + return self._stats_window_size(life_cycle.window_size) + + def _storage_pool_groups_by_window(self) -> dict[int, set[int]]: + pool_groups_by_window: dict[int, set[int]] = defaultdict(set) + for life_cycle_id, life_cycle in self.impl._life_cycles.attention_life_cycles(): + pool_group_id = self.impl._storage.get_pool_group_index(life_cycle_id) + pool_groups_by_window[self._stats_window_size(life_cycle.window_size)].add( + int(pool_group_id) + ) + return pool_groups_by_window + + @staticmethod + def _windows_by_pool_group( + pool_groups_by_window: dict[int, set[int]], + ) -> dict[int, tuple[int, ...]]: + windows_by_pool_group: dict[int, set[int]] = defaultdict(set) + for window_size, pool_group_ids in pool_groups_by_window.items(): + for pool_group_id in pool_group_ids: + windows_by_pool_group[pool_group_id].add(window_size) + return { + pool_group_id: tuple(sorted(window_sizes)) + for pool_group_id, window_sizes in windows_by_pool_group.items() + } + + @staticmethod + def _filter_iteration_stats_delta(delta, field_names) -> KVCacheIterationStatsDelta: + filtered = KVCacheIterationStatsDelta() + for field_name in field_names: + setattr(filtered, field_name, getattr(delta, field_name)) + return filtered + + @staticmethod + def _add_iteration_stats_delta( + bucket: dict[int, KVCacheIterationStatsDelta], key: int, delta: KVCacheIterationStatsDelta + ) -> None: + if delta.empty: + return + if key not in bucket: + bucket[key] = delta.copy() + return + bucket[key].add(delta) + + @staticmethod + def _iteration_cache_hit_rate(stats) -> float: + total = stats.iter_reused_blocks + stats.iter_missed_blocks + if stats.iter_reused_blocks == 0 or total == 0: + return 0.0 + return stats.iter_reused_blocks / total + + @staticmethod + def _apply_iteration_stats_delta( + stats, delta, field_names=KV_CACHE_ITERATION_STATS_DELTA_FIELDS + ) -> None: + if delta is None: + return + for field_name in field_names: + setattr(stats, field_name, getattr(delta, field_name)) + stats.iter_cache_hit_rate = KVCacheManagerV2._iteration_cache_hit_rate(stats) + + def _build_iteration_stats( + self, + pool_group_ids: Iterable[int], + primary_stats, + secondary_stats_by_level, + delta, + field_names=KV_CACHE_ITERATION_STATS_DELTA_FIELDS, + ): + pool_group_ids = tuple(pool_group_ids) + stats = KvCacheIterationStats() + stats.primary_max_num_blocks = sum( + primary_stats[pool_group_id].total for pool_group_id in pool_group_ids + ) + stats.primary_free_num_blocks = sum( + primary_stats[pool_group_id].available for pool_group_id in pool_group_ids + ) + stats.primary_used_num_blocks = stats.primary_max_num_blocks - stats.primary_free_num_blocks + stats.secondary_max_num_blocks = sum( + level_stats[pool_group_id].total + for level_stats in secondary_stats_by_level + for pool_group_id in pool_group_ids + ) + stats.secondary_free_num_blocks = sum( + level_stats[pool_group_id].available + for level_stats in secondary_stats_by_level + for pool_group_id in pool_group_ids + ) + stats.secondary_used_num_blocks = ( + stats.secondary_max_num_blocks - stats.secondary_free_num_blocks + ) + self._apply_iteration_stats_delta(stats, delta, field_names) + return stats + + def _collect_iteration_stats_deltas( + self, raw_iteration_stats, storage + ) -> tuple[dict, dict, dict, dict]: + reuse_deltas_by_window: dict[int, KVCacheIterationStatsDelta] = {} + reuse_deltas_by_life_cycle: dict[int, KVCacheIterationStatsDelta] = {} + pool_group_deltas_by_window: dict[int, KVCacheIterationStatsDelta] = {} + pool_group_deltas: dict[int, KVCacheIterationStatsDelta] = {} + + for life_cycle_id, delta in raw_iteration_stats.items(): + life_cycle = self.impl._life_cycles.get_life_cycle(life_cycle_id) + pool_group_id = int(storage.get_pool_group_index(life_cycle_id)) + window_size = self._stats_life_cycle_window_size(life_cycle) + + pool_group_delta = self._filter_iteration_stats_delta( + delta, KV_CACHE_ITERATION_STATS_POOL_GROUP_FIELDS + ) + self._add_iteration_stats_delta(pool_group_deltas, pool_group_id, pool_group_delta) + if window_size is not None: + self._add_iteration_stats_delta( + pool_group_deltas_by_window, window_size, pool_group_delta + ) + + reuse_delta = self._filter_iteration_stats_delta( + delta, KV_CACHE_ITERATION_STATS_REUSE_FIELDS + ) + if reuse_delta.empty: + continue + reuse_deltas_by_life_cycle[int(life_cycle_id)] = reuse_delta.copy() + if window_size is not None: + self._add_iteration_stats_delta(reuse_deltas_by_window, window_size, reuse_delta) + + return ( + reuse_deltas_by_window, + reuse_deltas_by_life_cycle, + pool_group_deltas_by_window, + pool_group_deltas, + ) + + def _build_window_iteration_stats( + self, + window_size: int, + pool_groups_by_window: dict[int, set[int]], + windows_by_pool_group: dict[int, tuple[int, ...]], + primary_stats, + secondary_stats_by_level, + pool_group_delta, + reuse_delta, + ): + pool_group_ids = tuple( + pool_group_id + for pool_group_id in pool_groups_by_window.get(window_size, set()) + if windows_by_pool_group.get(pool_group_id) == (window_size,) + ) + stats = self._build_iteration_stats( + pool_group_ids, + primary_stats, + secondary_stats_by_level, + pool_group_delta, + KV_CACHE_ITERATION_STATS_POOL_GROUP_FIELDS, + ) + self._apply_iteration_stats_delta(stats, reuse_delta, KV_CACHE_ITERATION_STATS_REUSE_FIELDS) + return stats + + def _build_pool_group_iteration_stats( + self, + pool_group_id: int, + windows_by_pool_group: dict[int, tuple[int, ...]], + primary_stats, + secondary_stats_by_level, + pool_group_delta, + ) -> KVCacheV2PoolGroupIterationStats: + return KVCacheV2PoolGroupIterationStats( + pool_group_id=pool_group_id, + slot_size=tuple(primary_stats[pool_group_id].slot_size), + window_sizes=windows_by_pool_group.get(pool_group_id, ()), + stats=self._build_iteration_stats( + (pool_group_id,), + primary_stats, + secondary_stats_by_level, + pool_group_delta, + KV_CACHE_ITERATION_STATS_POOL_GROUP_FIELDS, + ), + ) + + def _build_life_cycle_iteration_stats( + self, + life_cycle_id: int, + storage, + primary_stats, + secondary_stats_by_level, + reuse_delta, + ) -> KVCacheV2LifeCycleIterationStats: + typed_life_cycle_id = LifeCycleId(life_cycle_id) + life_cycle = self.impl._life_cycles.get_life_cycle(typed_life_cycle_id) + pool_group_id = int(storage.get_pool_group_index(typed_life_cycle_id)) + return KVCacheV2LifeCycleIterationStats( + life_cycle_id=life_cycle_id, + pool_group_id=pool_group_id, + window_size=self._stats_life_cycle_window_size(life_cycle), + kind="attention" if isinstance(life_cycle, AttnLifeCycle) else "ssm", + stats=self._build_iteration_stats( + (), + primary_stats, + secondary_stats_by_level, + reuse_delta, + KV_CACHE_ITERATION_STATS_REUSE_FIELDS, + ), + ) + def get_kv_cache_stats(self): kv_cache_stats = KvCacheStats() - kv_cache_stats.allocated_bytes = self.impl.get_quota(GPU_LEVEL) + storage_stats = self.impl._get_storage_level_stats(GPU_LEVEL) + pool_group_stats = storage_stats.pool_group_stats + committed_stats = self.impl.get_committed_stats() + + kv_cache_stats.max_num_blocks = storage_stats.max_num_blocks + kv_cache_stats.free_num_blocks = storage_stats.free_num_blocks + kv_cache_stats.used_num_blocks = storage_stats.used_num_blocks + kv_cache_stats.tokens_per_block = self.tokens_per_block + kv_cache_stats.alloc_total_blocks = committed_stats.alloc_total_blocks + kv_cache_stats.alloc_new_blocks = committed_stats.alloc_new_blocks + kv_cache_stats.reused_blocks = committed_stats.reused_blocks + kv_cache_stats.missed_blocks = committed_stats.missed_blocks + total = kv_cache_stats.reused_blocks + kv_cache_stats.missed_blocks + kv_cache_stats.cache_hit_rate = ( + 0.0 + if kv_cache_stats.reused_blocks == 0 or total == 0 + else kv_cache_stats.reused_blocks / total + ) + kv_cache_stats.num_free_blocks_per_window_size = { + window_size: sum( + pool_group_stats[pool_group_id].available for pool_group_id in pool_group_ids + ) + for window_size, pool_group_ids in self._storage_pool_groups_by_window().items() + } + kv_cache_stats.allocated_bytes = storage_stats.allocated_bytes return kv_cache_stats + def flush_iteration_events(self): + if self.event_manager is not None: + self.event_manager.flush_iteration_events() + + def get_latest_events(self, timeout_ms: Optional[float] = None): + if self.event_manager is None: + return [] + return self.event_manager.get_latest_events(timeout_ms) + def get_iteration_stats(self): - """V2 does not support per-iteration stats yet.""" - return None + if not self.enable_stats: + return None + + storage = self.impl._storage + pool_groups_by_window = self._storage_pool_groups_by_window() + windows_by_pool_group = self._windows_by_pool_group(pool_groups_by_window) + raw_iteration_stats = self.impl.get_and_reset_iteration_stats() + ( + reuse_deltas_by_window, + reuse_deltas_by_life_cycle, + pool_group_deltas_by_window, + pool_group_deltas, + ) = self._collect_iteration_stats_deltas(raw_iteration_stats, storage) + + windows = set(pool_groups_by_window) + windows.update(reuse_deltas_by_window) + windows.update(pool_group_deltas_by_window) + primary_stats = storage.get_statistics(GPU_LEVEL) + secondary_stats_by_level = [ + storage.get_statistics(CacheLevel(level)) + for level in range(1, int(storage.num_cache_levels)) + ] + + stats_by_window = { + window_size: self._build_window_iteration_stats( + window_size, + pool_groups_by_window, + windows_by_pool_group, + primary_stats, + secondary_stats_by_level, + pool_group_deltas_by_window.get(window_size), + reuse_deltas_by_window.get(window_size), + ) + for window_size in sorted(windows) + } + + pool_group_ids = sorted(set(windows_by_pool_group) | set(pool_group_deltas)) + stats_by_pool_group = { + pool_group_id: self._build_pool_group_iteration_stats( + pool_group_id, + windows_by_pool_group, + primary_stats, + secondary_stats_by_level, + pool_group_deltas.get(pool_group_id), + ) + for pool_group_id in pool_group_ids + } + + stats_by_life_cycle = { + life_cycle_id: self._build_life_cycle_iteration_stats( + life_cycle_id, + storage, + primary_stats, + secondary_stats_by_level, + reuse_delta, + ) + for life_cycle_id, reuse_delta in sorted(reuse_deltas_by_life_cycle.items()) + } + + return KVCacheV2IterationStatsReport( + stats_by_window, stats_by_pool_group, stats_by_life_cycle + ) def get_block_ids_per_seq(self, request_ids: List[int]) -> torch.Tensor: block_ids_per_seq = self.get_batch_cache_indices(request_ids) @@ -1553,7 +2104,14 @@ def release_resources( # writes to the radix tree, so the choice of branch does not # affect committed state. ``cache_salt`` is left defaulted # to None to avoid coupling synthetic data to any salted branch. - kv_cache = self._create_kv_cache(req.py_request_id, req.lora_task_id, input_tokens) + kv_cache = self._create_kv_cache( + req.py_request_id, req.lora_task_id, input_tokens, is_dummy=req.is_dummy + ) + # Saturated IndexMapper (e.g. disagg gen trans in progress) + # returns None; retry next iter. + if kv_cache is None: + release_resources(req) + return None assert kv_cache.num_committed_tokens == 0 success = kv_cache.resume(self._stream.cuda_stream) if not success: @@ -1570,9 +2128,12 @@ def release_resources( draft_kv_cache = None if draft_kv_cache_manager is not None: draft_kv_cache = draft_kv_cache_manager._create_kv_cache( - req.py_request_id, req.lora_task_id, input_tokens + req.py_request_id, req.lora_task_id, input_tokens, is_dummy=req.is_dummy ) # Dummy path: see comment above, no salt. + if draft_kv_cache is None: + release_resources(req) + return None success = draft_kv_cache.resume(draft_kv_cache_manager._stream.cuda_stream) if not success: release_resources(req, free_draft_resources=True) @@ -1637,9 +2198,12 @@ def free_resources(self, request: LlmRequest, pin_on_release: bool = False): 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: + self.impl.clear_stats_excluded(request.py_request_id) return + kv_cache.discard_pending_stats() self.try_commit_blocks_for_reuse(request, kv_cache) kv_cache.close() + self.impl.clear_stats_excluded(request.py_request_id) if request.py_request_id in self._early_freed_index_requests: self._early_freed_index_requests.discard(request.py_request_id) else: @@ -1661,7 +2225,7 @@ def _get_batch_cache_indices_by_pool_id( ) -> List[List[int]]: if is_kv_aggregate: # Div by kv_factor to index kv cache with size - # [num_blocks, kv_factor, tokens_per_block, num_kv_heads, head_dim]. + # [num_blocks, kv_factor, tokens_per_block, num_kv_heads, head_dim] div_factor = self.kv_factor else: div_factor = 1 @@ -1785,8 +2349,8 @@ def check_invalid_values_in_kv_cache(self, fill_with_zero: bool = False) -> bool if some_checks_unavailable: logger.warning( - "`torch.isnan` or `torch.isinf` is not implemented for current kv cache dtype, " - "related checks are skipped" + "`torch.isnan` or `torch.isinf` is not implemented for current " + "kv cache dtype, related checks are skipped" ) return bool(has_invalid_values) @@ -1852,25 +2416,25 @@ def get_cache_size_per_token( mem_per_token *= kv_factor return mem_per_token - def update_resources( - self, - scheduled_batch: ScheduledRequests, - attn_metadata: "AttentionMetadata" = None, - kv_cache_dtype_byte_size: float = None, - ): - if not self.is_draft: - _update_kv_cache_draft_token_location( - self, scheduled_batch, attn_metadata, kv_cache_dtype_byte_size - ) + def update_context_resources(self, scheduled_batch: ScheduledRequests): + """Update KV cache for context requests in the current batch. + + This is separated from update_resources (which handles generation + requests only) because the overlap executor needs context KV cache + updates to happen before next batch scheduling. Otherwise, the scheduler + would under-estimate available KV cache for sliding-window attention + layer. In non-overlap scheduler, you should call it together with + update_resources(). + """ for req in scheduled_batch.context_requests: if req.py_request_id not in self.kv_cache_map: continue kv_cache = self.kv_cache_map[req.py_request_id] # In the overlap scheduler, iteration N+1's eviction may # suspend a ctx request's KV cache while iteration N's - # update_resources still needs to process it. Skip the - # resize — the request will be resumed by the scheduler - # on the next iteration. + # update still needs to process it. Skip the resize — the + # request will be resumed by the scheduler on the next + # iteration. if not kv_cache.is_active: continue if self.enable_block_reuse and not self.is_draft and not req.is_dummy_request: @@ -1893,6 +2457,18 @@ def update_resources( "at context update" ) + def update_resources( + self, + scheduled_batch: ScheduledRequests, + attn_metadata: "AttentionMetadata" = None, + kv_cache_dtype_byte_size: float = None, + ): + if not self.is_draft: + _update_kv_cache_draft_token_location( + self, scheduled_batch, attn_metadata, kv_cache_dtype_byte_size + ) + # Context request KV cache updates are handled by + # update_context_resources, called separately from the executor loop. for req in scheduled_batch.generation_requests: if req.py_request_id not in self.kv_cache_map: continue @@ -1912,9 +2488,9 @@ def update_resources( success = kv_cache.resize(new_capacity, req.max_beam_num_tokens - 1) if not success: raise ValueError( - f"Failed to resize KV cache for request {req.py_request_id} to capacity " - f"{new_capacity} and history length {req.max_beam_num_tokens - 1} " - "tokens at generation update" + f"Failed to resize KV cache for request {req.py_request_id} " + f"to capacity {new_capacity} and history length " + f"{req.max_beam_num_tokens - 1} tokens at generation update" ) def copy_batch_block_offsets( @@ -1944,7 +2520,10 @@ def _create_kv_cache( request_id: int, lora_task_id: int | None, input_tokens: Sequence[TokenIdExt] | None, + *, cache_salt: str | None = None, + is_dummy: bool = False, + expected_prompt_length: int | None = None, ): assert request_id not in self.kv_cache_map, ( f"KV cache for request {request_id} already exists" @@ -1970,8 +2549,13 @@ def _create_kv_cache( kv_cache = self.impl.create_kv_cache( ReuseScope(lora_id=lora_task_id, salt=salt_int), input_tokens, + id=request_id, + expected_prompt_length=expected_prompt_length, ) self.kv_cache_map[request_id] = kv_cache + if is_dummy: + self.impl.mark_stats_excluded(request_id) + kv_cache.discard_pending_stats() index = self.index_mapper.add_new_sequence(request_id) for i in range(self.max_beam_width): for pool_idx in range(self.num_pools): diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_stats.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_stats.py new file mode 100644 index 000000000000..ff7a06643928 --- /dev/null +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_stats.py @@ -0,0 +1,139 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from dataclasses import dataclass, field +from typing import Any + +KV_CACHE_ITERATION_STATS_REUSE_KEYS = ( + "iterReusedBlocks", + "iterFullReusedBlocks", + "iterPartialReusedBlocks", + "iterMissedBlocks", + "iterCacheHitRate", +) + +KV_CACHE_ITERATION_STATS_POOL_GROUP_KEYS = ( + "primaryMaxNumBlocks", + "primaryFreeNumBlocks", + "primaryUsedNumBlocks", + "secondaryMaxNumBlocks", + "secondaryFreeNumBlocks", + "secondaryUsedNumBlocks", + "iterAllocTotalBlocks", + "iterAllocNewBlocks", + "iterGenAllocBlocks", + "iterOnboardBlocks", + "iterOnboardBytes", + "iterOffloadBlocks", + "iterOffloadBytes", + "iterIntraDeviceCopyBlocks", + "iterIntraDeviceCopyBytes", +) + + +@dataclass(slots=True) +class KVCacheV2PoolGroupIterationStats: + pool_group_id: int + slot_size: tuple[int, ...] + window_sizes: tuple[int, ...] + stats: Any + + +@dataclass(slots=True) +class KVCacheV2LifeCycleIterationStats: + life_cycle_id: int + pool_group_id: int + window_size: int | None + kind: str + stats: Any + + +@dataclass(slots=True) +class KVCacheV2IterationStatsReport: + by_window_size: dict[int, Any] + by_pool_group: dict[int, KVCacheV2PoolGroupIterationStats] + by_life_cycle: dict[int, KVCacheV2LifeCycleIterationStats] = field(default_factory=dict) + + +def serialize_kv_cache_iteration_stats(stats, keys: tuple[str, ...] | None = None) -> dict: + fields = { + "primaryMaxNumBlocks": stats.primary_max_num_blocks, + "primaryFreeNumBlocks": stats.primary_free_num_blocks, + "primaryUsedNumBlocks": stats.primary_used_num_blocks, + "secondaryMaxNumBlocks": stats.secondary_max_num_blocks, + "secondaryFreeNumBlocks": stats.secondary_free_num_blocks, + "secondaryUsedNumBlocks": stats.secondary_used_num_blocks, + "iterAllocTotalBlocks": stats.iter_alloc_total_blocks, + "iterAllocNewBlocks": stats.iter_alloc_new_blocks, + "iterReusedBlocks": stats.iter_reused_blocks, + "iterFullReusedBlocks": stats.iter_full_reused_blocks, + "iterPartialReusedBlocks": stats.iter_partial_reused_blocks, + "iterMissedBlocks": stats.iter_missed_blocks, + "iterCacheHitRate": stats.iter_cache_hit_rate, + "iterGenAllocBlocks": stats.iter_gen_alloc_blocks, + "iterOnboardBlocks": stats.iter_onboard_blocks, + "iterOnboardBytes": stats.iter_onboard_bytes, + "iterOffloadBlocks": stats.iter_offload_blocks, + "iterOffloadBytes": stats.iter_offload_bytes, + "iterIntraDeviceCopyBlocks": stats.iter_intra_device_copy_blocks, + "iterIntraDeviceCopyBytes": stats.iter_intra_device_copy_bytes, + } + if keys is None: + return fields + return {key: fields[key] for key in keys} + + +def append_kv_cache_iteration_stats(stats_dict: dict, kv_iter_stats) -> None: + if kv_iter_stats is None: + return + if isinstance(kv_iter_stats, KVCacheV2IterationStatsReport): + by_window_size = kv_iter_stats.by_window_size + by_pool_group = kv_iter_stats.by_pool_group + else: + by_window_size = kv_iter_stats + by_pool_group = None + + stats_dict["kvCacheIterationStats"] = { + str(window_size): serialize_kv_cache_iteration_stats(stats) + for window_size, stats in by_window_size.items() + } + if by_pool_group is None: + return + + stats_dict["kvCacheIterationStatsByPoolGroup"] = { + str(pool_group_id): { + "poolGroupId": stats.pool_group_id, + "slotSize": list(stats.slot_size), + "windowSizes": list(stats.window_sizes), + **serialize_kv_cache_iteration_stats( + stats.stats, KV_CACHE_ITERATION_STATS_POOL_GROUP_KEYS + ), + } + for pool_group_id, stats in by_pool_group.items() + } + + if not kv_iter_stats.by_life_cycle: + return + + stats_dict["kvCacheIterationStatsByLifecycle"] = { + str(life_cycle_id): { + "lifeCycleId": stats.life_cycle_id, + "poolGroupId": stats.pool_group_id, + "windowSize": stats.window_size, + "kind": stats.kind, + **serialize_kv_cache_iteration_stats(stats.stats, KV_CACHE_ITERATION_STATS_REUSE_KEYS), + } + for life_cycle_id, stats in kv_iter_stats.by_life_cycle.items() + } diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 3a680cbb919a..c04fa49087f6 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -1668,6 +1668,12 @@ def _create_cuda_graph_warmup_request( if requests is None: return None + def free_warmup_requests() -> None: + for r in requests: + kv_cache_manager.free_resources(r) + if draft_kv_cache_manager is not None: + draft_kv_cache_manager.free_resources(r) + # Add one dummy request with the maximum possible sequence length. max_seq_len = min( self.max_seq_len if max_seq_len is None else max_seq_len, @@ -1719,10 +1725,7 @@ def _create_cuda_graph_warmup_request( draft_kv_cache_manager=draft_kv_cache_manager) if max_seq_len_request is None: - for r in requests: - kv_cache_manager.free_resources(r) - if draft_kv_cache_manager is not None: - draft_kv_cache_manager.free_resources(r) + free_warmup_requests() return None else: max_seq_len_request = max_seq_len_request[0] diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 904964acb8a3..57a6f57ddac9 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -64,6 +64,7 @@ from .handle_logits import HandleLogits from .hang_detector import HangDetector from .kv_cache_manager_v2 import KVCacheManagerV2 +from .kv_cache_stats import append_kv_cache_iteration_stats from .kv_cache_transceiver import KvCacheTransceiver from .llm_request import (ATTENTION_DP_DUMMY_REQUEST_ID, MAX_SPEC_DECODE_POSITIONS, ExecutorRequest, @@ -1775,34 +1776,7 @@ def _append_iter_stats(self, local_dict["requestStats"] = [ _json.loads(r.to_json_str()) for r in req_stats ] - if kv_iter_stats is not None: - local_dict["kvCacheIterationStats"] = { - str(window_size): { - "primaryMaxNumBlocks": s.primary_max_num_blocks, - "primaryFreeNumBlocks": s.primary_free_num_blocks, - "primaryUsedNumBlocks": s.primary_used_num_blocks, - "secondaryMaxNumBlocks": s.secondary_max_num_blocks, - "secondaryFreeNumBlocks": s.secondary_free_num_blocks, - "secondaryUsedNumBlocks": s.secondary_used_num_blocks, - "iterAllocTotalBlocks": s.iter_alloc_total_blocks, - "iterAllocNewBlocks": s.iter_alloc_new_blocks, - "iterReusedBlocks": s.iter_reused_blocks, - "iterFullReusedBlocks": s.iter_full_reused_blocks, - "iterPartialReusedBlocks": s.iter_partial_reused_blocks, - "iterMissedBlocks": s.iter_missed_blocks, - "iterCacheHitRate": s.iter_cache_hit_rate, - "iterGenAllocBlocks": s.iter_gen_alloc_blocks, - "iterOnboardBlocks": s.iter_onboard_blocks, - "iterOnboardBytes": s.iter_onboard_bytes, - "iterOffloadBlocks": s.iter_offload_blocks, - "iterOffloadBytes": s.iter_offload_bytes, - "iterIntraDeviceCopyBlocks": - s.iter_intra_device_copy_blocks, - "iterIntraDeviceCopyBytes": - s.iter_intra_device_copy_bytes, - } - for window_size, s in kv_iter_stats.items() - } + append_kv_cache_iteration_stats(local_dict, kv_iter_stats) if host_step_time_ms is not None: local_dict["hostStepTimeMS"] = host_step_time_ms if prev_device_step_time_ms is not None: @@ -2479,6 +2453,9 @@ def _handle_executed_batch(self, executed_batch: Optional[BatchStatePP]): if self._disagg_pp_termination_handler is not None: self._disagg_pp_termination_handler.terminate_pending_requests() + if executed_batch is not None: + self._commit_kv_cache_stats(executed_batch.scheduled_requests) + if self.enable_iter_perf_stats and executed_batch is not None: self._process_iter_stats( finished_requests, @@ -2614,6 +2591,12 @@ def _prefetch_for_context_requests(self) -> None: if candidates: self.kv_cache_manager.prefetch_for_context_tokens(candidates) + def _commit_kv_cache_stats(self, + scheduled_batch: ScheduledRequests) -> None: + if self._is_kv_manager_v2: + self.kv_cache_manager.commit_scheduled_kv_cache_stats( + scheduled_batch) + def _prepare_and_schedule_batch(self): new_requests = self._fetch_and_activate_new_requests() if self.should_stop_processing: @@ -3012,6 +2995,7 @@ def _executor_loop(self): scheduled_batch_stats = ( self._collect_scheduled_batch_stats(scheduled_batch) if self.enable_iter_perf_stats else None) + self._commit_kv_cache_stats(scheduled_batch) # GPU and CPU timing for perf metrics gpu_forward_start, gpu_forward_end, gpu_sample_end = self.perf_manager.create_timing_events( @@ -3500,6 +3484,8 @@ def _executor_loop_overlap(self): self._update_request_states(scheduled_batch) if self.previous_batch is not None and should_process_previous_batch: + self._commit_kv_cache_stats( + self.previous_batch.scheduled_requests) self._process_previous_batch() self.perf_manager.compute_batch_gpu_times( self.previous_batch.scheduled_requests.all_requests()) @@ -5354,6 +5340,9 @@ def _handle_responses(self, emit_first_iter: bool = True): and not request.is_finished): should_emit = False if should_emit: + if request.return_perf_metrics: + # Response creation may finalize and copy scalar ctx GPU totals. + self.perf_manager.compute_batch_gpu_times([request]) response = request.create_response(False, self.dist.rank) if response: request_done = request.is_finished diff --git a/tensorrt_llm/_torch/pyexecutor/resource_manager.py b/tensorrt_llm/_torch/pyexecutor/resource_manager.py index e4ad48050606..1ef1c2843d8d 100644 --- a/tensorrt_llm/_torch/pyexecutor/resource_manager.py +++ b/tensorrt_llm/_torch/pyexecutor/resource_manager.py @@ -1,6 +1,17 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. - +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. import copy import enum import math @@ -26,6 +37,8 @@ from tensorrt_llm.lora_helper import LoraConfig from tensorrt_llm.lora_manager import LoraManager, LoraModelConfig from tensorrt_llm.runtime import ModelConfig as ModelConfigPython +from tensorrt_llm.runtime.kv_cache_hash import (KV_CACHE_HASH_ALGO_AUTO, + KV_CACHE_HASH_ALGO_V1) # isort: off # isort: on @@ -89,6 +102,15 @@ def compute_page_count(token_count: int, tokens_per_page: int) -> int: return (token_count + tokens_per_page) // tokens_per_page +def _warn_if_unsupported_v1_kv_cache_event_hash_algo(hash_algo: str) -> None: + if hash_algo in (KV_CACHE_HASH_ALGO_AUTO, KV_CACHE_HASH_ALGO_V1): + return + logger.warning( + f"KVCacheManager only supports kv_cache_event_hash_algo={KV_CACHE_HASH_ALGO_V1}; " + f"requested {hash_algo}. The V1 event manager will emit {KV_CACHE_HASH_ALGO_V1} " + "event hashes.") + + class BaseResourceManager(ABC): @abstractmethod @@ -573,6 +595,8 @@ def append_to_kv_heads_per_layer(num_kv_heads_per_layer: List[int], } if self.event_buffer_max_size > 0: + _warn_if_unsupported_v1_kv_cache_event_hash_algo( + kv_cache_config.kv_cache_event_hash_algo) if mapping.enable_attention_dp: kwargs['event_manager'] = KVCacheEventManagerCpp( max_kv_event_entries=self.event_buffer_max_size, diff --git a/tensorrt_llm/_utils.py b/tensorrt_llm/_utils.py index cff3383c8579..0781efe2f1e2 100644 --- a/tensorrt_llm/_utils.py +++ b/tensorrt_llm/_utils.py @@ -1175,6 +1175,12 @@ def to_json_str(cls, event): "data": event_serialize_func(event.data), "window_size": event.window_size, } + hash_algo = getattr(event, "hash_algo", None) + if hash_algo is not None: + json_str["hash_algo"] = hash_algo + layer_group_id = getattr(event, "layer_group_id", None) + if layer_group_id is not None: + json_str["layer_group_id"] = layer_group_id if event.attention_dp_rank is not None: json_str["attention_dp_rank"] = event.attention_dp_rank diff --git a/tensorrt_llm/executor/base_worker.py b/tensorrt_llm/executor/base_worker.py index fab7f898a70f..99bd28282e41 100644 --- a/tensorrt_llm/executor/base_worker.py +++ b/tensorrt_llm/executor/base_worker.py @@ -27,6 +27,7 @@ from tensorrt_llm.logger import logger +from .._torch.pyexecutor.kv_cache_stats import append_kv_cache_iteration_stats from .._torch.pyexecutor.llm_request import LlmResponse from .._utils import (global_mpi_rank, global_mpi_size, mpi_comm, mpi_rank, nvtx_range_debug) @@ -804,7 +805,6 @@ def get_disaggregated_params(self) -> dict: return {} return self.engine.kv_cache_transceiver.get_disaggregated_params() - # Define a Callable to join iteration and request stats @staticmethod def _stats_serializer(stats) -> str: # Per-rank path: stats is ("per_rank_dict", {..., "rank": N}). @@ -836,34 +836,7 @@ def _stats_serializer(stats) -> str: stats_dict["requestStats"].append( json.loads(req_stat.to_json_str())) - # Inject per-iteration KV cache stats (keyed by window size) - if kv_iter_stats is not None: - stats_dict["kvCacheIterationStats"] = { - str(window_size): { - "primaryMaxNumBlocks": s.primary_max_num_blocks, - "primaryFreeNumBlocks": s.primary_free_num_blocks, - "primaryUsedNumBlocks": s.primary_used_num_blocks, - "secondaryMaxNumBlocks": s.secondary_max_num_blocks, - "secondaryFreeNumBlocks": s.secondary_free_num_blocks, - "secondaryUsedNumBlocks": s.secondary_used_num_blocks, - "iterAllocTotalBlocks": s.iter_alloc_total_blocks, - "iterAllocNewBlocks": s.iter_alloc_new_blocks, - "iterReusedBlocks": s.iter_reused_blocks, - "iterFullReusedBlocks": s.iter_full_reused_blocks, - "iterPartialReusedBlocks": s.iter_partial_reused_blocks, - "iterMissedBlocks": s.iter_missed_blocks, - "iterCacheHitRate": s.iter_cache_hit_rate, - "iterGenAllocBlocks": s.iter_gen_alloc_blocks, - "iterOnboardBlocks": s.iter_onboard_blocks, - "iterOnboardBytes": s.iter_onboard_bytes, - "iterOffloadBlocks": s.iter_offload_blocks, - "iterOffloadBytes": s.iter_offload_bytes, - "iterIntraDeviceCopyBlocks": - s.iter_intra_device_copy_blocks, - "iterIntraDeviceCopyBytes": s.iter_intra_device_copy_bytes, - } - for window_size, s in kv_iter_stats.items() - } + append_kv_cache_iteration_stats(stats_dict, kv_iter_stats) # Per-loop CPU wall captured by profile_step() — always a clean # single-loop measurement, matching the log line's `host_step_time`. diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 609e11c5a27a..73bcc1936192 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -3091,6 +3091,16 @@ class KvCacheConfig(StrictBaseModel, PybindMirror): status="prototype", description="Whether to use the KV cache manager v2 (experimental).") + kv_cache_event_hash_algo: Literal[ + "auto", "v1_block_key", "v2_sha256", "v2_sha256_64"] = Field( + default="auto", + status="prototype", + description= + "The block hash algorithm used by KV cache manager events. " + "'auto' uses the native hash for each KV cache manager. " + "Explicit V2 hash choices are ignored with a warning by the V1 " + "KV cache manager.") + max_util_for_resume: float = Field( default=0.95, ge=0, diff --git a/tensorrt_llm/metrics/collector.py b/tensorrt_llm/metrics/collector.py index d876d2ddf0a0..cd5a91123ddf 100644 --- a/tensorrt_llm/metrics/collector.py +++ b/tensorrt_llm/metrics/collector.py @@ -797,9 +797,16 @@ def log_iteration_stats(self, iteration_stats: dict) -> None: self._log_gauge(self.spec_decode_draft_overhead, spec_stats["draftOverhead"]) - # Per-iteration KV cache stats (aggregated across window sizes) - if kv_iter := iteration_stats.get("kvCacheIterationStats"): - # Aggregate across all window sizes + # Per-iteration KV cache stats. V2 reports reuse/miss by lifecycle and + # storage/transfer counters by pool group; legacy V1 uses window stats. + kv_iter = iteration_stats.get("kvCacheIterationStats") + kv_iter_by_lifecycle = iteration_stats.get( + "kvCacheIterationStatsByLifecycle") + kv_iter_by_pool_group = iteration_stats.get( + "kvCacheIterationStatsByPoolGroup") + if kv_iter or kv_iter_by_lifecycle or kv_iter_by_pool_group: + reuse_stats = kv_iter_by_lifecycle or kv_iter or {} + pool_group_stats = kv_iter_by_pool_group or kv_iter or {} total_secondary_max = 0 total_secondary_used = 0 total_reused = 0 @@ -811,19 +818,19 @@ def log_iteration_stats(self, iteration_stats: dict) -> None: total_offload_bytes = 0 total_intra_device_copy_bytes = 0 - for ws_stats in kv_iter.values(): - total_secondary_max += ws_stats.get("secondaryMaxNumBlocks", 0) - total_secondary_used += ws_stats.get("secondaryUsedNumBlocks", - 0) - total_reused += ws_stats.get("iterReusedBlocks", 0) - total_full_reused += ws_stats.get("iterFullReusedBlocks", 0) - total_partial_reused += ws_stats.get("iterPartialReusedBlocks", - 0) - total_missed += ws_stats.get("iterMissedBlocks", 0) - total_gen_alloc += ws_stats.get("iterGenAllocBlocks", 0) - total_onboard_bytes += ws_stats.get("iterOnboardBytes", 0) - total_offload_bytes += ws_stats.get("iterOffloadBytes", 0) - total_intra_device_copy_bytes += ws_stats.get( + for stats in reuse_stats.values(): + total_reused += stats.get("iterReusedBlocks", 0) + total_full_reused += stats.get("iterFullReusedBlocks", 0) + total_partial_reused += stats.get("iterPartialReusedBlocks", 0) + total_missed += stats.get("iterMissedBlocks", 0) + + for stats in pool_group_stats.values(): + total_secondary_max += stats.get("secondaryMaxNumBlocks", 0) + total_secondary_used += stats.get("secondaryUsedNumBlocks", 0) + total_gen_alloc += stats.get("iterGenAllocBlocks", 0) + total_onboard_bytes += stats.get("iterOnboardBytes", 0) + total_offload_bytes += stats.get("iterOffloadBytes", 0) + total_intra_device_copy_bytes += stats.get( "iterIntraDeviceCopyBytes", 0) # Gauges diff --git a/tensorrt_llm/runtime/kv_cache_hash.py b/tensorrt_llm/runtime/kv_cache_hash.py new file mode 100644 index 000000000000..9bc65289ff8f --- /dev/null +++ b/tensorrt_llm/runtime/kv_cache_hash.py @@ -0,0 +1,94 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from collections.abc import Sequence +from typing import Optional + +KV_CACHE_HASH_ALGO_AUTO = "auto" +KV_CACHE_HASH_ALGO_V1 = "v1_block_key" +KV_CACHE_HASH_ALGO_V2 = "v2_sha256" +KV_CACHE_HASH_ALGO_V2_SHA256_64 = "v2_sha256_64" +KV_CACHE_HASH_ALGO_DEFAULT = KV_CACHE_HASH_ALGO_V1 + +_UINT32_MASK = (1 << 32) - 1 +_UINT64_MASK = (1 << 64) - 1 +_HASH32_CONST = 0x45D9F3B +_HASH64_CONST_1 = 0xBF58476D1CE4E5B9 +_HASH64_CONST_2 = 0x94D049BB133111EB +_HASH_COMBINE_CONST = 0x9E3779B9 +_PARENT_HASH_CONST = 0xBF58476D1CE4E5B9 + + +class NonTextTokenHashError(ValueError): + """Raised when v1-compatible hashing receives a non-text token.""" + + +def get_effective_kv_cache_event_hash_algo(hash_algo: str, use_kv_cache_manager_v2: bool) -> str: + if hash_algo != KV_CACHE_HASH_ALGO_AUTO: + return hash_algo + return KV_CACHE_HASH_ALGO_DEFAULT + + +def truncate_sha256_hash_to_int64(block_hash: bytes) -> int: + return int.from_bytes(block_hash[:8], "big", signed=False) + + +def get_cache_salt_id(cache_salt: str) -> int: + """Return the cache salt id used by request handling and cache-aware routing.""" + from blake3 import blake3 + + h = blake3(cache_salt.encode("utf-8")).digest(length=8) + return int.from_bytes(h, "little", signed=False) + + +def hash_v1_block_key( + tokens: Sequence[int], + parent_hash: int = 0, + lora_task_id: Optional[int] = None, + cache_salt_id: Optional[int] = None, +) -> int: + seed = (len(tokens) ^ ((parent_hash * _PARENT_HASH_CONST) & _UINT64_MASK)) & _UINT64_MASK + if parent_hash == 0 and cache_salt_id is not None: + seed = _hash64_mix(cache_salt_id, seed) + for token in tokens: + if type(token) is not int: + raise NonTextTokenHashError("v1-compatible hashing only supports text tokens") + seed = _hash32_mix(token, seed) + if lora_task_id is not None: + seed = _hash64_mix(lora_task_id, seed) + return seed + + +def _hash32_mix(value: int, seed: int) -> int: + value &= _UINT32_MASK + value = (((value >> 16) ^ value) * _HASH32_CONST) & _UINT32_MASK + value = (((value >> 16) ^ value) * _HASH32_CONST) & _UINT32_MASK + value = ((value >> 16) ^ value) & _UINT32_MASK + # In C++, value and _HASH_COMBINE_CONST are both 32-bit unsigned values, + # so this part wraps to uint32_t before size_t terms are added. + value = (value + _HASH_COMBINE_CONST) & _UINT32_MASK + combined = (value + ((seed << 6) & _UINT64_MASK) + (seed >> 2)) & _UINT64_MASK + return (seed ^ combined) & _UINT64_MASK + + +def _hash64_mix(value: int, seed: int) -> int: + value &= _UINT64_MASK + value = ((value ^ (value >> 30)) * _HASH64_CONST_1) & _UINT64_MASK + value = ((value ^ (value >> 27)) * _HASH64_CONST_2) & _UINT64_MASK + value = (value ^ (value >> 31)) & _UINT64_MASK + combined = ( + value + _HASH_COMBINE_CONST + ((seed << 6) & _UINT64_MASK) + (seed >> 2) + ) & _UINT64_MASK + return (seed ^ combined) & _UINT64_MASK diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py index 4d3607a9cc7c..8e118c11f013 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py @@ -50,7 +50,19 @@ ScratchDesc, _KVCache, ) +from ._event_manager import ( + KVCacheCreatedData, + KVCacheEvent, + KVCacheEventDiff, + KVCacheEventManager, + KVCacheRemovedData, + KVCacheStoredBlockData, + KVCacheStoredData, + KVCacheUpdatedData, + UniqueToken, +) from ._life_cycle_registry import LayerGroupId, LifeCycleId +from ._stats import KVCacheIterationStatsDelta, KVCacheStatsDelta from ._storage import BufferId __all__ = [ @@ -60,6 +72,15 @@ "TokenIdExt", "KVCacheManager", "_KVCache", + "KVCacheCreatedData", + "KVCacheEvent", + "KVCacheEventDiff", + "KVCacheEventManager", + "KVCacheRemovedData", + "KVCacheStoredBlockData", + "KVCacheStoredData", + "KVCacheUpdatedData", + "UniqueToken", "BeamIndex", "DEFAULT_BEAM_INDEX", "LayerId", @@ -89,4 +110,6 @@ "PageIndexConverter", "PageIndexMode", "ScratchDesc", + "KVCacheIterationStatsDelta", + "KVCacheStatsDelta", ] diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi index 451778f7e754..ee4a4badc218 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi @@ -63,6 +63,30 @@ MemAddress = NewType("MemAddress", int) Priority = NewType("Priority", int) PoolGroupIndex = NewType("PoolGroupIndex", int) +# From _stats.py +@dataclass(slots=True) +class KVCacheStatsDelta: + alloc_total_blocks: int = 0 + alloc_new_blocks: int = 0 + reused_blocks: int = 0 + missed_blocks: int = 0 + +@dataclass(slots=True) +class KVCacheIterationStatsDelta: + iter_alloc_total_blocks: int = 0 + iter_alloc_new_blocks: int = 0 + iter_reused_blocks: int = 0 + iter_full_reused_blocks: int = 0 + iter_partial_reused_blocks: int = 0 + iter_missed_blocks: int = 0 + iter_gen_alloc_blocks: int = 0 + iter_onboard_blocks: int = 0 + iter_onboard_bytes: int = 0 + iter_offload_blocks: int = 0 + iter_offload_bytes: int = 0 + iter_intra_device_copy_blocks: int = 0 + iter_intra_device_copy_bytes: int = 0 + # From _config.py DataRole = NewType("DataRole", str) @@ -149,10 +173,104 @@ class KVCacheManagerConfig: typical_step: BatchDesc | None = None ssm_reuse_interval: int = 512 swa_scratch_reuse: SwaScratchReuseConfig | None = None + enable_stats: bool = True helix_config: HelixConfig | None = None @property def enable_swa_scratch_reuse(self) -> bool: ... +# From _event_manager.py +EventBlockHash: TypeAlias = int | str +BlockHashLike: TypeAlias = bytes | EventBlockHash +BlockHashesLike: TypeAlias = BlockHashLike | Iterable[BlockHashLike] +EventTokenId: TypeAlias = int | str +MmKey: TypeAlias = tuple[bytes, int] | tuple[bytes, int, str | None] +AttentionDpGatherFn: TypeAlias = Callable[[list["KVCacheEvent"]], list[list["KVCacheEvent"]]] + +@dataclass(slots=True, frozen=True) +class UniqueToken: + token_id: EventTokenId + token_extra_id: int = ... + +@dataclass(slots=True, frozen=True) +class KVCacheCreatedData: + num_blocks_per_cache_level: list[int] + +@dataclass(slots=True, frozen=True) +class KVCacheStoredBlockData: + block_hash: EventBlockHash + tokens: list[UniqueToken] + cache_level: int + priority: int + mm_keys: list[MmKey] = ... + cache_salt: str | None = ... + +@dataclass(slots=True, frozen=True) +class KVCacheStoredData: + parent_hash: EventBlockHash | None + blocks: list[KVCacheStoredBlockData] + +@dataclass(slots=True, frozen=True) +class KVCacheRemovedData: + block_hashes: list[EventBlockHash] + +@dataclass(slots=True, frozen=True) +class KVCacheEventDiff: + old_value: int + new_value: int + +@dataclass(slots=True, frozen=True) +class KVCacheUpdatedData: + block_hash: EventBlockHash + cache_level: KVCacheEventDiff | None + priority: KVCacheEventDiff | None + +@dataclass(slots=True, frozen=True) +class KVCacheEvent: + event_id: int + data: KVCacheCreatedData | KVCacheStoredData | KVCacheRemovedData | KVCacheUpdatedData + window_size: int + hash_algo: str | None = None + attention_dp_rank: int | None = None + layer_group_id: int | None = None + +class KVCacheEventManager: + def __init__( + self, + max_kv_event_entries: int, + *, + window_size: int = ..., + attention_dp_rank: int | None = None, + attention_dp_gather: AttentionDpGatherFn | None = None, + hash_algo: str = ..., + window_size_by_layer_group: dict[int, int] | None = None, + ) -> None: ... + def add_created_event( + self, + num_blocks_per_cache_level: Sequence[int], + layer_group_ids: Sequence[int] | None = None, + ) -> None: ... + def set_layer_group_window_sizes(self, window_sizes: dict[int, int]) -> None: ... + def add_stored_event( + self, + parent_hash: EventBlockHash | None, + blocks: Sequence[KVCacheStoredBlockData], + layer_group_id: int | None = None, + ) -> None: ... + def add_stored_block_event_from_block(self, block: Any) -> None: ... + def add_stored_life_cycle_event_from_block(self, block: Any, life_cycle_id: int) -> None: ... + def add_removed_event(self, block_hashes: BlockHashesLike) -> None: ... + def add_removed_life_cycle_event(self, block_hash: bytes, life_cycle_id: int) -> None: ... + def add_updated_event( + self, + block_hash: BlockHashLike, + *, + cache_level: KVCacheEventDiff | None = None, + priority: KVCacheEventDiff | None = None, + layer_group_id: int | None = None, + ) -> None: ... + def flush_iteration_events(self) -> None: ... + def get_latest_events(self, timeout_ms: float | None = None) -> list[KVCacheEvent]: ... + # From _block_radix_tree.py def gen_multimodal_cache_key_tokens( id_offset: int, @@ -193,6 +311,8 @@ class _KVCache: def finish_event(self) -> Any: ... @property def num_blocks(self) -> int: ... + def commit_pending_stats(self) -> KVCacheStatsDelta: ... + def discard_pending_stats(self) -> None: ... def close(self) -> None: ... @property def beam_width(self) -> BeamIndex: ... @@ -295,7 +415,11 @@ class PageIndexConverter: ) -> list[int]: ... class KVCacheManager: - def __init__(self, config: KVCacheManagerConfig) -> None: ... + def __init__( + self, + config: KVCacheManagerConfig, + event_manager: KVCacheEventManager | None = None, + ) -> None: ... def __del__(self) -> None: ... def shutdown(self) -> None: ... def clear_reusable_blocks(self) -> None: ... @@ -314,6 +438,7 @@ class KVCacheManager: input_tokens: Sequence[TokenIdExt] | None = None, id: Any = None, custom_priority_callback: Callable[[int, Any], Priority] = ..., + expected_prompt_length: int | None = None, ) -> _KVCache: ... def probe_reuse( self, @@ -322,11 +447,21 @@ class KVCacheManager: ) -> int: ... def resize(self, cache_level: CacheLevel, quota: int, best_efforts: bool = False) -> bool: ... def get_quota(self, cache_level: CacheLevel) -> int: ... + def get_committed_stats(self) -> KVCacheStatsDelta: ... + def get_and_reset_iteration_stats(self) -> dict[LifeCycleId, KVCacheIterationStatsDelta]: ... + def mark_stats_dirty(self, kv_cache_id: int | None) -> None: ... + def clear_stats_dirty(self, kv_cache_id: int | None) -> None: ... + def get_dirty_stats_kv_cache_ids(self) -> set[int]: ... + def mark_stats_excluded(self, kv_cache_id: int | None) -> None: ... + def clear_stats_excluded(self, kv_cache_id: int | None) -> None: ... + def is_stats_excluded(self, kv_cache_id: int | None) -> bool: ... @property def cache_tier_list(self) -> Sequence[CacheTier]: ... @property def tokens_per_block(self) -> int: ... @property + def event_manager(self) -> Any | None: ... + @property def allow_seq_rebasing(self) -> bool: ... @property def enable_partial_match(self) -> bool: ... diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_block_radix_tree.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_block_radix_tree.py index 1024eca91575..ba0a759b36c7 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_block_radix_tree.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_block_radix_tree.py @@ -23,6 +23,7 @@ from ._utils import TypedIndexList, chunked, div_up, filled_list, find_index, unwrap_rawref if TYPE_CHECKING: + from ._event_manager import KVCacheEventManager from ._page import CommittedPage BlockKey = bytes @@ -130,23 +131,36 @@ def sequence_to_blockchain_keys( Children = dict[BlockKey, Child] -def get_tree(block: "RootBlock | Block") -> "BlockRadixTree": +def try_get_tree(block: "RootBlock | Block") -> "BlockRadixTree | None": node = block while not isinstance(node, BlockRadixTree): - node = node.prev + node = node._prev() + if node is None: + return None return node +def get_tree(block: "RootBlock | Block") -> "BlockRadixTree": + tree = try_get_tree(block) + if tree is None: + raise ValueError("Dereferencing a dangling rawref") + return tree + + def remove_subtree(root: "RootBlock | Block") -> list[rawref.ref["CommittedPage"]]: # taking O(1) space # remove leaf blocks one by one, in post-order ret: list[rawref.ref["CommittedPage"]] = [] + removed_block_hashes: list[BlockKey] = [] + tree = try_get_tree(root) + event_manager = tree.event_manager if tree is not None else None block: "RootBlock | Block" = root while True: if block.next: block = next(iter(block.next.values())) else: if isinstance(block, Block): + removed_block_hashes.append(block.key) ret.extend(p for p in block.storage if p is not None) block.storage = filled_list(None, block.num_life_cycles) assert isinstance(block, RootBlock) or all(page is None for page in block.storage), ( @@ -164,6 +178,8 @@ def remove_subtree(root: "RootBlock | Block") -> list[rawref.ref["CommittedPage" break assert not isinstance(prev_block, BlockRadixTree) block = prev_block + if event_manager is not None: + event_manager.add_removed_event(removed_block_hashes) return ret @@ -240,7 +256,7 @@ def _add_or_get_existing( class RootBlock: - __slots__ = ("_prev", "key", "next", "reuse_scope", "__rawref__") + __slots__ = ("__rawref__", "_prev", "key", "next", "reuse_scope") key: BlockKey reuse_scope: ReuseScope _prev: rawref.ref["BlockRadixTree"] @@ -285,7 +301,7 @@ class Block: A block of tokens. Manages data for all layers. """ - __slots__ = ("key", "tokens", "ordinal", "_prev", "next", "storage", "__rawref__") + __slots__ = ("__rawref__", "_prev", "key", "next", "ordinal", "storage", "tokens") key: BlockKey tokens: Sequence[TokenIdExt] ordinal: BlockOrdinal @@ -324,8 +340,11 @@ def __init__(self, tokens: Sequence[TokenIdExt], prev: "Block | RootBlock") -> N if len(b.tokens) < len(tokens) and tokens[: len(b.tokens)] == b.tokens: assert NDEBUG or (not b.is_full and b is not self and b.key == k and not b.next) to_remove.append(k) + event_manager = get_tree(prev).event_manager if to_remove else None for k in to_remove: b = prev.next.pop(k) + if event_manager is not None: + event_manager.add_removed_event(b.key) assert b.is_orphan # _KVCache may still hold it. # prev.next keeps a strong ref to this _Block, so no need to remove self from prev.next in __del__(). prev.next[self.key] = self @@ -363,6 +382,8 @@ def unset_page(self, lc_idx: LifeCycleId, lc: LifeCycle) -> None: return ordinal = self.ordinal self.storage[lc_idx] = None + tree = try_get_tree(self) + event_manager = tree.event_manager if tree is not None else None if type(lc) is AttnLifeCycle and (lc.window_size is None or ordinal < lc.num_sink_blocks): pages = remove_subtree(self) for r in pages: @@ -371,6 +392,8 @@ def unset_page(self, lc_idx: LifeCycleId, lc: LifeCycle) -> None: assert page.status == PageStatus.DROPPABLE if page.scheduled_for_eviction: page.manager.exclude_from_eviction(page) + elif event_manager is not None: + event_manager.add_removed_life_cycle_event(self.key, int(lc_idx)) # It's possible to implement more sophisticated logic to remove useless blocks for SWA, e.g. # check if consecutive available blocks is sufficient for window_size. (TRTLLM-8802) # But for simplicity, we leave it for now. @@ -382,6 +405,8 @@ def unset_page(self, lc_idx: LifeCycleId, lc: LifeCycle) -> None: ): if curr.key in curr.prev.next: curr.prev.next.pop(curr.key) + if event_manager is not None: + event_manager.add_removed_event(curr.key) curr = curr.prev @property @@ -400,15 +425,28 @@ def is_orphan(self) -> bool: class BlockRadixTree: - __slots__ = ("_life_cycles", "_tokens_per_block", "next", "__rawref__") + __slots__ = ( + "__rawref__", + "_event_manager", + "_life_cycles", + "_tokens_per_block", + "next", + ) _life_cycles: LifeCycleRegistry _tokens_per_block: int + _event_manager: "KVCacheEventManager | None" next: Children[RootBlock] __rawref__: rawref.ref["BlockRadixTree"] - def __init__(self, life_cycles: LifeCycleRegistry, tokens_per_block: int) -> None: + def __init__( + self, + life_cycles: LifeCycleRegistry, + tokens_per_block: int, + event_manager: "KVCacheEventManager | None" = None, + ) -> None: self._life_cycles = life_cycles self._tokens_per_block = tokens_per_block + self._event_manager = event_manager self.next = {} self.__rawref__ = rawref.NULL @@ -429,6 +467,10 @@ def tokens_per_block(self) -> int: def life_cycles(self) -> LifeCycleRegistry: return self._life_cycles + @property + def event_manager(self) -> "KVCacheEventManager | None": + return self._event_manager + @property def num_life_cycles(self) -> LifeCycleId: return self.life_cycles.size diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py index 5a325f801875..9341141a139a 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py @@ -243,6 +243,11 @@ class KVCacheManagerConfig: where the number of out-of-window blocks dominates memory usage. """ + enable_stats: bool = True + """ + Collect V2 KV cache allocation, reuse, and transfer statistics. + """ + # unsupported yet helix_config: HelixConfig | None = None 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 a7fec6eea3e6..d2c66568808a 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 @@ -61,6 +61,7 @@ _SharedPageLock, batched_lock_to_gpu, ) +from .._stats import KVCacheIterationStatsDelta, KVCacheStatsDelta from .._storage._core import Slot from .._storage_manager import StorageManager from .._utils import ( @@ -84,6 +85,7 @@ value_or, ) from ._moving_average import Average +from ._pending_stats import _PendingStats if TYPE_CHECKING: from ._kv_cache_manager import KVCacheManager, ScratchDesc @@ -177,6 +179,8 @@ class _KVCache: "_cuda_stream", "_status", "_beam_width", + "_expected_prompt_length", + "_generation_alloc_ready", "_capacity", "_history_length", "_commit_state", @@ -192,6 +196,7 @@ class _KVCache: "_never_resumed", "_enable_swa_scratch_reuse", "_scratch_slots", + "_pending_stats", "__rawref__", ) @@ -205,6 +210,8 @@ class _KVCache: _cuda_stream: CudaStream | None _status: _Status _beam_width: BeamIndex + _expected_prompt_length: int | None + _generation_alloc_ready: bool _capacity: int _history_length: int _commit_state: _CommitState @@ -237,6 +244,7 @@ class _KVCache: # Managed via delta in resize(): existing slots are reused across resize calls, # only the additional needed slots are allocated. Freed on teardown/suspend. _scratch_slots: TypedIndexList[LifeCycleId, list[ScratchSlotLock]] + _pending_stats: _PendingStats def __init__( self, @@ -245,6 +253,7 @@ def __init__( reuse_match: ReuseMatch | None, id: int | None, custom_priority_callback: Callable[[BlockOrdinal, LifeCycle], Priority], + expected_prompt_length: int | None = None, ): self.id = id self._manager = manager @@ -253,6 +262,10 @@ def __init__( self._cuda_stream = None self._status = self.Status.SUSPENDED self._beam_width = BeamIndex(1) + self._expected_prompt_length = ( + max(expected_prompt_length, 0) if expected_prompt_length is not None else None + ) + self._generation_alloc_ready = False self._capacity = 0 self._history_length = 0 self._commit_state = self.CommitState.ALLOWED @@ -274,9 +287,11 @@ def __init__( self._scratch_slots = make_typed( lambda _: list[ScratchSlotLock](), manager._storage.num_life_cycles ) + self._pending_stats = _PendingStats() self.__rawref__ = rawref.NULL if reuse_match is not None: self._setup_for_reuse(reuse_match) + self._refresh_generation_alloc_ready() self._avg_history_length = Average() self._avg_capacity = Average() self._avg_history_length.update(self.history_length) @@ -333,12 +348,149 @@ def finish_event(self) -> CachedCudaEvent: def num_blocks(self) -> int: return len(self._blocks) + def _should_record_stats(self) -> bool: + return self.manager._stats_enabled and not self.manager.is_stats_excluded(self.id) + + def commit_pending_stats(self) -> KVCacheStatsDelta: + if not self._should_record_stats(): + self.discard_pending_stats() + return KVCacheStatsDelta() + self.manager.commit_stats( + self._pending_stats.global_stats, self._pending_stats.iteration_stats_by_life_cycle + ) + request_stats = self._pending_stats.request_stats.copy() + self._pending_stats.clear() + self.manager.clear_stats_dirty(self.id) + return request_stats + + def discard_pending_stats(self) -> None: + self._pending_stats.clear() + self.manager.clear_stats_dirty(self.id) + + def _refresh_stats_dirty_state(self) -> None: + if not self._pending_stats.empty: + self.manager.mark_stats_dirty(self.id) + else: + self.manager.clear_stats_dirty(self.id) + + def _stats_life_cycle_key(self, life_cycle: LifeCycleId) -> LifeCycleId | None: + life_cycle_obj = self.manager._life_cycles.get_life_cycle(life_cycle) + if isinstance(life_cycle_obj, AttnLifeCycle): + return life_cycle + return None + + def _refresh_generation_alloc_ready(self) -> None: + expected_prompt_length = self._expected_prompt_length + if expected_prompt_length is not None and self._history_length >= expected_prompt_length: + self._generation_alloc_ready = True + + def _should_record_generation_alloc_stats(self, capacity: int) -> bool: + return self._generation_alloc_ready and capacity > self._capacity + + @staticmethod + def _block_ranges_excluding( + block_begin: BlockOrdinal, + block_end: BlockOrdinal, + excluded: HalfOpenRange[BlockOrdinal], + ) -> Iterator[HalfOpenRange[BlockOrdinal]]: + first_end = min(block_end, excluded.beg) + if block_begin < first_end: + yield HalfOpenRange(block_begin, first_end) + second_begin = max(block_begin, excluded.end) + if second_begin < block_end: + yield HalfOpenRange(second_begin, block_end) + + def _record_resize_pending_allocations( + self, + block_begin: BlockOrdinal, + block_end: BlockOrdinal, + beam_width: BeamIndex, + excluded_ranges: TypedIndexList[LifeCycleId, HalfOpenRange[BlockOrdinal]], + count_as_generation: bool, + ) -> None: + if not self._should_record_stats() or block_begin >= block_end: + return + # V2 includes generation allocations in per-request alloc_total/new + # metrics. This intentionally differs from the legacy V1 C++ manager, + # where addToken() only updates manager-level generation counters. + changed = False + for lc_idx, _ in self.manager._life_cycles.attention_life_cycles(): + for block_range in self._block_ranges_excluding( + block_begin, block_end, excluded_ranges[lc_idx] + ): + changed |= self._pending_stats.record_allocation_range( + lc_idx, + block_range.beg, + block_range.end, + beam_width=int(beam_width), + count_as_missed=not count_as_generation, + count_as_generation=count_as_generation, + ) + if changed: + self.manager.mark_stats_dirty(self.id) + + @staticmethod + def _has_reuse_source(page: BlockPage) -> bool: + if page is None or not isinstance(page.page, CommittedPage): + return False + return page.page.block() is not None + + def _subtract_pending_allocation_range( + self, block_begin: BlockOrdinal, block_end: BlockOrdinal + ) -> None: + if self._pending_stats.subtract_allocation_range(block_begin, block_end): + self._refresh_stats_dirty_state() + + def _record_direct_iteration_stats( + self, life_cycle: LifeCycleId, iteration_stats: KVCacheIterationStatsDelta + ) -> None: + life_cycle_key = self._stats_life_cycle_key(life_cycle) + if life_cycle_key is None or iteration_stats.empty or not self._should_record_stats(): + return + self.manager.commit_stats(KVCacheStatsDelta(), {life_cycle_key: iteration_stats}) + + def _record_migrated_slots( + self, + pages: Sequence[Page], + slots: Sequence[Slot], + src_level: CacheLevel, + dst_level: CacheLevel, + ) -> None: + if not self._should_record_stats(): + return + assert len(pages) == len(slots) + for page in pages: + life_cycle_key = self._stats_life_cycle_key(page.life_cycle) + if life_cycle_key is None: + continue + pg_idx = self.manager._storage.get_pool_group_index(page.life_cycle) + page_size = sum(self.manager._storage.slot_size(pg_idx)) + stats = KVCacheStatsDelta() + iteration_stats = KVCacheIterationStatsDelta() + if src_level == GPU_LEVEL and dst_level > GPU_LEVEL: + iteration_stats.iter_offload_blocks = 1 + iteration_stats.iter_offload_bytes = page_size + elif dst_level == GPU_LEVEL: + stats.alloc_total_blocks = 1 + stats.alloc_new_blocks = 1 + iteration_stats.iter_alloc_total_blocks = 1 + iteration_stats.iter_alloc_new_blocks = 1 + if src_level > GPU_LEVEL: + iteration_stats.iter_onboard_blocks = 1 + iteration_stats.iter_onboard_bytes = page_size + elif src_level == GPU_LEVEL: + iteration_stats.iter_intra_device_copy_blocks = 1 + iteration_stats.iter_intra_device_copy_bytes = page_size + if not stats.empty or not iteration_stats.empty: + self.manager.commit_stats(stats, {life_cycle_key: iteration_stats}) + # destroy ownership of memory blocks, so KV cache manager can decide to evict or drop them. After # close, uncommitted data in blocks for (beam_index >= beam_width) will be lost. def close(self) -> None: assert NDEBUG or self._check_sanity() if self.status == self.Status.CLOSED: return + self.discard_pending_stats() self.stop_committing() assert NDEBUG or self._check_sanity() manager = self.manager @@ -511,11 +663,13 @@ def resize(self, capacity: int | None, history_length: int | None = None) -> boo f"history_length ({history_length}) <= " f"old_capacity ({self._capacity})" ) + record_generation_alloc_stats = self._should_record_generation_alloc_stats(capacity) if ( not enable_scratch and self._shortcut_set_capacity(capacity) and self._shortcut_set_history_length(history_length) ): + self._refresh_generation_alloc_ready() return True ssm_lc_id = manager._life_cycles.ssm_life_cycle_id beam_width = self.beam_width @@ -525,6 +679,7 @@ def resize(self, capacity: int | None, history_length: int | None = None) -> boo num_life_cycles = manager._life_cycles.size if new_num_blocks < old_num_blocks: assert not self.has_scratch_slots, "Cannot shrink while scratch slots exist" + self._subtract_pending_allocation_range(new_num_blocks, old_num_blocks) with self._record_event(): del self._blocks[new_num_blocks:] for beam_indices in self._base_page_indices: @@ -577,7 +732,8 @@ def resize(self, capacity: int | None, history_length: int | None = None) -> boo if any(c > 0 for c in net_alloc_counts): try: new_slots = storage.new_gpu_slots( - make_typed(lambda lc: max(0, net_alloc_counts[lc]), num_life_cycles) + make_typed(lambda lc: max(0, net_alloc_counts[lc]), num_life_cycles), + self._record_migrated_slots, ) except OutOfPagesError: self._recover_excess_scratch_slots(excess_scratch_slots) @@ -627,6 +783,23 @@ def resize(self, capacity: int | None, history_length: int | None = None) -> boo else: if len(indices) < new_num_blocks: raise ValueError("User-provided base page indices is too short") + + stream_wait_events( + self.cuda_stream, (s.ready_event for s in chain.from_iterable(slots)) + ) + + # Scratch blocks use temporary shared SWA slots instead of normal + # per-request KV pages, so they are excluded from alloc/miss stats. + excluded_ranges = ( + scratch_ranges if enable_scratch else to_typed(LifeCycleId, stale_ranges) + ) + self._record_resize_pending_allocations( + old_num_blocks, + new_num_blocks, + beam_width, + excluded_ranges, + record_generation_alloc_stats, + ) for ordinal in typed_range(old_num_blocks, new_num_blocks): block = make_typed( lambda _: filled_list(cast(BlockPage, None), num_life_cycles), beam_width @@ -652,6 +825,7 @@ def resize(self, capacity: int | None, history_length: int | None = None) -> boo assert all(len(slots[lc]) == 0 for lc in typed_range(num_life_cycles)) self._capacity = capacity self._history_length = history_length + self._refresh_generation_alloc_ready() assert NDEBUG or self._check_sanity() return True @@ -822,7 +996,7 @@ def resume(self, cuda_stream: CudaStream | None = None) -> bool: if any(c > 0 for c in num_slots): try: - tmp_slots = storage.new_gpu_slots(num_slots) + tmp_slots = storage.new_gpu_slots(num_slots, self._record_migrated_slots) except OutOfPagesError: return False @@ -855,7 +1029,7 @@ def resume(self, cuda_stream: CudaStream | None = None) -> bool: page = expect_type(_PageHolder, beam_block[lc_idx]).page tasks.append(BatchedLockTarget(page, beam_idx, ordinal, lc_idx)) try: - locks = batched_lock_to_gpu(self, tasks) + locks = batched_lock_to_gpu(self, tasks, self._record_migrated_slots) except OutOfPagesError: for lc_idx, slot in typed_enumerate(deferred_slots): if slot is not None: @@ -897,6 +1071,9 @@ def resume(self, cuda_stream: CudaStream | None = None) -> bool: else: lock = self._block(last_ordinal, beam_idx)[lc_idx] assert type(lock) is _SharedPageLock + # V2 still copies a partial reuse into a private slot before writing to it. + # The copy allocates a block, but it is a miss only without a reusable source. + has_partial_reuse_source = self._has_reuse_source(lock) src_locks.append(lock) pg_idx = storage._life_cycle_grouping[lc_idx] slot_size = storage.slot_size(pg_idx) @@ -911,6 +1088,25 @@ def resume(self, cuda_stream: CudaStream | None = None) -> bool: [CopyTask(dst, src)], self.cuda_stream, ) + if lc_idx != ssm_lc_id: + life_cycle_key = self._stats_life_cycle_key(lc_idx) + if life_cycle_key is not None and self._should_record_stats(): + changed = self._pending_stats.record_allocation_range( + life_cycle_key, + last_ordinal, + BlockOrdinal(last_ordinal + 1), + beam_width=1, + count_as_missed=not has_partial_reuse_source, + ) + if changed: + self.manager.mark_stats_dirty(self.id) + self._record_direct_iteration_stats( + lc_idx, + KVCacheIterationStatsDelta( + iter_intra_device_copy_blocks=1, + iter_intra_device_copy_bytes=sum(storage.slot_size(pg_idx)), + ), + ) # Unlock source pages — _record_event captures all prior cuda work # so the original pages know when we're done reading from them. if src_locks: @@ -1152,6 +1348,9 @@ def _commit_block(self, ordinal: BlockOrdinal, is_last: bool) -> None: seq_block.tree_block = tree_block assert self._get_tree_block(ordinal) is tree_block self._num_committed_blocks = BlockOrdinal(ordinal + 1) + event_manager = self.manager.event_manager + if event_manager is not None: + event_manager.add_stored_block_event_from_block(tree_block) elif tree_block.is_full and self.manager.allow_seq_rebasing and is_full: # Happens when a concurrent request committed the same tokens before us. # Try to replace our pages with pages from the existing block to save memory. @@ -1168,6 +1367,9 @@ def _commit_block(self, ordinal: BlockOrdinal, is_last: bool) -> None: page = cast(UncommittedPage, cast(_SharedPageLock, beam_block[lc]).page) beam_block[lc] = None p = page.convert_to_committed(tree_block, self.finish_event) + event_manager = self.manager.event_manager + if event_manager is not None: + event_manager.add_stored_life_cycle_event_from_block(tree_block, int(lc)) # The page comes from uncommitted page of self, so safe to skip wait. beam_block[lc] = ( p.lock(self, beam_idx, ordinal, lc, skip_wait=True) if locked else p.hold() @@ -1177,7 +1379,9 @@ def _commit_block(self, ordinal: BlockOrdinal, is_last: bool) -> None: beam_block[lc] = cast(_SharedPageLock, beam_block[lc]).holder reuse_list.append((lc, existing_page)) locks = batched_lock_to_gpu( - self, [BatchedLockTarget(p, beam_idx, ordinal, lc) for lc, p in reuse_list] + self, + [BatchedLockTarget(p, beam_idx, ordinal, lc) for lc, p in reuse_list], + self._record_migrated_slots, ) for (lc, _), lock in zip(reuse_list, locks): beam_block[lc] = lock @@ -1267,6 +1471,7 @@ def _lock_held_blocks( BatchedLockTarget(holder.page, beam_idx, ordinal, lc) for ordinal, beam_idx, lc, holder in backup_holders ], + self._record_migrated_slots, ) for lock in locks: user = lock._user @@ -1507,6 +1712,8 @@ def _setup_for_reuse(self, match: ReuseMatch) -> None: self._committed_tokens = self._get_matched_tokens(match) self._history_length = num_tokens self._capacity = num_tokens + full_reused_end = BlockOrdinal(num_tokens // tokens_per_block) + has_partial_match = num_tokens % tokens_per_block != 0 # fill self._blocks self._blocks = to_typed( BlockOrdinalT, @@ -1524,10 +1731,13 @@ def _setup_for_reuse(self, match: ReuseMatch) -> None: beam_idx = DEFAULT_BEAM_INDEX + should_record_stats = self._should_record_stats() for lc_idx, lc in life_cycles.items(): if lc_idx == ssm_lc_id: continue # SSM is handled separately below stale_start, stale_end = _KVCache._get_stale_range(tokens_per_block, num_tokens, lc) + full_reused_blocks = 0 + partial_reused_blocks = 0 for ordinal in chain( typed_range(stale_start), typed_range(stale_end, BlockOrdinal(len(matched))) ): @@ -1536,6 +1746,23 @@ def _setup_for_reuse(self, match: ReuseMatch) -> None: # For partial blocks (last block, not full), we defer the copy to first resume(). # Just store the holder of the original committed page for now. block[lc_idx] = holder + if should_record_stats and isinstance(lc, AttnLifeCycle): + if ordinal < full_reused_end: + full_reused_blocks += 1 + elif ( + has_partial_match + and ordinal == full_reused_end + and self._has_reuse_source(holder) + ): + partial_reused_blocks = 1 + if should_record_stats and isinstance(lc, AttnLifeCycle): + changed = self._pending_stats.record_reuse( + lc_idx, + full_reused_blocks=full_reused_blocks, + partial_reused_blocks=partial_reused_blocks, + ) + if changed: + self.manager.mark_stats_dirty(self.id) # SSM reuse: hold the snapshot from the last matched block. Copy is deferred to first resume(). if ssm_lc_id is not None and matched: snapshot_block = matched[-1] diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_core/_kv_cache_manager.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_core/_kv_cache_manager.py index da065a366158..d358d064c358 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_core/_kv_cache_manager.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_core/_kv_cache_manager.py @@ -18,7 +18,7 @@ from collections.abc import Callable, Sequence from copy import deepcopy from dataclasses import dataclass -from typing import Iterable, Iterator, cast +from typing import TYPE_CHECKING, Iterable, Iterator, cast from .. import rawref from .._block_radix_tree import BlockRadixTree, ReuseMatch, ReuseScope @@ -39,9 +39,10 @@ from .._config import DataRole, KVCacheManagerConfig from .._life_cycle_registry import LayerGroupId, LifeCycle, LifeCycleId, LifeCycleRegistry from .._page import Page, _PageHolder +from .._stats import KVCacheIterationStatsDelta, KVCacheStatsDelta from .._storage._config import BufferId, create_storage_config from .._storage._core import PoolGroupIndex, PoolIndex, SlotId -from .._storage_manager import StorageManager +from .._storage_manager import StorageManager, StorageStatistics from .._utils import ( HalfOpenRange, HomoTuple, @@ -58,6 +59,9 @@ from ._kv_cache import _KVCache from ._moving_average import MovingAverage +if TYPE_CHECKING: + from .._event_manager import KVCacheEventManager + @dataclass(slots=True, frozen=True) class MemoryPoolDesc: @@ -178,6 +182,15 @@ def __call__( return result +@dataclass(slots=True, frozen=True) +class _StorageLevelStats: + pool_group_stats: TypedIndexList[PoolGroupIndex, StorageStatistics] + max_num_blocks: int + free_num_blocks: int + used_num_blocks: int + allocated_bytes: int + + class KVCacheManager: __slots__ = ( "_init_config", @@ -194,6 +207,12 @@ class KVCacheManager: "_num_sampled_kv_caches", "_last_adjustment_time", "_last_update_num_sampled_kv_caches", + "_event_manager", + "_stats_enabled", + "_committed_stats", + "_iteration_stats_by_life_cycle", + "_dirty_stats_kv_cache_ids", + "_stats_excluded_kv_cache_ids", ) _init_config: KVCacheManagerConfig _life_cycles: LifeCycleRegistry @@ -216,13 +235,23 @@ class KVCacheManager: _num_sampled_kv_caches: int _last_adjustment_time: float _last_update_num_sampled_kv_caches: int - - def __init__(self, config: KVCacheManagerConfig) -> None: + _event_manager: "KVCacheEventManager | None" + _stats_enabled: bool + _committed_stats: KVCacheStatsDelta + _iteration_stats_by_life_cycle: dict[LifeCycleId, KVCacheIterationStatsDelta] + _dirty_stats_kv_cache_ids: set[int] + _stats_excluded_kv_cache_ids: set[int] + + def __init__( + self, + config: KVCacheManagerConfig, + event_manager: "KVCacheEventManager | None" = None, + ) -> None: init_cuda_once() config = deepcopy(config) self._init_config = config self._life_cycles = LifeCycleRegistry(config) - self._radix_tree = BlockRadixTree(self._life_cycles, config.tokens_per_block) + self._radix_tree = BlockRadixTree(self._life_cycles, config.tokens_per_block, event_manager) storage_config = create_storage_config(config) self._storage = StorageManager( self._life_cycles, @@ -231,6 +260,7 @@ def __init__(self, config: KVCacheManagerConfig) -> None: config.swa_scratch_reuse, typical_batch=config.typical_step, constraints=config.constraints, + event_manager=event_manager, ) self._living_kv_caches = set[rawref.ref[_KVCache]]() decay = 0.9999 @@ -243,6 +273,12 @@ def __init__(self, config: KVCacheManagerConfig) -> None: self._num_sampled_kv_caches = 0 self._last_adjustment_time = time.monotonic() self._last_update_num_sampled_kv_caches = 0 + self._event_manager = event_manager + self._stats_enabled = config.enable_stats + self._committed_stats = KVCacheStatsDelta() + self._iteration_stats_by_life_cycle = {} + self._dirty_stats_kv_cache_ids = set() + self._stats_excluded_kv_cache_ids = set() def __del__(self) -> None: self.shutdown() @@ -346,11 +382,13 @@ def create_kv_cache( id: int | None = None, custom_priority_callback: Callable[[BlockOrdinal, LifeCycle], Priority] = lambda _, __: PRIORITY_DEFAULT, + expected_prompt_length: int | None = None, ) -> _KVCache: """ reuse_scope: namespace to match before matching any tokens. custom_priority_callback: takes block index and layer sliding window size, returns priority. If priority returned is higher than existing priority for reused blocks, the block priority is updated. + expected_prompt_length: optional prompt length hint used to size SWA scratch slots. Newly created KV cache is suspended. You need to call resume() with a cuda stream to make it active & ready in that stream. Returns None if suspended=False and we don't have enough resource. @@ -364,12 +402,15 @@ def create_kv_cache( reuse_match = ( self._match_reuse(reuse_scope, input_tokens) if input_tokens is not None else None ) + if expected_prompt_length is None and input_tokens is not None: + expected_prompt_length = len(input_tokens) return _KVCache( self, reuse_scope, reuse_match, id, custom_priority_callback, + expected_prompt_length, ) def _match_reuse( @@ -416,6 +457,71 @@ def resize(self, cache_level: CacheLevel, quota: int, best_efforts: bool = False def get_quota(self, cache_level: CacheLevel) -> int: return self._storage._levels[cache_level].storage.total_quota + def _get_storage_level_stats(self, cache_level: CacheLevel) -> _StorageLevelStats: + pool_group_stats = self._storage.get_statistics(cache_level) + max_num_blocks = sum(stat.total for stat in pool_group_stats) + free_num_blocks = sum(stat.available for stat in pool_group_stats) + return _StorageLevelStats( + pool_group_stats=pool_group_stats, + max_num_blocks=max_num_blocks, + free_num_blocks=free_num_blocks, + used_num_blocks=max_num_blocks - free_num_blocks, + allocated_bytes=self.get_quota(cache_level), + ) + + def commit_stats( + self, + stats: KVCacheStatsDelta, + iteration_stats_by_life_cycle: dict[LifeCycleId, KVCacheIterationStatsDelta] | None = None, + ) -> None: + if not self._stats_enabled: + return + self._committed_stats.add(stats) + if iteration_stats_by_life_cycle is None: + return + for life_cycle, iteration_stats in iteration_stats_by_life_cycle.items(): + if iteration_stats.empty: + continue + destination = self._iteration_stats_by_life_cycle.setdefault( + life_cycle, KVCacheIterationStatsDelta() + ) + destination.add(iteration_stats) + + def get_committed_stats(self) -> KVCacheStatsDelta: + return self._committed_stats.copy() + + def get_and_reset_iteration_stats(self) -> dict[LifeCycleId, KVCacheIterationStatsDelta]: + stats = { + life_cycle: delta.copy() + for life_cycle, delta in self._iteration_stats_by_life_cycle.items() + if not delta.empty + } + self._iteration_stats_by_life_cycle.clear() + return stats + + def mark_stats_dirty(self, kv_cache_id: int | None) -> None: + if kv_cache_id is not None: + self._dirty_stats_kv_cache_ids.add(kv_cache_id) + + def clear_stats_dirty(self, kv_cache_id: int | None) -> None: + if kv_cache_id is not None: + self._dirty_stats_kv_cache_ids.discard(kv_cache_id) + + def get_dirty_stats_kv_cache_ids(self) -> set[int]: + return self._dirty_stats_kv_cache_ids.copy() + + def mark_stats_excluded(self, kv_cache_id: int | None) -> None: + if kv_cache_id is not None: + self._stats_excluded_kv_cache_ids.add(kv_cache_id) + self.clear_stats_dirty(kv_cache_id) + + def clear_stats_excluded(self, kv_cache_id: int | None) -> None: + if kv_cache_id is not None: + self._stats_excluded_kv_cache_ids.discard(kv_cache_id) + + def is_stats_excluded(self, kv_cache_id: int | None) -> bool: + return kv_cache_id is not None and kv_cache_id in self._stats_excluded_kv_cache_ids + # sorted by CacheLevel from warm to cold @property def cache_tier_list(self) -> HomoTuple[CacheTier]: @@ -425,6 +531,10 @@ def cache_tier_list(self) -> HomoTuple[CacheTier]: def tokens_per_block(self) -> int: return self._radix_tree.tokens_per_block + @property + def event_manager(self) -> "KVCacheEventManager | None": + return self._event_manager + @property def allow_seq_rebasing(self) -> bool: """ @@ -577,7 +687,7 @@ def check_mismatch( b: TypedIndexList[PoolGroupIndex, float], thres: float, ) -> bool: - return any(not (1 / thres < x / y < thres) for x, y in zip(a, b)) + return any(not (1 / thres < x / y < thres) for x, y in zip(a, b, strict=True)) if level == GPU_LEVEL: return check_mismatch(self._target_ratio_list_gpu, self._current_gpu_ratio, 1.25) @@ -687,15 +797,19 @@ def get_num_slots(seq_len: int) -> TypedIndexList[PoolGroupIndex, int]: for pg in typed_range(num_pool_groups): remaining_slots[pg] -= get_num_slots(1)[pg] * (batch_size - 1) - assert remaining_slots[pg] >= 0 + if remaining_slots[pg] < 0: + return 0 def is_enough(num_blocks: int) -> bool: return all( cnt <= rem - for cnt, rem in zip(get_num_slots(num_blocks * tokens_per_block), remaining_slots) + for cnt, rem in zip( + get_num_slots(num_blocks * tokens_per_block), remaining_slots, strict=True + ) ) - assert is_enough(1) + if not is_enough(1): + return 0 lb = 1 ub = div_up(token_num_upper_bound, tokens_per_block) if is_enough(ub): diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_core/_pending_stats.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_core/_pending_stats.py new file mode 100644 index 000000000000..94d0957537be --- /dev/null +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_core/_pending_stats.py @@ -0,0 +1,190 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from dataclasses import dataclass, field + +from .._common import BlockOrdinal +from .._life_cycle_registry import LifeCycleId +from .._stats import KVCacheIterationStatsDelta, KVCacheStatsDelta + + +@dataclass(slots=True) +class _PendingAllocationSegment: + life_cycle: LifeCycleId + block_begin: BlockOrdinal + block_end: BlockOrdinal + beam_width: int + count_as_missed: bool + count_as_generation: bool + + +@dataclass(slots=True) +class _PendingStatsDelta: + global_stats: KVCacheStatsDelta + request_stats: KVCacheStatsDelta + iteration_stats: KVCacheIterationStatsDelta + life_cycle: LifeCycleId | None = None + + @property + def empty(self) -> bool: + return self.global_stats.empty and self.request_stats.empty and self.iteration_stats.empty + + +@dataclass(slots=True) +class _PendingStats: + request_stats: KVCacheStatsDelta = field(default_factory=KVCacheStatsDelta) + global_stats: KVCacheStatsDelta = field(default_factory=KVCacheStatsDelta) + iteration_stats_by_life_cycle: dict[LifeCycleId, KVCacheIterationStatsDelta] = field( + default_factory=dict + ) + allocation_segments: list[_PendingAllocationSegment] = field(default_factory=list) + + @property + def empty(self) -> bool: + return ( + self.request_stats.empty + and self.global_stats.empty + and not self.iteration_stats_by_life_cycle + ) + + def clear(self) -> None: + self.request_stats.clear() + self.global_stats.clear() + self.iteration_stats_by_life_cycle.clear() + self.allocation_segments.clear() + + def add(self, delta: _PendingStatsDelta) -> bool: + if delta.empty: + return False + if not delta.global_stats.empty: + self.global_stats.add(delta.global_stats) + if not delta.request_stats.empty: + self.request_stats.add(delta.request_stats) + if not delta.iteration_stats.empty: + assert delta.life_cycle is not None + pending = self.iteration_stats_by_life_cycle.setdefault( + delta.life_cycle, KVCacheIterationStatsDelta() + ) + pending.add(delta.iteration_stats) + return True + + def subtract(self, delta: _PendingStatsDelta) -> bool: + if delta.empty: + return False + if not delta.global_stats.empty: + self.global_stats.subtract(delta.global_stats) + if not delta.request_stats.empty: + self.request_stats.subtract(delta.request_stats) + if not delta.iteration_stats.empty: + assert delta.life_cycle is not None + pending = self.iteration_stats_by_life_cycle.get(delta.life_cycle) + if pending is not None: + pending.subtract(delta.iteration_stats) + if pending.empty: + del self.iteration_stats_by_life_cycle[delta.life_cycle] + return True + + @staticmethod + def _allocation_delta( + segment: _PendingAllocationSegment, + block_begin: BlockOrdinal, + block_end: BlockOrdinal, + ) -> _PendingStatsDelta: + num_blocks = max(0, int(block_end) - int(block_begin)) * segment.beam_width + stats = KVCacheStatsDelta( + alloc_total_blocks=num_blocks, + alloc_new_blocks=num_blocks, + missed_blocks=num_blocks if segment.count_as_missed else 0, + ) + request_stats = stats.copy() + iteration_stats = KVCacheIterationStatsDelta( + iter_alloc_total_blocks=num_blocks, + iter_alloc_new_blocks=num_blocks, + iter_missed_blocks=num_blocks if segment.count_as_missed else 0, + iter_gen_alloc_blocks=num_blocks if segment.count_as_generation else 0, + ) + return _PendingStatsDelta(stats, request_stats, iteration_stats, segment.life_cycle) + + def record_allocation_range( + self, + life_cycle: LifeCycleId, + block_begin: BlockOrdinal, + block_end: BlockOrdinal, + *, + beam_width: int, + count_as_missed: bool, + count_as_generation: bool = False, + ) -> bool: + if block_begin >= block_end: + return False + segment = _PendingAllocationSegment( + life_cycle=life_cycle, + block_begin=block_begin, + block_end=block_end, + beam_width=beam_width, + count_as_missed=count_as_missed, + count_as_generation=count_as_generation, + ) + if not self.add(self._allocation_delta(segment, block_begin, block_end)): + return False + self.allocation_segments.append(segment) + return True + + def record_reuse( + self, + life_cycle: LifeCycleId, + *, + full_reused_blocks: int, + partial_reused_blocks: int, + ) -> bool: + reused_blocks = full_reused_blocks + partial_reused_blocks + if reused_blocks == 0: + return False + return self.add( + _PendingStatsDelta( + global_stats=KVCacheStatsDelta(reused_blocks=reused_blocks), + request_stats=KVCacheStatsDelta(reused_blocks=reused_blocks), + iteration_stats=KVCacheIterationStatsDelta( + iter_reused_blocks=reused_blocks, + iter_full_reused_blocks=full_reused_blocks, + iter_partial_reused_blocks=partial_reused_blocks, + ), + life_cycle=life_cycle, + ) + ) + + def subtract_allocation_range(self, block_begin: BlockOrdinal, block_end: BlockOrdinal) -> bool: + if block_begin >= block_end or not self.allocation_segments: + return False + changed = False + idx = len(self.allocation_segments) - 1 + while idx >= 0: + segment = self.allocation_segments[idx] + if segment.block_end <= block_begin: + break + removed_begin = max(block_begin, segment.block_begin) + removed_end = min(block_end, segment.block_end) + if removed_begin >= removed_end: + idx -= 1 + continue + changed = True + self.subtract(self._allocation_delta(segment, removed_begin, removed_end)) + if removed_begin <= segment.block_begin: + del self.allocation_segments[idx] + else: + assert removed_end == segment.block_end + segment.block_end = removed_begin + idx -= 1 + return changed diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_event_manager.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_event_manager.py new file mode 100644 index 000000000000..86fbc3db1bb3 --- /dev/null +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_event_manager.py @@ -0,0 +1,638 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +import time +from collections import deque +from collections.abc import Iterable, Sequence +from dataclasses import dataclass, field, replace +from threading import Condition +from typing import Any, Callable + +from tensorrt_llm.logger import logger +from tensorrt_llm.runtime.kv_cache_hash import ( + KV_CACHE_HASH_ALGO_AUTO, + KV_CACHE_HASH_ALGO_DEFAULT, + KV_CACHE_HASH_ALGO_V1, + KV_CACHE_HASH_ALGO_V2, + KV_CACHE_HASH_ALGO_V2_SHA256_64, + NonTextTokenHashError, + hash_v1_block_key, + truncate_sha256_hash_to_int64, +) + +from ._common import GPU_LEVEL, PRIORITY_DEFAULT, CacheLevel, Priority, TokenIdExt + +EventBlockHash = int | str +BlockHashLike = bytes | EventBlockHash +BlockHashesLike = BlockHashLike | Iterable[BlockHashLike] +LayerGroupId = int | None +EventTokenId = int | str +MmKey = tuple[bytes, int] | tuple[bytes, int, str | None] +AttentionDpGatherFn = Callable[[list["KVCacheEvent"]], list[list["KVCacheEvent"]]] + + +@dataclass(slots=True, frozen=True) +class UniqueToken: + token_id: EventTokenId + token_extra_id: int = 0 + + +@dataclass(slots=True, frozen=True) +class KVCacheCreatedData: + num_blocks_per_cache_level: list[int] + + +@dataclass(slots=True, frozen=True) +class KVCacheStoredBlockData: + block_hash: EventBlockHash + tokens: list[UniqueToken] + cache_level: int + priority: int + mm_keys: list[MmKey] = field(default_factory=list) + cache_salt: str | None = None + + +@dataclass(slots=True, frozen=True) +class KVCacheStoredData: + parent_hash: EventBlockHash | None + blocks: list[KVCacheStoredBlockData] + + +@dataclass(slots=True, frozen=True) +class KVCacheRemovedData: + block_hashes: list[EventBlockHash] + + +@dataclass(slots=True, frozen=True) +class KVCacheEventDiff: + old_value: int + new_value: int + + +@dataclass(slots=True, frozen=True) +class KVCacheUpdatedData: + block_hash: EventBlockHash + cache_level: KVCacheEventDiff | None + priority: KVCacheEventDiff | None + + +@dataclass(slots=True, frozen=True) +class KVCacheEvent: + event_id: int + data: KVCacheCreatedData | KVCacheStoredData | KVCacheRemovedData | KVCacheUpdatedData + window_size: int + hash_algo: str | None = None + attention_dp_rank: int | None = None + layer_group_id: int | None = None + + +@dataclass(slots=True) +class _StoredBlockState: + block_hash: EventBlockHash + life_cycle_ids: set[int] + + +class KVCacheEventManager: + """Python event queue matching the C++ KV cache event serializer contract.""" + + def __init__( + self, + max_kv_event_entries: int, + *, + window_size: int = 0, + attention_dp_rank: int | None = None, + attention_dp_gather: AttentionDpGatherFn | None = None, + hash_algo: str = KV_CACHE_HASH_ALGO_V2, + window_size_by_layer_group: dict[int, int] | None = None, + ) -> None: + if hash_algo == KV_CACHE_HASH_ALGO_AUTO: + hash_algo = KV_CACHE_HASH_ALGO_DEFAULT + elif hash_algo not in ( + KV_CACHE_HASH_ALGO_V1, + KV_CACHE_HASH_ALGO_V2, + KV_CACHE_HASH_ALGO_V2_SHA256_64, + ): + raise ValueError(f"Unsupported V2 KV cache event hash algorithm: {hash_algo}") + self._max_kv_event_entries = max_kv_event_entries + self._window_size = window_size + self._window_size_by_layer_group = dict(window_size_by_layer_group or {}) + self._attention_dp_rank = attention_dp_rank + self._attention_dp_gather = attention_dp_gather + self._hash_algo = hash_algo + self._next_event_id = 0 + self._stored_blocks: dict[bytes, _StoredBlockState] = {} + self._latest_stored_events: dict[LayerGroupId, KVCacheEvent] = {} + self._latest_removed_block_hashes: dict[LayerGroupId, list[EventBlockHash]] = {} + self._pending_events: list[KVCacheEvent] = [] + self._events: deque[KVCacheEvent] = deque() + self._condition = Condition() + self._v1_hash_by_block_key: dict[bytes, int] = {} + self._v1_hash_compatible_keys: set[bytes] = set() + self._v1_root_attrs_by_block_key: dict[bytes, tuple[int | None, int | None]] = {} + self._warned_v1_hash_fallback = False + + def add_created_event( + self, + num_blocks_per_cache_level: Sequence[int], + layer_group_ids: Sequence[int] | None = None, + ) -> None: + data = KVCacheCreatedData(list(num_blocks_per_cache_level)) + if layer_group_ids is None: + self._add_event(data) + return + for layer_group_id in layer_group_ids: + self._add_event(data, layer_group_id=int(layer_group_id)) + + def set_layer_group_window_sizes(self, window_sizes: dict[int, int]) -> None: + with self._condition: + self._window_size_by_layer_group = dict(window_sizes) + + def add_stored_event( + self, + parent_hash: EventBlockHash | None, + blocks: Sequence[KVCacheStoredBlockData], + layer_group_id: int | None = None, + ) -> None: + if not blocks: + return + self._flush_removed_events(layer_group_id) + self._add_stored_event( + KVCacheStoredData(parent_hash, list(blocks)), + layer_group_id=layer_group_id, + ) + + def add_stored_block_event_from_block(self, block: Any) -> None: + life_cycle_ids = self._life_cycle_ids_from_radix_block(block) + if not life_cycle_ids: + return + parent_hash = self._parent_hash_from_radix_block(block) + self._stored_blocks[block.key] = _StoredBlockState( + block_hash=self._hash_from_radix_block(block), + life_cycle_ids=set(life_cycle_ids), + ) + for life_cycle_id in sorted(life_cycle_ids): + block_data = self._stored_block_from_radix_block(block, life_cycle_ids={life_cycle_id}) + if block_data is not None: + self.add_stored_event(parent_hash, [block_data], life_cycle_id) + + def add_stored_life_cycle_event_from_block(self, block: Any, life_cycle_id: int) -> None: + state = self._stored_blocks.get(block.key) + life_cycle_id = int(life_cycle_id) + if state is not None: + if life_cycle_id in state.life_cycle_ids: + return + block_data = self._stored_block_from_radix_block(block, life_cycle_ids={life_cycle_id}) + if block_data is None: + return + state.life_cycle_ids.add(life_cycle_id) + self.add_stored_event( + self._parent_hash_from_radix_block(block), + [block_data], + layer_group_id=life_cycle_id, + ) + return + self.add_stored_block_event_from_block(block) + + def add_removed_event(self, block_hashes: BlockHashesLike) -> None: + removed_block_hashes_by_layer_group: dict[int, list[EventBlockHash]] = {} + removed_block_hashes_without_layer_group: list[EventBlockHash] = [] + for block_hash in self._iter_block_hashes(block_hashes): + removed_state = self._pop_stored_block_state(block_hash) + if removed_state is None: + continue + normalized_hash, life_cycle_ids = removed_state + if life_cycle_ids: + for life_cycle_id in sorted(life_cycle_ids): + removed_block_hashes_by_layer_group.setdefault(life_cycle_id, []).append( + normalized_hash + ) + else: + removed_block_hashes_without_layer_group.append(normalized_hash) + + if removed_block_hashes_without_layer_group: + self._enqueue_removed_event(removed_block_hashes_without_layer_group) + for layer_group_id, removed_block_hashes in sorted( + removed_block_hashes_by_layer_group.items() + ): + self._enqueue_removed_event(removed_block_hashes, layer_group_id=layer_group_id) + + def add_removed_life_cycle_event(self, block_hash: bytes, life_cycle_id: int) -> None: + removed_state = self._pop_stored_life_cycle_block_state(block_hash, life_cycle_id) + if removed_state is None: + return + normalized_hash, removed_life_cycle_id, _ = removed_state + self._enqueue_removed_event( + [normalized_hash], + layer_group_id=removed_life_cycle_id, + ) + + def add_updated_event( + self, + block_hash: BlockHashLike, + *, + cache_level: KVCacheEventDiff | None = None, + priority: KVCacheEventDiff | None = None, + layer_group_id: int | None = None, + ) -> None: + if cache_level is None and priority is None: + return + normalized_block_hash = self._get_stored_block_hash(block_hash) + if normalized_block_hash is None: + return + self._add_event( + KVCacheUpdatedData( + block_hash=normalized_block_hash, + cache_level=cache_level, + priority=priority, + ), + layer_group_id=layer_group_id, + ) + + def flush_iteration_events(self) -> None: + if self._attention_dp_gather is not None: + with self._condition: + local_events = self._drain_pending_events_unlocked() + local_events = self._trim_events(local_events, self._max_kv_event_entries) + gathered_events = self._attention_dp_gather(local_events) + if self._attention_dp_rank != 0: + return + events = [ + event + for rank_events in gathered_events + for event in self._trim_events(rank_events, self._max_kv_event_entries) + ] + with self._condition: + self._publish_events_unlocked( + events, + max_kv_event_entries=( + self._max_kv_event_entries * max(1, len(gathered_events)) + ), + ) + self._condition.notify_all() + return + + with self._condition: + self._publish_events_unlocked(self._drain_pending_events_unlocked()) + self._condition.notify_all() + + def get_latest_events(self, timeout_ms: float | None = None) -> list[KVCacheEvent]: + with self._condition: + if not self._events and timeout_ms is None: + while not self._events: + self._condition.wait() + elif not self._events and timeout_ms > 0: + deadline = time.monotonic() + timeout_ms / 1000 + while not self._events: + remaining = deadline - time.monotonic() + if remaining <= 0: + break + self._condition.wait(timeout=remaining) + + events = list(self._events) + self._events.clear() + return events + + def _add_event( + self, + data: KVCacheCreatedData | KVCacheStoredData | KVCacheRemovedData | KVCacheUpdatedData, + layer_group_id: LayerGroupId = None, + ) -> None: + if self._max_kv_event_entries <= 0: + return + with self._condition: + self._add_event_unlocked(data, layer_group_id) + + def _add_stored_event( + self, + data: KVCacheStoredData, + layer_group_id: LayerGroupId = None, + ) -> None: + if self._max_kv_event_entries <= 0: + return + with self._condition: + has_pending_removed_events = bool(self._latest_removed_block_hashes) + latest_event = self._latest_stored_events.get(layer_group_id) + if ( + not has_pending_removed_events + and latest_event is not None + and isinstance(latest_event.data, KVCacheStoredData) + ): + latest_blocks = latest_event.data.blocks + if latest_blocks and latest_blocks[-1].block_hash == data.parent_hash: + merged_data = replace( + latest_event.data, + blocks=[*latest_blocks, *data.blocks], + ) + merged_event = replace(latest_event, data=merged_data) + self._replace_pending_event_unlocked(latest_event, merged_event) + self._latest_stored_events[layer_group_id] = merged_event + return + + event = self._add_event_unlocked(data, layer_group_id) + self._latest_stored_events[layer_group_id] = event + + def _replace_pending_event_unlocked( + self, + old_event: KVCacheEvent, + new_event: KVCacheEvent, + ) -> None: + for event_idx in range(len(self._pending_events) - 1, -1, -1): + if self._pending_events[event_idx] is old_event: + self._pending_events[event_idx] = new_event + return + raise RuntimeError("Stored event coalescing lost the pending event") + + def _enqueue_removed_event( + self, + block_hashes: Sequence[EventBlockHash], + layer_group_id: LayerGroupId = None, + ) -> None: + if not block_hashes or self._max_kv_event_entries <= 0: + return + with self._condition: + self._latest_removed_block_hashes.setdefault(layer_group_id, []).extend(block_hashes) + self._latest_stored_events.pop(layer_group_id, None) + + def _flush_removed_events(self, layer_group_id: LayerGroupId) -> None: + if self._max_kv_event_entries <= 0: + return + with self._condition: + self._flush_removed_events_unlocked(layer_group_id) + + def _flush_removed_events_unlocked(self, layer_group_id: LayerGroupId) -> None: + block_hashes = self._latest_removed_block_hashes.pop(layer_group_id, None) + if not block_hashes: + return + self._add_event_unlocked( + KVCacheRemovedData(block_hashes), + layer_group_id=layer_group_id, + ) + + def _flush_all_removed_events_unlocked(self) -> None: + layer_group_ids = list(self._latest_removed_block_hashes) + for layer_group_id in layer_group_ids: + self._flush_removed_events_unlocked(layer_group_id) + + def _add_event_unlocked( + self, + data: KVCacheCreatedData | KVCacheStoredData | KVCacheRemovedData | KVCacheUpdatedData, + layer_group_id: LayerGroupId = None, + ) -> KVCacheEvent: + if not isinstance(data, KVCacheRemovedData): + self._flush_all_removed_events_unlocked() + event = KVCacheEvent( + event_id=self._next_event_id, + data=data, + window_size=self._get_window_size(layer_group_id), + hash_algo=self._hash_algo, + attention_dp_rank=self._attention_dp_rank, + layer_group_id=layer_group_id, + ) + self._next_event_id += 1 + self._pending_events.append(event) + if not isinstance(data, KVCacheStoredData): + self._latest_stored_events.pop(layer_group_id, None) + return event + + def _drain_pending_events_unlocked(self) -> list[KVCacheEvent]: + self._flush_all_removed_events_unlocked() + events = self._pending_events + self._pending_events = [] + self._latest_stored_events.clear() + return events + + def _publish_events_unlocked( + self, + events: Sequence[KVCacheEvent], + *, + max_kv_event_entries: int | None = None, + ) -> None: + if not events: + return + if max_kv_event_entries is None: + max_kv_event_entries = self._max_kv_event_entries + self._events.extend(events) + while len(self._events) > max_kv_event_entries: + self._events.popleft() + + @staticmethod + def _trim_events( + events: Sequence[KVCacheEvent], max_kv_event_entries: int + ) -> list[KVCacheEvent]: + if max_kv_event_entries <= 0: + return [] + if len(events) <= max_kv_event_entries: + return list(events) + return list(events[-max_kv_event_entries:]) + + def _get_window_size(self, layer_group_id: LayerGroupId) -> int: + if layer_group_id is None: + return self._window_size + return self._window_size_by_layer_group.get(int(layer_group_id), self._window_size) + + @staticmethod + def _iter_block_hashes(block_hashes: BlockHashesLike) -> Iterable[BlockHashLike]: + if isinstance(block_hashes, (bytes, str, int)): + return (block_hashes,) + return block_hashes + + def _normalize_block_hash(self, block_hash: BlockHashLike) -> EventBlockHash: + if isinstance(block_hash, bytes): + if self._hash_algo == KV_CACHE_HASH_ALGO_V2_SHA256_64: + return truncate_sha256_hash_to_int64(block_hash) + return block_hash.hex() + return block_hash + + def _get_stored_block_hash(self, block_hash: BlockHashLike) -> EventBlockHash | None: + if isinstance(block_hash, bytes): + state = self._stored_blocks.get(block_hash) + return None if state is None else state.block_hash + return block_hash + + def _pop_stored_block_state( + self, block_hash: BlockHashLike + ) -> tuple[EventBlockHash, set[int]] | None: + if isinstance(block_hash, bytes): + state = self._stored_blocks.pop(block_hash, None) + if state is None: + return None + self._drop_hash_cache(block_hash) + return state.block_hash, set(state.life_cycle_ids) + return block_hash, set() + + def _pop_stored_life_cycle_block_state( + self, block_hash: bytes, life_cycle_id: int + ) -> tuple[EventBlockHash, int, bool] | None: + state = self._stored_blocks.get(block_hash) + if state is None or not state.life_cycle_ids: + return None + + life_cycle_id = int(life_cycle_id) + if life_cycle_id not in state.life_cycle_ids: + return None + + state.life_cycle_ids.remove(life_cycle_id) + is_last_life_cycle = not state.life_cycle_ids + if is_last_life_cycle: + self._stored_blocks.pop(block_hash, None) + self._drop_hash_cache(block_hash) + return state.block_hash, life_cycle_id, is_last_life_cycle + + def _drop_hash_cache(self, block_hash: bytes) -> None: + self._v1_hash_by_block_key.pop(block_hash, None) + self._v1_hash_compatible_keys.discard(block_hash) + self._v1_root_attrs_by_block_key.pop(block_hash, None) + + @staticmethod + def _normalize_token(token: TokenIdExt) -> UniqueToken: + if isinstance(token, bytes): + return UniqueToken(token.hex()) + return UniqueToken(int(token)) + + def _stored_block_from_radix_block( + self, block: Any, life_cycle_ids: set[int] | None = None + ) -> KVCacheStoredBlockData | None: + cache_level: CacheLevel = GPU_LEVEL + priority: Priority = PRIORITY_DEFAULT + found_page = False + for life_cycle_id, page_ref in enumerate(block.storage): + if life_cycle_ids is not None and life_cycle_id not in life_cycle_ids: + continue + if page_ref is None: + continue + page = page_ref() + if page is None: + continue + cache_level = page.cache_level + priority = page.priority + found_page = True + break + + if life_cycle_ids is not None and not found_page: + return None + + return KVCacheStoredBlockData( + block_hash=self._hash_from_radix_block(block), + tokens=[self._normalize_token(token) for token in block.tokens], + cache_level=int(cache_level), + priority=int(priority), + mm_keys=[], + ) + + @staticmethod + def _life_cycle_ids_from_radix_block(block: Any) -> set[int]: + return { + life_cycle_id + for life_cycle_id, page_ref in enumerate(block.storage) + if page_ref is not None and page_ref() is not None + } + + def _parent_hash_from_radix_block(self, block: Any) -> EventBlockHash | None: + parent = block.prev + if getattr(parent, "ordinal", -1) == -1: + return None + return self._hash_from_radix_block(parent) + + def _hash_from_radix_block(self, block: Any) -> EventBlockHash: + if self._hash_algo == KV_CACHE_HASH_ALGO_V1: + return self._v1_hash_from_radix_block(block) + return self._normalize_block_hash(block.key) + + def _v1_hash_from_radix_block(self, block: Any) -> int: + key = bytes(block.key) + cached = self._v1_hash_by_block_key.get(key) + if cached is not None: + return cached + + chain: list[Any] = [] + current = block + while self._is_radix_block(current): + current_key = bytes(current.key) + cached = self._v1_hash_by_block_key.get(current_key) + if cached is not None: + parent_hash = cached + parent_is_v1_compatible = current_key in self._v1_hash_compatible_keys + root_attrs = self._v1_root_attrs_by_block_key[current_key] + break + chain.append(current) + current = current.prev + + if not self._is_radix_block(current): + parent_hash = 0 + parent_is_v1_compatible = True + root_attrs = self._root_attrs_from_root_block(current) + + lora_task_id, cache_salt_id = root_attrs + for current in reversed(chain): + current_key = bytes(current.key) + if parent_is_v1_compatible: + try: + parent_hash = self._hash_block_key( + current.tokens, + parent_hash, + lora_task_id, + cache_salt_id, + ) + self._v1_hash_compatible_keys.add(current_key) + except NonTextTokenHashError: + parent_hash = self._fallback_v1_hash(current_key) + parent_is_v1_compatible = False + else: + parent_hash = self._fallback_v1_hash(current_key) + self._v1_hash_by_block_key[current_key] = parent_hash + self._v1_root_attrs_by_block_key[current_key] = root_attrs + + return parent_hash + + def _fallback_v1_hash(self, block_key: bytes) -> int: + if not self._warned_v1_hash_fallback: + logger.warning( + "V2 KV cache event hash algorithm %s only matches v1 for " + "text-token radix blocks. Falling back to truncated V2 block " + "hash for unsupported blocks.", + KV_CACHE_HASH_ALGO_V1, + ) + self._warned_v1_hash_fallback = True + return truncate_sha256_hash_to_int64(block_key) + + @staticmethod + def _is_radix_block(value: Any) -> bool: + return hasattr(value, "tokens") and hasattr(value, "key") + + @staticmethod + def _root_attrs_from_root_block(root: Any) -> tuple[int | None, int | None]: + # Read from the new location first; fall back to the legacy attribute + # names so V1-compat event hashes still resolve for any in-memory + # RootBlock predating the refactor. + scope = getattr(root, "reuse_scope", None) + if scope is not None: + return getattr(scope, "lora_id", None), getattr(scope, "salt", None) + return getattr(root, "lora_task_id", None), getattr(root, "cache_salt_id", None) + + @staticmethod + def _hash_block_key( + tokens: Sequence[TokenIdExt], + parent_hash: int, + lora_task_id: int | None, + cache_salt_id: int | None, + ) -> int: + return hash_v1_block_key( + tokens, + parent_hash=parent_hash, + lora_task_id=lora_task_id, + cache_salt_id=cache_salt_id, + ) diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_page.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_page.py index 8b3084267dd2..b276387e2385 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_page.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_page.py @@ -13,7 +13,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -from collections.abc import Sequence +from collections.abc import Callable, Sequence from dataclasses import dataclass, field from typing import TYPE_CHECKING, NamedTuple, cast @@ -441,7 +441,10 @@ class BatchedLockTarget(NamedTuple): def batched_lock_to_gpu( - kv_cache: "_KVCache", tasks: Sequence[BatchedLockTarget] + kv_cache: "_KVCache", + tasks: Sequence[BatchedLockTarget], + migration_recorder: Callable[[Sequence[Page], Sequence[Slot], CacheLevel, CacheLevel], None] + | None = None, ) -> list["_SharedPageLock"]: "Lock pages after migrating all pages to GPU. If migration fails, no locking happens." storage = kv_cache.manager._storage @@ -457,13 +460,18 @@ def batched_lock_to_gpu( requirements[lc2pg[t.life_cycle]] += 1 try: - storage.prepare_free_slots(GPU_LEVEL, requirements) + storage.prepare_free_slots(GPU_LEVEL, requirements, migration_recorder) partitioned = partition(tasks, lambda p: (p.page.cache_level, lc2pg[p.life_cycle])) for (lvl, pg_idx), part in partitioned.items(): if lvl == GPU_LEVEL: continue storage._batched_migrate( - pg_idx, GPU_LEVEL, lvl, [p.page for p in part], update_src=True + pg_idx, + GPU_LEVEL, + lvl, + [p.page for p in part], + update_src=True, + migration_recorder=migration_recorder, ) except Exception: for t, e in zip(tasks, scheduled_for_eviction): diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_stats.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_stats.py new file mode 100644 index 000000000000..73b104e7c02a --- /dev/null +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_stats.py @@ -0,0 +1,73 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from dataclasses import dataclass, fields + + +class _StatsDeltaMixin: + __slots__ = () + + def add(self, other) -> None: + for field in fields(self): + name = field.name + setattr(self, name, getattr(self, name) + getattr(other, name)) + + def subtract(self, other) -> None: + for field in fields(self): + name = field.name + setattr(self, name, getattr(self, name) - getattr(other, name)) + + def clear(self) -> None: + for field in fields(self): + setattr(self, field.name, 0) + + def copy(self): + return type(self)(**{field.name: getattr(self, field.name) for field in fields(self)}) + + @property + def empty(self) -> bool: + return all(getattr(self, field.name) == 0 for field in fields(self)) + + +@dataclass(slots=True) +class KVCacheStatsDelta(_StatsDeltaMixin): + alloc_total_blocks: int = 0 + alloc_new_blocks: int = 0 + reused_blocks: int = 0 + missed_blocks: int = 0 + + +@dataclass(slots=True) +class KVCacheIterationStatsDelta(_StatsDeltaMixin): + iter_alloc_total_blocks: int = 0 + iter_alloc_new_blocks: int = 0 + iter_reused_blocks: int = 0 + iter_full_reused_blocks: int = 0 + iter_partial_reused_blocks: int = 0 + iter_missed_blocks: int = 0 + iter_gen_alloc_blocks: int = 0 + iter_onboard_blocks: int = 0 + iter_onboard_bytes: int = 0 + iter_offload_blocks: int = 0 + iter_offload_bytes: int = 0 + iter_intra_device_copy_blocks: int = 0 + iter_intra_device_copy_bytes: int = 0 + + @property + def iter_cache_hit_rate(self) -> float: + total = self.iter_reused_blocks + self.iter_missed_blocks + if self.iter_reused_blocks == 0 or total == 0: + return 0.0 + return self.iter_reused_blocks / total diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_storage_manager.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_storage_manager.py index 9ade62d5dc99..7a32c732c3bf 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_storage_manager.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_storage_manager.py @@ -19,7 +19,7 @@ from collections import deque from dataclasses import dataclass from fractions import Fraction -from typing import Iterator, Sequence, cast +from typing import TYPE_CHECKING, Callable, Iterator, Sequence, cast from . import rawref from ._common import ( @@ -42,10 +42,11 @@ SwaScratchReuseConfig, ) from ._copy_engine import CopyTask, batched_copy +from ._event_manager import KVCacheEventDiff from ._eviction_controller import EvictablePage, PerLevelEvictionController from ._exceptions import OutOfPagesError from ._life_cycle_registry import LifeCycleId, LifeCycleRegistry, compute_scratch_range -from ._page import Page +from ._page import CommittedPage, Page from ._storage import CacheLevelStorage from ._storage._config import BufferAttr, BufferId, LayerAttr, SlotDesc, StorageConfig from ._storage._core import ( @@ -80,6 +81,9 @@ typed_range, ) +if TYPE_CHECKING: + from ._event_manager import KVCacheEventManager + class CacheLevelManager: __slots__ = ("cache_level", "storage", "controller") @@ -165,6 +169,9 @@ def unavailable(self) -> int: return self.total - self.available +MigrationRecorder = Callable[[Sequence[Page], Sequence[Slot], CacheLevel, CacheLevel], None] + + class StorageManager: __slots__ = ( "_life_cycles", @@ -177,6 +184,7 @@ class StorageManager: "_slot_desc_list", "_levels", "_min_slots", + "_event_manager", "__rawref__", ) _life_cycles: LifeCycleRegistry @@ -189,6 +197,7 @@ class StorageManager: _slot_desc_list: TypedIndexList[PoolGroupIndex, SlotDesc] _levels: TypedIndexList[CacheLevel, CacheLevelManager] _min_slots: TypedIndexList[PoolGroupIndex, int] + _event_manager: "KVCacheEventManager | None" __rawref__: rawref.ref["StorageManager"] def __init__( @@ -199,8 +208,10 @@ def __init__( swa_scratch_reuse: SwaScratchReuseConfig | None, typical_batch: BatchDesc | None = None, constraints: list[BatchDesc] | None = None, + event_manager: "KVCacheEventManager | None" = None, ) -> None: self.__rawref__ = rawref.NULL + self._event_manager = event_manager assert config.cache_tiers[GPU_LEVEL].tier == CacheTier.GPU_MEM, ( "The first cache tier must be GPU memory" ) @@ -277,12 +288,17 @@ def get_pool_group_index(self, life_cycle: LifeCycleId) -> PoolGroupIndex: return self._life_cycle_grouping[life_cycle] def new_gpu_slots( - self, num_slots: TypedIndexList[LifeCycleId, int] + self, + num_slots: TypedIndexList[LifeCycleId, int], + migration_recorder: MigrationRecorder | None = None, ) -> TypedIndexList[LifeCycleId, list[Slot]]: - return self.new_slots(GPU_LEVEL, num_slots) + return self.new_slots(GPU_LEVEL, num_slots, migration_recorder) def new_slots( - self, level: CacheLevel, num_slots: TypedIndexList[LifeCycleId, int] + self, + level: CacheLevel, + num_slots: TypedIndexList[LifeCycleId, int], + migration_recorder: MigrationRecorder | None = None, ) -> TypedIndexList[LifeCycleId, list[Slot]]: lc2pg = self._life_cycle_grouping pg_num_slots = filled_list(0, self.num_pool_groups) @@ -293,7 +309,7 @@ def new_slots( pg_num_slots[pg] > storage.get_num_free_slots(pg) for pg in typed_range(self.num_pool_groups) ): - self.prepare_free_slots(level, pg_num_slots) + self.prepare_free_slots(level, pg_num_slots, migration_recorder) assert all( pg_num_slots[pg] <= storage.get_num_free_slots(pg) for pg in typed_range(self.num_pool_groups) @@ -313,13 +329,17 @@ def new_slots( return ret def new_slots_for_pool_group( - self, level: CacheLevel, pg_idx: PoolGroupIndex, num_slots: int + self, + level: CacheLevel, + pg_idx: PoolGroupIndex, + num_slots: int, + migration_recorder: MigrationRecorder | None = None, ) -> list[Slot]: storage = self._levels[level].storage if num_slots > storage.get_num_free_slots(pg_idx): num_slots_list = filled_list(0, self.num_pool_groups) num_slots_list[pg_idx] = num_slots - self.prepare_free_slots(level, num_slots_list) + self.prepare_free_slots(level, num_slots_list, migration_recorder) assert num_slots <= storage.get_num_free_slots(pg_idx) try: return storage.allocate_multiple(pg_idx, num_slots) @@ -367,13 +387,16 @@ def is_evictable(self, page: EvictablePage, level: CacheLevel | None = None) -> ) def prepare_free_slots( - self, level: CacheLevel, requirements: TypedIndexList[PoolGroupIndex, int] + self, + level: CacheLevel, + requirements: TypedIndexList[PoolGroupIndex, int], + migration_recorder: MigrationRecorder | None = None, ) -> None: goals = filled_array2d(self.num_cache_levels, self.num_pool_groups, 0) for pg in typed_range(self.num_pool_groups): goals[level, pg] = requirements[pg] fallen_pages = make_typed(lambda _: list[Page](), self.num_pool_groups) - self._prepare_free_slots(goals, level, fallen_pages) + self._prepare_free_slots(goals, level, fallen_pages, migration_recorder) def force_evict( self, level: CacheLevel, min_num_pages: TypedIndexList[PoolGroupIndex, int] @@ -397,6 +420,7 @@ def _prepare_free_slots( goals: Array2D[CacheLevel, PoolGroupIndex, int], lvl_id: CacheLevel, fallen_pages: TypedIndexList[PoolGroupIndex, list[Page]], + migration_recorder: MigrationRecorder | None = None, ) -> None: assert NDEBUG or goals.rows == self.num_cache_levels and goals.cols == self.num_pool_groups assert NDEBUG or all( @@ -468,7 +492,12 @@ def _prepare_free_slots( if num_accepted > 0: accepted_pages[pg_idx] = fallen_pages[pg_idx][-num_accepted:] del fallen_pages[pg_idx][-num_accepted:] - self._prepare_free_slots(goals, CacheLevel(lvl_id + 1), fallen_pages) + self._prepare_free_slots( + goals, + CacheLevel(lvl_id + 1), + fallen_pages, + migration_recorder, + ) assert all(len(f) == 0 for f in fallen_pages) # migrate pages for pg_idx in typed_range(self.num_pool_groups): @@ -479,7 +508,14 @@ def _prepare_free_slots( accepted_pages[pg_idx].clear() for (src_lvl, pg_idx), pages in partitioned.items(): dst_lvl = lvl_id - self._batched_migrate(pg_idx, dst_lvl, src_lvl, pages, update_src=True) + self._batched_migrate( + pg_idx, + dst_lvl, + src_lvl, + pages, + update_src=True, + migration_recorder=migration_recorder, + ) for p in pages: if is_last_level and p.status == PageStatus.HELD: continue @@ -493,6 +529,7 @@ def _batched_migrate( src_level: CacheLevel, src_pages: Sequence[Page], update_src: bool, + migration_recorder: MigrationRecorder | None = None, defrag: bool = False, # we are doing defragmentation ) -> Sequence[Slot] | None: "Free slots must be prepared before calling this function." @@ -528,6 +565,15 @@ def _batched_migrate( for pool_idx, tasks in typed_enumerate(tasks_per_pool): batched_copy(dst_tier, src_tier, slot_sizes[pool_idx], tasks, stream.get()) finish_event = stream.take_finish_event() + emit_cache_level_updates = ( + update_src + and not defrag + and src_level != dst_level + and self._event_manager is not None + ) + emitted_update_keys: set[tuple[bytes, LifeCycleId]] = set() + if migration_recorder is not None and not defrag: + migration_recorder(src_pages, dst_slots, src_level, dst_level) for src, dst in zip(src_pages, dst_slots): dst.ready_event = finish_event src.ready_event = ( @@ -540,6 +586,10 @@ def _batched_migrate( src_pool_group.release(src) src.set_slot(dst) src.cache_level = dst_level + if emit_cache_level_updates: + self._emit_cache_level_updated_event( + src, src_level, dst_level, emitted_update_keys + ) if scheduled_for_eviction: self.schedule_for_eviction(src) return None if update_src else dst_slots @@ -548,6 +598,34 @@ def _batched_migrate( dst_pool_group.release(s) raise + def _emit_cache_level_updated_event( + self, + page: Page, + old_level: CacheLevel, + new_level: CacheLevel, + emitted_keys: set[tuple[bytes, LifeCycleId]], + ) -> None: + if self._event_manager is None or not isinstance(page, CommittedPage): + return + + block = page.block() + if block is None or block.is_orphan: + return + + event_key = (block.key, page.life_cycle) + if event_key in emitted_keys: + return + + emitted_keys.add(event_key) + self._event_manager.add_updated_event( + block.key, + cache_level=KVCacheEventDiff( + old_value=int(old_level), + new_value=int(new_level), + ), + layer_group_id=int(page.life_cycle), + ) + def _pool_group( self, cache_level: CacheLevel, pool_group_index: PoolGroupIndex ) -> PoolGroupBase: diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/setup_mypyc.py b/tensorrt_llm/runtime/kv_cache_manager_v2/setup_mypyc.py index cd80e4197bb7..f2f5d9e31c12 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/setup_mypyc.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/setup_mypyc.py @@ -76,6 +76,7 @@ "kv_cache_manager_v2/_config.py", "kv_cache_manager_v2/_copy_engine.py", "kv_cache_manager_v2/_cuda_virt_mem.py", + "kv_cache_manager_v2/_event_manager.py", "kv_cache_manager_v2/_exceptions.py", "kv_cache_manager_v2/_life_cycle_registry.py", "kv_cache_manager_v2/_page.py", @@ -85,6 +86,7 @@ "kv_cache_manager_v2/_core/__init__.py", "kv_cache_manager_v2/_core/_kv_cache_manager.py", "kv_cache_manager_v2/_core/_kv_cache.py", + "kv_cache_manager_v2/_core/_pending_stats.py", # _eviction_controller submodule "kv_cache_manager_v2/_eviction_controller/__init__.py", "kv_cache_manager_v2/_eviction_controller/_eviction_controller.py", diff --git a/tests/unittest/_torch/executor/test_resource_manager.py b/tests/unittest/_torch/executor/test_resource_manager.py index d45092117f6c..17278f9cb293 100644 --- a/tests/unittest/_torch/executor/test_resource_manager.py +++ b/tests/unittest/_torch/executor/test_resource_manager.py @@ -12,8 +12,9 @@ import tensorrt_llm import tensorrt_llm.bindings from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest -from tensorrt_llm._torch.pyexecutor.resource_manager import (KVCacheManager, - PeftCacheManager) +from tensorrt_llm._torch.pyexecutor.resource_manager import ( + KVCacheManager, PeftCacheManager, + _warn_if_unsupported_v1_kv_cache_event_hash_algo) from tensorrt_llm.bindings import LayerType from tensorrt_llm.bindings import ModelConfig as ModelConfigCpp from tensorrt_llm.bindings import executor as tllm @@ -24,6 +25,9 @@ from tensorrt_llm.llmapi.llm_args import KvCacheConfig, PeftCacheConfig from tensorrt_llm.lora_helper import LoraConfig from tensorrt_llm.mapping import Mapping +from tensorrt_llm.runtime.kv_cache_hash import (KV_CACHE_HASH_ALGO_AUTO, + KV_CACHE_HASH_ALGO_V1, + KV_CACHE_HASH_ALGO_V2) from tensorrt_llm.sampling_params import SamplingParams DataType = tensorrt_llm.bindings.DataType @@ -35,6 +39,33 @@ sys.path.append(str(root_dir / "tests" / "integration")) +def test_v1_kv_cache_event_hash_algo_warning_for_non_v1(): + with patch("tensorrt_llm._torch.pyexecutor.resource_manager.logger.warning" + ) as warning: + _warn_if_unsupported_v1_kv_cache_event_hash_algo(KV_CACHE_HASH_ALGO_V2) + + warning.assert_called_once() + assert KV_CACHE_HASH_ALGO_V1 in warning.call_args.args[0] + assert KV_CACHE_HASH_ALGO_V2 in warning.call_args.args[0] + + +def test_v1_kv_cache_event_hash_algo_no_warning_for_v1(): + with patch("tensorrt_llm._torch.pyexecutor.resource_manager.logger.warning" + ) as warning: + _warn_if_unsupported_v1_kv_cache_event_hash_algo(KV_CACHE_HASH_ALGO_V1) + + warning.assert_not_called() + + +def test_v1_kv_cache_event_hash_algo_no_warning_for_auto(): + with patch("tensorrt_llm._torch.pyexecutor.resource_manager.logger.warning" + ) as warning: + _warn_if_unsupported_v1_kv_cache_event_hash_algo( + KV_CACHE_HASH_ALGO_AUTO) + + warning.assert_not_called() + + class TestResourceManager(unittest.TestCase): CPP_RESOURCES_DIR = os.path.join(str(root_dir), "cpp", "tests", "resources") CPP_DATA_DIR = os.path.join(CPP_RESOURCES_DIR, "data") diff --git a/tests/unittest/bindings/test_bindings_ut.py b/tests/unittest/bindings/test_bindings_ut.py index 7585f787c8a9..210ea3378d3e 100644 --- a/tests/unittest/bindings/test_bindings_ut.py +++ b/tests/unittest/bindings/test_bindings_ut.py @@ -2,6 +2,7 @@ import pickle import tempfile import time +from datetime import timedelta from pathlib import Path import numpy as np @@ -416,6 +417,28 @@ def test_llm_request(): assert torch.equal(llm_request.draft_logits, logits) +def test_llm_request_kv_cache_transfer_metric_bindings(): + request = _tb.internal.batch_manager.LlmRequest( + request_id=0, + max_new_tokens=5, + sampling_config=_tb.SamplingConfig(1), + input_tokens=[0, 1, 2], + is_streaming=True, + ) + offset = _tb.internal.batch_manager.LlmRequest.global_steady_clock_offset + offset = offset if offset is not None else timedelta() + start = timedelta(seconds=1.25) + end = timedelta(seconds=2.5) + + request.set_kv_cache_transfer_start(start) + request.set_kv_cache_transfer_end(end) + request.set_kv_cache_size(128) + + assert request.kv_cache_transfer_start == start + offset + assert request.kv_cache_transfer_end == end + offset + assert request.kv_cache_size == 128 + + def test_Mpicomm(): size1 = _tb.MpiComm.size() rank1 = _tb.MpiComm.rank() diff --git a/tests/unittest/disaggregated/test_cache_transceiver_single_process.py b/tests/unittest/disaggregated/test_cache_transceiver_single_process.py index 6bb0bf2c1b8a..e4224ec03f4f 100644 --- a/tests/unittest/disaggregated/test_cache_transceiver_single_process.py +++ b/tests/unittest/disaggregated/test_cache_transceiver_single_process.py @@ -101,6 +101,7 @@ class KvCacheConfigV2: cross_kv_cache_fraction: Optional[float] = None secondary_offload_min_priority: Optional[int] = None event_buffer_max_size: int = 0 + kv_cache_event_hash_algo: str = "auto" max_gpu_total_bytes: Optional[int] = None enable_partial_reuse: bool = False copy_on_partial_reuse: bool = False diff --git a/tests/unittest/disaggregated/test_kv_transfer.py b/tests/unittest/disaggregated/test_kv_transfer.py index e87ce6fe7bcd..caf8dcb3f6d6 100644 --- a/tests/unittest/disaggregated/test_kv_transfer.py +++ b/tests/unittest/disaggregated/test_kv_transfer.py @@ -59,6 +59,7 @@ class KvCacheConfigV2: cross_kv_cache_fraction: Optional[float] = None secondary_offload_min_priority: Optional[int] = None event_buffer_max_size: int = 0 + kv_cache_event_hash_algo: str = "auto" max_gpu_total_bytes: Optional[int] = None enable_partial_reuse: bool = False diff --git a/tests/unittest/executor/test_stats_serializer.py b/tests/unittest/executor/test_stats_serializer.py index 2166af26de2d..07402ae0cde4 100644 --- a/tests/unittest/executor/test_stats_serializer.py +++ b/tests/unittest/executor/test_stats_serializer.py @@ -20,6 +20,11 @@ import pytest +from tensorrt_llm._torch.pyexecutor.kv_cache_stats import ( + KVCacheV2IterationStatsReport, + KVCacheV2LifeCycleIterationStats, + KVCacheV2PoolGroupIterationStats, +) from tensorrt_llm.executor.base_worker import BaseWorker @@ -237,3 +242,80 @@ def test_serializer_8_tuple_emits_new_timing_and_scheduler_mode(self): assert "prevDeviceStepTimeMS" not in d assert d["schedulerMode"] == "overlap" assert d["gpuForwardTimeMS"] == 4.25 + + def test_serializer_with_v2_pool_group_stats(self): + """KV cache manager V2 stats should include pool group breakdown.""" + iter_stats = _make_mock_iteration_stats() + by_window = _make_mock_kv_iter_stats( + window_size=16, + primary_used=10, + primary_max=20, + reused=5, + full_reused=4, + partial_reused=1, + missed=3, + gen_alloc=2, + ) + pool_group_stats = _make_mock_kv_iter_stats( + window_size=16, + primary_used=10, + primary_max=20, + reused=0, + full_reused=0, + partial_reused=0, + missed=0, + gen_alloc=2, + )[16] + life_cycle_stats = _make_mock_kv_iter_stats( + window_size=16, + primary_used=0, + primary_max=0, + reused=5, + full_reused=4, + partial_reused=1, + missed=3, + gen_alloc=0, + )[16] + kv_iter = KVCacheV2IterationStatsReport( + by_window, + { + 7: KVCacheV2PoolGroupIterationStats( + pool_group_id=7, + slot_size=(2 << 20,), + window_sizes=(16, 64), + stats=pool_group_stats, + ) + }, + { + 3: KVCacheV2LifeCycleIterationStats( + life_cycle_id=3, + pool_group_id=7, + window_size=16, + kind="attention", + stats=life_cycle_stats, + ) + }, + ) + + result = BaseWorker._stats_serializer((iter_stats, None, kv_iter)) + d = json.loads(result) + + assert d["kvCacheIterationStats"]["16"]["iterReusedBlocks"] == 5 + assert "kvCacheIterationStatsByPoolGroup" in d + pool_group = d["kvCacheIterationStatsByPoolGroup"]["7"] + assert pool_group["poolGroupId"] == 7 + assert pool_group["slotSize"] == [2 << 20] + assert pool_group["windowSizes"] == [16, 64] + assert pool_group["iterGenAllocBlocks"] == 2 + assert "iterReusedBlocks" not in pool_group + assert "iterMissedBlocks" not in pool_group + assert "iterCacheHitRate" not in pool_group + assert "kvCacheIterationStatsByLifecycle" in d + life_cycle = d["kvCacheIterationStatsByLifecycle"]["3"] + assert life_cycle["lifeCycleId"] == 3 + assert life_cycle["poolGroupId"] == 7 + assert life_cycle["windowSize"] == 16 + assert life_cycle["kind"] == "attention" + assert life_cycle["iterReusedBlocks"] == 5 + assert life_cycle["iterMissedBlocks"] == 3 + assert "iterGenAllocBlocks" not in life_cycle diff --git a/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_event_manager.py b/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_event_manager.py new file mode 100644 index 000000000000..807eb8fa69ea --- /dev/null +++ b/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_event_manager.py @@ -0,0 +1,1170 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import gc +import os +import threading +import time +from importlib.util import find_spec +from typing import TYPE_CHECKING, cast + +import pytest + +from tensorrt_llm._utils import KVCacheEventSerializer +from tensorrt_llm.runtime.kv_cache_hash import ( + KV_CACHE_HASH_ALGO_V1, + KV_CACHE_HASH_ALGO_V2_SHA256_64, + truncate_sha256_hash_to_int64, +) +from tensorrt_llm.runtime.kv_cache_manager_v2._event_manager import ( + KVCacheEventDiff, + KVCacheEventManager, + KVCacheStoredBlockData, + UniqueToken, +) + +if not TYPE_CHECKING and find_spec("kv_cache_manager_v2") is not None: + from kv_cache_manager_v2 import CacheLevel, CudaStream, KVCacheManager, TokenId + from kv_cache_manager_v2._block_radix_tree import Block, ReuseScope, RootBlock + from kv_cache_manager_v2._utils import CachedCudaStream, init_cuda_once, temporary_sys_path +else: + from tensorrt_llm.runtime.kv_cache_manager_v2 import ( + CacheLevel, + CudaStream, + KVCacheManager, + TokenId, + ) + from tensorrt_llm.runtime.kv_cache_manager_v2._block_radix_tree import ( + Block, + ReuseScope, + RootBlock, + ) + from tensorrt_llm.runtime.kv_cache_manager_v2._utils import ( + CachedCudaStream, + init_cuda_once, + temporary_sys_path, + ) + +try: + import torch +except ImportError: + torch = None + + +_DEFAULT_CACHE_LEVEL = CacheLevel(0) + + +class _FakePage: + def __init__(self, cache_level=_DEFAULT_CACHE_LEVEL, priority=0): + self.cache_level = cache_level + self.priority = priority + + +class _FakePageRef: + def __init__(self, page): + self._page = page + + def __call__(self): + return self._page + + +class _FakeRootBlock: + ordinal = -1 + + def __init__(self, lora_task_id=None, cache_salt_id=None, reuse_scope=None): + # Support both the new ReuseScope-based shape and the legacy flat + # (lora_task_id, cache_salt_id) shape so tests can exercise either. + self.lora_task_id = lora_task_id + self.cache_salt_id = cache_salt_id + if reuse_scope is not None: + self.reuse_scope = reuse_scope + + +class _FakeBlock: + def __init__(self, key, tokens, num_life_cycles=1, prev=None): + self.key = key + self.tokens = tokens + self.prev = prev or _FakeRootBlock() + self.ordinal = getattr(self.prev, "ordinal", -1) + 1 + self.storage = [_FakePageRef(_FakePage()) for _ in range(num_life_cycles)] + + +with temporary_sys_path(os.path.dirname(os.path.abspath(__file__))): + from test_kv_cache_manager_v2 import create_config + + +def _create_test_manager( + event_manager, + *, + tokens_per_block=4, + gpu_quota=16 << 20, + host_quota=0, + window_size=None, + kv_buf_size=8192, +): + return KVCacheManager( + create_config( + tokens_per_block=tokens_per_block, + gpu_quota=gpu_quota, + host_quota=host_quota, + disk_quota=0, + num_layers=2, + window_size=window_size, + sink_tokens=0, + kv_buf_size=kv_buf_size, + ), + event_manager=event_manager, + ) + + +def _token_ids(start, end): + return [TokenId(token_id) for token_id in range(start, end)] + + +def _flush_serialized_events(event_manager): + event_manager.flush_iteration_events() + return KVCacheEventSerializer.serialize(event_manager.get_latest_events(0)) + + +def _stored_events(events): + return [event for event in events if event["data"]["type"] == "stored"] + + +def _stored_block_hashes(events): + return [ + block["block_hash"] for event in _stored_events(events) for block in event["data"]["blocks"] + ] + + +def _stored_block_hashes_by_layer_group(events): + return { + event["layer_group_id"]: [block["block_hash"] for block in event["data"]["blocks"]] + for event in _stored_events(events) + } + + +def _stored_event_payloads(events): + return [event["data"] for event in _stored_events(events)] + + +def _commit_and_close(manager, stream, tokens, *, input_tokens=None, reuse_scope=None): + kv_cache = manager.create_kv_cache(input_tokens=input_tokens, reuse_scope=reuse_scope) + assert kv_cache.resume(stream) + kv_cache.capacity = len(input_tokens or []) + len(tokens) + kv_cache.commit(tokens) + kv_cache.close() + del kv_cache + gc.collect() + + +def test_v2_kv_cache_event_manager_serialization(): + event_manager = KVCacheEventManager(max_kv_event_entries=4, window_size=128) + event_manager.add_created_event([2, 3]) + event_manager.add_stored_event( + parent_hash=None, + blocks=[ + KVCacheStoredBlockData( + block_hash="abcd", + tokens=[UniqueToken(1), UniqueToken(2)], + cache_level=0, + priority=0, + ) + ], + ) + event_manager.add_removed_event("abcd") + + event_manager.flush_iteration_events() + events = KVCacheEventSerializer.serialize(event_manager.get_latest_events()) + + assert [event["event_id"] for event in events] == [0, 1, 2] + assert [event["hash_algo"] for event in events] == ["v2_sha256"] * 3 + assert events[0]["window_size"] == 128 + assert events[0]["data"] == { + "type": "created", + "num_blocks_per_cache_level": [2, 3], + } + assert events[1]["data"]["type"] == "stored" + assert events[1]["data"]["parent_hash"] is None + assert events[1]["data"]["blocks"][0]["block_hash"] == "abcd" + assert events[1]["data"]["blocks"][0]["tokens"][0] == { + "type": "unique_token", + "token_id": 1, + "token_extra_id": 0, + } + assert events[1]["data"]["blocks"][0]["mm_keys"] == [] + assert events[2]["data"] == { + "type": "removed", + "block_hashes": ["abcd"], + } + assert event_manager.get_latest_events(0) == [] + + +def test_v2_kv_cache_event_manager_default_get_latest_events_waits(): + event_manager = KVCacheEventManager(max_kv_event_entries=4, window_size=128) + events = [] + + def get_events(): + events.extend(event_manager.get_latest_events()) + + thread = threading.Thread(target=get_events) + thread.start() + time.sleep(0.05) + + assert thread.is_alive() + + event_manager.add_created_event([1]) + event_manager.flush_iteration_events() + thread.join(timeout=1) + + assert not thread.is_alive() + assert len(events) == 1 + assert events[0].data.num_blocks_per_cache_level == [1] + + +def test_v2_kv_cache_event_manager_uses_layer_group_window_size(): + event_manager = KVCacheEventManager( + max_kv_event_entries=4, + window_size=4096, + window_size_by_layer_group={ + 0: 128, + 1: 4096, + }, + ) + + event_manager.add_created_event([1], layer_group_ids=[0, 1]) + event_manager.flush_iteration_events() + events = event_manager.get_latest_events(0) + + assert [event.layer_group_id for event in events] == [0, 1] + assert [event.window_size for event in events] == [128, 4096] + + +def test_truncate_sha256_hash_to_int64_uses_unsigned_value(): + assert truncate_sha256_hash_to_int64(b"\x80" + b"\x00" * 31) == 1 << 63 + assert truncate_sha256_hash_to_int64(b"\xff" * 32) == (1 << 64) - 1 + + +def test_v2_kv_cache_event_diff_positional_order_matches_v1(): + diff = KVCacheEventDiff(3, 5) + + assert diff.old_value == 3 + assert diff.new_value == 5 + + +def test_v2_kv_cache_event_manager_drops_oldest_events(): + event_manager = KVCacheEventManager(max_kv_event_entries=2, window_size=128) + event_manager.add_created_event([1]) + event_manager.add_stored_event( + parent_hash=None, + blocks=[ + KVCacheStoredBlockData( + block_hash="stored", + tokens=[UniqueToken(1)], + cache_level=0, + priority=0, + ) + ], + ) + event_manager.add_removed_event("a") + event_manager.add_removed_event("b") + + event_manager.flush_iteration_events() + events = KVCacheEventSerializer.serialize(event_manager.get_latest_events()) + + assert [event["event_id"] for event in events] == [1, 2] + assert events[0]["data"]["blocks"][0]["block_hash"] == "stored" + assert events[1]["data"]["block_hashes"] == ["a", "b"] + + +def test_v2_kv_cache_event_manager_accepts_removed_iterables(): + event_manager = KVCacheEventManager(max_kv_event_entries=2, window_size=128) + event_manager.add_removed_event(["a", "b"]) + + event_manager.flush_iteration_events() + events = KVCacheEventSerializer.serialize(event_manager.get_latest_events()) + + assert len(events) == 1 + assert events[0]["data"] == { + "type": "removed", + "block_hashes": ["a", "b"], + } + + +def test_v2_kv_cache_event_manager_coalesces_contiguous_stored_events(): + event_manager = KVCacheEventManager(max_kv_event_entries=8, window_size=128) + block0 = _FakeBlock(b"\xab\xcd", [1, 2], num_life_cycles=2) + block1 = _FakeBlock(b"\xab\xce", [3, 4], num_life_cycles=2, prev=block0) + + event_manager.add_stored_block_event_from_block(block0) + event_manager.add_stored_block_event_from_block(block1) + + events = _flush_serialized_events(event_manager) + + assert [event["data"]["type"] for event in events] == ["stored", "stored"] + assert [event["layer_group_id"] for event in events] == [0, 1] + assert [block["block_hash"] for block in events[0]["data"]["blocks"]] == [ + "abcd", + "abce", + ] + assert [block["block_hash"] for block in events[1]["data"]["blocks"]] == [ + "abcd", + "abce", + ] + assert events[0]["data"]["parent_hash"] is None + assert events[1]["data"]["parent_hash"] is None + + +def test_v2_kv_cache_event_manager_serializes_layer_group_id(): + event_manager = KVCacheEventManager(max_kv_event_entries=4, window_size=128) + event_manager.add_created_event([2, 3], [0, 1]) + + events = _flush_serialized_events(event_manager) + + assert [event["layer_group_id"] for event in events] == [0, 1] + assert [event["data"] for event in events] == [ + { + "type": "created", + "num_blocks_per_cache_level": [2, 3], + }, + { + "type": "created", + "num_blocks_per_cache_level": [2, 3], + }, + ] + + +def test_v2_kv_cache_event_manager_sha256_64_compatibility_mode(): + event_manager = KVCacheEventManager( + max_kv_event_entries=8, + window_size=128, + hash_algo=KV_CACHE_HASH_ALGO_V2_SHA256_64, + ) + block0 = _FakeBlock(bytes.fromhex("8000000000000001" + "00" * 24), [1, 2]) + block1 = _FakeBlock(bytes.fromhex("0102030405060708" + "00" * 24), [3, 4], prev=block0) + + event_manager.add_stored_block_event_from_block(block0) + event_manager.add_stored_block_event_from_block(block1) + event_manager.add_removed_event(block0.key) + events = _flush_serialized_events(event_manager) + + expected_block0_hash = truncate_sha256_hash_to_int64(block0.key) + expected_block1_hash = truncate_sha256_hash_to_int64(block1.key) + assert [event["hash_algo"] for event in events] == [ + KV_CACHE_HASH_ALGO_V2_SHA256_64, + KV_CACHE_HASH_ALGO_V2_SHA256_64, + ] + assert events[0]["data"]["type"] == "stored" + assert events[0]["data"]["parent_hash"] is None + assert [block["block_hash"] for block in events[0]["data"]["blocks"]] == [ + expected_block0_hash, + expected_block1_hash, + ] + assert events[1]["data"] == { + "type": "removed", + "block_hashes": [expected_block0_hash], + } + assert all(isinstance(block_hash, int) for block_hash in events[1]["data"]["block_hashes"]) + + +def test_v2_kv_cache_event_manager_v1_hash_algo_matches_v1_block_key_hash(): + event_manager = KVCacheEventManager( + max_kv_event_entries=8, + window_size=128, + hash_algo=KV_CACHE_HASH_ALGO_V1, + ) + root = _FakeRootBlock() + block0 = _FakeBlock(b"block0", [1, 2, 3, 4], prev=root) + block1 = _FakeBlock(b"block1", [5, 6], prev=block0) + + event_manager.add_stored_block_event_from_block(block0) + event_manager.add_stored_block_event_from_block(block1) + events = _flush_serialized_events(event_manager) + + assert events[0]["hash_algo"] == KV_CACHE_HASH_ALGO_V1 + assert events[0]["data"]["parent_hash"] is None + assert [block["block_hash"] for block in events[0]["data"]["blocks"]] == [ + 924206229973855, + 6875034662206558884, + ] + assert all(isinstance(block["block_hash"], int) for block in events[0]["data"]["blocks"]) + + +def test_v2_root_key_distinguishes_lora_from_cache_salt_id(): + assert RootBlock.make_key(ReuseScope()) != RootBlock.make_key(ReuseScope(lora_id=123)) + assert RootBlock.make_key(ReuseScope()) != RootBlock.make_key(ReuseScope(salt=123)) + assert RootBlock.make_key(ReuseScope(lora_id=123)) != RootBlock.make_key(ReuseScope(salt=123)) + assert RootBlock.make_key(ReuseScope(lora_id=123, salt=456)) != RootBlock.make_key( + ReuseScope(lora_id=456, salt=123) + ) + + +def test_v2_kv_cache_event_manager_v1_hash_algo_mixes_cache_salt_id(): + event_manager = KVCacheEventManager( + max_kv_event_entries=8, + window_size=128, + hash_algo=KV_CACHE_HASH_ALGO_V1, + ) + root = _FakeRootBlock(cache_salt_id=123) + block0 = _FakeBlock(b"block0", [1, 2, 3, 4], prev=root) + block1 = _FakeBlock(b"block1", [5, 6], prev=block0) + + event_manager.add_stored_block_event_from_block(block0) + event_manager.add_stored_block_event_from_block(block1) + events = _flush_serialized_events(event_manager) + + assert events[0]["hash_algo"] == KV_CACHE_HASH_ALGO_V1 + assert [block["block_hash"] for block in events[0]["data"]["blocks"]] == [ + 6280297290684427985, + 18177682803760873588, + ] + + +def test_v2_kv_cache_event_manager_v1_hash_reads_root_reuse_scope(): + # Regression test: when ``RootBlock`` exposes its scope via a ReuseScope + # NamedTuple rather than direct ``lora_task_id`` / ``cache_salt_id`` + # attributes, ``_root_attrs_from_root_block`` must still recover the same + # (lora_id, salt) — otherwise V1-compat event hashes silently collapse to + # (None, None) for every LoRA/salt request and Dynamo routing degrades. + tokens0 = [1, 2, 3, 4] + tokens1 = [5, 6] + + def hashes_for(root): + event_manager = KVCacheEventManager( + max_kv_event_entries=8, + window_size=128, + hash_algo=KV_CACHE_HASH_ALGO_V1, + ) + block0 = _FakeBlock(b"block0", tokens0, prev=root) + block1 = _FakeBlock(b"block1", tokens1, prev=block0) + event_manager.add_stored_block_event_from_block(block0) + event_manager.add_stored_block_event_from_block(block1) + return _stored_block_hashes(_flush_serialized_events(event_manager)) + + # ReuseScope-shaped RootBlock and legacy-shape RootBlock must produce + # identical event hashes for the same scope. + scope_root = _FakeRootBlock(reuse_scope=ReuseScope(lora_id=11, salt=22)) + legacy_root = _FakeRootBlock(lora_task_id=11, cache_salt_id=22) + assert hashes_for(scope_root) == hashes_for(legacy_root) + + # Different scopes must still produce different hashes (no silent collapse). + other_scope_root = _FakeRootBlock(reuse_scope=ReuseScope(lora_id=99, salt=22)) + assert hashes_for(scope_root) != hashes_for(other_scope_root) + + # An unsalted ReuseScope must match an unsalted legacy root. + empty_scope_root = _FakeRootBlock(reuse_scope=ReuseScope()) + unsalted_legacy_root = _FakeRootBlock() + assert hashes_for(empty_scope_root) == hashes_for(unsalted_legacy_root) + + +def test_v2_kv_cache_event_manager_v1_hash_recomputes_removed_parent(): + event_manager = KVCacheEventManager( + max_kv_event_entries=8, + window_size=128, + hash_algo=KV_CACHE_HASH_ALGO_V1, + ) + root = _FakeRootBlock() + block0 = _FakeBlock(b"block0", [1, 2, 3, 4], prev=root) + block1 = _FakeBlock(b"block1", [5, 6], prev=block0) + + event_manager.add_stored_block_event_from_block(block0) + event_manager.add_removed_event(block0.key) + + assert event_manager._v1_hash_from_radix_block(block1) == 6875034662206558884 + + event_manager.add_stored_block_event_from_block(block1) + events = _flush_serialized_events(event_manager) + stored_events = _stored_events(events) + + assert stored_events[-1]["data"]["parent_hash"] == 924206229973855 + assert stored_events[-1]["data"]["blocks"][0]["block_hash"] == 6875034662206558884 + + +def test_v2_kv_cache_event_manager_v1_hash_algo_matches_cpp_hasher(): + _tb = pytest.importorskip("tensorrt_llm.bindings") + block_key = _tb.internal.batch_manager.BlockKey + block_key_hasher = _tb.internal.batch_manager.BlockKeyHasher + + parent_hash = block_key_hasher.hash(block_key([1, 2, 3, 4])) + child_hash = block_key_hasher.hash(block_key([5, 6]), parent_hash) + lora_hash = block_key_hasher.hash(block_key([1, 2, 3, 4], 123)) + + assert KVCacheEventManager._hash_block_key([1, 2, 3, 4], 0, None, None) == parent_hash + assert KVCacheEventManager._hash_block_key([5, 6], parent_hash, None, None) == child_hash + assert KVCacheEventManager._hash_block_key([1, 2, 3, 4], 0, 123, None) == lora_hash + + +def test_v2_kv_cache_event_manager_v1_hash_events_match_cpp_hasher(): + _tb = pytest.importorskip("tensorrt_llm.bindings") + block_key = _tb.internal.batch_manager.BlockKey + block_key_hasher = _tb.internal.batch_manager.BlockKeyHasher + event_manager = KVCacheEventManager( + max_kv_event_entries=8, + window_size=128, + hash_algo=KV_CACHE_HASH_ALGO_V1, + ) + root = _FakeRootBlock() + block0 = _FakeBlock(b"block0", [1, 2, 3, 4], prev=root) + block1 = _FakeBlock(b"block1", [5, 6], prev=block0) + + event_manager.add_stored_block_event_from_block(block0) + event_manager.add_stored_block_event_from_block(block1) + events = _flush_serialized_events(event_manager) + + parent_hash = block_key_hasher.hash(block_key([1, 2, 3, 4])) + child_hash = block_key_hasher.hash(block_key([5, 6]), parent_hash) + assert _stored_block_hashes(events) == [parent_hash, child_hash] + + +@pytest.mark.skipif(torch is None or not torch.cuda.is_available(), reason="requires CUDA") +def test_v1_and_v2_managers_emit_same_v1_hash_stored_events(): + import tensorrt_llm + from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest + from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager as KVCacheManagerV1 + from tensorrt_llm.bindings.internal.testing import ( + simulate_prefill_completion_only_use_for_testing, + ) + from tensorrt_llm.llmapi.llm_args import KvCacheConfig + from tensorrt_llm.mapping import Mapping + from tensorrt_llm.sampling_params import SamplingParams + + init_cuda_once() + gc.collect() + gc.disable() + + tokens_per_block = 4 + max_seq_len = 128 + prompt_tokens = list(range(1, 2 * tokens_per_block + 2)) + reusable_tokens = _token_ids(prompt_tokens[0], prompt_tokens[-1]) + event_buffer_max_size = 16 + manager_v1 = None + manager_v2 = None + try: + kv_cache_config_v1 = KvCacheConfig( + event_buffer_max_size=event_buffer_max_size, + enable_block_reuse=True, + max_tokens=256, + use_kv_cache_manager_v2=False, + ) + manager_v1 = KVCacheManagerV1( + kv_cache_config=kv_cache_config_v1, + kv_cache_type=tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF, + num_layers=2, + num_kv_heads=2, + head_dim=128, + tokens_per_block=tokens_per_block, + max_seq_len=max_seq_len, + max_batch_size=1, + mapping=Mapping(), + ) + manager_v1.flush_iteration_events() + manager_v1.get_latest_events(10) + + sampling_params = SamplingParams() + req = LlmRequest( + request_id=0, + max_new_tokens=1, + input_tokens=prompt_tokens, + sampling_config=tensorrt_llm.bindings.SamplingConfig( + sampling_params._get_sampling_config() + ), + is_streaming=False, + ) + manager_v1.impl.add_sequence_batch([(req.py_request_id, req.prompt_len, 1)], [req]) + simulate_prefill_completion_only_use_for_testing(req) + manager_v1.free_resources(req) + manager_v1.flush_iteration_events() + v1_events = KVCacheEventSerializer.serialize(manager_v1.get_latest_events(10)) + + event_manager_v2 = KVCacheEventManager( + max_kv_event_entries=event_buffer_max_size, + window_size=max_seq_len, + hash_algo=KV_CACHE_HASH_ALGO_V1, + ) + manager_v2 = _create_test_manager( + event_manager_v2, + tokens_per_block=tokens_per_block, + ) + stream_holder = CachedCudaStream() + stream = cast(CudaStream, stream_holder.handle) + _commit_and_close(manager_v2, stream, reusable_tokens) + v2_events = _flush_serialized_events(event_manager_v2) + + assert all( + event["hash_algo"] == KV_CACHE_HASH_ALGO_V1 for event in _stored_events(v2_events) + ) + assert _stored_event_payloads(v2_events) == _stored_event_payloads(v1_events) + finally: + gc.enable() + if manager_v1 is not None: + manager_v1.shutdown() + if manager_v2 is not None: + manager_v2.shutdown() + + +def test_v2_kv_cache_event_manager_unknown_hash_algo_raises(): + with pytest.raises(ValueError, match="Unsupported V2 KV cache event hash algorithm"): + KVCacheEventManager( + max_kv_event_entries=4, + window_size=128, + hash_algo="unknown_hash_algo", + ) + + +def test_v2_kv_cache_event_manager_gathers_attention_dp_events_on_rank_zero(): + remote_manager = KVCacheEventManager( + max_kv_event_entries=4, + window_size=128, + attention_dp_rank=1, + ) + remote_manager.add_created_event([3]) + remote_manager.flush_iteration_events() + remote_event_objects = remote_manager.get_latest_events() + remote_events = KVCacheEventSerializer.serialize(remote_event_objects) + assert remote_events[0]["attention_dp_rank"] == 1 + + gathered_local_events = [] + + def gather(local_events): + gathered_local_events.extend(local_events) + return [ + local_events, + remote_event_objects, + ] + + event_manager = KVCacheEventManager( + max_kv_event_entries=4, + window_size=128, + attention_dp_rank=0, + attention_dp_gather=gather, + ) + event_manager.add_created_event([2]) + + events = _flush_serialized_events(event_manager) + + assert len(gathered_local_events) == 1 + assert [event["attention_dp_rank"] for event in events] == [0, 1] + assert [event["data"]["num_blocks_per_cache_level"] for event in events] == [ + [2], + [3], + ] + + +def test_v2_kv_cache_event_manager_attention_dp_nonzero_rank_sends_without_publishing(): + gathered_local_events = [] + + def gather(local_events): + gathered_local_events.extend(local_events) + return [ + [], + local_events, + ] + + event_manager = KVCacheEventManager( + max_kv_event_entries=1, + window_size=128, + attention_dp_rank=1, + attention_dp_gather=gather, + ) + event_manager.add_created_event([2]) + event_manager.add_created_event([3]) + + events = _flush_serialized_events(event_manager) + + assert len(gathered_local_events) == 1 + assert gathered_local_events[0].attention_dp_rank == 1 + assert KVCacheEventSerializer.serialize(gathered_local_events)[0]["data"][ + "num_blocks_per_cache_level" + ] == [3] + assert events == [] + + +def test_v2_kv_cache_event_manager_attention_dp_rank_zero_uses_dp_capacity(): + remote_manager = KVCacheEventManager( + max_kv_event_entries=2, + window_size=128, + attention_dp_rank=1, + ) + remote_manager.add_created_event([10]) + remote_manager.add_created_event([11]) + remote_manager.flush_iteration_events() + remote_event_objects = remote_manager.get_latest_events() + + def gather(local_events): + return [ + local_events, + remote_event_objects, + ] + + event_manager = KVCacheEventManager( + max_kv_event_entries=1, + window_size=128, + attention_dp_rank=0, + attention_dp_gather=gather, + ) + event_manager.add_created_event([1]) + event_manager.add_created_event([2]) + + events = _flush_serialized_events(event_manager) + + assert [event["attention_dp_rank"] for event in events] == [0, 1] + assert [event["data"]["num_blocks_per_cache_level"] for event in events] == [ + [2], + [11], + ] + + +def test_v2_kv_cache_event_manager_serializes_updated_event(): + event_manager = KVCacheEventManager(max_kv_event_entries=2, window_size=128) + event_manager.add_updated_event( + "abcd", + cache_level=KVCacheEventDiff(old_value=0, new_value=1), + priority=KVCacheEventDiff(old_value=1, new_value=2), + ) + + event_manager.flush_iteration_events() + events = KVCacheEventSerializer.serialize(event_manager.get_latest_events()) + + assert len(events) == 1 + assert events[0]["hash_algo"] == "v2_sha256" + assert events[0]["data"] == { + "type": "updated", + "block_hash": "abcd", + "cache_level": { + "type": "event_diff", + "new_value": 1, + "old_value": 0, + }, + "priority": { + "type": "event_diff", + "new_value": 2, + "old_value": 1, + }, + } + + +def test_v2_kv_cache_event_manager_uses_stored_registry_for_removed_event(): + event_manager = KVCacheEventManager(max_kv_event_entries=8, window_size=128) + block = _FakeBlock(b"\xab\xcd", [1, 2]) + + event_manager.add_stored_block_event_from_block(block) + block.storage = [] + event_manager.add_removed_event(block.key) + event_manager.add_removed_event(block.key) + + events = _flush_serialized_events(event_manager) + + assert [event["data"]["type"] for event in events] == ["stored", "removed"] + assert [event["layer_group_id"] for event in events] == [0, 0] + assert events[0]["data"]["blocks"][0]["block_hash"] == "abcd" + assert "layer_groups" not in events[0]["data"]["blocks"][0] + assert events[1]["data"]["block_hashes"] == ["abcd"] + assert "layer_groups" not in events[1]["data"] + + +def test_v2_kv_cache_event_manager_emits_partial_life_cycle_removed_events(): + event_manager = KVCacheEventManager(max_kv_event_entries=8, window_size=128) + block = _FakeBlock(b"\xab\xcd", [1, 2], num_life_cycles=2) + + event_manager.add_stored_block_event_from_block(block) + event_manager.add_removed_life_cycle_event(block.key, 0) + + events = _flush_serialized_events(event_manager) + + assert [event["data"]["type"] for event in events] == [ + "stored", + "stored", + "removed", + ] + assert [event["layer_group_id"] for event in events] == [0, 1, 0] + assert events[0]["data"]["blocks"][0]["block_hash"] == "abcd" + assert events[1]["data"]["blocks"][0]["block_hash"] == "abcd" + assert "layer_groups" not in events[0]["data"]["blocks"][0] + assert "layer_groups" not in events[1]["data"]["blocks"][0] + assert events[2]["data"]["block_hashes"] == ["abcd"] + assert "layer_groups" not in events[2]["data"] + + event_manager.add_removed_life_cycle_event(block.key, 1) + event_manager.add_removed_life_cycle_event(block.key, 1) + + events = _flush_serialized_events(event_manager) + + assert [event["data"]["type"] for event in events] == ["removed"] + assert events[0]["layer_group_id"] == 1 + assert events[0]["data"]["block_hashes"] == ["abcd"] + assert "layer_groups" not in events[0]["data"] + + +def test_v2_kv_cache_event_manager_whole_block_removal_clears_life_cycle_state(): + event_manager = KVCacheEventManager(max_kv_event_entries=8, window_size=128) + block = _FakeBlock(b"\xab\xce", [1, 2], num_life_cycles=2) + + event_manager.add_stored_block_event_from_block(block) + event_manager.add_removed_life_cycle_event(block.key, 0) + event_manager.add_removed_event(block.key) + event_manager.add_removed_event(block.key) + + events = _flush_serialized_events(event_manager) + + assert [event["data"]["type"] for event in events] == [ + "stored", + "stored", + "removed", + "removed", + ] + assert [event["layer_group_id"] for event in events] == [0, 1, 0, 1] + assert events[2]["data"]["block_hashes"] == ["abce"] + assert events[3]["data"]["block_hashes"] == ["abce"] + + +def test_v2_kv_cache_event_manager_readds_life_cycle_emits_stored_event(): + event_manager = KVCacheEventManager(max_kv_event_entries=8, window_size=128) + block = _FakeBlock(b"\xab\xcf", [1, 2], num_life_cycles=2) + + event_manager.add_stored_block_event_from_block(block) + event_types = [event["data"]["type"] for event in _flush_serialized_events(event_manager)] + assert event_types == ["stored", "stored"] + + event_manager.add_removed_life_cycle_event(block.key, 0) + event_manager.add_stored_life_cycle_event_from_block(block, 0) + event_manager.add_removed_life_cycle_event(block.key, 1) + events = _flush_serialized_events(event_manager) + + assert [event["data"]["type"] for event in events] == [ + "removed", + "stored", + "removed", + ] + assert [event["layer_group_id"] for event in events] == [0, 0, 1] + assert events[0]["data"]["block_hashes"] == ["abcf"] + assert events[1]["data"]["blocks"][0]["block_hash"] == "abcf" + assert "layer_groups" not in events[1]["data"]["blocks"][0] + assert events[2]["data"]["block_hashes"] == ["abcf"] + + event_manager.add_removed_life_cycle_event(block.key, 0) + events = _flush_serialized_events(event_manager) + + assert [event["data"]["type"] for event in events] == ["removed"] + assert events[0]["layer_group_id"] == 0 + assert events[0]["data"]["block_hashes"] == ["abcf"] + assert "layer_groups" not in events[0]["data"] + + +def test_v2_kv_cache_event_manager_reemits_stored_after_all_life_cycles_were_removed(): + event_manager = KVCacheEventManager(max_kv_event_entries=8, window_size=128) + block = _FakeBlock(b"\xab\xd0", [1, 2]) + + event_manager.add_stored_block_event_from_block(block) + event_manager.add_removed_life_cycle_event(block.key, 0) + event_types = [event["data"]["type"] for event in _flush_serialized_events(event_manager)] + assert event_types == ["stored", "removed"] + + event_manager.add_stored_life_cycle_event_from_block(block, 0) + events = _flush_serialized_events(event_manager) + + assert [event["data"]["type"] for event in events] == ["stored"] + assert events[0]["layer_group_id"] == 0 + assert events[0]["data"]["blocks"][0]["block_hash"] == "abd0" + assert "layer_groups" not in events[0]["data"]["blocks"][0] + + +def test_v2_kv_cache_event_manager_flushes_removed_before_updated_event(): + event_manager = KVCacheEventManager(max_kv_event_entries=8, window_size=128) + removed_block = _FakeBlock(b"\xab\xd1", [1, 2]) + updated_block = _FakeBlock(b"\xab\xd2", [3, 4]) + + event_manager.add_stored_block_event_from_block(removed_block) + event_manager.add_stored_block_event_from_block(updated_block) + _flush_serialized_events(event_manager) + + event_manager.add_removed_event(removed_block.key) + event_manager.add_updated_event( + updated_block.key, + cache_level=KVCacheEventDiff(old_value=0, new_value=1), + layer_group_id=0, + ) + events = _flush_serialized_events(event_manager) + + assert [event["data"]["type"] for event in events] == ["removed", "updated"] + assert events[0]["data"]["block_hashes"] == ["abd1"] + assert events[1]["data"]["block_hash"] == "abd2" + + +@pytest.mark.skipif(torch is None or not torch.cuda.is_available(), reason="requires CUDA") +def test_v2_stored_events_match_block_hash_chain(): + init_cuda_once() + gc.collect() + gc.disable() + + event_manager = KVCacheEventManager(max_kv_event_entries=16, window_size=128) + manager = None + try: + tokens_per_block = 4 + manager = _create_test_manager(event_manager, tokens_per_block=tokens_per_block) + stream_holder = CachedCudaStream() + stream = cast(CudaStream, stream_holder.handle) + tokens = _token_ids(0, 2 * tokens_per_block) + + _commit_and_close(manager, stream, tokens) + events = _flush_serialized_events(event_manager) + + stored_events = _stored_events(events) + assert len(stored_events) == 1 + + root_key = RootBlock.make_key(ReuseScope()) + block0_key = Block.make_key(root_key, tokens[:tokens_per_block]) + block1_key = Block.make_key(block0_key, tokens[tokens_per_block:]) + expected_hashes = [block0_key.hex(), block1_key.hex()] + + assert _stored_block_hashes(stored_events) == expected_hashes + assert stored_events[0]["data"]["parent_hash"] is None + assert [ + block["block_hash"] for block in stored_events[0]["data"]["blocks"] + ] == expected_hashes + assert [ + token["token_id"] for token in stored_events[0]["data"]["blocks"][0]["tokens"] + ] == list(range(tokens_per_block)) + assert [ + token["token_id"] for token in stored_events[0]["data"]["blocks"][1]["tokens"] + ] == list(range(tokens_per_block, 2 * tokens_per_block)) + finally: + gc.enable() + if manager is not None: + manager.shutdown() + + +@pytest.mark.skipif(torch is None or not torch.cuda.is_available(), reason="requires CUDA") +def test_v2_v1_hash_events_include_cache_salt_from_kv_cache(): + init_cuda_once() + gc.collect() + gc.disable() + + event_manager = KVCacheEventManager( + max_kv_event_entries=16, + window_size=128, + hash_algo=KV_CACHE_HASH_ALGO_V1, + ) + manager = None + try: + tokens_per_block = 4 + manager = _create_test_manager(event_manager, tokens_per_block=tokens_per_block) + stream_holder = CachedCudaStream() + stream = cast(CudaStream, stream_holder.handle) + + _commit_and_close( + manager, + stream, + _token_ids(1, 7), + reuse_scope=ReuseScope(salt=123), + ) + events = _flush_serialized_events(event_manager) + + assert _stored_block_hashes(events) == [ + 6280297290684427985, + 18177682803760873588, + ] + finally: + gc.enable() + if manager is not None: + manager.shutdown() + + +@pytest.mark.skipif(torch is None or not torch.cuda.is_available(), reason="requires CUDA") +def test_v2_reused_prefix_does_not_emit_duplicate_stored_events(): + init_cuda_once() + gc.collect() + gc.disable() + + event_manager = KVCacheEventManager(max_kv_event_entries=16, window_size=128) + manager = None + try: + tokens_per_block = 4 + manager = _create_test_manager(event_manager, tokens_per_block=tokens_per_block) + stream_holder = CachedCudaStream() + stream = cast(CudaStream, stream_holder.handle) + prefix_tokens = _token_ids(0, 2 * tokens_per_block) + new_tokens = _token_ids(2 * tokens_per_block, 3 * tokens_per_block) + + _commit_and_close(manager, stream, prefix_tokens) + first_events = _flush_serialized_events(event_manager) + prefix_hashes = _stored_block_hashes(first_events) + assert len(prefix_hashes) == 2 + + _commit_and_close( + manager, + stream, + new_tokens, + input_tokens=prefix_tokens, + ) + reuse_events = _flush_serialized_events(event_manager) + reused_hashes = _stored_block_hashes(reuse_events) + + root_key = RootBlock.make_key(ReuseScope()) + block0_key = Block.make_key(root_key, prefix_tokens[:tokens_per_block]) + block1_key = Block.make_key(block0_key, prefix_tokens[tokens_per_block:]) + block2_key = Block.make_key(block1_key, new_tokens) + + assert prefix_hashes == [block0_key.hex(), block1_key.hex()] + assert reused_hashes == [block2_key.hex()] + assert not (set(prefix_hashes) & set(reused_hashes)) + finally: + gc.enable() + if manager is not None: + manager.shutdown() + + +@pytest.mark.skipif(torch is None or not torch.cuda.is_available(), reason="requires CUDA") +def test_v2_removed_events_match_stored_hashes(): + init_cuda_once() + gc.collect() + gc.disable() + + event_manager = KVCacheEventManager(max_kv_event_entries=16, window_size=128) + manager = None + try: + tokens_per_block = 4 + manager = _create_test_manager(event_manager, tokens_per_block=tokens_per_block) + stream_holder = CachedCudaStream() + stream = cast(CudaStream, stream_holder.handle) + tokens = _token_ids(0, 2 * tokens_per_block) + + _commit_and_close(manager, stream, tokens) + stored_events = _flush_serialized_events(event_manager) + stored_hashes = _stored_block_hashes(stored_events) + assert len(stored_hashes) == 2 + + manager.clear_reusable_blocks() + removal_events = _flush_serialized_events(event_manager) + removed_hashes = [ + block_hash + for event in removal_events + if event["data"]["type"] == "removed" + for block_hash in event["data"]["block_hashes"] + ] + + assert set(removed_hashes) == set(stored_hashes) + assert len(removed_hashes) == len(stored_hashes) + finally: + gc.enable() + if manager is not None: + manager.shutdown() + + +@pytest.mark.skipif(torch is None or not torch.cuda.is_available(), reason="requires CUDA") +def test_v2_removed_event_emitted_when_last_level_page_is_dropped(): + init_cuda_once() + gc.collect() + gc.disable() + + event_manager = KVCacheEventManager(max_kv_event_entries=16, window_size=4096) + manager = None + try: + tokens_per_block = 8 + manager = _create_test_manager( + event_manager, + tokens_per_block=tokens_per_block, + gpu_quota=8 << 20, + window_size=4096, + kv_buf_size=1 << 20, + ) + stream_holder = CachedCudaStream() + stream = cast(CudaStream, stream_holder.handle) + kv_cache = manager.create_kv_cache() + assert kv_cache.resume(stream) + kv_cache.capacity = tokens_per_block * 2 + kv_cache.commit(_token_ids(0, tokens_per_block * 2)) + + stored_events = _flush_serialized_events(event_manager) + stored_hashes_by_layer_group = _stored_block_hashes_by_layer_group(stored_events) + assert set(stored_hashes_by_layer_group) == {0, 1} + assert stored_hashes_by_layer_group[0] == stored_hashes_by_layer_group[1] + + kv_cache.close() + del kv_cache + gc.collect() + + assert manager.resize(CacheLevel(0), 4 << 20) + removal_events = _flush_serialized_events(event_manager) + removed_hashes_by_layer_group = { + event["layer_group_id"]: event["data"]["block_hashes"] + for event in removal_events + if event["data"]["type"] == "removed" + } + + assert set(removed_hashes_by_layer_group) == {0, 1} + assert removed_hashes_by_layer_group[0] == removed_hashes_by_layer_group[1] + for layer_group_id, removed_hashes in removed_hashes_by_layer_group.items(): + assert len(removed_hashes) == 1 + assert removed_hashes[0] in stored_hashes_by_layer_group[layer_group_id] + finally: + gc.enable() + if manager is not None: + manager.shutdown() + + +@pytest.mark.skipif(torch is None or not torch.cuda.is_available(), reason="requires CUDA") +def test_v2_kv_cache_event_manager_emits_updated_on_level_migration(): + init_cuda_once() + gc.collect() + gc.disable() + + event_manager = KVCacheEventManager(max_kv_event_entries=16, window_size=4096) + manager = None + try: + tokens_per_block = 8 + manager = _create_test_manager( + event_manager, + tokens_per_block=tokens_per_block, + gpu_quota=8 << 20, + host_quota=8 << 20, + window_size=4096, + kv_buf_size=1 << 20, + ) + stream_holder = CachedCudaStream() + stream = cast(CudaStream, stream_holder.handle) + kv_cache = manager.create_kv_cache() + assert kv_cache.resume(stream) + kv_cache.capacity = tokens_per_block * 2 + kv_cache.commit([TokenId(token_id) for token_id in range(tokens_per_block * 2)]) + + stored_events = _flush_serialized_events(event_manager) + stored_hashes_by_layer_group = _stored_block_hashes_by_layer_group(stored_events) + assert set(stored_hashes_by_layer_group) == {0, 1} + assert stored_hashes_by_layer_group[0] == stored_hashes_by_layer_group[1] + + kv_cache.close() + del kv_cache + gc.collect() + + assert manager.resize(CacheLevel(0), 4 << 20) + + event_manager.flush_iteration_events() + events = KVCacheEventSerializer.serialize(event_manager.get_latest_events()) + updated_events = [event for event in events if event["data"]["type"] == "updated"] + + assert len(updated_events) == 2 + assert {event["layer_group_id"] for event in updated_events} == {0, 1} + assert len({event["data"]["block_hash"] for event in updated_events}) == 1 + assert updated_events[0]["data"]["block_hash"] in stored_hashes_by_layer_group[0] + for event in updated_events: + assert event["data"]["cache_level"] == { + "type": "event_diff", + "old_value": 0, + "new_value": 1, + } + assert event["data"]["priority"] is None + finally: + gc.enable() + if manager is not None: + manager.shutdown() diff --git a/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_stats_behavior.py b/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_stats_behavior.py new file mode 100644 index 000000000000..aa20cb4f56e7 --- /dev/null +++ b/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_stats_behavior.py @@ -0,0 +1,647 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from dataclasses import dataclass, field + +import pytest +import torch + +from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import KVCacheManagerV2 +from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager as KVCacheManagerV1 +from tensorrt_llm._torch.pyexecutor.scheduler import ScheduledRequests +from tensorrt_llm.bindings import DataType, SamplingConfig +from tensorrt_llm.bindings.internal.batch_manager import CacheType +from tensorrt_llm.bindings.internal.testing import simulate_prefill_completion_only_use_for_testing +from tensorrt_llm.llmapi.llm_args import KvCacheConfig +from tensorrt_llm.mapping import Mapping +from tensorrt_llm.runtime.kv_cache_manager_v2 import DEFAULT_BEAM_INDEX +from tensorrt_llm.sampling_params import SamplingParams + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") + +TOKENS_PER_BLOCK = 4 +BYTES_PER_BLOCK = 2 << 20 + + +@dataclass +class _StatsRequest: + request_id: int + tokens: list[int] + context_remaining_length: int + py_request_id: int = field(init=False) + lora_task_id: int | None = None + cache_salt: str | None = None + cache_salt_id: int | None = None + is_first_context_chunk: bool = True + is_last_context_chunk: bool = True + is_encoder_init_state: bool = False + is_dummy_request: bool = False + is_attention_dp_dummy: bool = False + is_cuda_graph_dummy: bool = False + is_disagg_generation_transmission_complete: bool = False + context_phase_params: None = None + py_draft_tokens: list[int] = field(default_factory=list) + draft_tokens: list[int] = field(default_factory=list) + context_current_position: int = 0 + context_chunk_size: int = 0 + prepopulated_prompt: tuple[int, int] | None = None + kv_cache_perf_metric_calls: list[dict[str, int]] = field(default_factory=list) + multimodal_hashes: None = None + multimodal_positions: None = None + multimodal_lengths: None = None + + def __post_init__(self) -> None: + self.py_request_id = self.request_id + self.context_chunk_size = self.context_remaining_length + + @property + def prompt_len(self) -> int: + return len(self.tokens) + + @property + def is_dummy(self) -> bool: + return self.is_attention_dp_dummy or self.is_cuda_graph_dummy or self.is_dummy_request + + 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) + + @property + def prepopulated_prompt_len(self) -> int: + if self.prepopulated_prompt is None: + return 0 + return self.prepopulated_prompt[0] + + def update_kv_cache_perf_metrics( + self, + alloc_total_blocks: int, + alloc_new_blocks: int, + reused_blocks: int, + missed_blocks: int, + ) -> None: + self.kv_cache_perf_metric_calls.append( + { + "alloc_total_blocks": alloc_total_blocks, + "alloc_new_blocks": alloc_new_blocks, + "reused_blocks": reused_blocks, + "missed_blocks": missed_blocks, + } + ) + + +def _create_manager( + *, + gpu_bytes: int, + num_layers: int = 1, + max_attention_window: list[int] | None = None, + enable_block_reuse: bool = True, + enable_stats: bool = True, +) -> KVCacheManagerV2: + return KVCacheManagerV2( + KvCacheConfig( + enable_block_reuse=enable_block_reuse, + enable_partial_reuse=True, + max_gpu_total_bytes=gpu_bytes, + max_util_for_resume=1.0, + max_attention_window=max_attention_window, + ), + CacheType.SELF, + num_layers=num_layers, + num_kv_heads=128, + head_dim=1024, + tokens_per_block=TOKENS_PER_BLOCK, + max_seq_len=16, + 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=enable_stats, + ) + + +def _create_v1_manager( + *, + gpu_bytes: int, + enable_block_reuse: bool = True, +) -> KVCacheManagerV1: + max_gpu_blocks = gpu_bytes // BYTES_PER_BLOCK + return KVCacheManagerV1( + KvCacheConfig( + enable_block_reuse=enable_block_reuse, + enable_partial_reuse=True, + max_tokens=max_gpu_blocks * TOKENS_PER_BLOCK, + ), + CacheType.SELF, + num_layers=1, + num_kv_heads=128, + head_dim=1024, + tokens_per_block=TOKENS_PER_BLOCK, + max_seq_len=16, + max_batch_size=2, + mapping=Mapping(world_size=1, rank=0, tp_size=1, pp_size=1), + dtype=DataType.HALF, + ) + + +def _create_llm_request( + request_id: int, + tokens: list[int], +) -> LlmRequest: + sampling_params = SamplingParams() + return LlmRequest( + request_id=request_id, + max_new_tokens=1, + input_tokens=tokens, + sampling_config=SamplingConfig(sampling_params._get_sampling_config()), + is_streaming=False, + ) + + +def _context_batch(*requests) -> ScheduledRequests: + batch = ScheduledRequests() + for request in requests: + batch.append_context_request(request) + return batch + + +def _generation_batch(request: _StatsRequest) -> ScheduledRequests: + batch = ScheduledRequests() + batch.append_generation_request(request) + return batch + + +@pytest.fixture +def resource_guard(): + managers = [] + resources = [] + + def register(manager, *requests): + if manager not in managers: + managers.append(manager) + resources.extend((manager, request) for request in requests) + return manager + + yield register + + for manager, request in reversed(resources): + manager.free_resources(request) + for manager in reversed(managers): + manager.shutdown() + + +def _finish_context(manager: KVCacheManagerV2, request: _StatsRequest) -> None: + request.context_current_position = request.prompt_len + request.context_remaining_length = 0 + manager.update_context_resources(_context_batch(request)) + + +def _commit_and_get_stats(manager: KVCacheManagerV2, batch: ScheduledRequests): + manager.commit_scheduled_kv_cache_stats(batch) + stats_report = manager.get_iteration_stats() + assert stats_report is not None + assert manager.max_seq_len in stats_report.by_window_size + return stats_report.by_window_size[manager.max_seq_len] + + +def _assert_iteration_delta( + stats, + *, + alloc_total: int = 0, + alloc_new: int = 0, + reused: int = 0, + full_reused: int = 0, + partial_reused: int = 0, + missed: int = 0, + gen_alloc: int = 0, + intra_copy: int = 0, + intra_copy_bytes: int = 0, +) -> None: + assert stats.iter_alloc_total_blocks == alloc_total + assert stats.iter_alloc_new_blocks == alloc_new + assert stats.iter_reused_blocks == reused + assert stats.iter_full_reused_blocks == full_reused + assert stats.iter_partial_reused_blocks == partial_reused + assert stats.iter_missed_blocks == missed + assert stats.iter_gen_alloc_blocks == gen_alloc + assert stats.iter_intra_device_copy_blocks == intra_copy + assert stats.iter_intra_device_copy_bytes == intra_copy_bytes + + +def _metric_call( + *, + alloc_total: int = 0, + alloc_new: int = 0, + reused: int = 0, + missed: int = 0, +) -> dict[str, int]: + return { + "alloc_total_blocks": alloc_total, + "alloc_new_blocks": alloc_new, + "reused_blocks": reused, + "missed_blocks": missed, + } + + +def _assert_request_stats( + request: LlmRequest, + *, + alloc_total: int = 0, + alloc_new: int = 0, + reused: int = 0, + missed: int = 0, +) -> None: + assert request.alloc_total_blocks == alloc_total + assert request.alloc_new_blocks == alloc_new + assert request.reused_blocks == reused + assert request.missed_blocks == missed + + +def _run_v1_context(manager: KVCacheManagerV1, request: LlmRequest): + batch = _context_batch(request) + manager.prepare_resources(batch) + stats = manager.get_iteration_stats()[manager.max_seq_len] + simulate_prefill_completion_only_use_for_testing(request) + manager.update_resources(batch) + return stats + + +def _run_v1_generation(manager: KVCacheManagerV1, request: LlmRequest): + batch = _generation_batch(request) + manager.prepare_resources(batch) + return manager.get_iteration_stats()[manager.max_seq_len] + + +def _run_v2_context(manager: KVCacheManagerV2, request: LlmRequest): + assert manager.prepare_context(request) + assert manager.resize_context(request, num_tokens=request.context_remaining_length) + simulate_prefill_completion_only_use_for_testing(request) + manager.update_context_resources(_context_batch(request)) + return _commit_and_get_stats(manager, _context_batch(request)) + + +def _run_v2_generation(manager: KVCacheManagerV2, request: LlmRequest): + assert manager.try_allocate_generation(request) + return _commit_and_get_stats(manager, _generation_batch(request)) + + +def test_stats_disabled_suppresses_v2_accounting(resource_guard) -> None: + request = _StatsRequest(1, list(range(8)), context_remaining_length=8) + manager = resource_guard(_create_manager(gpu_bytes=8 << 20, enable_stats=False), request) + + assert not manager.kv_cache_manager_py_config.enable_stats + assert manager.prepare_context(request) + assert manager.resize_context(request, num_tokens=8) + _finish_context(manager, request) + + manager.commit_scheduled_kv_cache_stats(_context_batch(request)) + assert manager.get_iteration_stats() is None + + kv_stats = manager.get_kv_cache_stats() + assert kv_stats.alloc_total_blocks == 0 + assert kv_stats.alloc_new_blocks == 0 + assert kv_stats.reused_blocks == 0 + assert kv_stats.missed_blocks == 0 + assert request.kv_cache_perf_metric_calls == [] + + +def test_context_and_generation_stats_are_reported(resource_guard) -> None: + request = _StatsRequest(1, list(range(8)), context_remaining_length=8) + manager = resource_guard(_create_manager(gpu_bytes=8 << 20), request) + + assert manager.prepare_context(request) + assert manager.resize_context(request, num_tokens=8) + _finish_context(manager, request) + + context_stats = _commit_and_get_stats(manager, _context_batch(request)) + _assert_iteration_delta(context_stats, alloc_total=2, alloc_new=2, missed=2) + assert request.kv_cache_perf_metric_calls == [ + _metric_call(alloc_total=2, alloc_new=2, missed=2), + ] + + assert manager.try_allocate_generation(request) + generation_stats = _commit_and_get_stats(manager, _generation_batch(request)) + _assert_iteration_delta(generation_stats, alloc_total=1, alloc_new=1, gen_alloc=1) + assert request.kv_cache_perf_metric_calls == [ + _metric_call(alloc_total=2, alloc_new=2, missed=2), + _metric_call(alloc_total=1, alloc_new=1), + ] + + +def test_reverted_generation_allocation_does_not_report_stats(resource_guard) -> None: + request = _StatsRequest(1, list(range(8)), context_remaining_length=8) + manager = resource_guard(_create_manager(gpu_bytes=8 << 20), request) + + assert manager.prepare_context(request) + assert manager.resize_context(request, num_tokens=8) + _finish_context(manager, request) + _commit_and_get_stats(manager, _context_batch(request)) + + assert manager.try_allocate_generation(request) + manager.revert_allocate_generation(request) + manager.commit_scheduled_kv_cache_stats(_generation_batch(request)) + stats_report = manager.get_iteration_stats() + assert stats_report is not None + _assert_iteration_delta(stats_report.by_window_size[manager.max_seq_len]) + + +def test_reverted_context_allocation_does_not_report_pending_stats(resource_guard) -> None: + request = _StatsRequest(1, list(range(8)), context_remaining_length=8) + manager = resource_guard(_create_manager(gpu_bytes=8 << 20), request) + + assert manager.prepare_context(request) + assert manager.resize_context(request, num_tokens=4) + request.context_current_position = 4 + request.context_remaining_length = 4 + manager.update_context_resources(_context_batch(request)) + first_chunk_stats = _commit_and_get_stats(manager, _context_batch(request)) + _assert_iteration_delta(first_chunk_stats, alloc_total=1, alloc_new=1, missed=1) + assert request.kv_cache_perf_metric_calls == [ + _metric_call(alloc_total=1, alloc_new=1, missed=1), + ] + + request.is_first_context_chunk = False + assert manager.prepare_context(request) + assert manager.resize_context(request, num_tokens=4) + manager.revert_allocate_context(request) + manager.commit_scheduled_kv_cache_stats(_context_batch(request)) + + reverted_stats_report = manager.get_iteration_stats() + assert reverted_stats_report is not None + _assert_iteration_delta(reverted_stats_report.by_window_size[manager.max_seq_len]) + assert request.kv_cache_perf_metric_calls == [ + _metric_call(alloc_total=1, alloc_new=1, missed=1), + ] + kv_stats = manager.get_kv_cache_stats() + assert kv_stats.alloc_total_blocks == 1 + assert kv_stats.alloc_new_blocks == 1 + assert kv_stats.missed_blocks == 1 + + assert manager.prepare_context(request) + assert manager.resize_context(request, num_tokens=4) + _finish_context(manager, request) + second_chunk_stats = _commit_and_get_stats(manager, _context_batch(request)) + _assert_iteration_delta(second_chunk_stats, alloc_total=1, alloc_new=1, missed=1) + assert request.kv_cache_perf_metric_calls == [ + _metric_call(alloc_total=1, alloc_new=1, missed=1), + _metric_call(alloc_total=1, alloc_new=1, missed=1), + ] + + +def test_chunked_context_reports_generation_alloc_only_in_generation(resource_guard) -> None: + request = _StatsRequest(1, list(range(8)), context_remaining_length=8) + manager = resource_guard(_create_manager(gpu_bytes=8 << 20), request) + + assert manager.prepare_context(request) + assert manager.resize_context(request, num_tokens=4) + request.context_current_position = 4 + request.context_remaining_length = 4 + manager.update_context_resources(_context_batch(request)) + first_chunk_stats = _commit_and_get_stats(manager, _context_batch(request)) + assert first_chunk_stats.iter_gen_alloc_blocks == 0 + assert request.kv_cache_perf_metric_calls == [ + _metric_call(alloc_total=1, alloc_new=1, missed=1), + ] + + request.is_first_context_chunk = False + assert manager.prepare_context(request) + assert manager.resize_context(request, num_tokens=4) + _finish_context(manager, request) + second_chunk_stats = _commit_and_get_stats(manager, _context_batch(request)) + assert second_chunk_stats.iter_gen_alloc_blocks == 0 + assert request.kv_cache_perf_metric_calls == [ + _metric_call(alloc_total=1, alloc_new=1, missed=1), + _metric_call(alloc_total=1, alloc_new=1, missed=1), + ] + + assert manager.try_allocate_generation(request) + generation_stats = _commit_and_get_stats(manager, _generation_batch(request)) + _assert_iteration_delta(generation_stats, alloc_total=1, alloc_new=1, gen_alloc=1) + assert request.kv_cache_perf_metric_calls == [ + _metric_call(alloc_total=1, alloc_new=1, missed=1), + _metric_call(alloc_total=1, alloc_new=1, missed=1), + _metric_call(alloc_total=1, alloc_new=1), + ] + + +def test_v2_generation_alloc_updates_request_metrics_unlike_v1(resource_guard) -> None: + v1_request = _create_llm_request(101, list(range(8))) + v2_request = _create_llm_request(201, list(range(8))) + v1_manager = resource_guard(_create_v1_manager(gpu_bytes=8 << 20), v1_request) + v2_manager = resource_guard(_create_manager(gpu_bytes=8 << 20), v2_request) + + v1_context_stats = _run_v1_context(v1_manager, v1_request) + v2_context_stats = _run_v2_context(v2_manager, v2_request) + _assert_iteration_delta(v1_context_stats, alloc_total=2, alloc_new=2, missed=2) + _assert_iteration_delta(v2_context_stats, alloc_total=2, alloc_new=2, missed=2) + _assert_request_stats(v1_request, alloc_total=2, alloc_new=2, missed=2) + _assert_request_stats(v2_request, alloc_total=2, alloc_new=2, missed=2) + + v1_generation_stats = _run_v1_generation(v1_manager, v1_request) + v2_generation_stats = _run_v2_generation(v2_manager, v2_request) + _assert_iteration_delta(v1_generation_stats, alloc_total=1, alloc_new=1, gen_alloc=1) + _assert_iteration_delta(v2_generation_stats, alloc_total=1, alloc_new=1, gen_alloc=1) + # V2 records generation allocation in request-level alloc_total/new. + # Legacy V1 only reports it through iteration/global generation counters. + _assert_request_stats(v2_request, alloc_total=3, alloc_new=3, missed=2) + _assert_request_stats(v1_request, alloc_total=2, alloc_new=2, missed=2) + + +def test_v2_partial_prompt_reuse_classification_matches_v1(resource_guard) -> None: + v1_warmup_request = _create_llm_request(101, list(range(12))) + v2_warmup_request = _create_llm_request(201, list(range(12))) + v1_reuse_request = _create_llm_request(102, list(range(10))) + v2_reuse_request = _create_llm_request(202, list(range(10))) + v1_manager = resource_guard( + _create_v1_manager(gpu_bytes=8 << 20), v1_warmup_request, v1_reuse_request + ) + v2_manager = resource_guard( + _create_manager(gpu_bytes=8 << 20), v2_warmup_request, v2_reuse_request + ) + + _run_v1_context(v1_manager, v1_warmup_request) + _run_v2_context(v2_manager, v2_warmup_request) + v1_manager.free_resources(v1_warmup_request) + v2_manager.free_resources(v2_warmup_request) + + v1_reuse_stats = _run_v1_context(v1_manager, v1_reuse_request) + v2_reuse_stats = _run_v2_context(v2_manager, v2_reuse_request) + _assert_iteration_delta(v1_reuse_stats, reused=3, full_reused=2, partial_reused=1) + _assert_iteration_delta( + v2_reuse_stats, + alloc_total=1, + alloc_new=1, + reused=3, + full_reused=2, + partial_reused=1, + intra_copy=1, + intra_copy_bytes=BYTES_PER_BLOCK, + ) + _assert_request_stats(v1_reuse_request, reused=3) + # V2 copies the partially reused block into a private slot before writing to it. + _assert_request_stats(v2_reuse_request, alloc_total=1, alloc_new=1, reused=3) + + +def test_block_reuse_disabled_records_generation_alloc(resource_guard) -> None: + request = _StatsRequest(1, list(range(8)), context_remaining_length=8) + manager = resource_guard(_create_manager(gpu_bytes=8 << 20, enable_block_reuse=False), request) + + assert manager.prepare_context(request) + assert manager.resize_context(request, num_tokens=8) + _finish_context(manager, request) + + stats = _commit_and_get_stats(manager, _context_batch(request)) + _assert_iteration_delta(stats, alloc_total=2, alloc_new=2, missed=2) + assert request.kv_cache_perf_metric_calls == [ + _metric_call(alloc_total=2, alloc_new=2, missed=2), + ] + + assert manager.try_allocate_generation(request) + generation_stats = _commit_and_get_stats(manager, _generation_batch(request)) + _assert_iteration_delta(generation_stats, alloc_total=1, alloc_new=1, gen_alloc=1) + assert request.kv_cache_perf_metric_calls == [ + _metric_call(alloc_total=2, alloc_new=2, missed=2), + _metric_call(alloc_total=1, alloc_new=1), + ] + + +def test_v2_partial_leaf_reuse_counts_reuse_with_private_copy(resource_guard) -> None: + warmup_request = _StatsRequest(201, list(range(9)), context_remaining_length=9) + reuse_request = _StatsRequest(202, list(range(10)), context_remaining_length=10) + manager = resource_guard( + _create_manager(gpu_bytes=8 << 20), + warmup_request, + reuse_request, + ) + + assert manager.prepare_context(warmup_request) + assert manager.resize_context(warmup_request, num_tokens=9) + _finish_context(manager, warmup_request) + _commit_and_get_stats(manager, _context_batch(warmup_request)) + manager.free_resources(warmup_request) + + assert manager.prepare_context(reuse_request) + assert reuse_request.prepopulated_prompt == (9, TOKENS_PER_BLOCK) + assert manager.resize_context(reuse_request, num_tokens=1) + _finish_context(manager, reuse_request) + + reuse_stats = _commit_and_get_stats(manager, _context_batch(reuse_request)) + _assert_iteration_delta( + reuse_stats, + alloc_total=1, + alloc_new=1, + reused=3, + full_reused=2, + partial_reused=1, + intra_copy=1, + intra_copy_bytes=BYTES_PER_BLOCK, + ) + assert reuse_request.kv_cache_perf_metric_calls == [ + _metric_call(alloc_total=1, alloc_new=1, reused=3), + ] + + +def test_swa_context_reuse_stats_skip_stale_prefix_blocks(resource_guard) -> None: + warmup_request = _StatsRequest(1, list(range(16)), context_remaining_length=16) + reuse_request = _StatsRequest(2, list(range(16)), context_remaining_length=16) + manager = resource_guard( + _create_manager( + gpu_bytes=8 << 20, + num_layers=1, + max_attention_window=[8], + ), + warmup_request, + reuse_request, + ) + + assert manager.prepare_context(warmup_request) + assert manager.resize_context(warmup_request, num_tokens=16) + _finish_context(manager, warmup_request) + manager.commit_scheduled_kv_cache_stats(_context_batch(warmup_request)) + assert manager.get_iteration_stats() is not None + manager.free_resources(warmup_request) + + assert manager.prepare_context(reuse_request) + assert reuse_request.prepopulated_prompt == (15, TOKENS_PER_BLOCK) + assert manager.resize_context(reuse_request, num_tokens=1) + _finish_context(manager, reuse_request) + manager.commit_scheduled_kv_cache_stats(_context_batch(reuse_request)) + + stats_report = manager.get_iteration_stats() + assert stats_report is not None + swa_stats = stats_report.by_window_size[8] + _assert_iteration_delta( + swa_stats, + alloc_total=1, + alloc_new=1, + reused=2, + full_reused=1, + partial_reused=1, + intra_copy=1, + intra_copy_bytes=BYTES_PER_BLOCK, + ) + assert reuse_request.kv_cache_perf_metric_calls == [ + _metric_call(alloc_total=1, alloc_new=1, reused=2), + ] + + +def test_pool_group_stats_are_reported(resource_guard) -> None: + request = _StatsRequest(1, list(range(8)), context_remaining_length=8) + manager = resource_guard( + _create_manager( + gpu_bytes=16 << 20, + num_layers=2, + max_attention_window=[16, 8], + ), + request, + ) + + assert manager.prepare_context(request) + assert manager.resize_context(request, num_tokens=8) + _finish_context(manager, request) + manager.commit_scheduled_kv_cache_stats(_context_batch(request)) + + stats_report = manager.get_iteration_stats() + assert stats_report is not None + assert set(stats_report.by_window_size) == {manager.max_seq_len, 8} + assert set(stats_report.by_pool_group) == {0} + + pool_group = stats_report.by_pool_group[0] + assert pool_group.pool_group_id == 0 + assert set(pool_group.window_sizes) == {manager.max_seq_len, 8} + _assert_iteration_delta(pool_group.stats, alloc_total=4, alloc_new=4) + + life_cycle_stats = { + stats.window_size: stats.stats for stats in stats_report.by_life_cycle.values() + } + assert set(life_cycle_stats) == {manager.max_seq_len, 8} + _assert_iteration_delta(life_cycle_stats[manager.max_seq_len], missed=2) + _assert_iteration_delta(life_cycle_stats[8], missed=2) + + _assert_iteration_delta( + stats_report.by_window_size[manager.max_seq_len], + alloc_total=2, + alloc_new=2, + missed=2, + ) + _assert_iteration_delta( + stats_report.by_window_size[8], + alloc_total=2, + alloc_new=2, + missed=2, + ) diff --git a/tests/unittest/llmapi/test_llm_args.py b/tests/unittest/llmapi/test_llm_args.py index 5556f3bfbf04..41df9229ecc8 100644 --- a/tests/unittest/llmapi/test_llm_args.py +++ b/tests/unittest/llmapi/test_llm_args.py @@ -333,6 +333,8 @@ def get_model_defaults(cls, llm_args): def test_KvCacheConfig_declaration(): + assert KvCacheConfig().kv_cache_event_hash_algo == "auto" + config = KvCacheConfig(enable_block_reuse=True, max_tokens=1024, max_attention_window=[1024, 1024, 1024], @@ -343,6 +345,7 @@ def test_KvCacheConfig_declaration(): cross_kv_cache_fraction=0.5, secondary_offload_min_priority=1, event_buffer_max_size=0, + kv_cache_event_hash_algo="v2_sha256_64", enable_partial_reuse=True, copy_on_partial_reuse=True, attention_dp_events_gather_period_ms=10) @@ -358,6 +361,11 @@ def test_KvCacheConfig_declaration(): assert pybind_config.cross_kv_cache_fraction == 0.5 assert pybind_config.secondary_offload_min_priority == 1 assert pybind_config.event_buffer_max_size == 0 + assert config.kv_cache_event_hash_algo == "v2_sha256_64" + assert KvCacheConfig( + kv_cache_event_hash_algo="auto").kv_cache_event_hash_algo == "auto" + assert KvCacheConfig(kv_cache_event_hash_algo="v1_block_key" + ).kv_cache_event_hash_algo == "v1_block_key" 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 diff --git a/tests/unittest/llmapi/test_llm_kv_cache_events.py b/tests/unittest/llmapi/test_llm_kv_cache_events.py index ee002905d26c..b170cd01219c 100644 --- a/tests/unittest/llmapi/test_llm_kv_cache_events.py +++ b/tests/unittest/llmapi/test_llm_kv_cache_events.py @@ -24,6 +24,7 @@ serialize_item) from tensorrt_llm.llmapi import KvCacheConfig from tensorrt_llm.mapping import Mapping +from tensorrt_llm.runtime.kv_cache_hash import KV_CACHE_HASH_ALGO_V1 from tensorrt_llm.sampling_params import SamplingParams from tensorrt_llm.scheduling_params import SchedulingParams @@ -34,7 +35,8 @@ global_kvcache_config = KvCacheConfig(free_gpu_memory_fraction=0.4, event_buffer_max_size=1024, enable_block_reuse=True, - max_tokens=256) + max_tokens=256, + use_kv_cache_manager_v2=False) def create_kv_cache_manager(): @@ -66,6 +68,17 @@ def create_llm(tensor_parallel_size=1): enable_autotuner=False) +def create_v2_llm(): + v2_kvcache_config = KvCacheConfig(free_gpu_memory_fraction=0.4, + event_buffer_max_size=1024, + enable_block_reuse=True, + max_tokens=256, + use_kv_cache_manager_v2=True) + return LLM(model=llama_model_path, + kv_cache_config=v2_kvcache_config, + enable_autotuner=False) + + def create_llm_request(id, input_tokens, new_tokens=1): sampling_params = SamplingParams() req = LlmRequest(request_id=id, @@ -892,6 +905,38 @@ def test_expected_kv_cache_events(): assert event["data"]["type"] == "stored" +def test_expected_v2_kv_cache_events(): + with create_v2_llm() as llm: + sampling_params = SamplingParams(max_tokens=6, temperature=0.01) + prompt = list(range(127)) + + _ = llm.generate(prompt, sampling_params=sampling_params) + + events = llm.get_kv_cache_events(5) + assert events and len(events) >= 2 + assert all(event["hash_algo"] == KV_CACHE_HASH_ALGO_V1 + for event in events if event) + + created_events = [ + event for event in events + if event and event["data"]["type"] == "created" + ] + stored_events = [ + event for event in events + if event and event["data"]["type"] == "stored" + ] + assert created_events + assert stored_events + assert created_events[0]["event_id"] == 0 + + block_hashes = [ + block["block_hash"] for event in stored_events + for block in event["data"]["blocks"] + ] + assert block_hashes + assert all(isinstance(block_hash, int) for block_hash in block_hashes) + + def test_kv_cache_event_async_api(): llm = create_llm() sampling_params = SamplingParams(max_tokens=6, temperature=0.01) diff --git a/tests/unittest/metrics/test_collector.py b/tests/unittest/metrics/test_collector.py index aa1d931263a7..e66975ca6f8e 100644 --- a/tests/unittest/metrics/test_collector.py +++ b/tests/unittest/metrics/test_collector.py @@ -804,6 +804,48 @@ def test_counters_incremented(self): collector, "kv_cache_intra_device_copy_bytes_total" ) - before_intra_device == pytest.approx(16384) + def test_v2_lifecycle_and_pool_group_stats_are_aggregated(self): + """V2 split stats should aggregate reuse from lifecycle and storage from PG.""" + collector = _make_kv_iter_collector() + stats = { + "kvCacheIterationStatsByLifecycle": { + "0": { + "iterReusedBlocks": 5, + "iterFullReusedBlocks": 4, + "iterPartialReusedBlocks": 1, + "iterMissedBlocks": 3, + } + }, + "kvCacheIterationStatsByPoolGroup": { + "0": { + "secondaryMaxNumBlocks": 50, + "secondaryUsedNumBlocks": 20, + "iterGenAllocBlocks": 2, + "iterOnboardBytes": 4096, + "iterOffloadBytes": 2048, + "iterIntraDeviceCopyBytes": 8192, + } + }, + } + + before_reused = _get_counter_value(collector, "kv_cache_iter_reused_blocks") + before_gen_alloc = _get_counter_value(collector, "kv_cache_gen_alloc_blocks_total") + before_onboard = _get_counter_value(collector, "kv_cache_onboard_bytes_total") + + collector.log_iteration_stats(stats) + + assert _get_gauge_value(collector, "kv_cache_host_utilization") == pytest.approx(0.4) + assert _get_gauge_value(collector, "kv_cache_iter_reuse_rate") == pytest.approx(5 / 8) + assert _get_counter_value( + collector, "kv_cache_iter_reused_blocks" + ) - before_reused == pytest.approx(5) + assert _get_counter_value( + collector, "kv_cache_gen_alloc_blocks_total" + ) - before_gen_alloc == pytest.approx(2) + assert _get_counter_value( + collector, "kv_cache_onboard_bytes_total" + ) - before_onboard == pytest.approx(4096) + def test_multiple_windows_aggregated(self): """Stats from multiple window sizes should be summed.""" collector = _make_kv_iter_collector()