diff --git a/cpp/include/tensorrt_llm/batch_manager/llmRequest.h b/cpp/include/tensorrt_llm/batch_manager/llmRequest.h index 715bba11b394..79427f44a68a 100644 --- a/cpp/include/tensorrt_llm/batch_manager/llmRequest.h +++ b/cpp/include/tensorrt_llm/batch_manager/llmRequest.h @@ -1,5 +1,5 @@ /* - * Copyright (c) 2022-2025, NVIDIA CORPORATION. All rights reserved. + * Copyright (c) 2022-2026, NVIDIA CORPORATION. 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. @@ -1835,6 +1835,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 0af68c22a624..12ba0153f1aa 100644 --- a/cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2022-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2022-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); @@ -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/attention_backend/sparse/deepseek_v4/cache_manager.py b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/cache_manager.py index 7eb80eb3f99a..814070fd8ab0 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/cache_manager.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/cache_manager.py @@ -1,3 +1,18 @@ +# 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 import defaultdict from typing import Dict, List, Optional, Tuple @@ -21,6 +36,7 @@ AttentionLayerConfig, BatchDesc, BufferConfig, + DataRole, GpuCacheTierConfig, HostCacheTierConfig, KVCacheDesc, @@ -229,6 +245,22 @@ def __init__( device="cpu", ) + def _format_kv_cache_pool_lifecycle_entry(self, layer_id: LayerId, role: DataRole) -> str: + layer_semantics = self._manager_layer_id_to_layer_attn.get(layer_id) + if layer_semantics is None: + return super()._format_kv_cache_pool_lifecycle_entry(layer_id, role) + + model_layer_idx, attn_type = layer_semantics + 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"deepseek_role={attn_type.name}, " + f"compress_ratio={self._compress_ratios[model_layer_idx]}, " + f"pool_group_id={int(pool_group_id)}, " + f"lifecycle_id={int(attr.life_cycle_id)}, lifecycle={lifecycle}" + ) + def get_buffers(self, layer_idx: int, attn_type: DeepseekV4AttentionType) -> torch.Tensor: """ Get the buffers for a specific layer and attention type. @@ -365,6 +397,7 @@ def _build_cache_config( """ layers: List[AttentionLayerConfig] = [] layer_attn_to_layer_id: Dict[Tuple[int, DeepseekV4AttentionType], LayerId] = {} + manager_layer_id_to_layer_attn: Dict[LayerId, Tuple[int, DeepseekV4AttentionType]] = {} def _add_layer( layer_idx: int, attn_type: DeepseekV4AttentionType, sliding_window_size: int | None @@ -373,6 +406,7 @@ def _add_layer( layer_id = LayerId(len(layers)) # update the mapping from layer index and attention type to layer id layer_attn_to_layer_id[layer_idx, attn_type] = layer_id + manager_layer_id_to_layer_attn[layer_id] = (layer_idx, attn_type) # add the layer to the layers list layer_config = AttentionLayerConfig( layer_id=layer_id, @@ -433,6 +467,7 @@ def _add_layer( ) # the mapping from layer index and attention type to layer id self._layer_attn_to_layer_id = layer_attn_to_layer_id + self._manager_layer_id_to_layer_attn = manager_layer_id_to_layer_attn # number of layers in the KVCacheManagerPy self._num_manager_layers = len(layers) @@ -476,6 +511,7 @@ def _add_layer( 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=layers, typical_step=typical_step, constraints=constraints, diff --git a/tensorrt_llm/_torch/auto_deploy/shim/ad_executor.py b/tensorrt_llm/_torch/auto_deploy/shim/ad_executor.py index eec21a77f93c..8ffdc94b1108 100644 --- a/tensorrt_llm/_torch/auto_deploy/shim/ad_executor.py +++ b/tensorrt_llm/_torch/auto_deploy/shim/ad_executor.py @@ -274,6 +274,8 @@ def create_draft_kv_cache_manager_maybe( max_beam_width=ad_config.max_beam_width, kv_connector_manager=None, # KV connector manager not used in AutoDeploy (no disagg support) estimating_kv_cache=False, + enable_kv_cache_stats=ad_config.enable_iter_perf_stats + or getattr(ad_config, "return_perf_metrics", False), ) diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index 817616cf13d5..a79798013f73 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -470,11 +470,13 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): str(write_meta.slice_id).encode("ascii"), b"True", # is_last_slice — ensures receiver resolves its task future AgentResult.FAILED.value.encode("ascii"), + b"0", ] ) return agent_result = AgentResult.SUCCESS + transferred_bytes = int(write_meta.sizes.sum()) if write_meta.sizes.size > 0 else 0 if write_meta.src_ptrs.size > 0: request = Sender._make_agent_request(write_meta, device_id=self._device_id) if timer: @@ -510,9 +512,14 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): str(write_meta.slice_id).encode("ascii"), str(write_meta.is_last_slice).encode("ascii"), agent_result.value.encode("ascii"), + str(transferred_bytes if agent_result == AgentResult.SUCCESS else 0).encode( + "ascii" + ), ] ) + if agent_result == AgentResult.SUCCESS: + session.record_kv_transfer_bytes(transferred_bytes) task.transferred_count += 1 if timer: timer.record_task_end(write_meta.peer_rank) @@ -934,6 +941,7 @@ def _send_failed_result_to_receiver(self, info: RecvReqInfo): str(slice_id).encode("ascii"), b"True", # is_last_slice AgentResult.FAILED.value.encode("ascii"), + b"0", ] ) except Exception as e: @@ -1059,6 +1067,7 @@ def __init__( self.kv_tasks = [] self.aux_task = None self.lock = threading.Lock() + self._transferred_kv_bytes = 0 self._exception: Optional[Exception] = None self._closed = False @@ -1079,6 +1088,15 @@ def disagg_request_id(self) -> int: return params.ctx_request_id return self.request_id + def record_kv_transfer_bytes(self, transferred_bytes: int) -> None: + with self.lock: + self._transferred_kv_bytes += transferred_bytes + + @property + def transferred_kv_bytes(self) -> int: + with self.lock: + return self._transferred_kv_bytes + @property def status(self) -> SessionStatus: if self._terminal_status is not None: @@ -1497,12 +1515,26 @@ def _handle_cancel_session(self, message: list[bytes]): session.cancel() def _process_kv_agent_result(self, _send_id: bytes, message: list[bytes]): - msg_type, peer_rank, unique_rid, slice_id_str, is_last_slice_str, status = decode_message( - message - ) + decoded_message = decode_message(message) + if len(decoded_message) == 6: + msg_type, peer_rank, unique_rid, slice_id_str, is_last_slice_str, status = ( + decoded_message + ) + transferred_bytes = 0 + else: + ( + msg_type, + peer_rank, + unique_rid, + slice_id_str, + is_last_slice_str, + status, + transferred_bytes, + ) = decoded_message peer_rank = int(peer_rank) unique_rid = int(unique_rid) sender_slice_id = int(slice_id_str) + transferred_bytes = int(transferred_bytes) if msg_type.encode("ascii") != MessageType.KV_AGENT_RESULT: logger.error( f"_process_kv_agent_result: unexpected msg_type={msg_type!r}, expected KV_AGENT_RESULT" @@ -1515,7 +1547,11 @@ def _process_kv_agent_result(self, _send_id: bytes, message: list[bytes]): ) return session.process_kv_agent_result( - peer_rank, sender_slice_id, is_last_slice_str == "True", AgentResult(status) + peer_rank, + sender_slice_id, + is_last_slice_str == "True", + AgentResult(status), + transferred_bytes, ) def _process_aux_agent_result(self, _send_id: bytes, message: list[bytes]): @@ -1575,6 +1611,7 @@ def __init__( self._aux_status: TaskStatus = TaskStatus.INIT self._sender_endpoints: set[str] = set() self.lock = threading.Lock() + self._transferred_kv_bytes = 0 self._receiver.setup_session(self) @property @@ -1589,6 +1626,11 @@ def disagg_request_id(self) -> int: return params.ctx_request_id return self.request_id + @property + def transferred_kv_bytes(self) -> int: + with self.lock: + return self._transferred_kv_bytes + @property def status(self) -> SessionStatus: if self._terminal_status is not None: @@ -1623,7 +1665,12 @@ def receive(self, slice: KVSlice) -> None: self._receiver.dispatch_task(task) def process_kv_agent_result( - self, peer_rank: int, sender_slice_id: int, is_last_slice: bool, status: AgentResult + self, + peer_rank: int, + sender_slice_id: int, + is_last_slice: bool, + status: AgentResult, + transferred_bytes: int = 0, ): with self.lock: assert sender_slice_id < len(self._kv_tasks), ( @@ -1633,6 +1680,7 @@ def process_kv_agent_result( ) task = self._kv_tasks[sender_slice_id] if status == AgentResult.SUCCESS: + self._transferred_kv_bytes += transferred_bytes if is_last_slice: task.last_slice_count += 1 if task.last_slice_count == task.expected_transfers: diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index ec382c5aae11..1a2a8c519d1c 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -1,5 +1,7 @@ +import os import uuid from collections import defaultdict +from datetime import timedelta from itertools import chain from typing import Any, Callable, Dict, List, Optional, cast @@ -36,6 +38,7 @@ from tensorrt_llm.disaggregated_params import DisaggScheduleStyle from tensorrt_llm.llmapi.llm_args import CacheTransceiverConfig from tensorrt_llm.mapping import Mapping +from tensorrt_llm.serve.responses_utils import get_steady_clock_now_in_seconds def _find_consensus_request_ids(request_ids_all_ranks, sync_size): @@ -338,6 +341,82 @@ def _ctx_consensus_outcome(self, to_process, cancelled, failed, completed, timed c, f, d = self._consensus_outcome(to_process, c, f, d, pp_allgather, True) return c, f, d, timed_out + @staticmethod + def _clock_offset_seconds() -> float: + offset = LlmRequest.global_steady_clock_offset + if offset is None: + return 0.0 + if isinstance(offset, timedelta): + return offset.total_seconds() + return float(offset) + + def _record_transfer_start(self, req: LlmRequest) -> None: + local_start = get_steady_clock_now_in_seconds() + req.set_kv_cache_transfer_start(timedelta(seconds=local_start)) + req.set_kv_cache_size(0) + + def _record_transfer_end(self, req: LlmRequest, session: Optional[Any] = None) -> None: + local_end = get_steady_clock_now_in_seconds() + if session is not None: + req.set_kv_cache_size(int(getattr(session, "transferred_kv_bytes", 0))) + req.set_kv_cache_transfer_end(timedelta(seconds=local_end)) + + def _set_transfer_metrics( + self, req: LlmRequest, start_s: float, end_s: float, size_bytes: int + ) -> None: + offset_s = self._clock_offset_seconds() + req.set_kv_cache_transfer_start(timedelta(seconds=start_s - offset_s)) + req.set_kv_cache_transfer_end(timedelta(seconds=end_s - offset_s)) + req.set_kv_cache_size(size_bytes) + + @staticmethod + def _time_point_seconds(time_point: Any) -> float: + if time_point is None: + return 0.0 + if isinstance(time_point, timedelta): + return time_point.total_seconds() + return float(time_point) + + @staticmethod + def _get_transfer_metrics(req: LlmRequest) -> Optional[tuple[float, float, int]]: + start_s = KvCacheTransceiverV2._time_point_seconds(req.kv_cache_transfer_start) + end_s = KvCacheTransceiverV2._time_point_seconds(req.kv_cache_transfer_end) + if start_s <= 0.0 or end_s <= 0.0: + return None + return start_s, end_s, req.kv_cache_size + + @staticmethod + def _should_aggregate_gen_transfer_metrics() -> bool: + return bool(os.environ.get("TRTLLM_KVCACHE_TIME_OUTPUT_PATH", "")) + + def _publish_gen_transfer_metrics(self, completed: list[int], consensus: set[int]) -> None: + if not self._should_aggregate_gen_transfer_metrics(): + return + + metric_rids = [rid for rid in completed if not self._gen_need_sync or rid in consensus] + if not metric_rids: + return + + local_metrics = { + rid: self._get_transfer_metrics(self._recv_reqs[rid]) for rid in metric_rids + } + all_rank_metrics = ( + self._gen_allgather(local_metrics) if self._gen_need_sync else [local_metrics] + ) + + for rid in metric_rids: + metrics = [ + rank_metrics[rid] + for rank_metrics in all_rank_metrics + if rid in rank_metrics and rank_metrics[rid] is not None + ] + if not metrics: + continue + start_s = min(metric[0] for metric in metrics) + end_s = max(metric[1] for metric in metrics) + size_bytes = sum(metric[2] for metric in metrics) + self._set_transfer_metrics(self._recv_reqs[rid], start_s, end_s, size_bytes) + def _collect_done(self, sessions: dict, reqs: dict): """Scan sessions and return (completed_rids, failed_rids).""" completed, failed = [], [] @@ -416,7 +495,9 @@ def respond_and_send_async(self, req: LlmRequest): self._ever_had_send_session = True session = self._get_or_create_send_session(req) req.state = LlmRequestState.DISAGG_CONTEXT_TRANS_IN_PROGRESS - session.send(self._create_kv_slice(req)) + kv_slice = self._create_kv_slice(req) + self._record_transfer_start(req) + session.send(kv_slice) self._finalize_send(req, session) @nvtx_range("KvCacheTransceiverV2.request_and_receive_sync") @@ -433,10 +514,14 @@ def request_and_receive_sync(self, req: LlmRequest): session = self._transfer_worker.create_rx_session(req) self._recv_sessions[rid] = session self._recv_reqs[rid] = req - session.receive(self._create_kv_slice(req)) + kv_slice = self._create_kv_slice(req) + self._record_transfer_start(req) + session.receive(kv_slice) result = session.wait_complete(blocking=True) if result == WaitResult.COMPLETED: + self._record_transfer_end(req, session) + self._publish_gen_transfer_metrics([rid], {rid}) if self._need_aux_transfer(req): self._apply_aux(session, req) self._trim_kv_to_prompt_history(req) @@ -464,7 +549,9 @@ def request_and_receive_async(self, req: LlmRequest): req.state = LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS session = self._transfer_worker.create_rx_session(req) self._recv_sessions[rid] = session - session.receive(self._create_kv_slice(req)) + kv_slice = self._create_kv_slice(req) + self._record_transfer_start(req) + session.receive(kv_slice) self._recv_reqs[rid] = req def check_context_transfer_status( @@ -491,6 +578,7 @@ def check_context_transfer_status( if session.status == SessionStatus.CANCELLED: cancelled.append(rid) elif result == WaitResult.COMPLETED: + self._record_transfer_end(self._send_reqs[rid], session) completed.append(rid) elif result == WaitResult.TIMEOUT: logger.warning( @@ -551,6 +639,7 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]): # distinguish the two cases and set the appropriate state. cancelled.append(rid) elif result == WaitResult.COMPLETED: + self._record_transfer_end(self._recv_reqs[rid], session) completed.append(rid) elif result == WaitResult.FAILED: failed.append(rid) @@ -561,6 +650,7 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]): to_process, cancelled, failed, completed ) + self._publish_gen_transfer_metrics(completed, set(completed)) cancelled_reqs = [] for rid in cancelled: cancelled_reqs.append(self._recv_reqs[rid]) diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 62fb5c017d0e..78317b084e2a 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -235,6 +235,10 @@ def _warn_if_unsupported_kv_cache_manager_v2(self, kv_cache_manager_cls): "KVCacheManagerV2 is not supported with kv_connector_manager or beam width > 1." ) + def _enable_kv_cache_stats(self) -> bool: + return (self._llm_args.enable_iter_perf_stats + or getattr(self._llm_args, "return_perf_metrics", False)) + def _per_manager_cache_cost(self, manager_cls, model_config, **extra_kwargs) -> CacheCost: return CacheCost.from_raw( @@ -716,6 +720,8 @@ def _create_kv_cache_manager( max_input_len=self._llm_args.max_input_len, kv_connector_manager=self._kv_connector_manager, estimating_kv_cache=estimating_kv_cache, + enable_kv_cache_stats=self._enable_kv_cache_stats() + and not estimating_kv_cache, execution_stream=self._execution_stream, layer_mask=spec_dec_layer_mask, is_disagg=self._is_disagg, @@ -851,6 +857,8 @@ def _create_one_model_draft_kv_cache_manager( max_beam_width=self._max_beam_width, kv_connector_manager=self._kv_connector_manager, estimating_kv_cache=estimating_kv_cache, + enable_kv_cache_stats=self._enable_kv_cache_stats() + and not estimating_kv_cache, execution_stream=self._execution_stream, is_disagg=self._is_disagg, # One-model draft specific overrides @@ -1051,6 +1059,7 @@ def _create_kv_cache_manager( is_disagg: bool = False, max_input_len: Optional[int] = None, estimating_kv_cache: bool = False, + enable_kv_cache_stats: bool = False, execution_stream: Optional[torch.cuda.Stream] = None, # Optional overrides for one-model draft case (when model_engine is None) model_config: Optional[ModelConfig] = None, @@ -1168,6 +1177,9 @@ def _create_kv_cache_manager( per_layer_num_kv_heads = _build_per_layer_num_kv_heads( num_key_value_heads, num_hidden_layers, spec_config, draft_config_for_kv) + manager_extra_kwargs = {} + if issubclass(kv_cache_manager_cls, KVCacheManagerV2): + manager_extra_kwargs["enable_stats"] = enable_kv_cache_stats if is_mla(config): kv_cache_manager = kv_cache_manager_cls( @@ -1194,6 +1206,7 @@ def _create_kv_cache_manager( layer_mask=layer_mask, max_num_tokens=max_num_tokens, is_disagg=is_disagg, + **manager_extra_kwargs, ) elif is_nemotron_hybrid(config): if max_beam_width > 1: @@ -1275,6 +1288,7 @@ def _create_kv_cache_manager( execution_stream=execution_stream, model_type="nemotron_hybrid", use_replay_state_update=use_replay, + **manager_extra_kwargs, ) elif is_qwen3_hybrid(config): if max_beam_width > 1: @@ -1319,6 +1333,7 @@ def _create_kv_cache_manager( is_estimating_kv_cache=estimating_kv_cache, execution_stream=execution_stream, model_type="qwen3_next", + **manager_extra_kwargs, ) else: # NOTE: this is a workaround for VSWA to switch to calculate_max_num_blocks_for_vswa in KVCahceManager @@ -1361,6 +1376,7 @@ def _create_kv_cache_manager( execution_stream=execution_stream, layer_mask=layer_mask, is_disagg=is_disagg, + **manager_extra_kwargs, ) # Note: Gemma4 KV sharing cache remapping is handled in Gemma4Attention # via cache_layer_idx — shared layers use target layer's index for 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/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index fe39fe9d123c..42acbda93146 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -57,6 +57,7 @@ from .handle_additional_outputs import HandleAdditionalOutputs from .handle_logits import HandleLogits from .hang_detector import HangDetector +from .kv_cache_stats import append_kv_cache_iteration_stats from .kv_cache_transceiver import KvCacheTransceiver from .llm_request import (ExecutorRequest, LlmRequest, LlmRequestState, LlmResponse, get_draft_token_length) @@ -1428,34 +1429,8 @@ def _append_iter_stats(self, local_dict["requestStats"] = [ _json.loads(r.to_json_str()) for r in req_stats ] - if self._latest_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 self._latest_kv_iter_stats.items() - } + append_kv_cache_iteration_stats(local_dict, + self._latest_kv_iter_stats) local_dict["rank"] = self.dist.tp_rank gathered = self.dist.tp_allgather(local_dict) @@ -1721,6 +1696,8 @@ def _executor_loop_pp(self): # Return the first token to the client self._handle_first_token_response(scheduled_batch) + self._commit_kv_cache_stats(scheduled_batch) + # Stage 1.1: Async forward (all ranks) and decoding pass (last rank only) if not self.dist.is_last_pp_rank: with torch.cuda.nvtx.range( @@ -2138,6 +2115,13 @@ def _revert_ctx_alloc(self, dropped_context_requests): for req in dropped_context_requests: self.kv_cache_manager.revert_allocate_context(req) + def _commit_kv_cache_stats(self, + scheduled_batch: ScheduledRequests) -> None: + if self._scheduler_manages_kv_suspend and isinstance( + self.kv_cache_manager, KVCacheManagerV2): + 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: @@ -2509,6 +2493,8 @@ def _executor_loop(self): if hasattr(self.drafter, "guided_decoder"): self.guided_decoder.rollback_draft_tokens() + 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( ) @@ -2834,6 +2820,8 @@ def _executor_loop_overlap(self): else: previous_tensors_device = self.previous_batch and self.previous_batch.sample_state and self.previous_batch.sample_state.device + self._commit_kv_cache_stats(scheduled_batch) + # GPU timing for perf metrics gpu_forward_start, gpu_forward_end, gpu_sample_end = self.perf_manager.create_timing_events( ) @@ -2894,6 +2882,8 @@ def _executor_loop_overlap(self): 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()) @@ -4388,6 +4378,9 @@ def _handle_responses(self): request_done = False if request.py_decoding_iter == 1 or request.is_finished or \ request.py_decoding_iter % self.stream_interval == 0: + 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 d6c74b0693f6..68f346eb86a0 100644 --- a/tensorrt_llm/_torch/pyexecutor/resource_manager.py +++ b/tensorrt_llm/_torch/pyexecutor/resource_manager.py @@ -18,6 +18,7 @@ import os from abc import ABC, abstractmethod from collections import OrderedDict, defaultdict, deque +from dataclasses import fields from typing import (TYPE_CHECKING, Dict, Iterable, List, Optional, Sequence, Set, Tuple, Union) @@ -31,7 +32,8 @@ get_size_in_bytes, mpi_comm, mpi_disabled, prefer_pinned, torch_comm) from tensorrt_llm.bindings.internal.batch_manager import ( - KvCacheStats, LinearAttentionMetadata, LinearCacheType) + KvCacheIterationStats, KvCacheStats, LinearAttentionMetadata, + LinearCacheType) from tensorrt_llm.bindings.internal.batch_manager.kv_cache_manager_v2_utils import ( IndexMapper, copy_batch_block_offsets_to_device) from tensorrt_llm.bindings.internal.runtime import TaskLayerModuleConfig @@ -48,6 +50,7 @@ DEFAULT_BEAM_INDEX, AttentionLayerConfig, BufferConfig, CacheTierConfig, GpuCacheTierConfig, HostCacheTierConfig) # isort: on +from tensorrt_llm.runtime.kv_cache_manager_v2 import KVCacheIterationStatsDelta from tensorrt_llm.runtime.kv_cache_manager_v2 import \ KVCacheManager as KVCacheManagerPy from tensorrt_llm.runtime.kv_cache_manager_v2 import \ @@ -57,13 +60,16 @@ from tensorrt_llm.runtime.kv_cache_manager_v2._block_radix_tree import \ gen_multi_modal_tokens from tensorrt_llm.runtime.kv_cache_manager_v2._common import (BAD_PAGE_INDEX, - GPU_LEVEL) + 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 @@ -72,6 +78,9 @@ 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 .scheduler import ScheduledRequests @@ -93,6 +102,17 @@ BlocksPerWindow = Dict[int, Tuple[ int, int]] # window_size -> (blocks_in_primary_pool, blocks_in_secondary_pool) +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) def _warn_if_unsupported_v1_kv_cache_event_hash_algo(hash_algo: str) -> None: @@ -1866,6 +1886,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 @@ -1907,6 +1928,7 @@ def __init__( self.max_total_draft_tokens = spec_config.max_total_draft_tokens if spec_config is not None else 0 self.event_buffer_max_size = kv_cache_config.event_buffer_max_size + 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, @@ -2228,6 +2250,8 @@ def append_to_kv_heads_per_layer(num_kv_heads_per_layer: List[int], pin_memory=prefer_pinned(), 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()) @@ -2263,6 +2287,31 @@ def get_event_window_size(layer_id: int) -> int: 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([[ self.impl.get_mem_pool_base_address( @@ -2345,6 +2394,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=[ AttentionLayerConfig( layer_id=layer_id, @@ -2449,6 +2499,25 @@ 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: @@ -2645,10 +2714,10 @@ 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: @@ -2657,7 +2726,9 @@ def prepare_context(self, req: LlmRequest) -> bool: req.py_request_id, req.lora_task_id, tokens, - cache_salt_id=req.cache_salt_id) + cache_salt_id=req.cache_salt_id, + is_dummy=req.is_dummy, + expected_prompt_length=req.prompt_len - 1) if kv_cache is None: return False kv_cache.cuda_stream = self._stream.cuda_stream @@ -2703,7 +2774,8 @@ def resize_context(self, capacity = max(kv_cache.capacity, target) pre_cap = kv_cache.capacity - if not kv_cache.resize(capacity, history_length): + success = kv_cache.resize(capacity, history_length) + if not success: if req.is_first_context_chunk: kv_cache.suspend() return False @@ -2773,7 +2845,8 @@ def _prepare_draft_resources(self, scheduled_batch: ScheduledRequests): req.py_request_id, req.lora_task_id, None, - cache_salt_id=req.cache_salt_id) + cache_salt_id=req.cache_salt_id, + is_dummy=req.is_dummy) kv_cache.stop_committing() if not self._resume_and_restore(req.py_request_id, kv_cache): raise RuntimeError( @@ -2862,9 +2935,237 @@ def _augment_tokens_for_block_reuse( n] = mm_tokens[mm_offset:mm_offset + n] return result + 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 @@ -2878,8 +3179,72 @@ def get_latest_events(self, timeout_ms: Optional[float] = None): 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) @@ -2956,8 +3321,10 @@ def release_resources(current_request: LlmRequest, req.py_request_id, req.lora_task_id, input_tokens, - cache_salt_id=req.cache_salt_id) - # Saturated IndexMapper (e.g. disagg gen trans in progress) → None; retry next iter. + cache_salt_id=req.cache_salt_id, + 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 @@ -2981,7 +3348,8 @@ def release_resources(current_request: LlmRequest, req.py_request_id, req.lora_task_id, input_tokens, - cache_salt_id=req.cache_salt_id) + cache_salt_id=req.cache_salt_id, + is_dummy=req.is_dummy) if draft_kv_cache is None: release_resources(req) return None @@ -3051,9 +3419,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: @@ -3351,7 +3722,9 @@ def _create_kv_cache(self, lora_task_id: int | None, input_tokens: Sequence[TokenIdExt] | None, *, - cache_salt_id: int | None = None): + cache_salt_id: int | 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" if self.index_mapper.num_free_slots() == 0: logger.warning( @@ -3360,10 +3733,16 @@ def _create_kv_cache(self, "Skipping KV cache creation; request will retry next iteration.", request_id, self.index_mapper.size(), self.index_mapper.size()) return None - kv_cache = self.impl.create_kv_cache(lora_task_id, - input_tokens, - cache_salt_id=cache_salt_id) + kv_cache = self.impl._create_kv_cache( + lora_task_id, + input_tokens, + id=request_id, + cache_salt_id=cache_salt_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/executor/base_worker.py b/tensorrt_llm/executor/base_worker.py index 11e7c56cdf5f..998c4f9c1406 100644 --- a/tensorrt_llm/executor/base_worker.py +++ b/tensorrt_llm/executor/base_worker.py @@ -13,6 +13,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) @@ -684,7 +685,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}). @@ -711,34 +711,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) # Convert back to JSON string return json.dumps(stats_dict) diff --git a/tensorrt_llm/metrics/collector.py b/tensorrt_llm/metrics/collector.py index f10c622693bd..d4391339b2ca 100644 --- a/tensorrt_llm/metrics/collector.py +++ b/tensorrt_llm/metrics/collector.py @@ -677,9 +677,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 @@ -691,19 +698,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_manager_v2/__init__.py b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py index e84a14d4a87b..44b31ebaf9c2 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py @@ -61,6 +61,7 @@ UniqueToken, ) from ._life_cycle_registry import LayerGroupId, LifeCycleId +from ._stats import KVCacheIterationStatsDelta, KVCacheStatsDelta from ._storage import BufferId __all__ = [ @@ -106,4 +107,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 197da84e431e..bcc68b1068a9 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi @@ -57,6 +57,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) @@ -138,8 +162,9 @@ class KVCacheManagerConfig: constraints: list[BatchDesc] = ... typical_step: BatchDesc | None = None ssm_reuse_interval: int = 512 - helix_config: HelixConfig | None = None enable_swa_scratch_reuse: bool = False + enable_stats: bool = True + helix_config: HelixConfig | None = None # From _event_manager.py EventBlockHash: TypeAlias = int | str @@ -271,6 +296,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: ... @@ -399,6 +426,14 @@ class KVCacheManager: ) -> _KVCache: ... 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 diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py index 749216743c0a..b03b0e611444 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py @@ -223,6 +223,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 7550a5f0ee1c..76654d7d5139 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 @@ -54,12 +54,14 @@ BatchedLockTarget, BlockPage, CommittedPage, + Page, ScratchSlotLock, UncommittedPage, _PageHolder, _SharedPageLock, batched_lock_to_gpu, ) +from .._stats import KVCacheIterationStatsDelta, KVCacheStatsDelta from .._storage._core import Slot from .._storage_manager import StorageManager from .._utils import ( @@ -83,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 +180,8 @@ class _KVCache: "_cuda_stream", "_status", "_beam_width", + "_expected_prompt_length", + "_generation_alloc_ready", "_capacity", "_history_length", "_commit_state", @@ -192,6 +197,7 @@ class _KVCache: "_never_resumed", "_enable_swa_scratch_reuse", "_scratch_slots", + "_pending_stats", "__rawref__", ) @@ -206,6 +212,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 @@ -238,6 +246,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, @@ -247,6 +256,7 @@ def __init__( id: int | None, custom_priority_callback: Callable[[BlockOrdinal, LifeCycle], Priority], cache_salt_id: int | None, + expected_prompt_length: int | None = None, ): self.id = id self._manager = manager @@ -256,6 +266,12 @@ def __init__( self._cuda_stream = None self._status = self.Status.SUSPENDED self._beam_width = BeamIndex(1) + if expected_prompt_length is None and input_tokens is not None: + expected_prompt_length = len(input_tokens) + 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 @@ -277,9 +293,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 input_tokens is not None: self._setup_for_reuse(input_tokens) + self._refresh_generation_alloc_ready() self._avg_history_length = Average() self._avg_capacity = Average() self._avg_history_length.update(self.history_length) @@ -336,12 +354,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 @@ -509,11 +664,13 @@ def resize(self, capacity: int | None, history_length: int | None = None) -> boo f"SWA scratch requires 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 = self.manager._life_cycles.ssm_life_cycle_id beam_width = self.beam_width @@ -523,6 +680,7 @@ def resize(self, capacity: int | None, history_length: int | None = None) -> boo num_life_cycles = self.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: @@ -570,7 +728,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) @@ -639,6 +798,18 @@ def resize(self, capacity: int | None, history_length: int | None = None) -> boo for slot in normal_slots: slot.ready_event = CachedCudaEvent.NULL stream_wait_events(self.cuda_stream, new_slot_ready_events) + # 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 @@ -666,6 +837,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 @@ -836,7 +1008,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 @@ -869,7 +1041,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: @@ -911,6 +1083,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) @@ -925,6 +1100,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: @@ -1143,7 +1337,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 @@ -1226,6 +1422,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 @@ -1517,6 +1714,8 @@ def check_no_page_stale(b: tuple[Block, int]): self._committed_tokens = list(input_tokens[:num_tokens]) 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, @@ -1534,12 +1733,15 @@ def check_no_page_stale(b: tuple[Block, int]): 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, get_num_matched_tokens(matched), lc ) + full_reused_blocks = 0 + partial_reused_blocks = 0 for ordinal in chain( typed_range(stale_start), typed_range(stale_end, BlockOrdinal(len(matched))) ): @@ -1548,6 +1750,23 @@ def check_no_page_stale(b: tuple[Block, int]): # 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][0] 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 75f4acfbf4ec..f90d6da07bc9 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 @@ -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, @@ -181,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", @@ -198,6 +208,11 @@ class KVCacheManager: "_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 @@ -221,6 +236,11 @@ class KVCacheManager: _last_adjustment_time: float _last_update_num_sampled_kv_caches: int _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, @@ -254,6 +274,11 @@ def __init__( 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() @@ -371,6 +396,24 @@ def create_kv_cache( It's user responsibility to remove the last token from prompts if we need to re-compute the token generated by prefill. """ + return self._create_kv_cache( + lora_task_id, + input_tokens, + id, + custom_priority_callback, + cache_salt_id, + ) + + def _create_kv_cache( + self, + lora_task_id: int | None = None, + input_tokens: Sequence[TokenIdExt] | None = None, + id: int | None = None, + custom_priority_callback: Callable[[BlockOrdinal, LifeCycle], Priority] = lambda _, + __: PRIORITY_DEFAULT, + cache_salt_id: int | None = None, + expected_prompt_length: int | None = None, + ) -> _KVCache: return _KVCache( self, lora_task_id, @@ -378,6 +421,7 @@ def create_kv_cache( id, custom_priority_callback, cache_salt_id, + expected_prompt_length, ) def resize(self, cache_level: CacheLevel, quota: int, best_efforts: bool = False) -> bool: @@ -402,6 +446,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]: 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/_page.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_page.py index f0731cd555da..36e5e83d8f87 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 @@ -445,7 +445,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 @@ -461,13 +464,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 69037c1c72e7..30d73c340af6 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 TYPE_CHECKING, Iterator, Sequence, cast +from typing import TYPE_CHECKING, Callable, Iterator, Sequence, cast from . import rawref from ._common import ( @@ -162,6 +162,9 @@ def unavailable(self) -> int: return self.total - self.available +MigrationRecorder = Callable[[Sequence[Page], Sequence[Slot], CacheLevel, CacheLevel], None] + + class StorageManager: __slots__ = ( "_life_cycles", @@ -278,12 +281,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) @@ -294,7 +302,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) @@ -314,13 +322,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) @@ -368,13 +380,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] @@ -398,6 +413,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( @@ -469,7 +485,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): @@ -480,7 +501,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 @@ -494,6 +522,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." @@ -536,6 +565,8 @@ def _batched_migrate( 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 = ( 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 799520206298..f2f5d9e31c12 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/setup_mypyc.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/setup_mypyc.py @@ -86,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/integration/test_lists/test-db/l0_h100.yml b/tests/integration/test_lists/test-db/l0_h100.yml index 3e552fd6e92a..516fe31bc985 100644 --- a/tests/integration/test_lists/test-db/l0_h100.yml +++ b/tests/integration/test_lists/test-db/l0_h100.yml @@ -456,6 +456,7 @@ l0_h100: - unittest/kv_cache_manager_v2_tests/ # 4 min # ------------- KV Cache Iteration Stats --------------- - unittest/executor/test_stats_serializer.py + - unittest/disaggregated/test_kv_cache_transfer_perf_metrics.py - unittest/metrics/test_collector.py - kv_cache/test_kv_cache_iteration_stats.py::TestKvCacheIterationStats::test_cold_start - kv_cache/test_kv_cache_iteration_stats.py::TestKvCacheIterationStats::test_partial_block_reuse 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_kv_cache_transfer_perf_metrics.py b/tests/unittest/disaggregated/test_kv_cache_transfer_perf_metrics.py new file mode 100644 index 000000000000..d7266736e052 --- /dev/null +++ b/tests/unittest/disaggregated/test_kv_cache_transfer_perf_metrics.py @@ -0,0 +1,244 @@ +# 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 datetime import timedelta + +import numpy as np + +import tensorrt_llm._torch.disaggregation.transceiver as transceiver_module +from tensorrt_llm._torch.disaggregation.base.transfer import KVSlice, SessionStatus, WaitResult +from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 +from tensorrt_llm.bindings import LlmRequestState + + +class _Request: + def __init__(self, request_id: int = 0) -> None: + self.request_id = request_id + self.py_disaggregated_params = None + self.state = None + self.context_phase_params = None + self.kv_cache_transfer_start = None + self.kv_cache_transfer_end = None + self.kv_cache_size = None + + def set_kv_cache_transfer_start(self, time: timedelta) -> None: + self.kv_cache_transfer_start = time + + def set_kv_cache_transfer_end(self, time: timedelta) -> None: + self.kv_cache_transfer_end = time + + def set_kv_cache_size(self, size: int) -> None: + self.kv_cache_size = size + + +class _TxSession: + def __init__(self, rid: int, transferred_kv_bytes: int = 0) -> None: + self.disagg_request_id = rid + self.transferred_kv_bytes = transferred_kv_bytes + self.sent_slices = [] + self.closed = False + + def send(self, kv_slice: KVSlice) -> None: + self.sent_slices.append(kv_slice) + + def wait_complete(self) -> WaitResult: + return WaitResult.COMPLETED + + @property + def status(self) -> SessionStatus: + return SessionStatus.FULLY_TRANSFERRED + + def is_completed(self) -> bool: + return True + + def has_failed(self) -> bool: + return False + + def close(self) -> None: + self.closed = True + + +class _TransferWorker: + def __init__(self, session: _TxSession) -> None: + self.session = session + + def create_tx_session(self, _request: _Request) -> _TxSession: + return self.session + + def sweep_stale_req_infos(self) -> None: + pass + + +class _Dist: + tp_size = 1 + + def tp_allgather(self, value): + return [value] + + def pp_allgather(self, value): + return [value] + + +def _make_transceiver() -> KvCacheTransceiverV2: + transceiver = KvCacheTransceiverV2.__new__(KvCacheTransceiverV2) + transceiver._gen_need_sync = False + transceiver._gen_allgather = lambda metrics: [metrics] + transceiver._dist = _Dist() + transceiver._ctx_need_tp_sync = False + transceiver._ctx_need_pp_sync = False + transceiver._recv_reqs = {} + transceiver._clock_offset_seconds = lambda: 0.0 + transceiver._ctx_consensus = lambda ids: ids + return transceiver + + +def test_record_transfer_start_populates_request_perf_metrics(monkeypatch) -> None: + transceiver = _make_transceiver() + transceiver._clock_offset_seconds = lambda: 0.25 + request = _Request() + monkeypatch.setattr( + transceiver_module, + "get_steady_clock_now_in_seconds", + lambda: 100.0, + ) + + transceiver._record_transfer_start(request) + + assert request.kv_cache_transfer_start == timedelta(seconds=100.0) + assert request.kv_cache_size == 0 + + +def test_record_transfer_end_uses_actual_session_bytes(monkeypatch) -> None: + transceiver = _make_transceiver() + request = _Request() + session = _TxSession(rid=0, transferred_kv_bytes=123) + monkeypatch.setattr( + transceiver_module, + "get_steady_clock_now_in_seconds", + lambda: 12.5, + ) + + transceiver._record_transfer_end(request, session) + + assert request.kv_cache_transfer_end == timedelta(seconds=12.5) + assert request.kv_cache_size == 123 + + +def test_get_transfer_metrics_reads_request_perf_metrics() -> None: + request = _Request() + request.kv_cache_transfer_start = timedelta(seconds=10.0) + request.kv_cache_transfer_end = timedelta(seconds=14.0) + request.kv_cache_size = 64 + + assert KvCacheTransceiverV2._get_transfer_metrics(request) == (10.0, 14.0, 64) + + +def test_publish_gen_transfer_metrics_aggregates_completed_consensus_rids( + monkeypatch, +) -> None: + monkeypatch.setenv("TRTLLM_KVCACHE_TIME_OUTPUT_PATH", "/tmp/kvcache") + transceiver = _make_transceiver() + transceiver._gen_need_sync = True + transceiver._clock_offset_seconds = lambda: 2.0 + request = _Request() + request.kv_cache_transfer_start = timedelta(seconds=10.0) + request.kv_cache_transfer_end = timedelta(seconds=14.0) + request.kv_cache_size = 64 + transceiver._recv_reqs = {5: request} + transceiver._gen_allgather = lambda metrics: [ + metrics, + {5: (9.0, 16.0, 128)}, + ] + + transceiver._publish_gen_transfer_metrics([5], {5}) + + assert request.kv_cache_transfer_start == timedelta(seconds=7.0) + assert request.kv_cache_transfer_end == timedelta(seconds=14.0) + assert request.kv_cache_size == 192 + + +def test_publish_gen_transfer_metrics_respects_v1_env_gate(monkeypatch) -> None: + monkeypatch.delenv("TRTLLM_KVCACHE_TIME_OUTPUT_PATH", raising=False) + transceiver = _make_transceiver() + transceiver._gen_need_sync = True + request = _Request() + transceiver._recv_reqs = {5: request} + + def fail_allgather(_metrics): + raise AssertionError("aggregation should be gated by TRTLLM_KVCACHE_TIME_OUTPUT_PATH") + + transceiver._gen_allgather = fail_allgather + + transceiver._publish_gen_transfer_metrics([5], {5}) + + assert request.kv_cache_transfer_start is None + assert request.kv_cache_transfer_end is None + assert request.kv_cache_size is None + + +def test_publish_gen_transfer_metrics_skips_non_consensus_rids(monkeypatch) -> None: + monkeypatch.setenv("TRTLLM_KVCACHE_TIME_OUTPUT_PATH", "/tmp/kvcache") + transceiver = _make_transceiver() + transceiver._gen_need_sync = True + request = _Request() + transceiver._recv_reqs = {5: request} + + def fail_allgather(_metrics): + raise AssertionError("non-consensus request should not allgather") + + transceiver._gen_allgather = fail_allgather + + transceiver._publish_gen_transfer_metrics([5], set()) + + assert request.kv_cache_transfer_start is None + assert request.kv_cache_transfer_end is None + assert request.kv_cache_size is None + + +def test_context_transfer_metrics_cover_send_lifecycle(monkeypatch) -> None: + transceiver = _make_transceiver() + kv_slice = KVSlice( + block_ids_per_layer_groups=[ + np.array([0, 1], dtype=np.int64), + np.array([], dtype=np.int64), + ] + ) + session = _TxSession(rid=7, transferred_kv_bytes=321) + transceiver._transfer_worker = _TransferWorker(session) + transceiver._send_sessions = {} + transceiver._send_reqs = {} + transceiver._dp_rank = 1 + transceiver._context_info_endpoint = "tcp://ctx" + transceiver._create_kv_slice = lambda _request: kv_slice + now = iter([10.0, 12.5]) + monkeypatch.setattr( + transceiver_module, + "get_steady_clock_now_in_seconds", + lambda: next(now), + ) + request = _Request(request_id=7) + + transceiver.respond_and_send_async(request) + completed, failed = transceiver.check_context_transfer_status(1, mark_complete=True) + + assert session.sent_slices[0] is kv_slice + assert completed == [7] + assert failed == [] + assert request.kv_cache_transfer_start == timedelta(seconds=10.0) + assert request.kv_cache_transfer_end == timedelta(seconds=12.5) + assert request.kv_cache_size == 321 + assert request.state == LlmRequestState.DISAGG_CONTEXT_COMPLETE + assert request.context_phase_params.req_id == 7 + assert session.closed diff --git a/tests/unittest/executor/test_stats_serializer.py b/tests/unittest/executor/test_stats_serializer.py index 0c00ca9322f7..5c6c8f312ca4 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 @@ -196,3 +201,80 @@ def test_serializer_legacy_2_tuple(self): result = BaseWorker._stats_serializer((iter_stats, None)) d = json.loads(result) assert "kvCacheIterationStats" not in d + + 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_stats_behavior.py b/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_stats_behavior.py new file mode 100644 index 000000000000..efd778227aef --- /dev/null +++ b/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_stats_behavior.py @@ -0,0 +1,646 @@ +# 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.llm_request import LlmRequest +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager as KVCacheManagerV1 +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManagerV2 +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_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/metrics/test_collector.py b/tests/unittest/metrics/test_collector.py index c90f94545d07..f5e41b2968db 100644 --- a/tests/unittest/metrics/test_collector.py +++ b/tests/unittest/metrics/test_collector.py @@ -802,6 +802,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()