diff --git a/cpp/include/tensorrt_llm/batch_manager/kvCacheManager.h b/cpp/include/tensorrt_llm/batch_manager/kvCacheManager.h index 34c857980aaf..299bf1b570f4 100644 --- a/cpp/include/tensorrt_llm/batch_manager/kvCacheManager.h +++ b/cpp/include/tensorrt_llm/batch_manager/kvCacheManager.h @@ -327,8 +327,8 @@ struct KvCacheStats std::size_t allocatedBytes{}; }; -/// @brief Per-iteration KV cache statistics. All delta counters represent changes since the last call to -/// getIterationStats(). Gauges are instantaneous snapshots. +/// @brief Per-iteration KV cache statistics. All delta counters and peak gauges represent values since the last call +/// to getIterationStats(). Snapshot gauges are instantaneous. struct KvCacheIterationStats { // --- Instantaneous gauges --- @@ -336,10 +336,23 @@ struct KvCacheIterationStats SizeType32 primaryMaxNumBlocks{0}; SizeType32 primaryFreeNumBlocks{0}; SizeType32 primaryUsedNumBlocks{0}; + // Cached-but-unpinned blocks in the primary pool. Distinct from primaryUsedNumBlocks, + // which also counts blocks pinned during onboard memcpy windows. + SizeType32 primaryEvictableNumBlocks{0}; + SizeType32 primaryPeakFreeNumBlocks{0}; + SizeType32 primaryPeakUsedNumBlocks{0}; + SizeType32 primaryPeakEvictableNumBlocks{0}; // Secondary (host) pool SizeType32 secondaryMaxNumBlocks{0}; SizeType32 secondaryFreeNumBlocks{0}; SizeType32 secondaryUsedNumBlocks{0}; + // Cached-but-unpinned blocks in the secondary pool. Useful to gauge "how full is the + // host cache"; secondaryUsedNumBlocks only counts pinned blocks during the sub-ms + // onboard memcpy window so it cannot answer that question on its own. + SizeType32 secondaryEvictableNumBlocks{0}; + SizeType32 secondaryPeakFreeNumBlocks{0}; + SizeType32 secondaryPeakUsedNumBlocks{0}; + SizeType32 secondaryPeakEvictableNumBlocks{0}; // --- Per-iteration deltas (reset on each read) --- // Context phase: block allocation and reuse @@ -361,6 +374,11 @@ struct KvCacheIterationStats // Intra-device (GPU → GPU) block copies (e.g. partial reuse when source block has refs) SizeType32 iterIntraDeviceCopyBlocks{0}; std::size_t iterIntraDeviceCopyBytes{0}; + + // Pages released by LRU from the last cache tier without ever being onboarded back + // to GPU during their stay at that tier (i.e. fully dropped from the hierarchy). + SizeType32 iterHostDroppedBlocks{0}; + std::size_t iterHostDroppedBytes{0}; }; // Basic building block of a paged KV cache - a single diff --git a/cpp/include/tensorrt_llm/executor/executor.h b/cpp/include/tensorrt_llm/executor/executor.h index 3017737d2e37..825b8ad75959 100644 --- a/cpp/include/tensorrt_llm/executor/executor.h +++ b/cpp/include/tensorrt_llm/executor/executor.h @@ -1556,6 +1556,9 @@ class ExecutorConfig // Per request stats may have additional overhead due to going through all requests. Turned off by default. static constexpr SizeType32 kDefaultRequestStatsMaxIterations = 0; + // A value of -1 keeps all iteration/request stats until they are fetched. + static constexpr SizeType32 kUnlimitedStatsMaxIterations = -1; + explicit ExecutorConfig(SizeType32 maxBeamWidth = 1, SchedulerConfig schedulerConfig = SchedulerConfig(), KvCacheConfig kvCacheConfig = KvCacheConfig(), bool enableChunkedContext = true, bool normalizeLogProbs = false, SizeType32 iterStatsMaxIterations = kDefaultIterStatsMaxIterations, @@ -1661,9 +1664,11 @@ class ExecutorConfig bool mNormalizeLogProbs; /// @brief Controls the maximum number of iterations for which to keep statistics. + /// Set to -1 to keep all iteration statistics. Set to 0 to disable iteration statistics. SizeType32 mIterStatsMaxIterations; /// @brief Controls the maximum number of iterations for which to keep per-request statistics. + /// Set to -1 to keep all per-request statistics. Set to 0 to disable per-request statistics. SizeType32 mRequestStatsMaxIterations; /// @brief The type of batching strategy to use. See BatchingType. @@ -1944,12 +1949,12 @@ class Executor void shutdown(); /// @brief Returns the per-iterations statistics computed since last call to getLatestIterationStats. - /// Contains at most iterStatsMaxIterations iterations. + /// Contains at most iterStatsMaxIterations iterations, or all iterations when set to -1. /// @return Iteration stats std::deque getLatestIterationStats(); /// @brief Returns the request stats of each iteration computed since last call to getLatestRequestStats. - /// Contains at most requestStatsMaxIterations iterations. + /// Contains at most requestStatsMaxIterations iterations, or all iterations when set to -1. /// @return Request stats grouped by iterations std::deque getLatestRequestStats(); diff --git a/cpp/tensorrt_llm/executor/cacheTransceiverConfig.cpp b/cpp/tensorrt_llm/executor/cacheTransceiverConfig.cpp index 2d77299d3a24..45d994a46d21 100644 --- a/cpp/tensorrt_llm/executor/cacheTransceiverConfig.cpp +++ b/cpp/tensorrt_llm/executor/cacheTransceiverConfig.cpp @@ -28,8 +28,8 @@ CacheTransceiverConfig::CacheTransceiverConfig(std::optional backen , mMaxTokensInBuffer(maxNumTokens) , mKvTransferTimeoutMs(kvTransferTimeoutMs) , mKvTransferSenderFutureTimeoutMs(kvTransferSenderFutureTimeoutMs) - , mKvTransferPollIntervalMs(kvTransferPollIntervalMs) { + setKvTransferPollIntervalMs(kvTransferPollIntervalMs); } bool CacheTransceiverConfig::operator==(CacheTransceiverConfig const& other) const diff --git a/cpp/tensorrt_llm/executor/executorConfig.cpp b/cpp/tensorrt_llm/executor/executorConfig.cpp index 2dff78280f5a..25602a1a83d6 100644 --- a/cpp/tensorrt_llm/executor/executorConfig.cpp +++ b/cpp/tensorrt_llm/executor/executorConfig.cpp @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); @@ -65,8 +65,8 @@ ExecutorConfig::ExecutorConfig(SizeType32 maxBeamWidth, SchedulerConfig schedule , mEnableTrtOverlap(enableTrtOverlap) , mFailFastOnAttentionWindowTooLarge(failFastOnAttentionWindowTooLarge) { - TLLM_CHECK(iterStatsMaxIterations >= 0); - TLLM_CHECK(requestStatsMaxIterations >= 0); + TLLM_CHECK(iterStatsMaxIterations >= kUnlimitedStatsMaxIterations); + TLLM_CHECK(requestStatsMaxIterations >= kUnlimitedStatsMaxIterations); TLLM_CHECK(mMaxBeamWidth > 0); TLLM_CHECK(maxSeqIdleMicroseconds > 0); } @@ -271,13 +271,13 @@ void ExecutorConfig::setNormalizeLogProbs(bool normalizeLogProbs) void ExecutorConfig::setIterStatsMaxIterations(SizeType32 iterStatsMaxIterations) { mIterStatsMaxIterations = iterStatsMaxIterations; - TLLM_CHECK(mIterStatsMaxIterations >= 0); + TLLM_CHECK(mIterStatsMaxIterations >= kUnlimitedStatsMaxIterations); } void ExecutorConfig::setRequestStatsMaxIterations(SizeType32 requestStatsMaxIterations) { mRequestStatsMaxIterations = requestStatsMaxIterations; - TLLM_CHECK(mRequestStatsMaxIterations >= 0); + TLLM_CHECK(mRequestStatsMaxIterations >= kUnlimitedStatsMaxIterations); } void ExecutorConfig::setBatchingType(BatchingType batchingType) diff --git a/cpp/tensorrt_llm/executor/executorImpl.cpp b/cpp/tensorrt_llm/executor/executorImpl.cpp index 2fb20b7ea572..9f7fb654a2d5 100644 --- a/cpp/tensorrt_llm/executor/executorImpl.cpp +++ b/cpp/tensorrt_llm/executor/executorImpl.cpp @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); @@ -86,6 +86,16 @@ namespace return fixedExecutorConfig; } +[[nodiscard]] bool statsBufferIsEnabled(SizeType32 maxIterations) +{ + return maxIterations != 0; +} + +[[nodiscard]] bool statsBufferIsBounded(SizeType32 maxIterations) +{ + return maxIterations > 0; +} + SizeType32 getNumChildRequests(Request const& request) { auto samplingConfig = request.getSamplingConfig(); @@ -1908,9 +1918,13 @@ RequestStatsPerIteration Executor::Impl::getCurrentRequestStats( void Executor::Impl::appendCurrentIterStats(IterationStats&& currentIterStats) { std::scoped_lock lck(mIterStatsMtx); - if (mIterationStats.size() >= mIterStatsMaxIterations) + if (statsBufferIsBounded(mIterStatsMaxIterations)) { - mIterationStats.pop_front(); + auto const maxIterStats = static_cast(mIterStatsMaxIterations); + if (mIterationStats.size() >= maxIterStats) + { + mIterationStats.pop_front(); + } } mIterationStats.emplace_back(std::move(currentIterStats)); } @@ -1918,16 +1932,16 @@ void Executor::Impl::appendCurrentIterStats(IterationStats&& currentIterStats) void Executor::Impl::appendMultipleIterStats(std::vector&& currentIterStatsVec) { std::scoped_lock lck(mIterStatsMtx); - if (mIterationStats.size() + currentIterStatsVec.size() > mIterStatsMaxIterations) + mIterationStats.insert(mIterationStats.end(), std::make_move_iterator(currentIterStatsVec.begin()), + std::make_move_iterator(currentIterStatsVec.end())); + if (statsBufferIsBounded(mIterStatsMaxIterations)) { - size_t removeCount = mIterationStats.size() + currentIterStatsVec.size() - mIterStatsMaxIterations; - for (size_t i = 0; i < removeCount; i++) + auto const maxIterStats = static_cast(mIterStatsMaxIterations); + while (mIterationStats.size() > maxIterStats) { mIterationStats.pop_front(); } } - mIterationStats.insert(mIterationStats.end(), std::make_move_iterator(currentIterStatsVec.begin()), - std::make_move_iterator(currentIterStatsVec.end())); } void Executor::Impl::updateIterationStats(RequestList const& activeRequests, double iterLatencyMS, @@ -1935,7 +1949,7 @@ void Executor::Impl::updateIterationStats(RequestList const& activeRequests, dou bool flushToOrchestrator) { NVTX3_SCOPED_RANGE(updateIterationStats); - if (mIterStatsMaxIterations > 0 && mIsLeader) + if (statsBufferIsEnabled(mIterStatsMaxIterations) && mIsLeader) { auto currentIterStats = getCurrentIterationStats( activeRequests, iterLatencyMS, numNewActiveRequests, newActiveRequestsQueueLatencyMS, numCompletedRequests); @@ -1972,9 +1986,13 @@ void Executor::Impl::updateIterationStats(RequestList const& activeRequests, dou void Executor::Impl::appendCurrentRequestStats(RequestStatsPerIteration&& currentRequestStats) { std::scoped_lock lck(mRequestStatsMtx); - if (mRequestStats.size() >= mRequestStatsMaxIterations) + if (statsBufferIsBounded(mRequestStatsMaxIterations)) { - mRequestStats.pop_front(); + auto const maxRequestStats = static_cast(mRequestStatsMaxIterations); + if (mRequestStats.size() >= maxRequestStats) + { + mRequestStats.pop_front(); + } } mRequestStats.emplace_back(std::move(currentRequestStats)); } @@ -1982,23 +2000,23 @@ void Executor::Impl::appendCurrentRequestStats(RequestStatsPerIteration&& curren void Executor::Impl::appendMultipleRequestStats(std::vector&& currentRequestStatsVec) { std::scoped_lock lck(mRequestStatsMtx); - if (mRequestStats.size() + currentRequestStatsVec.size() > mRequestStatsMaxIterations) + mRequestStats.insert(mRequestStats.end(), std::make_move_iterator(currentRequestStatsVec.begin()), + std::make_move_iterator(currentRequestStatsVec.end())); + if (statsBufferIsBounded(mRequestStatsMaxIterations)) { - size_t removeCount = mRequestStats.size() + currentRequestStatsVec.size() - mRequestStatsMaxIterations; - for (size_t i = 0; i < removeCount; i++) + auto const maxRequestStats = static_cast(mRequestStatsMaxIterations); + while (mRequestStats.size() > maxRequestStats) { mRequestStats.pop_front(); } } - mRequestStats.insert(mRequestStats.end(), std::make_move_iterator(currentRequestStatsVec.begin()), - std::make_move_iterator(currentRequestStatsVec.end())); } void Executor::Impl::updateRequestStats( RequestList const& activeRequests, RequestList const& finishedRequests, bool flushToOrchestrator) { NVTX3_SCOPED_RANGE(updateRequestStats); - if (mRequestStatsMaxIterations > 0 && mIsLeader) + if (statsBufferIsEnabled(mRequestStatsMaxIterations) && mIsLeader) { // Add current iteration request stats auto currentRequestStats = getCurrentRequestStats(activeRequests, finishedRequests); diff --git a/cpp/tensorrt_llm/executor/executorImpl.h b/cpp/tensorrt_llm/executor/executorImpl.h index 6e545dbf6def..f812b55a3fa0 100644 --- a/cpp/tensorrt_llm/executor/executorImpl.h +++ b/cpp/tensorrt_llm/executor/executorImpl.h @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); @@ -310,12 +310,12 @@ class Executor::Impl std::unordered_map> mChildReqIdsMap; // Iteration stats - IterationType mIterStatsMaxIterations; + SizeType32 mIterStatsMaxIterations; std::mutex mIterStatsMtx; std::deque mIterationStats; // Request stats - IterationType mRequestStatsMaxIterations; + SizeType32 mRequestStatsMaxIterations; std::mutex mRequestStatsMtx; std::deque mRequestStats; diff --git a/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManager.cpp b/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManager.cpp index 12b29d4981e2..1496ebbb883d 100644 --- a/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManager.cpp +++ b/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManager.cpp @@ -379,9 +379,17 @@ void tb::kv_cache_manager::KVCacheManagerBindings::initBindings(nb::module_& m) .def_rw("primary_max_num_blocks", &tbk::KvCacheIterationStats::primaryMaxNumBlocks) .def_rw("primary_free_num_blocks", &tbk::KvCacheIterationStats::primaryFreeNumBlocks) .def_rw("primary_used_num_blocks", &tbk::KvCacheIterationStats::primaryUsedNumBlocks) + .def_rw("primary_evictable_num_blocks", &tbk::KvCacheIterationStats::primaryEvictableNumBlocks) + .def_rw("primary_peak_free_num_blocks", &tbk::KvCacheIterationStats::primaryPeakFreeNumBlocks) + .def_rw("primary_peak_used_num_blocks", &tbk::KvCacheIterationStats::primaryPeakUsedNumBlocks) + .def_rw("primary_peak_evictable_num_blocks", &tbk::KvCacheIterationStats::primaryPeakEvictableNumBlocks) .def_rw("secondary_max_num_blocks", &tbk::KvCacheIterationStats::secondaryMaxNumBlocks) .def_rw("secondary_free_num_blocks", &tbk::KvCacheIterationStats::secondaryFreeNumBlocks) .def_rw("secondary_used_num_blocks", &tbk::KvCacheIterationStats::secondaryUsedNumBlocks) + .def_rw("secondary_evictable_num_blocks", &tbk::KvCacheIterationStats::secondaryEvictableNumBlocks) + .def_rw("secondary_peak_free_num_blocks", &tbk::KvCacheIterationStats::secondaryPeakFreeNumBlocks) + .def_rw("secondary_peak_used_num_blocks", &tbk::KvCacheIterationStats::secondaryPeakUsedNumBlocks) + .def_rw("secondary_peak_evictable_num_blocks", &tbk::KvCacheIterationStats::secondaryPeakEvictableNumBlocks) .def_rw("iter_alloc_total_blocks", &tbk::KvCacheIterationStats::iterAllocTotalBlocks) .def_rw("iter_alloc_new_blocks", &tbk::KvCacheIterationStats::iterAllocNewBlocks) .def_rw("iter_reused_blocks", &tbk::KvCacheIterationStats::iterReusedBlocks) @@ -395,7 +403,9 @@ void tb::kv_cache_manager::KVCacheManagerBindings::initBindings(nb::module_& m) .def_rw("iter_offload_blocks", &tbk::KvCacheIterationStats::iterOffloadBlocks) .def_rw("iter_offload_bytes", &tbk::KvCacheIterationStats::iterOffloadBytes) .def_rw("iter_intra_device_copy_blocks", &tbk::KvCacheIterationStats::iterIntraDeviceCopyBlocks) - .def_rw("iter_intra_device_copy_bytes", &tbk::KvCacheIterationStats::iterIntraDeviceCopyBytes); + .def_rw("iter_intra_device_copy_bytes", &tbk::KvCacheIterationStats::iterIntraDeviceCopyBytes) + .def_rw("iter_host_dropped_blocks", &tbk::KvCacheIterationStats::iterHostDroppedBlocks) + .def_rw("iter_host_dropped_bytes", &tbk::KvCacheIterationStats::iterHostDroppedBytes); nb::class_(m, "BlockKey") .def(nb::init<>()) diff --git a/cpp/tensorrt_llm/nanobind/executor/request.cpp b/cpp/tensorrt_llm/nanobind/executor/request.cpp index 3eca465cb544..a4effbd2e942 100644 --- a/cpp/tensorrt_llm/nanobind/executor/request.cpp +++ b/cpp/tensorrt_llm/nanobind/executor/request.cpp @@ -26,6 +26,7 @@ #include "tensorrt_llm/runtime/cudaStream.h" #include +#include #include #include #include @@ -36,6 +37,7 @@ #include #include +#include #include #include @@ -604,7 +606,15 @@ void initRequestBindings(nb::module_& m) auto requestGetstate = [](tle::Request const& self) { - return nb::make_tuple(self.getInputTokenIds(), self.getMaxTokens(), self.getStreaming(), + // Serialize input_token_ids as a raw int32 byte buffer instead of a Python + // list[int]: nanobind casts VecTokens element-by-element (a PyLong storm, + // ISL-proportional) on every Request pickle -- the request broadcast and the + // RPC IPC submit. A bytes blob is a memcpy, ~order-of-magnitude cheaper. + // Paired with requestSetstate. + auto const& inputTokenIds = self.getInputTokenIds(); + auto inputTokenIdsBytes = nb::bytes( + reinterpret_cast(inputTokenIds.data()), inputTokenIds.size() * sizeof(VecTokens::value_type)); + return nb::make_tuple(std::move(inputTokenIdsBytes), self.getMaxTokens(), self.getStreaming(), self.getSamplingConfig(), self.getOutputConfig(), self.getEndId(), self.getPadId(), self.getPositionIds(), self.getBadWords(), self.getStopWords(), self.getEmbeddingBias(), self.getExternalDraftTokensConfig(), self.getPromptTuningConfig(), self.getMultimodalInput(), self.getMultimodalEmbedding(), @@ -621,8 +631,21 @@ void initRequestBindings(nb::module_& m) { throw std::runtime_error("Invalid Request state!"); } - new (&self) tle::Request(nb::cast(state[0]), nb::cast(state[1]), - nb::cast(state[2]), nb::cast(state[3]), nb::cast(state[4]), + // input_token_ids is a raw int32 byte buffer (see requestGetstate). + auto const inputTokenIdsBytes = nb::cast(state[0]); + auto constexpr kTokenByteSize = sizeof(VecTokens::value_type); + auto const inputTokenIdsByteSize = inputTokenIdsBytes.size(); + if (inputTokenIdsByteSize % kTokenByteSize != 0) + { + throw std::runtime_error("Invalid Request state: input_token_ids byte buffer has invalid size!"); + } + VecTokens inputTokenIds(inputTokenIdsByteSize / kTokenByteSize); + if (inputTokenIdsByteSize > 0) + { + std::memcpy(inputTokenIds.data(), inputTokenIdsBytes.c_str(), inputTokenIdsByteSize); + } + new (&self) tle::Request(std::move(inputTokenIds), nb::cast(state[1]), nb::cast(state[2]), + nb::cast(state[3]), nb::cast(state[4]), nb::cast>(state[5]), nb::cast>(state[6]), nb::cast>>(state[7]), nb::cast>>(state[8]), @@ -645,47 +668,69 @@ void initRequestBindings(nb::module_& m) nb::cast>(state[33]), nb::cast>(state[34])); }; + // Convert input_token_ids to VecTokens. Fast path: a 1-D contiguous int32 + // ndarray is memcpy'd into the vector (no per-element PyLong cast, which is + // O(ISL) on the GIL-held submit path). Anything else (list[int], etc.) falls + // back to the default nanobind sequence cast, so this is fully back-compatible. + // This complements PR #15134 (which bytes-encodes Request *pickling*); here we + // target Request *construction* on the RpcWorker.submit / _enqueue_request path. + auto toVecTokens = [](nb::handle ids) -> tle::VecTokens + { + nb::ndarray, nb::c_contig> arr; + if (nb::try_cast(ids, arr, /*convert=*/false)) + { + tle::VecTokens out(arr.shape(0)); + if (arr.shape(0) > 0) + { + std::memcpy(out.data(), arr.data(), arr.shape(0) * sizeof(int32_t)); + } + return out; + } + return nb::cast(ids); + }; + nb::class_ request(m, "Request", nb::dynamic_attr()); request - .def(nb::init const&, // endId - std::optional const&, // padId - std::optional>, // positionIds - std::optional>, // badWords - std::optional>, // stopWords - std::optional, // embeddingBias - std::optional, // externalDraftTokensConfig - std::optional, // pTuningConfig - std::optional, // multimodalInput - std::optional, // multimodalEmbedding - std::optional, // mRopeConfig - std::optional, // loraConfig - std::optional, // lookaheadConfig - std::optional, // kvCacheRetentionConfig - std::optional, // logitsPostProcessorName - std::optional, // logitsPostProcessor - std::optional, // encoderInputTokenIds - std::optional, // clientId - bool, // returnAllGeneratedTokens - tle::PriorityType, // priority - tle::RequestType, // type - std::optional, // contextPhaseParams - std::optional, // encoderInputFeatures - std::optional, // encoderOutputLength - std::optional, // crossAttentionMask - SizeType32, // numReturnSequences - std::optional, // eagleConfig - std::optional, // skipCrossAttnBlocks - std::optional, // guidedDecodingParams - std::optional, // languageAdapterUid - std::optional, // allottedTimeMs - std::optional, // disaggRequestId - std::optional // cacheSalt - >(), + .def( + "__init__", + [toVecTokens](tle::Request* self, + nb::handle input_token_ids, // list[int] or int32 ndarray + tle::SizeType32 max_tokens, bool streaming, tle::SamplingConfig const& sampling_config, + tle::OutputConfig const& output_config, std::optional const& end_id, + std::optional const& pad_id, std::optional> position_ids, + std::optional> bad_words, std::optional> stop_words, + std::optional embedding_bias, + std::optional external_draft_tokens_config, + std::optional prompt_tuning_config, + std::optional multimodal_input, std::optional multimodal_embedding, + std::optional mrope_config, std::optional lora_config, + std::optional lookahead_config, + std::optional kv_cache_retention_config, + std::optional logits_post_processor_name, + std::optional logits_post_processor, + std::optional encoder_input_token_ids, std::optional client_id, + bool return_all_generated_tokens, tle::PriorityType priority, tle::RequestType type, + std::optional context_phase_params, + std::optional encoder_input_features, std::optional encoder_output_length, + std::optional cross_attention_mask, SizeType32 num_return_sequences, + std::optional eagle_config, std::optional skip_cross_attn_blocks, + std::optional guided_decoding_params, + std::optional language_adapter_uid, + std::optional allotted_time_ms, std::optional disagg_request_id, + std::optional cache_salt) + { + new (self) tle::Request(toVecTokens(input_token_ids), max_tokens, streaming, sampling_config, + output_config, end_id, pad_id, std::move(position_ids), std::move(bad_words), std::move(stop_words), + std::move(embedding_bias), std::move(external_draft_tokens_config), std::move(prompt_tuning_config), + std::move(multimodal_input), std::move(multimodal_embedding), std::move(mrope_config), + std::move(lora_config), std::move(lookahead_config), std::move(kv_cache_retention_config), + std::move(logits_post_processor_name), std::move(logits_post_processor), + std::move(encoder_input_token_ids), client_id, return_all_generated_tokens, priority, type, + std::move(context_phase_params), std::move(encoder_input_features), encoder_output_length, + std::move(cross_attention_mask), num_return_sequences, std::move(eagle_config), + std::move(skip_cross_attn_blocks), std::move(guided_decoding_params), language_adapter_uid, + allotted_time_ms, disagg_request_id, std::move(cache_salt)); + }, // clang-format off nb::arg("input_token_ids"), nb::arg("max_tokens"), @@ -726,7 +771,7 @@ void initRequestBindings(nb::module_& m) nb::arg("allotted_time_ms") = nb::none(), nb::arg("disagg_request_id") = nb::none(), nb::arg("cache_salt") = nb::none() - ) // clang-format on + ) // clang-format on .def_prop_ro("input_token_ids", &tle::Request::getInputTokenIds) .def_prop_ro("num_input_tokens", &tle::Request::getNumInputTokens) .def_prop_ro("max_tokens", &tle::Request::getMaxTokens) diff --git a/cpp/tests/unit_tests/executor/executorConfigTest.cpp b/cpp/tests/unit_tests/executor/executorConfigTest.cpp index f4bd9959a253..1378e15fcd7f 100644 --- a/cpp/tests/unit_tests/executor/executorConfigTest.cpp +++ b/cpp/tests/unit_tests/executor/executorConfigTest.cpp @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2023-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); @@ -27,6 +27,17 @@ using ::testing::Invoke; using namespace tensorrt_llm::executor; using namespace tensorrt_llm::common; +TEST(CacheTransceiverConfigTest, validatesKvTransferPollInterval) +{ + auto makeConfig = [](std::optional pollIntervalMs) + { return CacheTransceiverConfig(std::nullopt, std::nullopt, std::nullopt, std::nullopt, pollIntervalMs); }; + + EXPECT_EQ(makeConfig(std::nullopt).getKvTransferPollIntervalMs(), std::nullopt); + EXPECT_EQ(makeConfig(1).getKvTransferPollIntervalMs(), 1); + EXPECT_THROW(makeConfig(0), TllmException); + EXPECT_THROW(makeConfig(-1), TllmException); +} + TEST(ExecutorConfigTest, ctorValidInputs) { SchedulerConfig schedulerConfig; @@ -37,6 +48,12 @@ TEST(ExecutorConfigTest, ctorValidInputs) { auto executorConfig = ExecutorConfig(1, schedulerConfig, kvCacheConfig, true, true, 0); } + { + auto executorConfig = ExecutorConfig(1, schedulerConfig, kvCacheConfig, true, true, + ExecutorConfig::kUnlimitedStatsMaxIterations, ExecutorConfig::kUnlimitedStatsMaxIterations); + EXPECT_EQ(executorConfig.getIterStatsMaxIterations(), ExecutorConfig::kUnlimitedStatsMaxIterations); + EXPECT_EQ(executorConfig.getRequestStatsMaxIterations(), ExecutorConfig::kUnlimitedStatsMaxIterations); + } { auto executorConfig = ExecutorConfig(1, schedulerConfig, kvCacheConfig, true, true, 1000); } @@ -87,9 +104,24 @@ TEST(ExecutorConfigTest, ctorInvalidInputs) FAIL() << "Expected TllmException"; } - // iter stats negative + // iter stats below the unlimited sentinel ParallelConfig parallelConfigValid; - testInvalid(1, schedulerConfig, kvCacheConfig, true, true, -1, BatchingType::kINFLIGHT, parallelConfigValid); + testInvalid(1, schedulerConfig, kvCacheConfig, true, true, -2, BatchingType::kINFLIGHT, parallelConfigValid); + + // request stats below the unlimited sentinel + try + { + auto executorConfig = ExecutorConfig(1, schedulerConfig, kvCacheConfig, true, true, 1000, -2); + FAIL() << "Expected TllmException"; + } + catch (TllmException& e) + { + EXPECT_THAT(e.what(), testing::HasSubstr("Assertion failed")); + } + catch (std::exception const& e) + { + FAIL() << "Expected TllmException"; + } } TEST(ExecutorConfigTest, extendedRuntimePerfKnobConfigTest) diff --git a/examples/disaggregated/slurm/cache_transceiver_test/run_cache_transceiver_test.py b/examples/disaggregated/slurm/cache_transceiver_test/run_cache_transceiver_test.py index fc8423d68d0d..6a6775d5c08e 100644 --- a/examples/disaggregated/slurm/cache_transceiver_test/run_cache_transceiver_test.py +++ b/examples/disaggregated/slurm/cache_transceiver_test/run_cache_transceiver_test.py @@ -104,6 +104,8 @@ class KvCacheConfigV2: dtype: str = "auto" pool_ratio: Optional[List[float]] = None avg_seq_len: Optional[int] = None + block_reuse_policy: str = "all_reusable" + enable_swa_scratch_reuse: bool = False disk_prefetch_num_reqs: int = 4 max_util_for_resume: float = 0.95 @@ -171,12 +173,13 @@ def build_kv_cache_manager(cfg_kv, mapping, use_v2): def add_sequence(mgr, req, prompt_len, use_v2): """Allocate KV blocks for the request. Returns a handle to close (V2) or None.""" if use_v2: - kv = mgr._create_kv_cache(req.py_request_id, None, None) - ok = kv.resume(torch.cuda.current_stream().cuda_stream) + if req.is_disagg_generation_init_state: + ok = mgr.prepare_disagg_gen_init(req) + else: + ok = mgr.prepare_context(req) and mgr.resize_context(req, prompt_len) if not ok: - raise RuntimeError(f"V2 resume failed for request {req.py_request_id}") - kv.resize(prompt_len) - return kv + raise RuntimeError(f"V2 KV cache allocation failed for request {req.py_request_id}") + return None mgr.impl.add_sequence_batch([(req.py_request_id, prompt_len, 1)], [req]) return None diff --git a/tensorrt_llm/_torch/attention_backend/interface.py b/tensorrt_llm/_torch/attention_backend/interface.py index e91348f9a4e3..04fc5a771924 100644 --- a/tensorrt_llm/_torch/attention_backend/interface.py +++ b/tensorrt_llm/_torch/attention_backend/interface.py @@ -25,6 +25,7 @@ from ..pyexecutor.kv_cache_manager_v2 import KVCacheManagerV2 from ..pyexecutor.mamba_cache_manager import BaseMambaCacheManager from ..pyexecutor.resource_manager import KVCacheManager +from ..pyexecutor.trace_log_utils import log_tensor_size from ..utils import get_model_extra_attrs from .sparse.params import SparseMetadataParams @@ -763,6 +764,14 @@ def create_rope_const_params(self, interleave: bool = True): if rope_inv_freq is not None else None, weakref.ref(rope_cos_sin), ) + # One-shot log on cache miss (typically 2-4 times per model load). + log_tensor_size("rope/new_table", + rope_cos_sin, + max_pos=self.max_positions, + dim=self.dim, + theta=self.theta, + scale_type=self.scale_type, + interleave=interleave) return rope_inv_freq, rope_cos_sin 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 a88f45148b5e..6ac39598ff14 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 @@ -19,11 +19,13 @@ import torch from tensorrt_llm._torch.pyexecutor import llm_request -from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import GPU_LEVEL, KVCacheManagerV2, Role +from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import GPU_LEVEL, KVCacheManagerV2 +from tensorrt_llm._torch.utils import maybe_compile from tensorrt_llm._utils import ( TensorWrapper, convert_to_torch_tensor, get_size_in_bytes, + nvtx_range_debug, prefer_pinned, ) from tensorrt_llm.bindings import DataType @@ -36,51 +38,107 @@ AttentionLayerConfig, BatchDesc, BufferConfig, + DataRole, GpuCacheTierConfig, HostCacheTierConfig, KVCacheDesc, LayerId, + PageIndexMode, + ScratchDesc, + SwaScratchReuseConfig, ) from tensorrt_llm.runtime.kv_cache_manager_v2 import KVCacheManagerConfig as KVCacheManagerConfigPy +from tensorrt_llm.runtime.kv_cache_manager_v2._common import BAD_PAGE_INDEX from .compressor import KVCacheDtype from .deepseek_v4 import ( + DEEPSEEK_V4_NON_SLIDING_ATTENTION, + DEEPSEEK_V4_SLIDING_ATTENTION, DEEPSEEK_V4_SPARSE_RATIO, DeepseekV4AttentionType, compress_ratio_has_attention, get_attn_dim, get_token_bytes, - is_compress_layer, is_overlap_compressor, - is_sparse_layer, ) -def _estimate_bytes_per_token( +def _estimate_non_sliding_attn_size_per_token( head_dim: int, index_head_dim: int, compress_ratios: List[int], has_fp8_kv_cache, - attn_types: set[DeepseekV4AttentionType] | None = None, indexer_k_dtype: str = "fp8", ) -> int: total_bytes = 0 - for ratio in compress_ratios: - for attn in DeepseekV4AttentionType: - if attn_types is not None and attn not in attn_types: - continue - if compress_ratio_has_attention(ratio, attn): + for compress_ratio in compress_ratios: + for attn_type in DEEPSEEK_V4_NON_SLIDING_ATTENTION: + if compress_ratio_has_attention(compress_ratio, attn_type): total_bytes += _get_attn_bytes_per_token( head_dim, index_head_dim, - ratio, - attn, + compress_ratio, + attn_type, has_fp8_kv_cache, indexer_k_dtype=indexer_k_dtype, ) return total_bytes +def _estimate_swa_cache_size( + head_dim: int, + index_head_dim: int, + compress_ratios: List[int], + has_fp8_kv_cache, + tokens_per_block: int, + swa_window_size: int | None, + *, + context: bool, + scratch: bool, + indexer_k_dtype: str = "fp8", +) -> Tuple[int, int]: + tokens_per_block = int(tokens_per_block) + size_per_token = 0 + size_per_request = 0 + scratch_keys = set() + for compress_ratio in compress_ratios: + for attn_type in DEEPSEEK_V4_SLIDING_ATTENTION: + if not compress_ratio_has_attention(compress_ratio, attn_type): + continue + if attn_type == DeepseekV4AttentionType.SWA: + if swa_window_size is None: + continue + window_size = swa_window_size + else: + state_factor = 2 if is_overlap_compressor(compress_ratio) else 1 + window_size = state_factor * compress_ratio + if window_size <= 0: + continue + window_tokens = ( + (int(window_size) + tokens_per_block - 1) // tokens_per_block + ) * tokens_per_block + token_bytes = _get_attn_bytes_per_token( + head_dim, + index_head_dim, + compress_ratio, + attn_type, + has_fp8_kv_cache, + indexer_k_dtype=indexer_k_dtype, + ) + if not context: + size_per_request += window_tokens * token_bytes + elif not scratch: + size_per_token += token_bytes + else: + scratch_key = (attn_type, compress_ratio) + if scratch_key in scratch_keys: + size_per_request += window_tokens * token_bytes + else: + scratch_keys.add(scratch_key) + size_per_token += token_bytes + return size_per_token, size_per_request + + def _get_attn_bytes_per_token( head_dim: int, index_head_dim: int, @@ -102,18 +160,88 @@ def _get_attn_bytes_per_token( return token_bytes +def _get_index_mode(attn_type: DeepseekV4AttentionType) -> PageIndexMode: + if attn_type in DEEPSEEK_V4_SLIDING_ATTENTION: + return PageIndexMode.PER_LAYER + else: + return PageIndexMode.SHARED + + +@maybe_compile(options={"max-autotune": True}) +def _compute_sliding_block_tables_compiled( + block_offsets: torch.Tensor, + copy_idx: torch.Tensor, + pool_ids: torch.Tensor, + valid_pool: torch.Tensor, + scales: torch.Tensor, + layer_offsets: torch.Tensor, + output: torch.Tensor, +) -> None: + base = block_offsets[pool_ids[:, :, None], copy_idx[None, None, :], 0, :] + scaled_base = torch.where( + (base == BAD_PAGE_INDEX) | ~(valid_pool[:, :, None, None]), + BAD_PAGE_INDEX, + base * scales[:, :, None, None] + layer_offsets[:, :, None, None], + ) + output.copy_(scaled_base) + + +@maybe_compile(options={"max-autotune": True}) +def _compute_sliding_block_tables_with_scratch_compiled( + block_offsets: torch.Tensor, + copy_idx: torch.Tensor, + pool_ids: torch.Tensor, + valid_pool: torch.Tensor, + scales: torch.Tensor, + layer_offsets: torch.Tensor, + block_positions: torch.Tensor, + scratch_pages: torch.Tensor, + scratch_begs: torch.Tensor, + scratch_ends: torch.Tensor, + scratch_slots: torch.Tensor, + num_contexts: torch.Tensor, + output: torch.Tensor, +) -> None: + base = block_offsets[pool_ids[:, :, None], copy_idx[None, None, :], 0, :] + scaled_base = torch.where( + (base == BAD_PAGE_INDEX) | ~(valid_pool[:, :, None, None]), + BAD_PAGE_INDEX, + base * scales[:, :, None, None] + layer_offsets[:, :, None, None], + ) + output.copy_(scaled_base) + + context_positions = torch.arange( + scratch_begs.shape[1], + dtype=torch.int32, + device=scratch_begs.device, + ) + active_context = context_positions < num_contexts + mask = ( + (block_positions >= scratch_begs[:, :, None]) + & (block_positions < scratch_ends[:, :, None]) + & active_context[None, :, None] + ) + range_index = torch.where(mask, block_positions - scratch_begs[:, :, None], 0) + total_offset = range_index[pool_ids] * scratch_pages[:, :, None, None] + slot_idx = (total_offset // scales[:, :, None, None]).clamp( + max=scratch_slots.shape[-1] - 1, + ) + slot_id = scratch_slots[pool_ids].gather(-1, slot_idx.long()) + offset = total_offset % scales[:, :, None, None] + scratch_index = ( + slot_id * scales[:, :, None, None] + + (offset + layer_offsets[:, :, None, None]) % scales[:, :, None, None] + ) + scratch_capacity = scratch_begs.shape[1] + scratch_rows = scaled_base[:, :, :scratch_capacity, :] + mask = mask[pool_ids] & valid_pool[:, :, None, None] + output[:, :, :scratch_capacity, :].copy_(torch.where(mask, scratch_index, scratch_rows)) + + class DeepseekV4CacheManager(KVCacheManagerV2): - fixed_size_attention = { - DeepseekV4AttentionType.SWA, - DeepseekV4AttentionType.COMPRESSOR_STATE, - DeepseekV4AttentionType.COMPRESSOR_SCORE, - DeepseekV4AttentionType.INDEXER_COMPRESSOR_STATE, - DeepseekV4AttentionType.INDEXER_COMPRESSOR_SCORE, - } # This tensor is for compatibility with AttentionOp, it only contains swa attention. - # kv_cache_pool_pointers contains pool pointers swa pool, shape: [1, 2] - # It assume the KVCacheManagerPy has only one pool for swa attention. - # The second column is always 0. + # kv_cache_pool_pointers contains one virtual attention-op pool per local + # SWA layer, shape: [num_local_layers, 2]. The second column is always 0. kv_cache_pool_pointers: torch.Tensor # This tensor is for compatibility with AttentionOp, it only contains swa attention. # kv_cache_pool_mapping contains pool id and layer offset for each layer's swa attention, @@ -223,7 +351,9 @@ def __init__( # Use first PP layer instead of hardcoded 0 for pipeline parallelism. first_pp_layer = self.pp_layers[0] self.swa_pool_ptr = self.impl.get_mem_pool_base_address( - self._layer_attn_to_layer_id[first_pp_layer, DeepseekV4AttentionType.SWA], Role.KEY + self._layer_attn_to_layer_id[first_pp_layer, DeepseekV4AttentionType.SWA], + DeepseekV4AttentionType.SWA.role, + PageIndexMode.PER_LAYER, ) self.compress_pool_ptrs = {} @@ -233,7 +363,8 @@ def __init__( first_layer_with_4 = self.pp_layers[pp_compress_ratios.index(4)] self.compress_pool_ptrs[4] = self.impl.get_mem_pool_base_address( self._layer_attn_to_layer_id[first_layer_with_4, DeepseekV4AttentionType.COMPRESS], - Role.KEY, + DeepseekV4AttentionType.COMPRESS.role, + PageIndexMode.SHARED, ) if 128 in pp_compress_ratios: # compressor first_layer_with_128 = self.pp_layers[pp_compress_ratios.index(128)] @@ -241,17 +372,24 @@ def __init__( self._layer_attn_to_layer_id[ first_layer_with_128, DeepseekV4AttentionType.COMPRESS ], - Role.KEY, + DeepseekV4AttentionType.COMPRESS.role, + PageIndexMode.SHARED, ) - # Use pinned staging buffer to avoid pageable H2D memcpy - max_num_sequences = max_batch_size * mapping.pp_size - self._host_block_offsets_staging = torch.empty( - (max_num_sequences + 1) * max_beam_width, - 2, # key and value - self.max_blocks_per_seq, - dtype=torch.int32, - pin_memory=prefer_pinned(), - 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, role)) + 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: @@ -267,7 +405,9 @@ def get_buffers(self, layer_idx: int, attn_type: DeepseekV4AttentionType) -> tor For blockwise FP8 layers, shape is [num_blocks, tokens_per_block, attn_dim + scale_size] """ layer_id = self._layer_attn_to_layer_id[(layer_idx, attn_type)] - addr = self.impl.get_mem_pool_base_address(layer_id, Role.KEY) + data_role = attn_type.role + page_index_mode = _get_index_mode(attn_type) + addr = self.impl.get_mem_pool_base_address(layer_id, data_role, page_index_mode) block_size = self.tokens_per_block if attn_type in [ @@ -285,18 +425,20 @@ def get_buffers(self, layer_idx: int, attn_type: DeepseekV4AttentionType) -> tor else: dim_per_token = attn_dim - shape = ( - self.impl.get_page_index_upper_bound(layer_id, Role.KEY), - block_size, - dim_per_token, - ) + page_index_upper_bound = self.impl.get_page_index_upper_bound(layer_id, data_role) + if page_index_mode == PageIndexMode.PER_LAYER: + converter = self.impl.get_page_index_converter(layer_id, data_role) + if converter.layer_offset is not None: + page_index_upper_bound += converter.layer_offset * converter.expansion + + shape = (page_index_upper_bound, block_size, dim_per_token) dtype = self.dtype - # (indexer) compressor state and score use compressor_dtype + # (indexer) compressor kv and score use compressor_dtype if attn_type in [ - DeepseekV4AttentionType.COMPRESSOR_STATE, + DeepseekV4AttentionType.COMPRESSOR_KV, DeepseekV4AttentionType.COMPRESSOR_SCORE, - DeepseekV4AttentionType.INDEXER_COMPRESSOR_STATE, + DeepseekV4AttentionType.INDEXER_COMPRESSOR_KV, DeepseekV4AttentionType.INDEXER_COMPRESSOR_SCORE, ]: dtype = self._compressor_dtype @@ -305,48 +447,260 @@ def get_buffers(self, layer_idx: int, attn_type: DeepseekV4AttentionType) -> tor return convert_to_torch_tensor(TensorWrapper(addr, dtype, shape)) - def _get_window_size(self, compress_ratio: int, attn_type: DeepseekV4AttentionType) -> int: + def _get_window_size( + self, compress_ratio: int, attn_type: DeepseekV4AttentionType + ) -> int | None: if attn_type == DeepseekV4AttentionType.SWA: base_window_size = self._swa_window_size - elif attn_type in self.fixed_size_attention: + elif attn_type in ( + DeepseekV4AttentionType.COMPRESSOR_KV, + DeepseekV4AttentionType.COMPRESSOR_SCORE, + DeepseekV4AttentionType.INDEXER_COMPRESSOR_KV, + DeepseekV4AttentionType.INDEXER_COMPRESSOR_SCORE, + ): state_factor = 2 if is_overlap_compressor(compress_ratio) else 1 base_window_size = state_factor * compress_ratio else: - raise ValueError(f"Unsupported fixed-size attention type: {attn_type}") + return None return base_window_size + self._max_draft_len - def _build_pool_mapping_tensors(self) -> Tuple[torch.Tensor, torch.Tensor]: - first_pp_layer = self.pp_layers[0] - swa_bytes_per_block = self._get_attn_bytes_per_block( - DeepseekV4AttentionType.SWA, first_pp_layer - ) - swa_pool_ptr = self.impl.get_mem_pool_base_address( - self._layer_attn_to_layer_id[first_pp_layer, DeepseekV4AttentionType.SWA], Role.KEY - ) - - def _get_layer_offset(pp_layer: int) -> int: - buffer_ptr = self.impl.get_mem_pool_base_address( - self._layer_attn_to_layer_id[pp_layer, DeepseekV4AttentionType.SWA], Role.KEY - ) - return (buffer_ptr - swa_pool_ptr) // swa_bytes_per_block - + def _prepare_page_table_tensor(self, index_mapper_capacity: int) -> None: # Tensors for compatibility with AttentionOp, only contains swa attention. - # Assume the SWA of all layers share the same pool. - # shape: [1, 2] - kv_cache_pool_pointers = torch.tensor( - [[swa_pool_ptr, 0]], + # SWA uses per-layer page indices, so each SWA layer has a virtual + # attention-op pool. + # shape: [num_local_layers, 2] + self.num_attention_op_pools = self.num_local_layers + self.kv_cache_pool_pointers = torch.tensor( + [ + [ + self.impl.get_mem_pool_base_address( + self._layer_attn_to_layer_id[pp_layer, DeepseekV4AttentionType.SWA], + DeepseekV4AttentionType.SWA.role, + PageIndexMode.PER_LAYER, + ), + 0, + ] + for pp_layer in self.pp_layers + ], dtype=torch.int64, device="cpu", pin_memory=prefer_pinned(), ) # shape: [num_local_layers, 2] - kv_cache_pool_mapping = torch.tensor( - [[0, _get_layer_offset(pp_layer)] for pp_layer in self.pp_layers], + self.kv_cache_pool_mapping = torch.tensor( + [[local_layer_idx, 0] for local_layer_idx in range(self.num_local_layers)], dtype=torch.int32, device="cpu", pin_memory=prefer_pinned(), ) - return kv_cache_pool_pointers, kv_cache_pool_mapping + self.host_kv_cache_block_offsets = torch.empty( + self.num_pools, + index_mapper_capacity * self.max_beam_width, + 2, # key and value + self.max_blocks_per_seq, + dtype=torch.int32, + pin_memory=prefer_pinned(), + device="cpu", + ) + staging_capacity = self.max_batch_size * self.max_beam_width + self._host_compress_block_tables_staging = { + compress_ratio: torch.full( + (staging_capacity, self.max_blocks_per_seq), + BAD_PAGE_INDEX, + dtype=torch.int32, + pin_memory=prefer_pinned(), + device="cpu", + ) + for compress_ratio in set(self._compress_ratios) + if compress_ratio_has_attention(compress_ratio, DeepseekV4AttentionType.COMPRESS) + } + + # layer offsets per layer and attn, shape [num_local_layers, len(DEEPSEEK_V4_SLIDING_ATTENTION)]. + self._layer_offsets = torch.full( + (self.num_local_layers, len(DEEPSEEK_V4_SLIDING_ATTENTION)), + -1, + dtype=torch.int32, + device="cpu", + ) + # Pool ids per layer and sliding attention type, shape [num_local_layers, num_sliding_attention_types]. + self._layer_attn_pool_ids = torch.full( + (self.num_local_layers, len(DEEPSEEK_V4_SLIDING_ATTENTION)), + -1, + dtype=torch.int32, + device="cpu", + ) + # Scales per layer and sliding attention type, shape [num_local_layers, num_sliding_attention_types]. + self._layer_attn_scales = torch.ones( + (self.num_local_layers, len(DEEPSEEK_V4_SLIDING_ATTENTION)), + dtype=torch.int32, + device="cpu", + ) + # Scratch pages per block per layer and sliding attention type, shape + # [num_local_layers, num_sliding_attention_types]. + self._scratch_pages = torch.zeros( + (self.num_local_layers, len(DEEPSEEK_V4_SLIDING_ATTENTION)), + dtype=torch.int32, + device="cpu", + ) + self._csa_compress_pool_id = None + self._csa_compress_scale = None + self._csa_indexer_compress_pool_id = None + self._csa_indexer_compress_scale = None + self._hca_compress_pool_id = None + self._hca_compress_scale = None + + for layer_idx in self.pp_layers: + compress_ratio = self._compress_ratios[layer_idx] + if compress_ratio == 4: + compress_layer_id = self._layer_attn_to_layer_id[ + layer_idx, DeepseekV4AttentionType.COMPRESS + ] + compress_converter = self.impl.get_page_index_converter( + compress_layer_id, DeepseekV4AttentionType.COMPRESS.role + ) + self._csa_compress_pool_id = self.layer_to_pool_mapping_dict[compress_layer_id] + self._csa_compress_scale = int(compress_converter.scale) + indexer_layer_id = self._layer_attn_to_layer_id[ + layer_idx, DeepseekV4AttentionType.INDEXER_COMPRESS + ] + indexer_converter = self.impl.get_page_index_converter( + indexer_layer_id, DeepseekV4AttentionType.INDEXER_COMPRESS.role + ) + self._csa_indexer_compress_pool_id = self.layer_to_pool_mapping_dict[ + indexer_layer_id + ] + self._csa_indexer_compress_scale = int(indexer_converter.scale) + elif compress_ratio == 128: + compress_layer_id = self._layer_attn_to_layer_id[ + layer_idx, DeepseekV4AttentionType.COMPRESS + ] + compress_converter = self.impl.get_page_index_converter( + compress_layer_id, DeepseekV4AttentionType.COMPRESS.role + ) + self._hca_compress_pool_id = self.layer_to_pool_mapping_dict[compress_layer_id] + self._hca_compress_scale = int(compress_converter.scale) + + local_layer_idx = self.layer_offsets[layer_idx] + for attn_type in DEEPSEEK_V4_SLIDING_ATTENTION: + if not compress_ratio_has_attention(self._compress_ratios[layer_idx], attn_type): + continue + layer_id = self._layer_attn_to_layer_id[layer_idx, attn_type] + pool_id = self.layer_to_pool_mapping_dict[layer_id] + converter = self.impl.get_page_index_converter(layer_id, attn_type.role) + self._layer_attn_pool_ids[local_layer_idx, attn_type.value] = pool_id + self._layer_attn_scales[local_layer_idx, attn_type.value] = converter.scale + self._layer_offsets[local_layer_idx, attn_type.value] = converter.layer_offset + self._scratch_pages[local_layer_idx, attn_type.value] = ( + converter.scratch_pages_per_block + ) + + device = torch.device("cuda", torch.cuda.current_device()) + self._device_kv_cache_block_offsets_input = torch.empty_like( + self.host_kv_cache_block_offsets, + device=device, + ) + self._precomputed_sliding_block_tables = torch.empty( + ( + self.num_local_layers, + len(DEEPSEEK_V4_SLIDING_ATTENTION), + self.host_kv_cache_block_offsets.size(1), + self.max_blocks_per_seq, + ), + dtype=torch.int32, + device=device, + ) + self._device_copy_idx_staging = torch.zeros( + self.host_kv_cache_block_offsets.size(1), + dtype=torch.int32, + device=device, + ) + self._device_num_contexts = torch.empty((), dtype=torch.int32, device=device) + self._device_layer_offsets = self._layer_offsets.to(device=device) + self._device_layer_attn_pool_ids = self._layer_attn_pool_ids.to( + device=device, + dtype=torch.long, + ) + self._device_layer_attn_scales = self._layer_attn_scales.to(device=device) + self._device_scratch_pages = self._scratch_pages.to(device=device) + self._device_valid_sliding_pool = self._device_layer_attn_pool_ids >= 0 + self._device_block_positions = torch.arange( + self.max_blocks_per_seq, + dtype=torch.int32, + device=device, + ) + + if self.enable_swa_scratch_reuse: + valid_scales = self._layer_attn_scales[self._layer_attn_pool_ids >= 0] + min_scale = int(valid_scales.min().item()) if valid_scales.numel() > 0 else 1 + max_scratch_pages = int(self._scratch_pages.max().item()) + self._max_scratch_slots = max( + 1, + (self.max_blocks_per_seq * max_scratch_pages + min_scale - 1) // min_scale, + ) + scratch_slots_shape = ( + self.num_pools, + staging_capacity, + self._max_scratch_slots, + ) + self._host_scratch_begs_staging = torch.empty( + self.num_pools, + staging_capacity, + dtype=torch.int32, + pin_memory=prefer_pinned(), + device="cpu", + ) + self._host_scratch_ends_staging = torch.empty( + self.num_pools, + staging_capacity, + dtype=torch.int32, + pin_memory=prefer_pinned(), + device="cpu", + ) + self._host_scratch_slots_staging = torch.empty( + scratch_slots_shape, + dtype=torch.int32, + pin_memory=prefer_pinned(), + device="cpu", + ) + self._device_scratch_begs_staging = torch.empty( + self.num_pools, + staging_capacity, + dtype=torch.int32, + device=device, + ) + self._device_scratch_ends_staging = torch.empty( + self.num_pools, + staging_capacity, + dtype=torch.int32, + device=device, + ) + self._device_scratch_slots_staging = torch.empty( + scratch_slots_shape, + dtype=torch.int32, + device=device, + ) + + @property + def blocks_in_primary_pool(self) -> int: + first_pp_layer = self.pp_layers[0] + swa_layer_id = self._layer_attn_to_layer_id[first_pp_layer, DeepseekV4AttentionType.SWA] + return self.impl.get_page_index_upper_bound(swa_layer_id, DeepseekV4AttentionType.SWA.role) + + def get_num_free_blocks(self) -> int: + # This method reports primary-pool capacity while the manager is empty. + # DSV4 does not allocate the generic Role.KEY buffer, so use SWA's + # model-specific DataRole for warmup capacity estimation. + assert len(self.kv_cache_map) == 0, ( + "get_num_free_blocks is only used when the kv cache manager is empty" + ) + max_num_pages = max( + self.impl.get_page_index_upper_bound( + self._layer_attn_to_layer_id[layer_idx, DeepseekV4AttentionType.SWA], + DeepseekV4AttentionType.SWA.role, + ) + for layer_idx in self.pp_layers + ) + return max_num_pages def get_cache_indices( self, @@ -366,16 +720,125 @@ def get_cache_indices( The cache block indices, shape (max_blocks_per_seq,) """ layer_id = self._layer_attn_to_layer_id[(layer_idx, attn_type)] + data_role = attn_type.role pool_id = self.layer_to_pool_mapping_dict[layer_id] - base_indices = self.kv_cache_map[request_id].get_base_page_indices(pool_id).tolist() - converter = self.impl.get_page_index_converter(layer_id, Role.KEY) - return converter(base_indices) + kv_cache = self.kv_cache_map[request_id] + base_indices = kv_cache.get_base_page_indices(pool_id).tolist() + converter = self.impl.get_page_index_converter(layer_id, data_role) + page_index_mode = _get_index_mode(attn_type) + return converter( + base_indices, + page_index_mode, + kv_cache.get_scratch_desc(pool_id), + ) + + def _get_extra_quota_padding(self) -> int: + """Ensure each attention type has minimal space when max_tokens is small.""" + return len(DeepseekV4AttentionType) * (2 << 20) - def _get_cache_quota(self, max_tokens: int) -> int: - quota = int(max_tokens * self.get_cache_bytes_per_token()) - # Add extra quota to ensure sufficient space for small max_tokens cases. - quota += len(DeepseekV4AttentionType) * (2 << 20) - return quota + def _get_quota_from_max_tokens(self, max_tokens: int) -> int: + compress_ratios = [self._compress_ratios[layer] for layer in self.pp_layers] + has_fp8_kv_cache = self.dtype == DataType.FP8 + non_sliding_attn_size_per_token = _estimate_non_sliding_attn_size_per_token( + self.head_dim, + self.index_head_dim, + compress_ratios, + has_fp8_kv_cache, + indexer_k_dtype=self._indexer_k_dtype, + ) + ( + context_swa_size_per_token, + _, + ) = _estimate_swa_cache_size( + self.head_dim, + self.index_head_dim, + compress_ratios, + has_fp8_kv_cache, + self.tokens_per_block, + self._swa_window_size, + context=True, + scratch=self.enable_swa_scratch_reuse, + indexer_k_dtype=self._indexer_k_dtype, + ) + ( + generation_swa_size_per_token, + generation_swa_size_per_request, + ) = _estimate_swa_cache_size( + self.head_dim, + self.index_head_dim, + compress_ratios, + has_fp8_kv_cache, + self.tokens_per_block, + self._swa_window_size, + context=False, + scratch=False, + indexer_k_dtype=self._indexer_k_dtype, + ) + max_context_tokens = ( + self._max_num_tokens if self._max_num_tokens is not None else max_tokens + ) + context_tokens = min(max_tokens, max_context_tokens) + generation_tokens = max_tokens - context_tokens + generation_quota = ( + max_tokens * non_sliding_attn_size_per_token + + generation_tokens * generation_swa_size_per_token + + self.max_batch_size * generation_swa_size_per_request + ) + context_extra_quota = context_tokens * context_swa_size_per_token + padding = self._get_extra_quota_padding() + return int(generation_quota + context_extra_quota + padding) + + def _get_max_tokens_from_quota(self, quota: int) -> float: + compress_ratios = [self._compress_ratios[layer] for layer in self.pp_layers] + has_fp8_kv_cache = self.dtype == DataType.FP8 + non_sliding_attn_size_per_token = _estimate_non_sliding_attn_size_per_token( + self.head_dim, + self.index_head_dim, + compress_ratios, + has_fp8_kv_cache, + indexer_k_dtype=self._indexer_k_dtype, + ) + context_swa_size_per_token, _ = _estimate_swa_cache_size( + self.head_dim, + self.index_head_dim, + compress_ratios, + has_fp8_kv_cache, + self.tokens_per_block, + self._swa_window_size, + context=True, + scratch=self.enable_swa_scratch_reuse, + indexer_k_dtype=self._indexer_k_dtype, + ) + ( + generation_swa_size_per_token, + generation_swa_size_per_request, + ) = _estimate_swa_cache_size( + self.head_dim, + self.index_head_dim, + compress_ratios, + has_fp8_kv_cache, + self.tokens_per_block, + self._swa_window_size, + context=False, + scratch=False, + indexer_k_dtype=self._indexer_k_dtype, + ) + padding = self._get_extra_quota_padding() + size_per_batch = self.max_batch_size * generation_swa_size_per_request + padding + if quota < size_per_batch: + return 0 + context_size_per_token = non_sliding_attn_size_per_token + context_swa_size_per_token + if self._max_num_tokens is None: + return (quota - size_per_batch) / context_size_per_token + + context_limit_quota = self._max_num_tokens * context_size_per_token + size_per_batch + if quota <= context_limit_quota: + return (quota - size_per_batch) / context_size_per_token + + generation_size_per_token = non_sliding_attn_size_per_token + generation_swa_size_per_token + if generation_size_per_token <= 0: + return float("inf") + return self._max_num_tokens + (quota - context_limit_quota) / generation_size_per_token def _build_cache_config( self, @@ -390,21 +853,30 @@ def _build_cache_config( """ layers: List[AttentionLayerConfig] = [] layer_attn_to_layer_id: Dict[Tuple[int, DeepseekV4AttentionType], LayerId] = {} + manager_layer_id_to_layer_attn: Dict[ + Tuple[LayerId, DataRole], Tuple[int, DeepseekV4AttentionType] + ] = {} def _add_layer( - layer_idx: int, attn_type: DeepseekV4AttentionType, sliding_window_size: int | None - ): - nonlocal layers, layer_attn_to_layer_id + layer_idx: int, + attention_types: List[DeepseekV4AttentionType], + sliding_window_size: int | None, + ) -> None: 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 - # add the layer to the layers list + for attn_type in attention_types: + layer_attn_to_layer_id[layer_idx, attn_type] = layer_id + manager_layer_id_to_layer_attn[layer_id, attn_type.role] = ( + layer_idx, + attn_type, + ) layer_config = AttentionLayerConfig( layer_id=layer_id, buffers=[ BufferConfig( - role=Role.KEY, size=self._get_attn_bytes_per_block(attn_type, layer_idx) + role=attn_type.role, + size=self._get_attn_bytes_per_block(attn_type, layer_idx), ) + for attn_type in attention_types ], sliding_window_size=sliding_window_size, num_sink_tokens=None, @@ -414,50 +886,53 @@ def _add_layer( # create the layer config for DeepSeek-V4 for layer in self.pp_layers: compress_ratio = self._compress_ratios[layer] - is_compress = is_compress_layer(compress_ratio) - is_sparse = is_sparse_layer(compress_ratio) - - # sliding window attention pool - _add_layer( - layer, - DeepseekV4AttentionType.SWA, - self._get_window_size(compress_ratio, DeepseekV4AttentionType.SWA), - ) - - if is_compress: - # compressed attention pool - _add_layer(layer, DeepseekV4AttentionType.COMPRESS, None) - # compressor state, managed as a sliding window attention cache, - # including compressor kv states and compressor score states. - # Add max_draft_len so rewind after rejected draft tokens can - # still reach past states. - compressor_window = self._get_window_size( - compress_ratio, DeepseekV4AttentionType.COMPRESSOR_STATE + if compress_ratio == 1: + _add_layer( + layer, + [DeepseekV4AttentionType.SWA], + self._get_window_size(compress_ratio, DeepseekV4AttentionType.SWA), + ) + elif compress_ratio == DEEPSEEK_V4_SPARSE_RATIO: + _add_layer( + layer, + [DeepseekV4AttentionType.SWA], + self._get_window_size(compress_ratio, DeepseekV4AttentionType.SWA), ) - _add_layer(layer, DeepseekV4AttentionType.COMPRESSOR_STATE, compressor_window) - _add_layer(layer, DeepseekV4AttentionType.COMPRESSOR_SCORE, compressor_window) - - # sparse attention layer has indexer - if is_sparse: - # indexer kv cache pool, dim is indexer_head_dim - _add_layer(layer, DeepseekV4AttentionType.INDEXER_COMPRESS, None) - # indexer has its own compressor, so a separate compressor state - # similarly, indexer compressor state is managed as a sliding window attention cache - indexer_compressor_window = self._get_window_size( - compress_ratio, DeepseekV4AttentionType.INDEXER_COMPRESSOR_STATE + _add_layer( + layer, + [ + DeepseekV4AttentionType.COMPRESS, + DeepseekV4AttentionType.INDEXER_COMPRESS, + ], + None, ) _add_layer( layer, - DeepseekV4AttentionType.INDEXER_COMPRESSOR_STATE, - indexer_compressor_window, + [ + DeepseekV4AttentionType.COMPRESSOR_KV, + DeepseekV4AttentionType.COMPRESSOR_SCORE, + DeepseekV4AttentionType.INDEXER_COMPRESSOR_KV, + DeepseekV4AttentionType.INDEXER_COMPRESSOR_SCORE, + ], + self._get_window_size(compress_ratio, DeepseekV4AttentionType.COMPRESSOR_KV), ) + elif compress_ratio == 128: _add_layer( layer, - DeepseekV4AttentionType.INDEXER_COMPRESSOR_SCORE, - indexer_compressor_window, + [ + DeepseekV4AttentionType.SWA, + DeepseekV4AttentionType.COMPRESSOR_KV, + DeepseekV4AttentionType.COMPRESSOR_SCORE, + ], + self._get_window_size(compress_ratio, DeepseekV4AttentionType.SWA), ) + _add_layer(layer, [DeepseekV4AttentionType.COMPRESS], None) + else: + raise ValueError(f"Unsupported DeepSeek-V4 compress ratio {compress_ratio}.") + # 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) @@ -466,45 +941,74 @@ def _add_layer( max_seq_len = self.max_seq_len max_num_tokens = self._max_num_tokens max_draft_len = self._max_draft_len + typical_step = None + constraints = [] + if kv_cache_config.pool_ratio is None: + typical_seq_len = ( + kv_cache_config.avg_seq_len + if kv_cache_config.avg_seq_len is not None + else max_seq_len + ) + if typical_seq_len > max_seq_len: + raise ValueError( + f"kv_cache_config.avg_seq_len ({typical_seq_len}) must be less than or " + f"equal to max_seq_len ({max_seq_len})" + ) - # For aggregated serving in large batch size: - # Use 1 context request + (max_batch_size - 1) generation requests as - # the typical step. An all-generation typical_step over-provisions the - # compressed-cache pool at the expense of the SWA pool, starving the - # SWA pool and artificially capping the achievable batch size. - ctx_capacity = max_num_tokens if max_num_tokens is not None else max_seq_len - typical_step = BatchDesc( - kv_caches=[ - KVCacheDesc(capacity=ctx_capacity, history_length=0), - ] - + [KVCacheDesc(capacity=max_seq_len, history_length=max_seq_len - max_draft_len - 1)] - * (max_batch_size - 1), - ) + # For aggregated serving in large batch size: + # Use 1 context request + (max_batch_size - 1) generation requests as + # the typical step. An all-generation typical_step over-provisions the + # compressed-cache pool at the expense of the SWA pool, starving the + # SWA pool and artificially capping the achievable batch size. + ctx_capacity = max_num_tokens if max_num_tokens is not None else typical_seq_len + generation_history_length = max(0, typical_seq_len - max_draft_len - 1) + typical_step = BatchDesc( + kv_caches=[ + KVCacheDesc(capacity=ctx_capacity, history_length=0), + ] + + [ + KVCacheDesc( + capacity=typical_seq_len, + history_length=generation_history_length, + ) + ] + * (max_batch_size - 1), + ) - constraints = [] - # Constraint 1: cuda graph generation warmup — one decode request that has - # accumulated to the tail of max_seq_len. Using history_length=max_seq_len-1 - # (instead of 0) lets SWA / SSM pools collapse to their windowed working set, - # while full-cache pools still need max_seq_len/tokens_per_block blocks - # because they don't age. - constraints.append( - BatchDesc([KVCacheDesc(capacity=max_seq_len, history_length=max_seq_len - 1)]) - ) + # Constraint 1: cuda graph generation warmup — one decode request that has + # accumulated to the tail of max_seq_len. Using history_length=max_seq_len-1 + # (instead of 0) lets SWA / SSM pools collapse to their windowed working set, + # while full-cache pools still need max_seq_len/tokens_per_block blocks + # because they don't age. + constraints.append( + BatchDesc([KVCacheDesc(capacity=max_seq_len, history_length=max_seq_len - 1)]) + ) + + # Constraint 2: general / chunked-prefill warmup — one fresh context request + # at max_num_tokens (the per-iteration token budget). + if max_num_tokens is not None: + constraints.append( + BatchDesc([KVCacheDesc(capacity=max_num_tokens, history_length=0)]) + ) - # Constraint 2: general / chunked-prefill warmup — one fresh context request - # at max_num_tokens (the per-iteration token budget). - if max_num_tokens is not None: - constraints.append(BatchDesc([KVCacheDesc(capacity=max_num_tokens, history_length=0)])) + scratch_reuse_config = None + if self.enable_swa_scratch_reuse: + # Context requests will allocate num_extra_kv_tokens tokens for spec decoding. + # Cache manager should not take them into account when calculating scratch range. + # Therefore set max_rewind_len to num_extra_kv_tokens. + scratch_reuse_config = SwaScratchReuseConfig(max_rewind_len=self.num_extra_kv_tokens) return KVCacheManagerConfigPy( tokens_per_block=tokens_per_block, vocab_size=vocab_size, cache_tiers=cache_tiers, max_util_for_resume=kv_cache_config.max_util_for_resume, + swa_scratch_reuse=scratch_reuse_config, layers=layers, typical_step=typical_step, constraints=constraints, enable_stats=self.enable_stats, + initial_pool_ratio=kv_cache_config.pool_ratio, ) def _init_indexer_dtype(self, sparse_attn_config: DeepSeekV4SparseAttentionConfig) -> None: @@ -564,7 +1068,8 @@ def _assert_layer_pool_scale(self) -> None: compress_ratio = self._compress_ratios[layer_idx] layer_id = self._layer_attn_to_layer_id[layer_idx, attn_type] pool_id = self.layer_to_pool_mapping_dict[layer_id] - converter = self.impl.get_page_index_converter(layer_id, Role.KEY) + converter = self.impl.get_page_index_converter(layer_id, attn_type.role) + assert converter.expansion == 1, "DeepSeek-V4 page index expansion must be 1" scale = converter.scale # check if the pool id is consistent @@ -609,16 +1114,13 @@ def _assert_layer_pool_scale(self) -> None: attn_ratio_to_pool_id[DeepseekV4AttentionType.SWA][ratio] = swa_pool_id attn_ratio_to_scale[DeepseekV4AttentionType.SWA][ratio] = swa_scale - self._attn_ratio_to_pool_id = attn_ratio_to_pool_id - self._attn_ratio_to_scale = attn_ratio_to_scale - def _get_attn_bytes_per_block( self, attn_type: DeepseekV4AttentionType, layer_idx: int, ) -> int: """ - Get the cache bytes per token for a specific attention type and layer. + Get the cache bytes per block for a specific attention type and layer. """ has_fp8_kv_cache = self.dtype == DataType.FP8 token_bytes = get_token_bytes( @@ -643,7 +1145,7 @@ def get_cache_bytes_per_token(self) -> int: """Get the average cache bytes per token for DeepSeek-V4.""" has_fp8_kv_cache = self.dtype == DataType.FP8 compress_ratios = [self._compress_ratios[layer] for layer in self.pp_layers] - return _estimate_bytes_per_token( + return _estimate_non_sliding_attn_size_per_token( self.head_dim, self.index_head_dim, compress_ratios, @@ -656,51 +1158,58 @@ def get_max_resource_count(self) -> int: return int(self.impl.get_quota(GPU_LEVEL)) def _is_context_request(self, request: llm_request.LlmRequest) -> bool: - if getattr(request, "is_context_init_state", False): + if request.is_context_init_state: return True - return getattr(request, "state", None) == llm_request.LlmRequestState.CONTEXT_INIT + return request.state == llm_request.LlmRequestState.CONTEXT_INIT def _is_generation_request(self, request: llm_request.LlmRequest) -> bool: if ( - getattr(request, "is_generation_in_progress_state", False) - or getattr(request, "is_generation_to_complete_state", False) - or getattr(request, "is_disagg_generation_init_state", False) + request.is_generation_in_progress_state + or request.is_generation_to_complete_state + or request.is_disagg_generation_init_state ): return True - return getattr(request, "state", None) in ( + return request.state in ( llm_request.LlmRequestState.GENERATION_IN_PROGRESS, llm_request.LlmRequestState.GENERATION_TO_COMPLETE, ) def _get_context_bytes(self, request: llm_request.LlmRequest) -> int: - prompt_len = max(0, getattr(request, "prompt_len", request.orig_prompt_len)) + prompt_len = max(0, request.prompt_len) total_tokens = prompt_len + self.num_extra_kv_tokens - return total_tokens * self.get_cache_bytes_per_token() + return self._get_cache_bytes_for_tokens(total_tokens, context=True) def _get_generation_bytes(self, request: llm_request.LlmRequest) -> int: - prompt_len = max(0, getattr(request, "prompt_len", request.orig_prompt_len)) + prompt_len = max(0, request.prompt_len) max_new_tokens = max(0, request.max_new_tokens) total_tokens = prompt_len + max_new_tokens + self.num_extra_kv_tokens + return self._get_cache_bytes_for_tokens(total_tokens, context=False) + + def _get_cache_bytes_for_tokens(self, total_tokens: int, *, context: bool) -> int: has_fp8_kv_cache = self.dtype == DataType.FP8 - total_bytes = 0 - for layer in self.pp_layers: - compress_ratio = self._compress_ratios[layer] - for attn_type in DeepseekV4AttentionType: - if not compress_ratio_has_attention(compress_ratio, attn_type): - continue - token_bytes = _get_attn_bytes_per_token( - self.head_dim, - self.index_head_dim, - compress_ratio, - attn_type, - has_fp8_kv_cache, - indexer_k_dtype=self._indexer_k_dtype, - ) - attn_tokens = total_tokens - if attn_type in self.fixed_size_attention: - attn_tokens = self._get_window_size(compress_ratio, attn_type) - total_bytes += attn_tokens * token_bytes - return total_bytes + compress_ratios = [self._compress_ratios[layer] for layer in self.pp_layers] + non_sliding_attn_size_per_token = _estimate_non_sliding_attn_size_per_token( + self.head_dim, + self.index_head_dim, + compress_ratios, + has_fp8_kv_cache, + indexer_k_dtype=self._indexer_k_dtype, + ) + swa_size_per_token, swa_size_per_request = _estimate_swa_cache_size( + self.head_dim, + self.index_head_dim, + compress_ratios, + has_fp8_kv_cache, + self.tokens_per_block, + self._swa_window_size, + context=context, + scratch=self.enable_swa_scratch_reuse, + indexer_k_dtype=self._indexer_k_dtype, + ) + return int( + total_tokens * (non_sliding_attn_size_per_token + swa_size_per_token) + + swa_size_per_request + ) def get_needed_resource_to_completion(self, request: llm_request.LlmRequest) -> int: if self._is_generation_request(request): @@ -712,7 +1221,7 @@ def get_needed_resource_to_completion(self, request: llm_request.LlmRequest) -> def get_layer_bytes_per_token( self, local_layer_idx: int, - data_role: Role, + data_role: DataRole, ) -> int: raise NotImplementedError( "DeepSeek-V4 doesn't support get_layer_bytes_per_token, use _get_attn_bytes_per_block" @@ -725,20 +1234,129 @@ def get_indexer_k_cache_buffers(self, layer_idx: int) -> torch.Tensor: buffer = self.get_buffers(layer_idx, DeepseekV4AttentionType.INDEXER_COMPRESS).unsqueeze(2) return buffer.view(torch.uint8) - def get_batch_indexer_k_cache_indices(self, request_ids: List[int]) -> List[List[int]]: + def _compute_shared_block_table( + self, pool_id: int, scale: int, copy_idx: torch.Tensor + ) -> torch.Tensor: """ - Get the indices for the indexer k cache for a batch of requests. + Get the shared offset for one pool and copy index. + Return shape: [num_seqs, max_blocks_per_seq] """ - return self.get_batch_attn_offset( - request_ids, - # use beam_width=1 and num_contexts=0 since we don't support beam search - 1, - 0, - len(request_ids), - DeepseekV4AttentionType.INDEXER_COMPRESS, - DEEPSEEK_V4_SPARSE_RATIO, - ).tolist() + base = self.host_kv_cache_block_offsets[pool_id, copy_idx, 0, :] + return torch.where(base == BAD_PAGE_INDEX, BAD_PAGE_INDEX, base * scale) + def _copy_idx_to_device(self, copy_idx: torch.Tensor) -> torch.Tensor: + num_tables = copy_idx.size(0) + device_copy_idx = self._device_copy_idx_staging[:num_tables] + device_copy_idx.copy_(copy_idx, non_blocking=True) + # Keep the compiled graph independent of the active table count. + return self._device_copy_idx_staging + + def _copy_scratch_metadata_to_device( + self, + scratch_descs_by_pool: list[list[ScratchDesc | None]], + num_contexts: int, + host_begs_staging: torch.Tensor, + host_ends_staging: torch.Tensor, + host_slots_staging: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + # shape: [num_pools, num_contexts] + host_begs = host_begs_staging[:, :num_contexts] + # shape: [num_pools, num_contexts] + host_ends = host_ends_staging[:, :num_contexts] + # shape: [num_pools, num_contexts, num_slots] + host_slots = host_slots_staging[:, :num_contexts, :] + host_begs.zero_() + host_ends.zero_() + host_slots.zero_() + + for pool_idx, scratch_descs in enumerate(scratch_descs_by_pool): + for context_idx, desc in enumerate(scratch_descs): + if desc is None: + continue + slot_ids = desc.slot_ids + if len(slot_ids) > self._max_scratch_slots: + raise RuntimeError( + f"Scratch slot count {len(slot_ids)} exceeds staging capacity " + f"{self._max_scratch_slots}" + ) + host_begs[pool_idx, context_idx] = int(desc.range.beg) + host_ends[pool_idx, context_idx] = int(desc.range.end) + for slot_idx, slot_id in enumerate(slot_ids): + host_slots[pool_idx, context_idx, slot_idx] = int(slot_id) + + self._device_scratch_begs_staging.copy_(host_begs_staging, non_blocking=True) + self._device_scratch_ends_staging.copy_(host_ends_staging, non_blocking=True) + self._device_scratch_slots_staging.copy_(host_slots_staging, non_blocking=True) + # Keep scratch tensor shapes fixed; the device-side mask gates num_contexts. + return ( + self._device_scratch_begs_staging, + self._device_scratch_ends_staging, + self._device_scratch_slots_staging, + ) + + @nvtx_range_debug("dsv4_compute_sliding_block_tables") + def compute_sliding_block_tables( + self, + request_ids: List[int], + num_contexts: int, + ) -> None: + """Compute all per-layer sliding-window block tables for this batch.""" + copy_idx = self.index_mapper.get_copy_index(request_ids, num_contexts, 1) + num_tables = copy_idx.size(0) + self._num_tables = num_tables + + scratch_descs_by_pool = None + if self.enable_swa_scratch_reuse and num_contexts > 0: + scratch_descs_by_pool = [ + [ + self.kv_cache_map[req].get_scratch_desc(pool_id) + for req in request_ids[:num_contexts] + ] + for pool_id in range(self.num_pools) + ] + + device_copy_idx = self._copy_idx_to_device(copy_idx) + self._device_kv_cache_block_offsets_input.copy_( + self.host_kv_cache_block_offsets, + non_blocking=True, + ) + + if scratch_descs_by_pool is not None: + scratch_begs, scratch_ends, scratch_slots = self._copy_scratch_metadata_to_device( + scratch_descs_by_pool, + num_contexts, + self._host_scratch_begs_staging, + self._host_scratch_ends_staging, + self._host_scratch_slots_staging, + ) + self._device_num_contexts.fill_(num_contexts) + _compute_sliding_block_tables_with_scratch_compiled( + self._device_kv_cache_block_offsets_input, + device_copy_idx, + self._device_layer_attn_pool_ids, + self._device_valid_sliding_pool, + self._device_layer_attn_scales, + self._device_layer_offsets, + self._device_block_positions, + self._device_scratch_pages, + scratch_begs, + scratch_ends, + scratch_slots, + self._device_num_contexts, + self._precomputed_sliding_block_tables, + ) + else: + _compute_sliding_block_tables_compiled( + self._device_kv_cache_block_offsets_input, + device_copy_idx, + self._device_layer_attn_pool_ids, + self._device_valid_sliding_pool, + self._device_layer_attn_scales, + self._device_layer_offsets, + self._precomputed_sliding_block_tables, + ) + + @nvtx_range_debug("dsv4_copy_batch_block_offsets") def copy_batch_block_offsets( self, dst_tensor: torch.Tensor, @@ -748,90 +1366,83 @@ def copy_batch_block_offsets( num_seqs: int, ) -> None: """For compatibility with AttentionOp, copy only the SWA block offsets.""" - offsets = self.get_batch_attn_offset( - request_ids, - beam_width, - num_contexts, - num_seqs, - DeepseekV4AttentionType.SWA, - # all compress ratios have SWA attention and they are in the same pool - self._compress_ratios[self.pp_layers[0]], - ) - self._host_block_offsets_staging[:num_seqs, :, :] = offsets[:, None, :] - dst_tensor[0, :num_seqs, :, :].copy_( - self._host_block_offsets_staging[:num_seqs, :, :], non_blocking=True + assert beam_width == 1, "DSV4 only supports beam width 1 now" + assert dst_tensor.is_cuda, "copy_batch_block_offsets expects a CUDA destination" + dst_tensor.fill_(BAD_PAGE_INDEX) + dst_tensor[:, : self._num_tables, 0, :].copy_( + self._precomputed_sliding_block_tables[ + :, DeepseekV4AttentionType.SWA.value, : self._num_tables, : + ], + non_blocking=True, ) - def get_batch_attn_offset( + @nvtx_range_debug("dsv4_copy_batch_sliding_block_tables") + def copy_batch_sliding_block_tables( self, + dst_tensor: torch.Tensor, request_ids: List[int], - beam_width: int, num_contexts: int, num_seqs: int, - attn_type: DeepseekV4AttentionType, - compress_ratio: int, - ) -> torch.Tensor: + ) -> None: """ - Get the block offsets for a specific attention type for a batch of requests. - - Args: - request_ids: The request ids - beam_width: The beam width - num_contexts: The number of context requests - num_seqs: The number of sequence requests - attn_type: The attention type - compress_ratio: The compress ratio. Used for non-SWA attention types. - - Returns: - The block offsets, shape (num_seqs, max_blocks_per_seq) + Copy the per-layer block tables for attentions managed in sliding-window mode to the GPU tensor. """ - assert beam_width == 1, "beam_width must be 1 for KVCacheManagerV2" - assert attn_type == DeepseekV4AttentionType.SWA or compress_ratio is not None, ( - "compress_ratio must be provided for non-SWA attention types" + assert dst_tensor.is_cuda, "copy_batch_sliding_block_tables expects a CUDA destination" + dst_tensor.fill_(BAD_PAGE_INDEX) + dst_tensor[:, :, : self._num_tables, :].copy_( + self._precomputed_sliding_block_tables[:, :, : self._num_tables, :], + non_blocking=True, ) + @nvtx_range_debug("dsv4_copy_batch_compress_block_tables") + def copy_batch_compress_block_tables( + self, + dst_tensor: torch.Tensor, + request_ids: List[int], + compress_ratio: int, + beam_width: int, + num_contexts: int, + num_seqs: int, + ) -> None: + """Build the COMPRESS block table for one compression ratio and copy it to the destination.""" + assert beam_width == 1, "DSV4 only supports beam width 1 now" copy_idx = self.index_mapper.get_copy_index(request_ids, num_contexts, beam_width) - assert copy_idx.shape[0] == num_seqs + staging = self._host_compress_block_tables_staging[compress_ratio] + if compress_ratio == 4: + pool_id = self._csa_compress_pool_id + scale = self._csa_compress_scale + elif compress_ratio == 128: + pool_id = self._hca_compress_pool_id + scale = self._hca_compress_scale + else: + raise ValueError( + f"Unsupported compress ratio {compress_ratio} for copy_batch_compress_block_tables" + ) - pool_id = self._attn_ratio_to_pool_id[attn_type][compress_ratio] - scale = self._attn_ratio_to_scale[attn_type][compress_ratio] - offsets = self.host_kv_cache_block_offsets[pool_id, copy_idx, 0] * scale - offsets[offsets == -scale] = -1 - return offsets + if pool_id is None or scale is None: + raise RuntimeError( + f"Missing COMPRESS pool metadata for compress ratio {compress_ratio}" + ) + staging[:num_seqs] = self._compute_shared_block_table(pool_id, scale, copy_idx) + dst_tensor[:num_seqs].copy_(staging[:num_seqs], non_blocking=True) - def get_batch_block_offsets( + @nvtx_range_debug("dsv4_copy_batch_indexer_compress_block_tables") + def copy_batch_indexer_compress_block_tables( self, + host_block_table: torch.Tensor, request_ids: List[int], + beam_width: int, num_contexts: int, - attention_type_set: set, - ) -> Dict[Tuple[int, "DeepseekV4AttentionType"], torch.Tensor]: - """Get block offsets for all attention types in a single call. - - Calls get_copy_index once and deduplicates offset computation by - (pool_id, scale) to avoid redundant work. - - Args: - request_ids: The request ids. - num_contexts: The number of context requests. - attention_type_set: Set of (compress_ratio, attention_type) tuples. - - Returns: - Dict mapping (compress_ratio, attention_type) -> offset tensor. - """ - copy_idx = self.index_mapper.get_copy_index(request_ids, num_contexts, 1) - - offset_cache = {} # (pool_id, scale) -> offsets tensor - result = {} - for compress_ratio, attention_type in attention_type_set: - pool_id = self._attn_ratio_to_pool_id[attention_type][compress_ratio] - scale = self._attn_ratio_to_scale[attention_type][compress_ratio] - cache_key = (pool_id, scale) - if cache_key not in offset_cache: - offsets = self.host_kv_cache_block_offsets[pool_id, copy_idx, 0] * scale - offsets[offsets == -scale] = -1 - offset_cache[cache_key] = offsets - result[(compress_ratio, attention_type)] = offset_cache[cache_key] - return result + num_seqs: int, + ) -> None: + """Build the shared INDEXER_COMPRESS compatibility block table.""" + assert beam_width == 1, "DSV4 only supports beam width 1 now" + copy_idx = self.index_mapper.get_copy_index(request_ids, num_contexts, beam_width) + pool_id = self._csa_indexer_compress_pool_id + scale = self._csa_indexer_compress_scale + if pool_id is None or scale is None: + raise RuntimeError("Missing INDEXER_COMPRESS pool metadata") + host_block_table[:num_seqs] = self._compute_shared_block_table(pool_id, scale, copy_idx) @staticmethod def get_cache_size_per_token(model_config: ModelConfig, mapping: Mapping, **kwargs): @@ -848,25 +1459,43 @@ def get_cache_size_per_token(model_config: ModelConfig, mapping: Mapping, **kwar else: has_fp8_kv_cache = False indexer_k_dtype = model_config.sparse_attention_config.indexer_k_dtype - return _estimate_bytes_per_token( + non_sliding_attn_size_per_token = _estimate_non_sliding_attn_size_per_token( head_dim, index_head_dim, compress_ratios, has_fp8_kv_cache, indexer_k_dtype=indexer_k_dtype, ) + swa_size_per_token, swa_size_per_request = _estimate_swa_cache_size( + head_dim, + index_head_dim, + compress_ratios, + has_fp8_kv_cache, + kwargs["tokens_per_block"], + model_config.sparse_attention_config.window_size, + context=False, + scratch=False, + indexer_k_dtype=indexer_k_dtype, + ) + max_batch_size = int(kwargs.get("max_batch_size") or 0) + return ( + non_sliding_attn_size_per_token + swa_size_per_token, + swa_size_per_request * max_batch_size, + ) def check_invalid_values_in_kv_cache(self, fill_with_zero: bool = False) -> bool: some_checks_unavailable = False has_invalid_values = torch.tensor( [False], dtype=torch.bool, device=torch.cuda.current_device() ) - pool_handled = set() + buffers_handled = set() - # Handle each layer from start to end to traverse the whole KV cache. + # Handle each attention buffer from start to end to traverse the whole + # KV cache. Multiple attention buffers can now share one cache layer. for (layer, attn), layer_id in self._layer_attn_to_layer_id.items(): - pool_id = self.layer_to_pool_mapping_dict[layer_id] - if pool_id in pool_handled: + data_role = attn.role + buffer_key = (layer_id, data_role) + if buffer_key in buffers_handled: continue buffer = self.get_buffers(layer, attn) # process in chunks of 256 pages to avoid OoM @@ -879,7 +1508,7 @@ def check_invalid_values_in_kv_cache(self, fill_with_zero: bool = False) -> bool some_checks_unavailable = True if fill_with_zero: buffer.zero_() - pool_handled.add(pool_id) + buffers_handled.add(buffer_key) torch.cuda.synchronize() if some_checks_unavailable: diff --git a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/compressor.py b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/compressor.py index afc4302c8710..6dd402369d9f 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/compressor.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/compressor.py @@ -156,25 +156,29 @@ def forward( # Determine attention types based on whether this is an indexer compressor if self.is_indexer: compress_type = DeepseekV4AttentionType.INDEXER_COMPRESS - state_type = DeepseekV4AttentionType.INDEXER_COMPRESSOR_STATE + kv_type = DeepseekV4AttentionType.INDEXER_COMPRESSOR_KV score_type = DeepseekV4AttentionType.INDEXER_COMPRESSOR_SCORE else: compress_type = DeepseekV4AttentionType.COMPRESS - state_type = DeepseekV4AttentionType.COMPRESSOR_STATE + kv_type = DeepseekV4AttentionType.COMPRESSOR_KV score_type = DeepseekV4AttentionType.COMPRESSOR_SCORE # Get cache buffers kv_cache = metadata.kv_cache_manager.get_buffers(self.layer_idx, compress_type) - paged_kv_state = metadata.kv_cache_manager.get_buffers(self.layer_idx, state_type) + paged_kv_state = metadata.kv_cache_manager.get_buffers(self.layer_idx, kv_type) paged_score_state = metadata.kv_cache_manager.get_buffers(self.layer_idx, score_type) # Get block tables - block_table = metadata.block_tables[(self.compress_ratio, compress_type)] - block_table_kv_state = metadata.block_tables[(self.compress_ratio, state_type)] - block_table_score_state = metadata.block_tables[(self.compress_ratio, score_type)] + local_layer_idx = metadata.kv_cache_manager.layer_offsets[self.layer_idx] + if self.is_indexer: + block_table = metadata.indexer_k_cache_block_offsets + else: + block_table = metadata.compress_block_tables[self.compress_ratio] + block_table_kv_state = metadata.sliding_block_tables[local_layer_idx, kv_type.value] + block_table_score_state = metadata.sliding_block_tables[local_layer_idx, score_type.value] # Get tokens_per_block from cache manager - # state_tokens_per_block: for state/score caches (used in compress kernels) + # state_tokens_per_block: for compressor kv/score state caches (used in compress kernels) # compress_tokens_per_block: for compressed KV cache (used in scatter) state_tokens_per_block = metadata.kv_cache_manager.tokens_per_block compress_tokens_per_block = metadata.kv_cache_manager.compressed_block_sizes[self.layer_idx] diff --git a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/deepseek_v4.py b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/deepseek_v4.py index 109910d5a8cc..5f4b04c9ea34 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/deepseek_v4.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/deepseek_v4.py @@ -35,6 +35,7 @@ from tensorrt_llm._utils import prefer_pinned from tensorrt_llm.models.modeling_utils import QuantConfig from tensorrt_llm.quantization.utils import fp8_utils +from tensorrt_llm.runtime.kv_cache_manager_v2 import DataRole from ..dsa import ( HAS_FAST_HADAMARD, @@ -55,13 +56,37 @@ class DeepseekV4AttentionType(Enum): + # attentions managed in sliding-window mode are put in the front on purpose SWA = 0 - COMPRESS = 1 - COMPRESSOR_STATE = 2 - COMPRESSOR_SCORE = 3 - INDEXER_COMPRESS = 4 - INDEXER_COMPRESSOR_STATE = 5 - INDEXER_COMPRESSOR_SCORE = 6 + COMPRESSOR_KV = 1 + COMPRESSOR_SCORE = 2 + INDEXER_COMPRESSOR_KV = 3 + INDEXER_COMPRESSOR_SCORE = 4 + + # attentions not managed in sliding-window mode + COMPRESS = 5 + INDEXER_COMPRESS = 6 + + @property + def role(self) -> DataRole: + return DataRole(f"deepseek_v4_{self.name.lower()}") + + +DEEPSEEK_V4_SLIDING_ATTENTION = ( + DeepseekV4AttentionType.SWA, + DeepseekV4AttentionType.COMPRESSOR_KV, + DeepseekV4AttentionType.COMPRESSOR_SCORE, + DeepseekV4AttentionType.INDEXER_COMPRESSOR_KV, + DeepseekV4AttentionType.INDEXER_COMPRESSOR_SCORE, +) +assert tuple(attn_type.value for attn_type in DEEPSEEK_V4_SLIDING_ATTENTION) == tuple( + range(len(DEEPSEEK_V4_SLIDING_ATTENTION)) +) + +DEEPSEEK_V4_NON_SLIDING_ATTENTION = ( + DeepseekV4AttentionType.COMPRESS, + DeepseekV4AttentionType.INDEXER_COMPRESS, +) def is_overlap_compressor(compress_ratio: int) -> bool: @@ -84,13 +109,13 @@ def compress_ratio_has_attention(compress_ratio: int, attn_type: DeepseekV4Atten return True if attn_type == DeepseekV4AttentionType.COMPRESS: return is_compress - if attn_type == DeepseekV4AttentionType.COMPRESSOR_STATE: + if attn_type == DeepseekV4AttentionType.COMPRESSOR_KV: return is_compress if attn_type == DeepseekV4AttentionType.COMPRESSOR_SCORE: return is_compress if attn_type == DeepseekV4AttentionType.INDEXER_COMPRESS: return is_sparse - if attn_type == DeepseekV4AttentionType.INDEXER_COMPRESSOR_STATE: + if attn_type == DeepseekV4AttentionType.INDEXER_COMPRESSOR_KV: return is_sparse if attn_type == DeepseekV4AttentionType.INDEXER_COMPRESSOR_SCORE: return is_sparse @@ -105,13 +130,13 @@ def get_attn_dim( return head_dim if attn_type == DeepseekV4AttentionType.COMPRESS: return head_dim - if attn_type == DeepseekV4AttentionType.COMPRESSOR_STATE: + if attn_type == DeepseekV4AttentionType.COMPRESSOR_KV: return state_factor * head_dim if attn_type == DeepseekV4AttentionType.COMPRESSOR_SCORE: return state_factor * head_dim if attn_type == DeepseekV4AttentionType.INDEXER_COMPRESS: return index_head_dim - if attn_type == DeepseekV4AttentionType.INDEXER_COMPRESSOR_STATE: + if attn_type == DeepseekV4AttentionType.INDEXER_COMPRESSOR_KV: return state_factor * index_head_dim if attn_type == DeepseekV4AttentionType.INDEXER_COMPRESSOR_SCORE: return state_factor * index_head_dim @@ -134,14 +159,18 @@ def get_token_bytes( attn_dim = get_attn_dim(head_dim, index_head_dim, compress_ratio, attn_type) dtype_bytes = 1 if has_fp8_kv_cache else 2 + # (indexer) compressor kv and score always use float32 if attn_type in [ - DeepseekV4AttentionType.COMPRESSOR_STATE, + DeepseekV4AttentionType.COMPRESSOR_KV, DeepseekV4AttentionType.COMPRESSOR_SCORE, - DeepseekV4AttentionType.INDEXER_COMPRESSOR_STATE, + DeepseekV4AttentionType.INDEXER_COMPRESSOR_KV, DeepseekV4AttentionType.INDEXER_COMPRESSOR_SCORE, ]: - dtype_bytes = 4 - + dtype_bytes = 4 # (indexer) compressor kv and score use float32 + # Indexer cache always packs data + per-block scales into one row. Only + # the two indexer presets ("fp8" blockwise / "fp4" mxfp4) are valid + # here — bf16 / fp8_pertensor are reserved for the main-attention + # compressor. if attn_type == DeepseekV4AttentionType.INDEXER_COMPRESS: if indexer_k_dtype == "fp8": return attn_dim + index_head_dim // 128 * 4 @@ -275,11 +304,11 @@ def __post_init__(self): elif compress_ratio == 4: attention_types.append((1, DeepseekV4AttentionType.SWA)) attention_types.append((compress_ratio, DeepseekV4AttentionType.COMPRESS)) - attention_types.append((compress_ratio, DeepseekV4AttentionType.COMPRESSOR_STATE)) + attention_types.append((compress_ratio, DeepseekV4AttentionType.COMPRESSOR_KV)) attention_types.append((compress_ratio, DeepseekV4AttentionType.COMPRESSOR_SCORE)) attention_types.append((compress_ratio, DeepseekV4AttentionType.INDEXER_COMPRESS)) attention_types.append( - (compress_ratio, DeepseekV4AttentionType.INDEXER_COMPRESSOR_STATE) + (compress_ratio, DeepseekV4AttentionType.INDEXER_COMPRESSOR_KV) ) attention_types.append( (compress_ratio, DeepseekV4AttentionType.INDEXER_COMPRESSOR_SCORE) @@ -287,7 +316,7 @@ def __post_init__(self): else: attention_types.append((1, DeepseekV4AttentionType.SWA)) attention_types.append((compress_ratio, DeepseekV4AttentionType.COMPRESS)) - attention_types.append((compress_ratio, DeepseekV4AttentionType.COMPRESSOR_STATE)) + attention_types.append((compress_ratio, DeepseekV4AttentionType.COMPRESSOR_KV)) attention_types.append((compress_ratio, DeepseekV4AttentionType.COMPRESSOR_SCORE)) self.attention_type_set = set(attention_types) @@ -424,21 +453,38 @@ def __post_init__(self): capture_graph=capture_graph, ) - self.block_tables = { - attention_type: self.get_empty( + # Sliding-window caches use per-layer page indices and are indexed by + # [local_layer_idx, attention_type.value, sequence, block]. COMPRESS + # uses ratio-shared page indices, and INDEXER_COMPRESS keeps a separate + # compatibility table for the generic DSA indexer path. + block_table_shape = ( + self.kv_cache_manager.num_local_layers, + len(DEEPSEEK_V4_SLIDING_ATTENTION), + self.max_num_sequences, + self.kv_cache_manager.max_blocks_per_seq, + ) + self.sliding_block_tables = self.get_empty( + self.cuda_graph_buffers, + block_table_shape, + cache_name="sliding_block_tables", + dtype=torch.int32, + capture_graph=capture_graph, + ) + + compress_block_table_shape = ( + self.max_num_sequences, + self.kv_cache_manager.max_blocks_per_seq, + ) + self.compress_block_tables = { + compress_ratio: self.get_empty( self.cuda_graph_buffers, - (self.max_num_sequences, self.kv_cache_manager.max_blocks_per_seq), - cache_name=f"block_tables_{attention_type}", + compress_block_table_shape, + cache_name=f"compress_block_tables_{compress_ratio}", dtype=torch.int32, capture_graph=capture_graph, ) - for attention_type in self.attention_type_set - } - self.host_block_tables = { - attention_type: torch.empty_like( - self.block_tables[attention_type], device="cpu", pin_memory=prefer_pinned() - ) - for attention_type in self.attention_type_set + for compress_ratio in self._compress_ratios_sorted + if is_compress_layer(compress_ratio) } # sparse_mla_topk_lens: actual token count per token for each compress_ratio (SWA + compressed) @@ -471,47 +517,37 @@ def __post_init__(self): self._init_cache_buffer_data_pointers() def prepare_for_indexer_k_cache(self): - """Optimized: bulk tensor copy instead of per-row Python loop. - - Note: must use num_contexts=0 for get_batch_attn_offset because the - indexer k cache always uses generation-style copy indices regardless - of request type. - """ - num_seqs = self.num_seqs - offsets = self.kv_cache_manager.get_batch_attn_offset( + """Prepare the shared indexer K-cache decode table for DSA kernels.""" + # INDEXER_COMPRESS uses shared page indices, so the generic DSA + # indexer path only needs one 2D block table. + self.kv_cache_manager.copy_batch_indexer_compress_block_tables( + self.host_indexer_k_cache_block_offsets, self.request_ids, - 1, # beam_width - 0, # num_contexts=0 (indexer always uses gen-style copy index) - num_seqs, - DeepseekV4AttentionType.INDEXER_COMPRESS, - DEEPSEEK_V4_SPARSE_RATIO, + beam_width=self.beam_width, + num_contexts=self.num_contexts, + num_seqs=self.num_seqs, ) - num_cols = offsets.shape[1] - self.host_indexer_k_cache_block_offsets[:num_seqs, :num_cols] = offsets[:num_seqs] - self.indexer_k_cache_block_offsets[:num_seqs].copy_( - self.host_indexer_k_cache_block_offsets[:num_seqs], + self.indexer_k_cache_block_offsets[: self.num_seqs].copy_( + self.host_indexer_k_cache_block_offsets[: self.num_seqs], non_blocking=True, ) def prepare_for_block_tables(self): - """Prepare block tables for all attention types. - - Delegates offset computation to DeepseekV4CacheManager.get_batch_block_offsets - (single get_copy_index call, deduplicated by pool), then copies to device. - """ - num_seqs = self.num_seqs - offsets_map = self.kv_cache_manager.get_batch_block_offsets( - self.request_ids, self.num_contexts, self.attention_type_set + """Prepare block tables for sliding-window and compressed attention.""" + self.kv_cache_manager.copy_batch_sliding_block_tables( + self.sliding_block_tables, + self.request_ids, + self.num_contexts, + self.num_seqs, ) - - # Phase 1: write host buffers. - for key, offsets in offsets_map.items(): - self.host_block_tables[key][:num_seqs] = offsets[:num_seqs] - - # Phase 2: batch H2D copies. - for key in self.attention_type_set: - self.block_tables[key][:num_seqs].copy_( - self.host_block_tables[key][:num_seqs], non_blocking=True + for compress_ratio, compress_block_table in self.compress_block_tables.items(): + self.kv_cache_manager.copy_batch_compress_block_tables( + compress_block_table, + self.request_ids, + compress_ratio=compress_ratio, + beam_width=self.beam_width, + num_contexts=self.num_contexts, + num_seqs=self.num_seqs, ) def prepare_for_deepseek_v4_indices(self, token_positions=None): @@ -612,10 +648,17 @@ def _init_cache_buffer_data_pointers(self): extend_compress_ratios = self.compress_ratios + [self.compress_ratios[-1]] * ( self.max_draft_tokens - 1 ) + # SWA uses PER_LAYER indices; COMPRESS uses SHARED indices. The sparse + # MLA conversion kernel receives a representative base pointer per pool + # and a per-layer buffer pointer so it can account for any layer offset. + self.sparse_mla_base_ptrs = { + 1: self.kv_cache_manager.swa_pool_ptr, + } + for ratio, compress_pool_ptr in self.kv_cache_manager.compress_pool_ptrs.items(): + self.sparse_mla_base_ptrs[ratio] = compress_pool_ptr + self.swa_buffer_ptrs = { - layer_idx: self.kv_cache_manager.get_buffers( - layer_idx, DeepseekV4AttentionType.SWA - ).data_ptr() + layer_idx: self.kv_cache_manager.swa_pool_ptr for layer_idx in self.kv_cache_manager.pp_layers } self.compressed_buffer_ptrs = { @@ -626,15 +669,15 @@ def _init_cache_buffer_data_pointers(self): if is_compress_layer(extend_compress_ratios[layer_idx]) } - # Per-ratio base pointer for sparse MLA global-index conversion. - # Ratio 1 is the SWA pool; ratios greater than 1 are compressed pools. - self.sparse_mla_base_ptrs = { - 1: self.kv_cache_manager.swa_pool_ptr, - } - for ratio, compress_pool_ptr in self.kv_cache_manager.compress_pool_ptrs.items(): - self.sparse_mla_base_ptrs[ratio] = compress_pool_ptr - def prepare(self): + assert self.kv_cache_manager is not None + assert self.request_ids is not None + + self.kv_cache_manager.compute_sliding_block_tables( + self.request_ids, + self.num_contexts, + ) + TrtllmAttentionMetadata.prepare(self) num_requests = self.num_contexts + self.num_generations @@ -662,6 +705,9 @@ def prepare(self): has_sparse_layers = DEEPSEEK_V4_SPARSE_RATIO in self.compress_ratio_set + # For block offsets + self.prepare_for_block_tables() + # For indexer k cache (only needed when sparse layers exist) if has_sparse_layers: self.prepare_for_indexer_k_cache() @@ -669,9 +715,6 @@ def prepare(self): # For spec decode self.prepare_for_spec_decode(kv_lens) - # For block offsets - self.prepare_for_block_tables() - # For DeepSeek-V4 indices self.prepare_for_deepseek_v4_indices() @@ -1193,6 +1236,16 @@ def _run_serial_indexer_prepare( self._update_k_cache_if_needed(k_fp8, k_scale, metadata) return q_fp8, q_scale, k_fp8, k_scale, weights + def _update_k_cache( + self, + k_fp8: torch.Tensor, + k_scale: torch.Tensor, + metadata: DSAtrtllmAttentionMetadata, + ) -> None: + # DSV4's indexer compressor already scatters INDEXER_COMPRESS. The + # shared DSA scatter would duplicate that write. + return + def forward( self, qr: torch.Tensor, @@ -1401,14 +1454,15 @@ def sparse_attn_predict( # Use global req_id directly req_id = metadata.req_idx_per_token[start_idx:end_idx] swa_local_indices = metadata.swa_local_indices_cuda[start_idx:end_idx] - block_table_swa = metadata.block_tables[(1, DeepseekV4AttentionType.SWA)] + local_layer_idx = kv_cache_manager.layer_offsets[layer_idx] + block_table_swa = metadata.sliding_block_tables[ + local_layer_idx, DeepseekV4AttentionType.SWA.value + ] if self.compress_ratio > 1: compressed_buffer_ptr = metadata.compressed_buffer_ptrs[layer_idx] compress_pool_base_ptr = metadata.sparse_mla_base_ptrs[self.compress_ratio] - block_table_compressed = metadata.block_tables[ - (self.compress_ratio, DeepseekV4AttentionType.COMPRESS) - ] + block_table_compressed = metadata.compress_block_tables[self.compress_ratio] if self.compress_ratio == 4: topk_indices = forward_args.topk_indices assert topk_indices is not None, "topk_indices is required when compress_ratio=4" diff --git a/tensorrt_llm/_torch/attention_backend/trtllm.py b/tensorrt_llm/_torch/attention_backend/trtllm.py index 685472d66813..36791838d9d7 100644 --- a/tensorrt_llm/_torch/attention_backend/trtllm.py +++ b/tensorrt_llm/_torch/attention_backend/trtllm.py @@ -306,10 +306,13 @@ def _post_init_with_buffers(self, buffers) -> None: ) if self.kv_cache_manager is not None: + num_attention_op_pools = getattr(self.kv_cache_manager, + "num_attention_op_pools", + self.kv_cache_manager.num_pools) self.kv_cache_block_offsets = self.get_empty( buffers, [ - self.kv_cache_manager.num_pools, self.max_num_sequences, 2, + num_attention_op_pools, self.max_num_sequences, 2, self.kv_cache_manager.max_blocks_per_seq ], cache_name="kv_cache_block_offsets", @@ -323,11 +326,13 @@ def _post_init_with_buffers(self, buffers) -> None: # Allocate separate block offset tensors for draft KV cache manager # Used in one-model speculative decoding with different KV cache layouts if self.draft_kv_cache_manager is not None: + num_draft_attention_op_pools = getattr( + self.draft_kv_cache_manager, "num_attention_op_pools", + self.draft_kv_cache_manager.num_pools) self.draft_kv_cache_block_offsets = self.get_empty( buffers, [ - self.draft_kv_cache_manager.num_pools, - self.max_num_sequences, 2, + num_draft_attention_op_pools, self.max_num_sequences, 2, self.draft_kv_cache_manager.max_blocks_per_seq ], cache_name="draft_kv_cache_block_offsets", diff --git a/tensorrt_llm/_torch/disaggregation/base/region.py b/tensorrt_llm/_torch/disaggregation/base/region.py index 249bf742bc2b..62896250173e 100644 --- a/tensorrt_llm/_torch/disaggregation/base/region.py +++ b/tensorrt_llm/_torch/disaggregation/base/region.py @@ -39,15 +39,6 @@ class MemRegionGroup(NamedTuple): bytes_per_region: int -class DataRole(IntFlag): - """Logical role(s) a memory region plays. Supports combinations.""" - - KEY = auto() - VALUE = auto() - BLOCK_QUANT = auto() - INDEXER = auto() - - class DataLayout(IntFlag): """Possible orders for storing data in memory.""" @@ -71,7 +62,6 @@ class KVRegionSpec(RegionSpec): Specifies a region within the Key/Value cache, with optional axes. """ - role: DataRole = DataRole.KEY | DataRole.VALUE heads: Optional[IndexRange] = None tokens: Optional[IndexRange] = None diff --git a/tensorrt_llm/_torch/disaggregation/native/mixers/attention/peer.py b/tensorrt_llm/_torch/disaggregation/native/mixers/attention/peer.py index ed6a8d84a887..12d7cadfee4f 100644 --- a/tensorrt_llm/_torch/disaggregation/native/mixers/attention/peer.py +++ b/tensorrt_llm/_torch/disaggregation/native/mixers/attention/peer.py @@ -8,7 +8,7 @@ SpecRegionPair, ) from tensorrt_llm._torch.disaggregation.native.rank_info import RankInfo -from tensorrt_llm._torch.disaggregation.resource.utils import PoolRole +from tensorrt_llm._torch.disaggregation.resource.page import MapperKind from tensorrt_llm._utils import nvtx_range @@ -404,7 +404,7 @@ def build_kv_mapper( self, *, peer_ri: RankInfo, - pool_role: PoolRole, + mapper_kind: MapperKind, transfer_layers: int, self_layer_offset: int, peer_layer_offset: int, @@ -419,7 +419,7 @@ def build_kv_mapper( return IdentityMapper() if head_match: - if pool_role == PoolRole.INDEXER: + if mapper_kind == MapperKind.FLAT: block_size_per_layer = self_pool_slot_bytes // self_pool_num_layers return IndexerKCacheHeadMatchMapper( transfer_layers=transfer_layers, @@ -445,7 +445,7 @@ def build_kv_mapper( slot_size_per_layer=slot_size_per_layer, ) - if pool_role == PoolRole.INDEXER: + if mapper_kind == MapperKind.FLAT: raise ValueError("IndexerKCacheHeadMatchMapper is not supported for head mismatch case") return HeadMismatchMapper( diff --git a/tensorrt_llm/_torch/disaggregation/native/peer.py b/tensorrt_llm/_torch/disaggregation/native/peer.py index 80eaa30bb59e..1902c9d14f4f 100644 --- a/tensorrt_llm/_torch/disaggregation/native/peer.py +++ b/tensorrt_llm/_torch/disaggregation/native/peer.py @@ -6,14 +6,12 @@ from tensorrt_llm._torch.disaggregation.native.mixers.attention.peer import AttentionPolicy from tensorrt_llm._torch.disaggregation.native.rank_info import RankInfo from tensorrt_llm._torch.disaggregation.resource.kv_extractor import KVRegionExtractorV1 -from tensorrt_llm._torch.disaggregation.resource.page import AttentionLayerGroup +from tensorrt_llm._torch.disaggregation.resource.page import AttentionLayerGroup, MapperKind from tensorrt_llm._torch.disaggregation.resource.utils import ( - PoolRole, get_global_layer_ids, get_layer_group_num_layers, get_layer_to_layer_group, get_physical_pool, - get_pool_role, get_pool_view_global_layer_ids, get_pool_view_num_layers, ) @@ -113,8 +111,13 @@ def get_pool_mapping(self, peer_ri: RankInfo) -> Dict[LGPoolKey, LGPoolKey]: Two-step matching: 1. Find peer layer_group via layer_to_layer_group (global_layer_id -> lg_idx). - 2. Within the matched peer layer_group, find the peer pool by matching - pool_role AND global_layer_ids overlap. + 2. Within the matched peer layer_group, find the unique peer pool whose + ``PoolView.pool_role`` equals self's and whose global_layer_ids + overlap. + + Layer-overlap is required: a peer pool with the same pool_role but + zero layer overlap with self is *not* a match — the two pools cover + disjoint layers and have nothing to transfer. """ key = self._unique_key(peer_ri.instance_name, peer_ri.instance_rank) if key in self._lg_pool_mapping_cache: @@ -132,55 +135,68 @@ def get_pool_mapping(self, peer_ri: RankInfo) -> Dict[LGPoolKey, LGPoolKey]: return mapping peer_layer_to_group = get_layer_to_layer_group(peer_pt) - assert self._ri.attention is not None - kv_factor = self._ri.attention.kv_factor for self_lg_idx, self_lg in enumerate(self_pt.layer_groups): if not isinstance(self_lg, AttentionLayerGroup): continue for self_pi, self_pv in enumerate(self_lg.pool_views): - is_indexer = len(self_pv.buffer_entries) == 0 - # For INDEXER (empty buffer_entries), use group-level IDs for step-1 lookup + # The only place mapper_kind affects pool matching: + # INDEXED → pool may cover a subset of the LG; read + # buffer_entries to find the exact layer set. + # FLAT → pool covers the entire LG by convention; + # use the LG's layer ids directly. + self_is_flat = self_pv.mapper_kind == MapperKind.FLAT pv_global_ids = ( get_global_layer_ids(self_lg) - if is_indexer + if self_is_flat else get_pool_view_global_layer_ids(self_pv, self_lg) ) if not pv_global_ids: continue - # Step 1: find peer layer_group via any overlapping global_layer_id - peer_lg_idx = None - for glid in pv_global_ids: - if glid in peer_layer_to_group: - peer_lg_idx = peer_layer_to_group[glid] - break + # Step 1: find peer layer_group via any overlapping global_layer_id. + peer_lg_idx = next( + (peer_layer_to_group[g] for g in pv_global_ids if g in peer_layer_to_group), + None, + ) if peer_lg_idx is None: continue peer_lg = peer_pt.layer_groups[peer_lg_idx] - # Step 2: find peer pool within group by matching pool_role + layer overlap - self_pool_role = ( - PoolRole.INDEXER if is_indexer else get_pool_role(self_pv, kv_factor=kv_factor) - ) + # Step 2: pick the first peer pool with the same pool_role + # whose layers overlap self's (zero-overlap pools cover + # disjoint layers — nothing to transfer). + # + # Uniqueness assumption: at most one peer pool can match on + # both ``pool_role`` (frozenset equality) and layer overlap. + # We do *not* assume ``pool_role`` is unique within a peer LG + # — V2 may split an LG into multiple same-role pools by + # buffer-size class (e.g. VSWA). What we rely on is that both + # peers run the same pool-grouping logic, so for every self_pv + # there is exactly one peer pool with the same role *and* an + # overlapping layer set; other same-role peer pools cover + # disjoint layers and fall out via the overlap filter. self_layer_set = set(pv_global_ids) matched_peer_pi = None for peer_pi, peer_pv in enumerate(peer_lg.pool_views): - peer_is_indexer = len(peer_pv.buffer_entries) == 0 - peer_pool_role = ( - PoolRole.INDEXER - if peer_is_indexer - else get_pool_role(peer_pv, kv_factor=kv_factor) + if peer_pv.pool_role != self_pv.pool_role: + continue + peer_global_ids = ( + get_global_layer_ids(peer_lg) + if peer_pv.mapper_kind == MapperKind.FLAT + else get_pool_view_global_layer_ids(peer_pv, peer_lg) ) - if peer_pool_role != self_pool_role: + if not set(peer_global_ids) & self_layer_set: continue - if is_indexer: - # INDEXER pools match by role alone - matched_peer_pi = peer_pi - break - if set(get_pool_view_global_layer_ids(peer_pv, peer_lg)) & self_layer_set: - matched_peer_pi = peer_pi - break + if peer_pv.mapper_kind != self_pv.mapper_kind: + raise ValueError( + "PeerRegistrar.get_pool_mapping: incompatible mapper " + f"kinds for pool role {sorted(self_pv.pool_role)} " + f"(local={self_pv.mapper_kind.name}, " + f"peer={peer_pv.mapper_kind.name}, peer_pool={peer_pi})" + ) + matched_peer_pi = peer_pi + break if matched_peer_pi is not None: mapping[(self_lg_idx, self_pi)] = (peer_lg_idx, matched_peer_pi) @@ -218,21 +234,27 @@ def get_kv_map( peer_pv = peer_lg.pool_views[peer_pi] assert self._ri.attention is not None - kv_factor = self._ri.attention.kv_factor - is_indexer = len(self_pv.buffer_entries) == 0 - self_pool_role = ( - PoolRole.INDEXER if is_indexer else get_pool_role(self_pv, kv_factor=kv_factor) - ) + if self_pv.mapper_kind != peer_pv.mapper_kind: + raise ValueError( + "PeerRegistrar.get_kv_map: incompatible mapper kinds " + f"(local={self_pv.mapper_kind.name}, peer={peer_pv.mapper_kind.name})" + ) - # For INDEXER (empty buffer_entries), use group-level global layer IDs - if is_indexer: - self_global_ids = get_global_layer_ids(self_lg) - peer_global_ids = get_global_layer_ids(peer_lg) + # FLAT pools carry no per-buffer layer info, so layer ids and + # layer count come from the layer_group itself. + # + # Sort by global_layer_id so that ``.index(first_overlap_layer)`` + # below returns the layer's slot position. This relies on the + # convention that managers (V1 / V2 / DSv4) assign global_layer_id + # monotonically with the layer's byte offset in the slot. + if self_pv.mapper_kind == MapperKind.FLAT: + self_global_ids = sorted(get_global_layer_ids(self_lg)) + peer_global_ids = sorted(get_global_layer_ids(peer_lg)) self_num_layers = get_layer_group_num_layers(self_lg) peer_num_layers = get_layer_group_num_layers(peer_lg) else: - self_global_ids = get_pool_view_global_layer_ids(self_pv, self_lg) - peer_global_ids = get_pool_view_global_layer_ids(peer_pv, peer_lg) + self_global_ids = sorted(get_pool_view_global_layer_ids(self_pv, self_lg)) + peer_global_ids = sorted(get_pool_view_global_layer_ids(peer_pv, peer_lg)) self_num_layers = get_pool_view_num_layers(self_pv) peer_num_layers = get_pool_view_num_layers(peer_pv) @@ -252,7 +274,7 @@ def get_kv_map( mapper = self._attention_policy.build_kv_mapper( peer_ri=peer_ri, - pool_role=self_pool_role, + mapper_kind=self_pv.mapper_kind, transfer_layers=transfer_layers, self_layer_offset=self_layer_offset, peer_layer_offset=peer_layer_offset, diff --git a/tensorrt_llm/_torch/disaggregation/resource/cache_reuse.py b/tensorrt_llm/_torch/disaggregation/resource/cache_reuse.py index 3a0f59a87651..3fbbe070c44f 100644 --- a/tensorrt_llm/_torch/disaggregation/resource/cache_reuse.py +++ b/tensorrt_llm/_torch/disaggregation/resource/cache_reuse.py @@ -142,12 +142,7 @@ def get_block_ids(self, req, group_idx, lg): # noqa: ARG002 ) def commit_blocks_for_reuse(self, req: LlmRequest) -> None: - if not self.enable_block_reuse: - return - kv_cache = self._mgr.kv_cache_map.get(req.py_request_id) - if kv_cache is None: - return - self._mgr.try_commit_blocks_for_reuse(req, kv_cache) + self._mgr.try_commit_blocks(req) def create_cache_reuse_adapter( diff --git a/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py b/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py index 7a6674cd8e48..0e1c0f7274de 100644 --- a/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py +++ b/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py @@ -4,7 +4,6 @@ from tensorrt_llm._torch.disaggregation.base.region import ( DataLayout, - DataRole, MemRegionGroup, RegionExtractorBase, SpecRegion, @@ -16,6 +15,7 @@ LayerGroup, LocalLayer, MambaLayerGroup, + MapperKind, PhysicalPool, PhysicalPoolGroup, PoolView, @@ -183,16 +183,24 @@ def build_page_table(kv_cache_manager: KVCacheManager) -> KVCachePageTable: slot_bytes = stride * len(local_layer_ids) entries = [] + kv_role_names: set[str] = {"key"} + if not is_key_only: + kv_role_names.add("value") for i, lid in enumerate(local_layer_ids): base_offset = i * stride - entries.append((lid, int(DataRole.KEY), base_offset, buffer_size)) + entries.append((lid, base_offset, buffer_size)) if not is_key_only: - entries.append((lid, int(DataRole.VALUE), base_offset + buffer_size, buffer_size)) + entries.append((lid, base_offset + buffer_size, buffer_size)) kv_physical = PhysicalPool( base_address=base_addr, slot_bytes=slot_bytes, num_slots=num_blocks ) - kv_view = PoolView(pool_idx=0, buffer_entries=np.array(entries, dtype=BUFFER_ENTRY_DTYPE)) + kv_view = PoolView( + pool_idx=0, + buffer_entries=np.array(entries, dtype=BUFFER_ENTRY_DTYPE), + pool_role=frozenset(kv_role_names), + mapper_kind=MapperKind.INDEXED, + ) physical_pools = [kv_physical] pool_views = [kv_view] @@ -211,7 +219,10 @@ def build_page_table(kv_cache_manager: KVCacheManager) -> KVCachePageTable: num_slots=num_blocks, ) indexer_view = PoolView( - pool_idx=1, buffer_entries=np.array([], dtype=BUFFER_ENTRY_DTYPE) + pool_idx=1, + buffer_entries=np.array([], dtype=BUFFER_ENTRY_DTYPE), + pool_role=frozenset({"indexer_k"}), + mapper_kind=MapperKind.FLAT, ) physical_pools.append(indexer_physical) pool_views.append(indexer_view) @@ -285,8 +296,9 @@ def _build_page_table_v2(manager) -> KVCachePageTable: """Build a KVCachePageTable from a KVCacheManagerV2. Uses the V2 storage layer APIs (pool.slot_address, pool.slot_size, - pool.num_slots) for accurate pool metadata, and determines PoolRole - from the DataRole of buffers in each pool. + pool.num_slots) for accurate pool metadata, and stamps each PoolView + with the manager's native role-name strings (``pool_role``) plus the + closed-set ``mapper_kind`` discriminator used by ``build_kv_mapper``. Important: iterates over life cycles (layer groups), not storage pool groups. Multiple life cycles with different sliding-window sizes may @@ -296,22 +308,8 @@ def _build_page_table_v2(manager) -> KVCachePageTable: """ from collections import defaultdict - from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import Role from tensorrt_llm.runtime.kv_cache_manager_v2 import CacheTier - _ROLE_STR_TO_ENUM: dict[str, DataRole] = { - Role.KEY: DataRole.KEY, - Role.VALUE: DataRole.VALUE, - Role.KEY_BLOCK_SCALE: DataRole.KEY | DataRole.BLOCK_QUANT, - Role.VALUE_BLOCK_SCALE: DataRole.VALUE | DataRole.BLOCK_QUANT, - } - - def _role_str_to_enum(role: str) -> DataRole: - if role not in _ROLE_STR_TO_ENUM: - valid_roles = list(_ROLE_STR_TO_ENUM.keys()) - raise ValueError(f"Invalid role: '{role}'. Valid roles: {valid_roles}") - return _ROLE_STR_TO_ENUM[role] - storage = manager.impl._storage config = manager.impl._init_config @@ -322,17 +320,20 @@ def _role_str_to_enum(role: str) -> DataRole: gpu_level = level_idx break - # Collect buffer entries keyed by (life_cycle_id, pool_idx) + # Collect buffer entries keyed by (life_cycle_id, pool_idx). + # Also collect the set of native role-name strings per pool — used as + # ``PoolView.pool_role``, the manager-supplied equivalence label that + # disagg uses to match pools across peers without enumerating roles. buffer_by_lc_pool: Dict[tuple, list] = defaultdict(list) + native_roles_by_pool: Dict[tuple, set] = defaultdict(set) for buffer_id, attr in storage._buffer_attr.items(): layer_id, role = buffer_id lc_id = attr.life_cycle_id pool_idx = attr.pool_index pool_key = (int(lc_id), pool_idx) - buffer_by_lc_pool[pool_key].append( - (layer_id, _role_str_to_enum(role), attr.offset, attr.size) - ) + buffer_by_lc_pool[pool_key].append((layer_id, attr.offset, attr.size)) + native_roles_by_pool[pool_key].add(str(role)) # Iterate over life cycles (layer groups), not storage pool groups. # Multiple layer_groups can share the same storage pool_group when their @@ -398,6 +399,8 @@ def _role_str_to_enum(role: str) -> DataRole: PoolView( pool_idx=pool_idx, buffer_entries=np.array(buffers_info, dtype=BUFFER_ENTRY_DTYPE), + pool_role=frozenset(native_roles_by_pool[pool_key]), + mapper_kind=MapperKind.INDEXED, ) ) diff --git a/tensorrt_llm/_torch/disaggregation/resource/page.py b/tensorrt_llm/_torch/disaggregation/resource/page.py index 8e96e3889745..81514c06d6a0 100644 --- a/tensorrt_llm/_torch/disaggregation/resource/page.py +++ b/tensorrt_llm/_torch/disaggregation/resource/page.py @@ -1,20 +1,48 @@ from __future__ import annotations from dataclasses import dataclass, field -from typing import Dict, List, Optional +from enum import IntEnum +from typing import Dict, FrozenSet, List, Optional import numpy as np BUFFER_ENTRY_DTYPE = np.dtype( [ ("local_layer_id", np.uint32), - ("role", np.uint32), ("offset", np.uint32), ("size", np.uint32), ] ) +class MapperKind(IntEnum): + """Slot metadata shape — selects how disagg derives the pool's layer set. + + INDEXED: PoolView.buffer_entries lists ``(local_layer_id, offset, size)`` + per buffer. Disagg reads ``local_layer_id`` to know *which* layers + from the LG live in this pool (a pool may cover a subset when V2 + splits an LG into multiple pools by buffer-size class). The + ``offset`` / ``size`` columns are carried for future use but are not + currently consumed at byte-transfer time. + FLAT: PoolView.buffer_entries is empty. Disagg assumes the pool + covers *all* layers of the LG, packed equal-sized in + ``local_layers`` order. Used today by the DSA (DeepSeek Sparse + Attention, v3.2) indexer K cache pool, whose slot layout is a dense + ``(numLayers, kvFactor, blockSize)`` array. + + Byte arithmetic is the same for both kinds: per-layer stride is + ``slot_bytes // num_layers``. The kind only affects how disagg discovers + the pool's layer set during pool matching. + + Mamba state pools do not use this enum: Mamba's transfer is dispatched + through :class:`MambaPolicy` which hard-codes the ``is_conv`` switch and + bypasses the attention pool-matching path entirely. + """ + + INDEXED = 0 + FLAT = 1 + + @dataclass class PhysicalPool: base_address: int # uint64 @@ -74,15 +102,32 @@ def from_dict(data: dict) -> "LocalLayer": class PoolView: """ Per-layer-group view of a physical pool (slot layout for this life cycle). + + Fields: + pool_idx: Index of the physical pool within its pool group. + buffer_entries: Structured array using ``BUFFER_ENTRY_DTYPE``. Each + entry records a buffer's ``local_layer_id`` and its byte ``offset`` + and ``size`` within the pool slot. FLAT pools have no entries. + pool_role: Set of native role-name strings (whatever the cache manager + uses, e.g. ``"key"`` / ``"value"`` / ``"deepseek_v4_swa"``) that + live in this pool. Used as the *equivalence label* for peer-to-peer + pool matching: two pools match iff their ``pool_role`` frozensets + are equal. Disagg never enumerates the role-name vocabulary — + adding a new role on the manager side requires no disagg change. + mapper_kind: Closed-set discriminator for picking the Mapper family. """ pool_idx: int buffer_entries: np.ndarray # dtype=BUFFER_ENTRY_DTYPE + pool_role: FrozenSet[str] = field(default_factory=frozenset) + mapper_kind: MapperKind = MapperKind.INDEXED def to_dict(self) -> dict: return { "pool_idx": int(self.pool_idx), "buffer_entries": self.buffer_entries.tolist(), + "pool_role": sorted(self.pool_role), + "mapper_kind": int(self.mapper_kind), } @staticmethod @@ -96,6 +141,8 @@ def from_dict(data: dict) -> "PoolView": [tuple(row) for row in raw], dtype=BUFFER_ENTRY_DTYPE, ), + pool_role=frozenset(data["pool_role"]), + mapper_kind=MapperKind(int(data["mapper_kind"])), ) diff --git a/tensorrt_llm/_torch/disaggregation/resource/utils.py b/tensorrt_llm/_torch/disaggregation/resource/utils.py index c97dd2721325..ee03e2171a22 100644 --- a/tensorrt_llm/_torch/disaggregation/resource/utils.py +++ b/tensorrt_llm/_torch/disaggregation/resource/utils.py @@ -1,21 +1,9 @@ from __future__ import annotations -from enum import Enum, auto from typing import Dict, List, Set -from tensorrt_llm._torch.disaggregation.base.region import DataRole as RegionDataRole - from .page import AttentionLayerGroup, KVCachePageTable, MambaLayerGroup, PhysicalPool, PoolView - -class PoolRole(Enum): - """Logical role of a memory pool within a layer group.""" - - KV_CACHE = auto() - KV_BLOCK_SCALE = auto() - INDEXER = auto() - - # ------------------------------------------------------------------------- # PhysicalPool helpers # ------------------------------------------------------------------------- @@ -43,11 +31,6 @@ def get_unique_layers(pool_view: PoolView) -> Set[int]: return {int(e["local_layer_id"]) for e in pool_view.buffer_entries} -def get_unique_roles(pool_view: PoolView) -> Set[int]: - """Unique role values in *pool_view*.""" - return {int(e["role"]) for e in pool_view.buffer_entries} - - def get_num_buffer_entries(pool_view: PoolView) -> int: """Number of buffer entries.""" return len(pool_view.buffer_entries) @@ -75,42 +58,6 @@ def get_pool_view_global_layer_ids( ] -def get_pool_role(pool_view: PoolView, *, kv_factor: int) -> PoolRole: - """ - Infer :class:`PoolRole` from the DataRole values in *pool_view* - - Raises ``ValueError`` if *pool_view* has no buffer entries — the caller - must handle INDEXER pools (``len(pool_view.buffer_entries) == 0``) before - invoking this function. - """ - entries = pool_view.buffer_entries - if entries is None or len(entries) == 0: - raise ValueError( - "get_pool_role called on a PoolView with empty buffer_entries. " - "Check for INDEXER pools (len(pool_view.buffer_entries) == 0) " - "before calling this function." - ) - - roles = {int(entry["role"]) for entry in entries} - - has_key = int(RegionDataRole.KEY) in roles - has_value = int(RegionDataRole.VALUE) in roles - has_key_bq = int(RegionDataRole.KEY | RegionDataRole.BLOCK_QUANT) in roles - has_value_bq = int(RegionDataRole.VALUE | RegionDataRole.BLOCK_QUANT) in roles - - if has_key_bq or has_value_bq: - return PoolRole.KV_BLOCK_SCALE - if has_key and has_value: - return PoolRole.KV_CACHE - if has_key and not has_value: - if int(kv_factor) == 1: - return PoolRole.KV_CACHE - raise ValueError("kv_factor != 1 but pool has only KEY without VALUE") - if has_value and not has_key: - raise ValueError("pool has only VALUE without KEY") - raise ValueError(f"Unrecognized role combination in pool buffer_entries: {roles}") - - # ------------------------------------------------------------------------- # LayerGroup helpers # ------------------------------------------------------------------------- @@ -143,27 +90,6 @@ def get_physical_pool(page_table: KVCachePageTable, lg_idx: int, pool_idx: int) return page_table.pool_groups[int(lg.pool_group_idx)].pools[int(pool_idx)] -def get_device_pointer( - page_table: KVCachePageTable, - *, - lg_idx: int, - pool_view: PoolView, - slot_id: int, - local_layer_id: int, - role: int, -) -> int: - """ - Compute the device pointer for a specific buffer entry - """ - pool = get_physical_pool(page_table, lg_idx, int(pool_view.pool_idx)) - if slot_id >= pool.num_slots: - raise ValueError(f"slot_id {slot_id} >= num_slots {pool.num_slots}") - for e in pool_view.buffer_entries: - if int(e["local_layer_id"]) == int(local_layer_id) and int(e["role"]) == int(role): - return int(pool.base_address) + int(slot_id) * int(pool.slot_bytes) + int(e["offset"]) - raise ValueError(f"Buffer not found: local_layer_id={local_layer_id}, role={role}") - - # ------------------------------------------------------------------------- # NIXL memory registration helpers # ------------------------------------------------------------------------- diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index a5f11093e17b..6de5c0a64222 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -183,7 +183,7 @@ def _create_kv_slice( ) if token_range is None and req.prompt_len > 0: - # Align with KV cache allocation (resize_context / + # Align with KV cache allocation (prepare_disagg_gen_init / # _get_context_bytes), which reserves prompt_len + # num_extra_kv_tokens slots for speculative decoding methods # (e.g. EAGLE3) that consume extra KV positions per request. @@ -488,7 +488,7 @@ def request_and_receive_sync(self, req: LlmRequest): if result == WaitResult.COMPLETED: if self._need_aux_transfer(req): self._apply_aux(session, req) - self._trim_kv_to_prompt_history(req) + self._assert_disagg_history_declared(req) req.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE else: req.state = LlmRequestState.DISAGG_TRANS_ERROR @@ -627,7 +627,7 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]): req = self._recv_reqs[rid] if self._need_aux_transfer(req): self._apply_aux(session, req) - self._trim_kv_to_prompt_history(req) + self._assert_disagg_history_declared(req) req.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE session.close() del self._recv_reqs[rid] @@ -658,35 +658,39 @@ def _poll_gen_sessions_for_poll_interval(self, wait_num: int) -> None: def check_gen_transfer_complete(self): return len(self._recv_sessions) == 0 - def _trim_kv_to_prompt_history(self, req: LlmRequest) -> None: - """Mark received KV as historic so SWA pools release pre-window blocks. - - Call right before the TRANS_COMPLETE state transition. The cache - was sized to hold the full prompt by ``resize_context`` and just - got fully populated by the transfer; setting ``history_length`` to - the prompt length triggers ``_unlock_stale_blocks`` inside V2's - ``resize()`` for any sliding-window life cycle, releasing blocks - before the window back to their pool group. - - This closes the gap between transfer completion and - ``update_resources`` (which only runs after the *first* forward - pass and would otherwise be the first thing to update - ``history_length``). In benchmark fill-phase the first forward - is gated until every disagg-gen request is ready, so without - this trim the SWA / sparse-attn pool groups stay 100% occupied - with pre-window prompt blocks and the V2 scheduler deadlocks - on the next ``resize(+1)``. - - No-op for V1 managers and for V2 caches with only full-context - life cycles. + def _assert_disagg_history_declared(self, req: LlmRequest) -> None: + """Verify the V2 scheduler pre-declared prompt_len as history. + + Call right before the TRANS_COMPLETE state transition. The V2 + scheduler's ``_try_schedule_disagg_gen_init`` calls + ``prepare_disagg_gen_init``, which sets ``kv_cache.history_length`` + to ``prompt_len`` at allocation time so SWA stale computation + skips pre-window blocks. If that contract is violated, SWA / + sparse-attn pools may fill with pre-window prompt blocks and the + V2 scheduler can deadlock under high concurrency (e.g., benchmark + fill-phase). + + No-op for V1 managers (which lack ``get_history_length``) and for + V2 caches with only full-context life cycles (where the watermark + has no allocation effect). """ - trim = getattr(self._kv_cache_manager, "trim_to_history", None) - if trim is None: + get_history = getattr(self._kv_cache_manager, "get_history_length", None) + if get_history is None: return prompt_len = getattr(req, "prompt_len", None) if not prompt_len or prompt_len <= 0: return - trim(req, prompt_len) + history = get_history(req) + if history is None: + # Cache was already released (e.g., cancelled mid-transfer); nothing to verify. + return + if history < prompt_len: + raise RuntimeError( + f"req {req.py_request_id}: kv_cache.history_length={history} " + f"< prompt_len={prompt_len} at TRANS_COMPLETE boundary. " + f"V2 scheduler must call prepare_disagg_gen_init() in " + f"_try_schedule_disagg_gen_init." + ) def cancel_request(self, req: LlmRequest) -> bool: """Cancel the transfer for the given request. diff --git a/tensorrt_llm/_torch/models/modeling_deepseekv4.py b/tensorrt_llm/_torch/models/modeling_deepseekv4.py index 022cc877b796..2fb611ecb999 100644 --- a/tensorrt_llm/_torch/models/modeling_deepseekv4.py +++ b/tensorrt_llm/_torch/models/modeling_deepseekv4.py @@ -2453,7 +2453,13 @@ def forward( class DeepseekV4ForCausalLM(SpecDecOneEngineForCausalLM[DeepseekV4Model, PretrainedConfig]): @classmethod def get_model_defaults(cls, llm_args: "TorchLlmArgs") -> dict: - return {"kv_cache_config": {"tokens_per_block": 128}} + return { + "kv_cache_config": { + "tokens_per_block": 128, + "use_kv_cache_manager_v2": True, + "enable_swa_scratch_reuse": True, + } + } def __init__(self, model_config: ModelConfig[PretrainedConfig]): self.mapping_with_cp = None diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index fc8ce77fe9fe..b2d45dec50aa 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -1,3 +1,17 @@ +# Copyright (c) 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. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + import copy import dataclasses import os @@ -148,9 +162,10 @@ def get_kv_cache_manager_cls( # KVCacheManager.get_cache_size_per_token may return either an ``int`` # (legacy proportional model ``bytes = slope * tokens``) or an affine # ``(slope, intercept)`` tuple (CppMambaHybridCacheManager, where mamba -# state introduces a per-batch fixed cost). CacheCost normalizes both -# shapes so the rest of the file does plain attribute access and method -# calls instead of branching on type. +# state introduces a per-batch fixed cost). KVCacheManagerV2 reports +# sliding-window attention fixed cost in the tuple intercept. CacheCost +# normalizes the combined shape so the rest of the file does plain attribute +# access and method calls instead of branching on type. @dataclasses.dataclass(frozen=True) @@ -233,6 +248,7 @@ def __init__( self._mapping = mapping self._kv_cache_config = kv_cache_config self._max_kv_tokens_in = self._kv_cache_config.max_tokens + self._max_gpu_total_bytes_in = self._kv_cache_config.max_gpu_total_bytes self._max_num_tokens = max_num_tokens self._max_beam_width = max_beam_width self._kv_connector_manager = kv_connector_manager @@ -251,6 +267,8 @@ def __init__( self._execution_stream = execution_stream self._kv_cache_manager_cls = self._get_model_kv_cache_manager_cls( model_engine) + self._is_kv_cache_manager_v2 = issubclass(self._kv_cache_manager_cls, + KVCacheManagerV2) self._draft_config = draft_config self._skip_est = skip_est @@ -335,6 +353,10 @@ def _fallback_if_unsupported_kv_cache_manager_v2( return KVCacheManager return kv_cache_manager_cls + 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, @@ -347,8 +369,10 @@ def _per_manager_cache_cost(self, model_config, self._mapping, tokens_per_block=self._tokens_per_block, + max_seq_len=self._max_seq_len, max_batch_size=self._max_batch_size, kv_cache_config=kv_cache_config, + spec_config=self._speculative_config, **extra_kwargs)) def _get_kv_size_per_token(self, @@ -601,7 +625,7 @@ def _get_token_num_for_estimation(self) -> int: # heterogeneous layer_types) uses MambaHybridCacheManager and would # have its max_tokens estimate inflated incorrectly otherwise. num_pool_groups = 1 - if self._kv_cache_manager_cls == KVCacheManagerV2: + if self._is_kv_cache_manager_v2: model_cfg = self._model_engine.model.model_config.pretrained_config layer_types = getattr(model_cfg, "layer_types", None) if isinstance(layer_types, (list, tuple)): @@ -614,18 +638,22 @@ def _get_token_num_for_estimation(self) -> int: set(self._kv_cache_config.max_attention_window)) num_cache_blocks *= num_pool_groups - free_mem, total_mem = torch.cuda.mem_get_info() + # Multiply by beam width, to prevent rescaling of the max_seq_len caused by the influence of beam width during the preparation for kv_cache_estimation + max_num_tokens_for_estimation = ( + num_cache_blocks * self._tokens_per_block * + self._dummy_reqs[0].sampling_config.beam_width) + # V2 capacity is controlled by max_gpu_total_bytes; max_tokens only + # describes the dummy workload needed for estimation. + if self._is_kv_cache_manager_v2: + return max_num_tokens_for_estimation + + free_mem, _ = torch.cuda.mem_get_info() max_memory = self._kv_cache_config.free_gpu_memory_fraction * free_mem kv_size_per_token = self._get_kv_size_per_token() max_num_tokens_in_memory = ( kv_size_per_token.tokens_for_budget(max_memory) // self._tokens_per_block * self._tokens_per_block) - - # Multiply by beam width, to prevent rescaling of the max_seq_len caused by the influence of beam width during the preparation for kv_cache_estimation - return min( - num_cache_blocks * self._tokens_per_block * - self._dummy_reqs[0].sampling_config.beam_width, - max_num_tokens_in_memory) + return min(max_num_tokens_for_estimation, max_num_tokens_in_memory) def try_prepare_estimation(self) -> bool: """Prepare for possible KV cache capacity estimation. @@ -635,19 +663,37 @@ def try_prepare_estimation(self) -> bool: """ if self._skip_est: return False - estimating_kv_cache = False - if 'cp_type' not in self._mapping.cp_config: - estimating_kv_cache = True - estimate_max_tokens = self._get_token_num_for_estimation() - self._kv_cache_config.max_tokens = min( - estimate_max_tokens, self._kv_cache_config.max_tokens - ) if self._kv_cache_config.max_tokens is not None else estimate_max_tokens + + estimating_kv_cache = True + if 'cp_type' in self._mapping.cp_config: + estimating_kv_cache = False + logger.info( + "KV cache size estimation is not supported for context parallelism, disable it." + ) model_config = self._model_engine.model.model_config if model_config.attn_backend == "VANILLA": + estimating_kv_cache = False logger.info( "KV cache size estimation is not supported for Vanilla attention backend, disable it." ) - estimating_kv_cache = False + + if estimating_kv_cache: + estimate_max_tokens = self._get_token_num_for_estimation() + max_tokens = min( + estimate_max_tokens, self._kv_cache_config.max_tokens + ) if self._kv_cache_config.max_tokens is not None else estimate_max_tokens + if self._is_kv_cache_manager_v2: + free_mem, _ = torch.cuda.mem_get_info() + max_gpu_total_bytes = int( + self._kv_cache_config.free_gpu_memory_fraction * free_mem) + if (self._max_gpu_total_bytes_in is not None + and self._max_gpu_total_bytes_in > 0): + max_gpu_total_bytes = min(max_gpu_total_bytes, + self._max_gpu_total_bytes_in) + self._kv_cache_config.max_gpu_total_bytes = max_gpu_total_bytes + self._kv_cache_config.max_tokens = max_tokens + else: + self._kv_cache_config.max_tokens = max_tokens return estimating_kv_cache def configure_kv_cache_capacity(self, @@ -768,7 +814,7 @@ def configure_kv_cache_capacity(self, # This leaves max_tokens as a user-defined constraint. # ---------------------------handle max_tokens--------------------------------- - if issubclass(self._kv_cache_manager_cls, KVCacheManagerV2): + if self._is_kv_cache_manager_v2: # KVCacheManagerV2 doesn't rely on max_tokens to control capacity, so restore user provided value self._kv_cache_config.max_tokens = self._max_kv_tokens_in else: @@ -799,11 +845,12 @@ def configure_kv_cache_capacity(self, # ---------------------------handle max_gpu_total_bytes--------------------------------- # if user provided max_gpu_total_bytes, set max memory from max_gpu_total_bytes - if self._kv_cache_config.max_gpu_total_bytes > 0: + if (self._max_gpu_total_bytes_in is not None + and self._max_gpu_total_bytes_in > 0): kv_cache_max_memory = min(kv_cache_max_memory, - self._kv_cache_config.max_gpu_total_bytes) + self._max_gpu_total_bytes_in) logger.info( - f"max_gpu_total_bytes={self._kv_cache_config.max_gpu_total_bytes / (GB):.2f} GiB is provided. New max memory is {kv_cache_max_memory / (GB):.2f} GiB" + f"max_gpu_total_bytes={self._max_gpu_total_bytes_in / (GB):.2f} GiB is provided. New max memory is {kv_cache_max_memory / (GB):.2f} GiB" ) logger.info( @@ -852,6 +899,8 @@ def _create_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, layer_mask=spec_dec_layer_mask, is_disagg=self._is_disagg, @@ -981,6 +1030,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, # One-model draft specific overrides model_config=effective_draft_config, @@ -1361,7 +1412,7 @@ def _needs_gpu_kv_cache_budget_split( kv_cache_config: Optional[KvCacheConfig] = None, ) -> bool: """Whether max_gpu_total_bytes must be split per manager.""" - if issubclass(self._kv_cache_manager_cls, KVCacheManagerV2): + if self._is_kv_cache_manager_v2: return self._should_create_separate_draft_kv_cache() kv_cache_config = (kv_cache_config if kv_cache_config is not None else self._kv_cache_config) @@ -1401,8 +1452,7 @@ def build_managers(self, "max_gpu_total_bytes", self_kv_cache_config, draft_kv_cache_config)) # KVCacheManagerV2 does not support two-model draft budget splitting. - v2_two_model = (issubclass(self._kv_cache_manager_cls, - KVCacheManagerV2) + v2_two_model = (self._is_kv_cache_manager_v2 and self._draft_model_engine is not None) if not v2_two_model: # Each manager sizes its host pool from host_cache_size directly. @@ -1428,7 +1478,7 @@ def build_managers(self, # Two-model speculative decoding: draft model has separate engine if self._draft_model_engine is not None: - if issubclass(self._kv_cache_manager_cls, KVCacheManagerV2): + if self._is_kv_cache_manager_v2: assert draft_kv_cache_config is None, ( "KVCacheManagerV2 does not support two-model speculative " "decoding with separate draft KV cache budget splitting.") @@ -1519,6 +1569,7 @@ def _create_kv_cache_manager( max_beam_width: int, kv_connector_manager: Optional[KvCacheConnectorManager], 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, @@ -1644,6 +1695,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( @@ -1669,6 +1723,7 @@ def _create_kv_cache_manager( execution_stream=execution_stream, layer_mask=layer_mask, is_disagg=is_disagg, + **manager_extra_kwargs, ) elif is_nemotron_hybrid(config): if max_beam_width > 1: @@ -1774,6 +1829,7 @@ def _create_kv_cache_manager( model_type="nemotron_hybrid", use_replay_state_update=use_replay, mamba_ssm_stochastic_rounding=mamba_ssm_stochastic_rounding, + **manager_extra_kwargs, ) elif is_qwen3_hybrid(config): if max_beam_width > 1: @@ -1818,6 +1874,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 @@ -1860,6 +1917,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_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py index 3ec74e1fd72b..295372c4034a 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py @@ -20,6 +20,7 @@ from typing import TYPE_CHECKING, Dict, Iterable, List, NamedTuple, Optional, Sequence, Tuple, Union import torch +from strenum import StrEnum from tensorrt_llm._torch.distributed.communicator import Distributed, ReduceOp from tensorrt_llm._utils import ( @@ -45,7 +46,9 @@ HostCacheTierConfig, KVCacheIterationStatsDelta, LayerId, + PageIndexMode, ReuseScope, + SwaScratchReuseConfig, TokenIdExt, _KVCache, ) @@ -73,6 +76,7 @@ from ..._utils import binding_to_torch_dtype, mpi_rank, nvtx_range, str_dtype_to_torch from ...logger import logger from ...mapping import CpType, Mapping +from ..utils import maybe_compile from .connectors.kv_cache_connector import KvCacheConnectorManager from .kv_cache_stats import ( KVCacheV2IterationStatsReport, @@ -126,6 +130,109 @@ class Role: ALL = DataRole("all") +class BlockReusePolicy(StrEnum): + ALL_REUSABLE = "all_reusable" + PER_REQUEST = "per_request" + + +def _estimate_full_attn_size_per_token( + layer_sizes: Sequence[int], attention_windows: Sequence[Optional[int]] +) -> int: + return sum( + layer_size + for layer_size, window_size in zip(layer_sizes, attention_windows) + if window_size is None or window_size <= 0 + ) + + +def _estimate_swa_cache_size( + layer_sizes: Sequence[int], + attention_windows: Sequence[Optional[int]], + tokens_per_block: int, + *, + context: bool, + scratch: bool, +) -> tuple[int, int]: + tokens_per_block = int(tokens_per_block) + size_per_token = 0 + size_per_request = 0 + scratch_keys = set() + for layer_size, window_size in zip(layer_sizes, attention_windows): + if window_size is not None and window_size > 0: + window_tokens = math.ceil(window_size / tokens_per_block) * tokens_per_block + if not context: + size_per_request += window_tokens * layer_size + elif not scratch: + size_per_token += layer_size + else: + scratch_key = (int(window_size), layer_size) + if scratch_key in scratch_keys: + size_per_request += window_tokens * layer_size + else: + scratch_keys.add(scratch_key) + size_per_token += layer_size + return size_per_token, size_per_request + + +def _get_static_cache_size_layer_components( + model_config: ModelConfigPython, + mapping: Mapping, + num_layers: Optional[int] = None, + **kwargs, +) -> tuple[List[int], List[Optional[int]]]: + config = model_config.pretrained_config + max_seq_len = kwargs.get("max_seq_len") + kv_cache_config = kwargs.get("kv_cache_config") + + num_key_value_heads = getattr(config, "num_key_value_heads", config.num_attention_heads) + if isinstance(num_key_value_heads, Iterable): + num_key_value_heads = sum(num_key_value_heads) / len(num_key_value_heads) + + mla = hasattr(config, "kv_lora_rank") and config.kv_lora_rank is not None + if mla: + head_dim = config.kv_lora_rank + config.qk_rope_head_dim + kv_factor = 1 + else: + tp_size = 1 if mapping.enable_attention_dp else mapping.tp_size + head_dim = getattr(config, "head_dim", None) + if not isinstance(head_dim, int): + head_dim = config.hidden_size // config.num_attention_heads + head_dim = head_dim * num_key_value_heads // tp_size + kv_factor = 2 + + cache_size_per_token = kv_factor * head_dim + quant_config = model_config.quant_config + if quant_config is not None and quant_config.quant_mode.has_fp8_kv_cache(): + layer_size = cache_size_per_token + elif quant_config is not None and quant_config.quant_mode.has_fp4_kv_cache(): + layer_size = math.ceil(cache_size_per_token / 2) + math.ceil(cache_size_per_token / 16) + else: + assert quant_config is None or (not quant_config.quant_mode.has_kv_cache_quant()), ( + "Quantized kv cache is not expected" + ) + layer_size = cache_size_per_token * 2 + + num_attention_layers = KVCacheManager._resolve_num_attention_layers( + model_config, mapping, num_layers + ) + layer_sizes = [layer_size] * num_attention_layers + window_pattern = kv_cache_config.max_attention_window if kv_cache_config is not None else None + + def get_window_size(layer_idx: int) -> Optional[int]: + if window_pattern is None or not isinstance(window_pattern, (list, tuple)): + return None + window_size = window_pattern[layer_idx % len(window_pattern)] + if window_size is None or window_size <= 0: + return None + window_size = int(window_size) + if max_seq_len is not None and window_size == int(max_seq_len): + return None + return window_size + + attention_windows = [get_window_size(layer_idx) for layer_idx in range(num_attention_layers)] + return layer_sizes, attention_windows + + class _MmRunMetadata(NamedTuple): """CPU tensors for exact multimodal run lookup. @@ -439,6 +546,58 @@ def _update_kv_cache_draft_token_location( ) +@maybe_compile(options={"max-autotune": True}) +def _copy_swa_block_offsets_with_scratch_compiled( + block_offsets: torch.Tensor, + copy_idx: torch.Tensor, + pool_ids: torch.Tensor, + scales: torch.Tensor, + layer_offsets: torch.Tensor, + scratch_pages: torch.Tensor, + block_positions: torch.Tensor, + scratch_begs: torch.Tensor, + scratch_ends: torch.Tensor, + scratch_slots: torch.Tensor, + num_contexts: torch.Tensor, + output: torch.Tensor, +) -> None: + base = block_offsets[pool_ids[:, :, None], copy_idx[None, None, :], 0, :] + converted = torch.where( + base == BAD_PAGE_INDEX, + BAD_PAGE_INDEX, + base * scales[:, :, None, None] + layer_offsets[:, :, None, None], + ) + + context_positions = torch.arange( + scratch_begs.shape[1], + dtype=torch.int32, + device=scratch_begs.device, + ) + active_context = context_positions < num_contexts + scratch_mask_by_pool = ( + (block_positions >= scratch_begs[:, :, None]) + & (block_positions < scratch_ends[:, :, None]) + & active_context[None, :, None] + ) + range_index = torch.where( + scratch_mask_by_pool, + block_positions - scratch_begs[:, :, None], + 0, + ) + total_offset = range_index[pool_ids] * scratch_pages[:, :, None, None] + slot_idx = (total_offset // scales[:, :, None, None]).clamp(max=scratch_slots.shape[-1] - 1) + slot_id = scratch_slots[pool_ids].gather(-1, slot_idx.long()) + offset = total_offset % scales[:, :, None, None] + scratch_index = ( + slot_id * scales[:, :, None, None] + + (offset + layer_offsets[:, :, None, None]) % scales[:, :, None, None] + ) + scratch_mask = scratch_mask_by_pool[pool_ids] + converted = torch.where(scratch_mask, scratch_index, converted) + + output.copy_(converted.permute(0, 2, 1, 3)) + + class KVCacheManagerV2(BaseResourceManager): def __init__( self, @@ -488,6 +647,10 @@ def __init__( layer_mask=layer_mask, ) self.is_draft = is_draft + self.enable_swa_scratch_reuse = ( + kv_cache_config.enable_swa_scratch_reuse and not self.is_draft + ) + self.block_reuse_policy = BlockReusePolicy(kv_cache_config.block_reuse_policy) self.num_local_layers = len(self.pp_layers) self.layer_offsets = {idx: offset for offset, idx in enumerate(self.pp_layers)} self.max_beam_width = max_beam_width @@ -501,6 +664,7 @@ def __init__( self.tokens_per_block = tokens_per_block self.max_seq_len = max_seq_len self.max_batch_size = max_batch_size + self.max_num_tokens = max_num_tokens self.kv_factor = 1 if kv_cache_type == CacheTypeCpp.SELFKONLY else 2 from ..speculative import get_num_extra_kv_tokens @@ -654,10 +818,12 @@ def append_to_kv_heads_per_layer( # bytes_per_token varies across PP ranks (different local layers). if mapping.world_size > 1: dist = Distributed.get(mapping) - bytes_per_token = self.get_cache_bytes_per_token() - max_tokens = quota / bytes_per_token + max_tokens = self._get_max_tokens_from_quota(quota) max_tokens = dist.allreduce(max_tokens, op=ReduceOp.MIN) - quota = max_tokens * bytes_per_token + # inf max_tokens means all layers are SWA and every rank quota can + # fit all SWA fixed cache. + if not math.isinf(max_tokens): + quota = self._get_quota_from_max_tokens(max_tokens) logger.info(f"KV cache manager v2 device quota set to {quota / (1 << 30)}GiB") @@ -749,6 +915,14 @@ def append_to_kv_heads_per_layer( ) self.num_pools = len(self.impl.layer_grouping) + # num_pools is the physical pool count owned by the KV cache manager. + # With SWA scratch reuse, scratch slot IDs are only valid with + # per-layer page indices, so the attention op sees one virtual pool per + # local layer while the underlying manager can still group layers. + if self.enable_swa_scratch_reuse: + self.num_attention_op_pools = self.num_local_layers + else: + self.num_attention_op_pools = self.num_pools num_layers = len(config.layers) self.layer_to_pool_mapping_dict: dict[int, int] = { @@ -756,10 +930,6 @@ def append_to_kv_heads_per_layer( for layer_id in typed_range(LayerId(num_layers)) } - (self.kv_cache_pool_pointers, self.kv_cache_pool_mapping) = ( - self._build_pool_mapping_tensors() - ) - self.kv_cache_map: dict[int, _KVCache] = {} # Tracks the draft length allocated by try_allocate_generation per @@ -822,6 +992,120 @@ def append_to_kv_heads_per_layer( ) self.index_mapper = IndexMapper(index_mapper_capacity, max_beam_width) self._early_freed_index_requests: set[int] = set() + self._prepare_page_table_tensor(index_mapper_capacity) + + self._log_kv_cache_pool_lifecycle_mapping() + + def _prepare_page_table_tensor(self, index_mapper_capacity: int) -> None: + kv_cache_pool_pointers_list = [] + kv_cache_pool_mapping_list = [] + block_scale_pool_pointers_list = [] + if self.enable_swa_scratch_reuse: + for layer_id in typed_range(LayerId(self.num_local_layers)): + kv_cache_pool_pointers_list.append( + [ + self.impl.get_mem_pool_base_address( + layer_id, Role.KEY, PageIndexMode.PER_LAYER + ), + 0, + ] + ) + if self.dtype == DataType.NVFP4: + block_scale_pool_pointers_list.append( + [ + self.impl.get_mem_pool_base_address( + layer_id, Role.KEY_BLOCK_SCALE, PageIndexMode.PER_LAYER + ), + 0, + ] + ) + kv_cache_pool_mapping_list.append([int(layer_id), 0]) + else: + for pool_id in range(self.num_pools): + layer_id = self.impl.layer_grouping[pool_id][0] + kv_cache_pool_pointers_list.append( + [ + self.impl.get_mem_pool_base_address( + layer_id, Role.KEY, PageIndexMode.SHARED + ), + 0, + ] + ) + if self.dtype == DataType.NVFP4: + block_scale_pool_pointers_list.append( + [ + self.impl.get_mem_pool_base_address( + layer_id, Role.KEY_BLOCK_SCALE, PageIndexMode.SHARED + ), + 0, + ] + ) + + for layer_id in typed_range(LayerId(self.num_local_layers)): + layer_group_id = self.impl.get_layer_group_id(layer_id) + if self.dtype != DataType.NVFP4: + key_base_addr = kv_cache_pool_pointers_list[layer_group_id][0] + addr_offset = ( + self.impl.get_mem_pool_base_address( + layer_id, Role.KEY, PageIndexMode.SHARED + ) + - key_base_addr + ) + else: + key_base_addr = kv_cache_pool_pointers_list[layer_group_id][0] + block_scale_base_addr = block_scale_pool_pointers_list[layer_group_id][0] + addr_offset = ( + self.impl.get_mem_pool_base_address( + layer_id, Role.KEY, PageIndexMode.SHARED + ) + - key_base_addr + ) + block_scale_addr_offset = ( + self.impl.get_mem_pool_base_address( + layer_id, Role.KEY_BLOCK_SCALE, PageIndexMode.SHARED + ) + - block_scale_base_addr + ) + block_scale_offset = exact_div( + block_scale_addr_offset, + self.get_layer_bytes_per_token(layer_id, Role.KEY_BLOCK_SCALE) + * self.kv_factor + * self.tokens_per_block, + ) + offset = exact_div( + addr_offset, + self.get_layer_bytes_per_token(layer_id, Role.KEY) + * self.kv_factor + * self.tokens_per_block, + ) + + if self.dtype == DataType.NVFP4: + assert block_scale_offset == offset, ( + "Block scale offset and offset should be the same" + ) + + kv_cache_pool_mapping_list.append([layer_group_id, offset]) + + if self.dtype == DataType.NVFP4: + for pool_id, block_scale_pool_pointers in enumerate(block_scale_pool_pointers_list): + pool_pointers = kv_cache_pool_pointers_list[pool_id] + kv_cache_pool_pointers_list[pool_id] = [ + [pool_pointers[0], block_scale_pool_pointers[0]], + [pool_pointers[1], block_scale_pool_pointers[1]], + ] + + self.kv_cache_pool_pointers = torch.tensor( + kv_cache_pool_pointers_list, + dtype=torch.int64, + device="cpu", + pin_memory=prefer_pinned(), + ) + self.kv_cache_pool_mapping = torch.tensor( + kv_cache_pool_mapping_list, + dtype=torch.int32, + device="cpu", + pin_memory=prefer_pinned(), + ) self.index_scales = torch.empty( self.num_pools, dtype=torch.int32, pin_memory=prefer_pinned(), device="cpu" ) @@ -833,8 +1117,8 @@ def append_to_kv_heads_per_layer( self.index_scales[pool_id] = self.impl.get_page_index_scale(layer_id, Role.KEY) if self.kv_cache_type != CacheTypeCpp.SELFKONLY: self.kv_offset[pool_id] = exact_div( - self.impl.get_mem_pool_base_address(layer_id, Role.VALUE) - - self.impl.get_mem_pool_base_address(layer_id, Role.KEY), + self.impl.get_mem_pool_base_address(layer_id, Role.VALUE, PageIndexMode.SHARED) + - self.impl.get_mem_pool_base_address(layer_id, Role.KEY, PageIndexMode.SHARED), self.impl.get_page_stride(layer_id, Role.KEY), ) else: @@ -843,18 +1127,92 @@ def append_to_kv_heads_per_layer( # Keep unused block offsets as safe block index 0. self.host_kv_cache_block_offsets = torch.zeros( self.num_pools, - index_mapper_capacity * max_beam_width, + index_mapper_capacity * self.max_beam_width, 2, # key and value self.max_blocks_per_seq, dtype=torch.int32, pin_memory=prefer_pinned(), device="cpu", ) + if self.enable_swa_scratch_reuse: + self._prepare_swa_scratch_copy_tensors(index_mapper_capacity) + + def _get_runtime_cache_size_layer_components(self) -> tuple[List[int], List[Optional[int]]]: + layer_sizes = [] + attention_windows = [] + pattern_len = len(self.max_attention_window_vec) + for local_layer_idx in range(self.num_local_layers): + layer_sizes.append( + self.get_layer_bytes_per_token(local_layer_idx=local_layer_idx, data_role=Role.ALL) + ) + attention_windows.append( + self.max_attention_window_vec[self.pp_layers[local_layer_idx] % pattern_len] + ) + return layer_sizes, attention_windows - self._log_kv_cache_pool_lifecycle_mapping() + def _get_max_tokens_from_quota(self, quota: int) -> float: + layer_sizes, attention_windows = self._get_runtime_cache_size_layer_components() + full_attn_size_per_token = _estimate_full_attn_size_per_token( + layer_sizes, attention_windows + ) + context_swa_size_per_token, _ = _estimate_swa_cache_size( + layer_sizes, + attention_windows, + self.tokens_per_block, + context=True, + scratch=self.enable_swa_scratch_reuse, + ) + ( + generation_swa_size_per_token, + generation_swa_size_per_request, + ) = _estimate_swa_cache_size( + layer_sizes, attention_windows, self.tokens_per_block, context=False, scratch=False + ) + size_per_batch = self.max_batch_size * generation_swa_size_per_request + if quota < size_per_batch: + return 0 + context_size_per_token = full_attn_size_per_token + context_swa_size_per_token + context_limit_quota = self.max_num_tokens * context_size_per_token + size_per_batch + if quota <= context_limit_quota: + if context_size_per_token <= 0: + return float("inf") + return (quota - size_per_batch) / context_size_per_token + + generation_size_per_token = full_attn_size_per_token + generation_swa_size_per_token + if generation_size_per_token <= 0: + return float("inf") + return self.max_num_tokens + (quota - context_limit_quota) / generation_size_per_token def _get_quota_from_max_tokens(self, max_tokens: int) -> int: - return int(max_tokens * self.get_cache_bytes_per_token()) + layer_sizes, attention_windows = self._get_runtime_cache_size_layer_components() + full_attn_size_per_token = _estimate_full_attn_size_per_token( + layer_sizes, attention_windows + ) + ( + context_swa_size_per_token, + _, + ) = _estimate_swa_cache_size( + layer_sizes, + attention_windows, + self.tokens_per_block, + context=True, + scratch=self.enable_swa_scratch_reuse, + ) + ( + generation_swa_size_per_token, + generation_swa_size_per_request, + ) = _estimate_swa_cache_size( + layer_sizes, attention_windows, self.tokens_per_block, context=False, scratch=False + ) + context_tokens = min(max_tokens, self.max_num_tokens) + generation_tokens = max_tokens - context_tokens + generation_quota = ( + max_tokens * full_attn_size_per_token + + generation_tokens * generation_swa_size_per_token + + self.max_batch_size * generation_swa_size_per_request + ) + context_extra_quota = context_tokens * context_swa_size_per_token + return int(generation_quota + context_extra_quota) def _get_event_num_blocks_per_cache_level( self, @@ -909,82 +1267,208 @@ def _log_kv_cache_pool_lifecycle_mapping(self) -> None: 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( - self.impl.layer_grouping[pool_id][0], Role.KEY - ), - 0, - ] - for pool_id in range(self.num_pools) - ], - dtype=torch.int64, + def _prepare_swa_scratch_copy_tensors(self, index_mapper_capacity: int) -> None: + pool_ids = torch.empty( + self.num_attention_op_pools, + 2, + dtype=torch.long, device="cpu", - pin_memory=prefer_pinned(), + ) + scales = torch.empty( + self.num_attention_op_pools, + 2, + dtype=torch.int32, + device="cpu", + ) + layer_offsets = torch.empty_like(scales) + scratch_pages = torch.empty_like(scales) + for local_layer_idx in range(self.num_local_layers): + layer_id = LayerId(local_layer_idx) + pool_id = self.layer_to_pool_mapping_dict[layer_id] + roles = [Role.KEY, Role.VALUE] + if self.kv_cache_type == CacheTypeCpp.SELFKONLY: + roles[1] = Role.KEY + for role_idx, role in enumerate(roles): + converter = self.impl.get_page_index_converter(layer_id, role) + if converter.expansion != 1: + raise NotImplementedError( + "SWA scratch block-table conversion does not support " + f"expanded page indices yet: layer={layer_id}, role={role}, " + f"expansion={converter.expansion}" + ) + pool_ids[local_layer_idx, role_idx] = pool_id + scales[local_layer_idx, role_idx] = int(converter.scale) + layer_offsets[local_layer_idx, role_idx] = int(converter.layer_offset) + scratch_pages[local_layer_idx, role_idx] = int(converter.scratch_pages_per_block) + + staging_capacity = index_mapper_capacity * self.max_beam_width + device = torch.device("cuda", torch.cuda.current_device()) + self._device_kv_cache_block_offsets_input = torch.empty_like( + self.host_kv_cache_block_offsets, + device=device, + ) + self._device_attention_op_block_offsets_staging = torch.empty( + self.num_attention_op_pools, + staging_capacity, + 2, + self.max_blocks_per_seq, + dtype=torch.int32, + device=device, + ) + self._device_copy_idx_staging = torch.zeros( + staging_capacity, + dtype=torch.long, + device=device, + ) + self._device_num_contexts = torch.empty((), dtype=torch.int32, device=device) + self._device_attention_op_pool_ids = pool_ids.to(device=device) + self._device_attention_op_scales = scales.to(device=device) + self._device_attention_op_layer_offsets = layer_offsets.to(device=device) + self._device_attention_op_scratch_pages = scratch_pages.to(device=device) + self._device_block_positions = torch.arange( + self.max_blocks_per_seq, + dtype=torch.int32, + device=device, ) - if self.dtype == DataType.NVFP4: - kv_cache_pool_pointers = torch.stack( - [ - kv_cache_pool_pointers, - torch.tensor( - [ - [ - self.impl.get_mem_pool_base_address( - self.impl.layer_grouping[pool_id][0], Role.KEY_BLOCK_SCALE - ), - 0, - ] - for pool_id in range(self.num_pools) - ], - dtype=torch.int64, - device="cpu", - pin_memory=prefer_pinned(), - ), - ], - dim=-1, - ) + min_scale = int(scales.min().item()) if scales.numel() > 0 else 1 + max_scratch_pages = int(scratch_pages.max().item()) if scratch_pages.numel() > 0 else 1 + self._max_scratch_slots = max( + 1, + (self.max_blocks_per_seq * max_scratch_pages + min_scale - 1) // min_scale, + ) + scratch_slots_shape = ( + self.num_pools, + staging_capacity, + self._max_scratch_slots, + ) + self._host_scratch_begs_staging = torch.zeros( + self.num_pools, + staging_capacity, + dtype=torch.int32, + pin_memory=prefer_pinned(), + device="cpu", + ) + self._host_scratch_ends_staging = torch.zeros( + self.num_pools, + staging_capacity, + dtype=torch.int32, + pin_memory=prefer_pinned(), + device="cpu", + ) + self._host_scratch_slots_staging = torch.zeros( + scratch_slots_shape, + dtype=torch.int32, + pin_memory=prefer_pinned(), + device="cpu", + ) + self._device_scratch_begs_staging = torch.zeros( + self.num_pools, + staging_capacity, + dtype=torch.int32, + device=device, + ) + self._device_scratch_ends_staging = torch.zeros( + self.num_pools, + staging_capacity, + dtype=torch.int32, + device=device, + ) + self._device_scratch_slots_staging = torch.zeros( + scratch_slots_shape, + dtype=torch.int32, + device=device, + ) - kv_cache_pool_mapping_list = [] - for layer_id in typed_range(LayerId(self.num_local_layers)): - layer_group_id = self.impl.get_layer_group_id(layer_id) - if self.dtype != DataType.NVFP4: - addr_offset = self.impl.get_mem_pool_base_address(layer_id, Role.KEY) - int( - kv_cache_pool_pointers[layer_group_id][0] - ) - else: - addr_offset = self.impl.get_mem_pool_base_address(layer_id, Role.KEY) - int( - kv_cache_pool_pointers[layer_group_id][0][0] - ) - block_scale_addr_offset = self.impl.get_mem_pool_base_address( - layer_id, Role.KEY_BLOCK_SCALE - ) - int(kv_cache_pool_pointers[layer_group_id][0][1]) - block_scale_offset = exact_div( - block_scale_addr_offset, - self.get_layer_bytes_per_token(layer_id, Role.KEY_BLOCK_SCALE) - * self.kv_factor - * self.tokens_per_block, - ) - offset = exact_div( - addr_offset, - self.get_layer_bytes_per_token(layer_id, Role.KEY) - * self.kv_factor - * self.tokens_per_block, - ) + def _copy_idx_to_device(self, copy_idx: torch.Tensor) -> torch.Tensor: + self._device_copy_idx_staging[: copy_idx.size(0)].copy_(copy_idx, non_blocking=True) + return self._device_copy_idx_staging - if self.dtype == DataType.NVFP4: - assert block_scale_offset == offset, ( - "Block scale offset and offset should be the same" - ) + def _copy_scratch_metadata_to_device( + self, + request_ids: List[int], + num_contexts: int, + ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + host_begs = self._host_scratch_begs_staging[:, :num_contexts] + host_ends = self._host_scratch_ends_staging[:, :num_contexts] + host_slots = self._host_scratch_slots_staging[:, :num_contexts, :] + host_begs.zero_() + host_ends.zero_() + host_slots.zero_() - kv_cache_pool_mapping_list.append([layer_group_id, offset]) + for pool_id in range(self.num_pools): + for context_idx, request_id in enumerate(request_ids[:num_contexts]): + desc = self.kv_cache_map[request_id].get_scratch_desc(pool_id) + if desc is None: + continue + slot_ids = desc.slot_ids + if len(slot_ids) > self._max_scratch_slots: + raise RuntimeError( + f"Scratch slot count {len(slot_ids)} exceeds staging capacity " + f"{self._max_scratch_slots}" + ) + host_begs[pool_id, context_idx] = int(desc.range.beg) + host_ends[pool_id, context_idx] = int(desc.range.end) + for slot_idx, slot_id in enumerate(slot_ids): + host_slots[pool_id, context_idx, slot_idx] = int(slot_id) + + self._device_scratch_begs_staging.copy_( + self._host_scratch_begs_staging, + non_blocking=True, + ) + self._device_scratch_ends_staging.copy_( + self._host_scratch_ends_staging, + non_blocking=True, + ) + self._device_scratch_slots_staging.copy_( + self._host_scratch_slots_staging, + non_blocking=True, + ) + return ( + self._device_scratch_begs_staging, + self._device_scratch_ends_staging, + self._device_scratch_slots_staging, + ) - kv_cache_pool_mapping = torch.tensor( - kv_cache_pool_mapping_list, dtype=torch.int32, device="cpu", pin_memory=prefer_pinned() + def _copy_batch_block_offsets_per_layer( + self, + dst_tensor: torch.Tensor, + request_ids: List[int], + copy_idx: torch.Tensor, + num_contexts: int, + num_seqs: int, + ) -> None: + device_copy_idx = self._copy_idx_to_device(copy_idx) + self._device_kv_cache_block_offsets_input.copy_( + self.host_kv_cache_block_offsets, + non_blocking=True, + ) + scratch_begs, scratch_ends, scratch_slots = self._copy_scratch_metadata_to_device( + request_ids, + num_contexts, + ) + self._device_num_contexts.fill_(num_contexts) + _copy_swa_block_offsets_with_scratch_compiled( + self._device_kv_cache_block_offsets_input, + device_copy_idx, + self._device_attention_op_pool_ids, + self._device_attention_op_scales, + self._device_attention_op_layer_offsets, + self._device_attention_op_scratch_pages, + self._device_block_positions, + scratch_begs, + scratch_ends, + scratch_slots, + self._device_num_contexts, + self._device_attention_op_block_offsets_staging, + ) + dst_tensor[: self.num_attention_op_pools, :num_seqs].copy_( + self._device_attention_op_block_offsets_staging[ + : self.num_attention_op_pools, + :num_seqs, + ], + non_blocking=True, ) - return kv_cache_pool_pointers, kv_cache_pool_mapping def _build_cache_config( self, @@ -1006,6 +1490,12 @@ def _build_cache_config( if self.kv_cache_type != CacheTypeCpp.SELFKONLY: buffer_type.append(Role.VALUE_BLOCK_SCALE) + scratch_reuse_config = None + if self.enable_swa_scratch_reuse: + # Context requests allocate num_extra_kv_tokens for spec decoding. + # They should not count toward the scratch range. + scratch_reuse_config = SwaScratchReuseConfig(max_rewind_len=self.num_extra_kv_tokens) + # Subclasses (e.g. MiniMax-M3 sparse cache) can register additional # per-layer BufferConfig entries — for example a sparse index-K # buffer — without overriding the K/V/NVFP4 scale wiring above. @@ -1050,6 +1540,8 @@ def _build_cache_config( cache_tiers=cache_tiers, max_util_for_resume=kv_cache_config.max_util_for_resume, enable_stats=self.enable_stats, + swa_scratch_reuse=scratch_reuse_config, + initial_pool_ratio=kv_cache_config.pool_ratio, layers=layer_configs, ) @@ -1079,9 +1571,11 @@ def blocks_in_primary_pool(self) -> int: def get_buffers(self, layer_idx: int, kv_layout: str = "NHD") -> Optional[torch.Tensor]: layer_offset = self.layer_offsets[layer_idx] - addr_key = self.impl.get_mem_pool_base_address(layer_offset, Role.KEY) + addr_key = self.impl.get_mem_pool_base_address(layer_offset, Role.KEY, PageIndexMode.SHARED) if self.kv_cache_type != CacheTypeCpp.SELFKONLY: - addr_value = self.impl.get_mem_pool_base_address(layer_offset, Role.VALUE) + addr_value = self.impl.get_mem_pool_base_address( + layer_offset, Role.VALUE, PageIndexMode.SHARED + ) page_size_key = self.impl.get_page_stride(layer_offset, Role.KEY) page_size_value = self.impl.get_page_stride(layer_offset, Role.VALUE) @@ -1336,46 +1830,6 @@ def try_allocate_generation(self, req: LlmRequest) -> bool: self._allocated_draft_lens[req.py_request_id] = draft_len return kv_cache.resize(self._required_gen_capacity(req, kv_cache.capacity)) - def trim_to_history(self, req: LlmRequest, history_length: int) -> bool: - """Mark *history_length* tokens of this request's KV as historic. - - For sliding-window-style life cycles (AttnLifeCycle with non-None - window_size), this triggers ``_unlock_stale_blocks`` inside V2's - ``resize()`` so blocks before ``(history_length + 1 - window) // - tokens_per_block`` get released back to their pool group. For - full-context life cycles (``window_size=None``) and SSM cycles, it - is a no-op — the stale range stays empty. - - Used by the disagg-gen transceiver right after KV transfer - completes: at that moment the cache has the entire prompt KV - written, so ``history_length=prompt_len`` correctly classifies - every prompt token as historic. Without this call, ``history_length`` - stays 0 until ``update_resources`` runs after the first forward - pass — and in benchmark fill-phase the first forward never fires - until every disagg-gen request is ready, so SWA / sparse-attn - pool groups would otherwise stay 100% occupied with pre-window - prompt blocks and the V2 scheduler would deadlock on the next - ``resize(+1)``. - - Returns True on success (or no-op), False if the underlying - ``kv_cache.resize`` rejected the call. - """ - kv_cache = self.kv_cache_map.get(req.py_request_id) - if kv_cache is None or not kv_cache.is_active: - return True - if history_length <= kv_cache.history_length: - return True - # resize() requires capacity >= history_length; clamp for safety. - target_capacity = max(kv_cache.capacity, history_length) - try: - return kv_cache.resize(target_capacity, history_length=history_length) - except Exception as e: - logger.warning( - f"trim_to_history failed for req {req.py_request_id} " - f"(capacity={kv_cache.capacity}, target_history={history_length}): {e}" - ) - return False - def revert_allocate_generation(self, req: LlmRequest) -> None: """Undo the capacity growth from try_allocate_generation. @@ -1406,33 +1860,25 @@ def revert_allocate_generation(self, req: LlmRequest) -> None: ) def revert_allocate_context(self, req: LlmRequest) -> None: - """Undo the capacity growth from this iter's ``resize_context``. - - When delay batching (``_balance_adp_requests`` / - ``_waiting_requests``) defers a context request after V2 - scheduling, the forward pass is skipped for that request but the - scheduler already grew its KV cache capacity to cover the chunk. - This shrinks capacity back to the pre-resize value so the - freshly-allocated pages can be reused during the wait window — - important for long contexts where one deferred request can hold - GBs of KV. - """ + """Undo the capacity growth from this iteration's context resize.""" pre_cap = getattr(req, "py_ctx_pre_resize_cap", None) if pre_cap is None: return - # Mark as consumed even if the resize below is skipped, so a - # later iter does not see a stale snapshot. req.py_ctx_pre_resize_cap = None kv_cache = self.kv_cache_map.get(req.py_request_id) if kv_cache is None or not kv_cache.is_active: return if pre_cap >= kv_cache.capacity: return - if not kv_cache.resize(pre_cap): + if kv_cache.history_length > pre_cap: + self.free_resources(req) + return + history_length = min(kv_cache.history_length, pre_cap) + if not kv_cache.resize(pre_cap, history_length): raise RuntimeError( f"Failed to revert KV cache capacity for context " - f"request {req.py_request_id} from " - f"{kv_cache.capacity} to {pre_cap}" + f"request {req.py_request_id} from {kv_cache.capacity} " + f"to {pre_cap}" ) if pre_cap > 0: kv_cache.suspend() @@ -1475,6 +1921,12 @@ def prepare_context(self, req: LlmRequest) -> bool: For subsequent chunks: verifies existing cache is active. Returns True on success, False if preparation failed. """ + assert not req.is_disagg_generation_init_state, ( + f"req {req.py_request_id}: use prepare_disagg_gen_init" + ) + return self._prepare_context_impl(req) + + def _prepare_context_impl(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: @@ -1507,6 +1959,10 @@ def prepare_context(self, req: LlmRequest) -> bool: kv_cache.num_committed_tokens, self.tokens_per_block ) + if req.is_disagg_generation_init_state: + # Disagg generation receives prompt KV from the context worker; + # scratch blocks are only valid for local prefill chunks. + kv_cache.enable_swa_scratch_reuse = False return self._resume_and_restore(req.py_request_id, kv_cache) else: # Subsequent chunk: cache must exist from first chunk. @@ -1518,20 +1974,15 @@ def prepare_context(self, req: LlmRequest) -> bool: ) return self._resume_and_restore(req.py_request_id, kv_cache) - def resize_context( - self, req: LlmRequest, num_tokens: int, history_length: int | None = None - ) -> bool: + def resize_context(self, req: LlmRequest, num_tokens: int) -> bool: """Resize KV cache to cover context_current_position + num_tokens. - history_length, when set, lets SWA life cycles compute their stale - range at allocation time so pre-window blocks are never allocated. Returns True on success, False if resize failed (first chunk is suspended on failure). - - Snapshots the pre-resize capacity on ``req.py_ctx_pre_resize_cap`` - when growth happens so ``revert_allocate_context`` can undo it if - delay batching defers the request. """ + assert not req.is_disagg_generation_init_state, ( + f"req {req.py_request_id}: use prepare_disagg_gen_init" + ) kv_cache = self.kv_cache_map.get(req.py_request_id) if kv_cache is None: return False @@ -1540,17 +1991,55 @@ def resize_context( capacity = max(kv_cache.capacity, target) pre_cap = kv_cache.capacity - success = kv_cache.resize(capacity, history_length) + success = kv_cache.resize(capacity) if not success: if req.is_first_context_chunk: kv_cache.suspend() return False + req.py_ctx_pre_resize_cap = pre_cap if capacity > pre_cap else None + return True + + def prepare_disagg_gen_init(self, req: LlmRequest) -> bool: + """Prepare KV cache for a disagg generation init request. - # None means "no growth this iter, nothing to revert"; this also - # invalidates a stale snapshot from a prior iter on the same req. + Allocates capacity for the full prompt (+ draft) and sets + ``kv_cache.history_length`` to ``prompt_len``. Returns True on + success, False if preparation or resize failed (cache is suspended + on resize failure). + """ + if not self._prepare_context_impl(req): + return False + + kv_cache = self.kv_cache_map.get(req.py_request_id) + if kv_cache is None: + return False + + # prompt_len is the full incoming prompt length, robust to block + # reuse (which may leave a non-zero context_current_position). + target = req.prompt_len + get_draft_token_length(req) + self.num_extra_kv_tokens + capacity = max(kv_cache.capacity, target) + pre_cap = kv_cache.capacity + + success = kv_cache.resize(capacity, req.prompt_len) + if not success: + if req.is_first_context_chunk: + kv_cache.suspend() + return False req.py_ctx_pre_resize_cap = pre_cap if capacity > pre_cap else None return True + def get_history_length(self, req: LlmRequest) -> int | None: + """Return the cache's current history_length, or None if no cache. + + Exposes the per-request SWA history watermark so callers + (e.g., the disagg transceiver) can verify scheduler/cache contracts + without reaching into ``kv_cache_map`` directly. + """ + kv_cache = self.kv_cache_map.get(req.py_request_id) + if kv_cache is None: + return None + return kv_cache.history_length + def extend_capacity_for_tokens(self, request: LlmRequest) -> None: """Extend KV cache capacity for the CUDA-graph padding delta. @@ -1785,6 +2274,8 @@ def _build_iteration_stats( pool_group_ids: Iterable[int], primary_stats, secondary_stats_by_level, + primary_peak_stats, + secondary_peak_stats_by_level, delta, field_names=KV_CACHE_ITERATION_STATS_DELTA_FIELDS, ): @@ -1797,6 +2288,18 @@ def _build_iteration_stats( 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.primary_evictable_num_blocks = sum( + primary_stats[pool_group_id].evictable for pool_group_id in pool_group_ids + ) + stats.primary_peak_free_num_blocks = sum( + primary_peak_stats[pool_group_id].available for pool_group_id in pool_group_ids + ) + stats.primary_peak_used_num_blocks = sum( + primary_peak_stats[pool_group_id].unavailable for pool_group_id in pool_group_ids + ) + stats.primary_peak_evictable_num_blocks = sum( + primary_peak_stats[pool_group_id].evictable for pool_group_id in pool_group_ids + ) stats.secondary_max_num_blocks = sum( level_stats[pool_group_id].total for level_stats in secondary_stats_by_level @@ -1810,6 +2313,26 @@ def _build_iteration_stats( stats.secondary_used_num_blocks = ( stats.secondary_max_num_blocks - stats.secondary_free_num_blocks ) + stats.secondary_evictable_num_blocks = sum( + level_stats[pool_group_id].evictable + for level_stats in secondary_stats_by_level + for pool_group_id in pool_group_ids + ) + stats.secondary_peak_free_num_blocks = sum( + peak_stats[pool_group_id].available + for peak_stats in secondary_peak_stats_by_level + for pool_group_id in pool_group_ids + ) + stats.secondary_peak_used_num_blocks = sum( + peak_stats[pool_group_id].unavailable + for peak_stats in secondary_peak_stats_by_level + for pool_group_id in pool_group_ids + ) + stats.secondary_peak_evictable_num_blocks = sum( + peak_stats[pool_group_id].evictable + for peak_stats in secondary_peak_stats_by_level + for pool_group_id in pool_group_ids + ) self._apply_iteration_stats_delta(stats, delta, field_names) return stats @@ -1858,6 +2381,8 @@ def _build_window_iteration_stats( windows_by_pool_group: dict[int, tuple[int, ...]], primary_stats, secondary_stats_by_level, + primary_peak_stats, + secondary_peak_stats_by_level, pool_group_delta, reuse_delta, ): @@ -1870,6 +2395,8 @@ def _build_window_iteration_stats( pool_group_ids, primary_stats, secondary_stats_by_level, + primary_peak_stats, + secondary_peak_stats_by_level, pool_group_delta, KV_CACHE_ITERATION_STATS_POOL_GROUP_FIELDS, ) @@ -1882,6 +2409,8 @@ def _build_pool_group_iteration_stats( windows_by_pool_group: dict[int, tuple[int, ...]], primary_stats, secondary_stats_by_level, + primary_peak_stats, + secondary_peak_stats_by_level, pool_group_delta, ) -> KVCacheV2PoolGroupIterationStats: return KVCacheV2PoolGroupIterationStats( @@ -1892,6 +2421,8 @@ def _build_pool_group_iteration_stats( (pool_group_id,), primary_stats, secondary_stats_by_level, + primary_peak_stats, + secondary_peak_stats_by_level, pool_group_delta, KV_CACHE_ITERATION_STATS_POOL_GROUP_FIELDS, ), @@ -1903,6 +2434,8 @@ def _build_life_cycle_iteration_stats( storage, primary_stats, secondary_stats_by_level, + primary_peak_stats, + secondary_peak_stats_by_level, reuse_delta, ) -> KVCacheV2LifeCycleIterationStats: typed_life_cycle_id = LifeCycleId(life_cycle_id) @@ -1917,6 +2450,8 @@ def _build_life_cycle_iteration_stats( (), primary_stats, secondary_stats_by_level, + primary_peak_stats, + secondary_peak_stats_by_level, reuse_delta, KV_CACHE_ITERATION_STATS_REUSE_FIELDS, ), @@ -1924,13 +2459,14 @@ def _build_life_cycle_iteration_stats( def get_kv_cache_stats(self): kv_cache_stats = KvCacheStats() - storage_stats = self.impl._get_storage_level_stats(GPU_LEVEL) - pool_group_stats = storage_stats.pool_group_stats + pool_group_stats = self.impl._storage.get_statistics(GPU_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) 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.max_num_blocks = max_num_blocks + kv_cache_stats.free_num_blocks = free_num_blocks + kv_cache_stats.used_num_blocks = max_num_blocks - free_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 @@ -1948,7 +2484,7 @@ def get_kv_cache_stats(self): ) for window_size, pool_group_ids in self._storage_pool_groups_by_window().items() } - kv_cache_stats.allocated_bytes = storage_stats.allocated_bytes + kv_cache_stats.allocated_bytes = self.impl.get_quota(GPU_LEVEL) return kv_cache_stats @@ -1969,6 +2505,11 @@ def get_iteration_stats(self): 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() + primary_peak_stats = self.impl.get_and_reset_iteration_peak_block_stats(GPU_LEVEL) + secondary_peak_stats_by_level = [ + self.impl.get_and_reset_iteration_peak_block_stats(CacheLevel(level)) + for level in range(1, int(storage.num_cache_levels)) + ] ( reuse_deltas_by_window, reuse_deltas_by_life_cycle, @@ -1992,6 +2533,8 @@ def get_iteration_stats(self): windows_by_pool_group, primary_stats, secondary_stats_by_level, + primary_peak_stats, + secondary_peak_stats_by_level, pool_group_deltas_by_window.get(window_size), reuse_deltas_by_window.get(window_size), ) @@ -2005,6 +2548,8 @@ def get_iteration_stats(self): windows_by_pool_group, primary_stats, secondary_stats_by_level, + primary_peak_stats, + secondary_peak_stats_by_level, pool_group_deltas.get(pool_group_id), ) for pool_group_id in pool_group_ids @@ -2016,6 +2561,8 @@ def get_iteration_stats(self): storage, primary_stats, secondary_stats_by_level, + primary_peak_stats, + secondary_peak_stats_by_level, reuse_delta, ) for life_cycle_id, reuse_delta in sorted(reuse_deltas_by_life_cycle.items()) @@ -2122,6 +2669,8 @@ def release_resources( return None kv_cache.stop_committing() dummy_capacity = token_num + self.num_extra_kv_tokens + num_extra_decoding_steps + if is_gen: + kv_cache.enable_swa_scratch_reuse = False # Need to hint the committed history to activate stale-block # optimization and match the solver's pool budget. success = kv_cache.resize(dummy_capacity, history_length=history_hint) @@ -2170,13 +2719,18 @@ def release_resources( return requests - def try_commit_blocks_for_reuse(self, request: LlmRequest, kv_cache) -> None: - if ( - self.enable_block_reuse - and not self.is_draft - and not request.is_dummy_request - and request.context_current_position > kv_cache.num_committed_tokens - ): + def try_commit_blocks(self, request: LlmRequest) -> None: + should_block_reuse = ( + self.enable_block_reuse and not self.is_draft and not request.is_dummy_request + ) + if not should_block_reuse: + return + + kv_cache = self.kv_cache_map.get(request.py_request_id) + if kv_cache is None: + return + + if request.context_current_position > kv_cache.num_committed_tokens: tokens = self._augment_tokens_for_block_reuse( request.get_tokens(DEFAULT_BEAM_INDEX), request, @@ -2184,6 +2738,7 @@ def try_commit_blocks_for_reuse(self, request: LlmRequest, kv_cache) -> None: end=request.context_current_position, ) kv_cache.commit(tokens) + if request.context_remaining_length == 0: kv_cache.stop_committing() def release_index_slot(self, request_id: int) -> None: @@ -2194,6 +2749,11 @@ def release_index_slot(self, request_id: int) -> None: needed. Releasing it early allows new requests to be scheduled while the KV cache blocks are still being transferred via NIXL/UCX. """ + kv_cache = self.kv_cache_map.get(request_id) + if kv_cache is not None: + for i in range(self.max_beam_width): + for pool_idx in range(self.num_pools): + kv_cache.set_base_page_index_buf(i, pool_idx, None) self.index_mapper.remove_sequence(request_id) self._early_freed_index_requests.add(request_id) @@ -2204,7 +2764,6 @@ def free_resources(self, request: LlmRequest, pin_on_release: bool = False): 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: @@ -2385,39 +2944,24 @@ def get_cache_size_per_token( num_layers: Optional[int] = None, **kwargs, ): - # get kv cache dtype bytes - mem_per_token = 2 - quant_config = model_config.quant_config - if quant_config is not None and quant_config.quant_mode.has_fp8_kv_cache(): - mem_per_token = 1 - - # get num key value heads - config = model_config.pretrained_config - num_key_value_heads = getattr(config, "num_key_value_heads", config.num_attention_heads) - if isinstance(num_key_value_heads, Iterable): - num_key_value_heads = sum(num_key_value_heads) / len(num_key_value_heads) - - # get head dim - mla = hasattr(config, "kv_lora_rank") and config.kv_lora_rank is not None - if mla: - head_dim = config.kv_lora_rank + config.qk_rope_head_dim - kv_factor = 1 - else: - tp_size = 1 if mapping.enable_attention_dp else mapping.tp_size - head_dim = getattr(config, "head_dim", None) - if not isinstance(head_dim, int): - head_dim = config.hidden_size // config.num_attention_heads - head_dim = head_dim * num_key_value_heads // tp_size - kv_factor = 2 - - num_attention_layers = KVCacheManager._resolve_num_attention_layers( - model_config, mapping, num_layers + layer_sizes, attention_windows = _get_static_cache_size_layer_components( + model_config, mapping, num_layers=num_layers, **kwargs + ) + full_attn_size_per_token = _estimate_full_attn_size_per_token( + layer_sizes, attention_windows + ) + swa_size_per_token, swa_size_per_request = _estimate_swa_cache_size( + layer_sizes, + attention_windows, + kwargs["tokens_per_block"], + context=False, + scratch=False, + ) + max_batch_size = int(kwargs.get("max_batch_size") or 0) + return ( + full_attn_size_per_token + swa_size_per_token, + swa_size_per_request * max_batch_size, ) - mem_per_token *= num_attention_layers * head_dim - - # K and V - mem_per_token *= kv_factor - return mem_per_token def update_context_resources(self, scheduled_batch: ScheduledRequests): """Update KV cache for context requests in the current batch. @@ -2440,18 +2984,14 @@ def update_context_resources(self, scheduled_batch: ScheduledRequests): # iteration. if not kv_cache.is_active: continue - if self.enable_block_reuse and not self.is_draft and not req.is_dummy_request: - if req.context_current_position > kv_cache.num_committed_tokens: - tokens = self._augment_tokens_for_block_reuse( - req.get_tokens(DEFAULT_BEAM_INDEX), - req, - start=kv_cache.num_committed_tokens, - end=req.context_current_position, - ) - kv_cache.commit(tokens) - if req.context_remaining_length == 0: - kv_cache.stop_committing() - else: + should_block_reuse = ( + self.enable_block_reuse and not self.is_draft and not req.is_dummy_request + ) + is_all_reusable = self.block_reuse_policy == BlockReusePolicy.ALL_REUSABLE + should_resize = not should_block_reuse or not is_all_reusable + should_commit = is_all_reusable or req.context_remaining_length == 0 + + if should_resize: success = kv_cache.resize(None, req.context_current_position) if not success: raise ValueError( @@ -2459,6 +2999,13 @@ def update_context_resources(self, scheduled_batch: ScheduledRequests): f"{req.py_request_id} to {req.context_current_position} tokens " "at context update" ) + if should_commit: + self.try_commit_blocks(req) + if req.context_remaining_length == 0: + # Scratch blocks are only for prefill chunks. Disable them at + # the context/generation boundary so generation uses normal KV + # pages before the first generation allocation. + kv_cache.enable_swa_scratch_reuse = False def update_resources( self, @@ -2509,6 +3056,12 @@ def copy_batch_block_offsets( copy_idx = self.index_mapper.get_copy_index(request_ids, num_contexts, beam_width) assert copy_idx.shape[0] == num_seqs + if self.enable_swa_scratch_reuse: + self._copy_batch_block_offsets_per_layer( + dst_tensor, request_ids, copy_idx, num_contexts, num_seqs + ) + return + copy_batch_block_offsets_to_device( self.host_kv_cache_block_offsets, dst_tensor, diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_stats.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_stats.py index ff7a06643928..d9841ccbfe86 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_stats.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_stats.py @@ -28,9 +28,17 @@ "primaryMaxNumBlocks", "primaryFreeNumBlocks", "primaryUsedNumBlocks", + "primaryEvictableNumBlocks", + "primaryPeakFreeNumBlocks", + "primaryPeakUsedNumBlocks", + "primaryPeakEvictableNumBlocks", "secondaryMaxNumBlocks", "secondaryFreeNumBlocks", "secondaryUsedNumBlocks", + "secondaryEvictableNumBlocks", + "secondaryPeakFreeNumBlocks", + "secondaryPeakUsedNumBlocks", + "secondaryPeakEvictableNumBlocks", "iterAllocTotalBlocks", "iterAllocNewBlocks", "iterGenAllocBlocks", @@ -40,6 +48,8 @@ "iterOffloadBytes", "iterIntraDeviceCopyBlocks", "iterIntraDeviceCopyBytes", + "iterHostDroppedBlocks", + "iterHostDroppedBytes", ) @@ -72,9 +82,17 @@ def serialize_kv_cache_iteration_stats(stats, keys: tuple[str, ...] | None = Non "primaryMaxNumBlocks": stats.primary_max_num_blocks, "primaryFreeNumBlocks": stats.primary_free_num_blocks, "primaryUsedNumBlocks": stats.primary_used_num_blocks, + "primaryEvictableNumBlocks": stats.primary_evictable_num_blocks, + "primaryPeakFreeNumBlocks": stats.primary_peak_free_num_blocks, + "primaryPeakUsedNumBlocks": stats.primary_peak_used_num_blocks, + "primaryPeakEvictableNumBlocks": stats.primary_peak_evictable_num_blocks, "secondaryMaxNumBlocks": stats.secondary_max_num_blocks, "secondaryFreeNumBlocks": stats.secondary_free_num_blocks, "secondaryUsedNumBlocks": stats.secondary_used_num_blocks, + "secondaryEvictableNumBlocks": stats.secondary_evictable_num_blocks, + "secondaryPeakFreeNumBlocks": stats.secondary_peak_free_num_blocks, + "secondaryPeakUsedNumBlocks": stats.secondary_peak_used_num_blocks, + "secondaryPeakEvictableNumBlocks": stats.secondary_peak_evictable_num_blocks, "iterAllocTotalBlocks": stats.iter_alloc_total_blocks, "iterAllocNewBlocks": stats.iter_alloc_new_blocks, "iterReusedBlocks": stats.iter_reused_blocks, @@ -89,6 +107,8 @@ def serialize_kv_cache_iteration_stats(stats, keys: tuple[str, ...] | None = Non "iterOffloadBytes": stats.iter_offload_bytes, "iterIntraDeviceCopyBlocks": stats.iter_intra_device_copy_blocks, "iterIntraDeviceCopyBytes": stats.iter_intra_device_copy_bytes, + "iterHostDroppedBlocks": stats.iter_host_dropped_blocks, + "iterHostDroppedBytes": stats.iter_host_dropped_bytes, } if keys is None: return fields diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index e492548fd449..208a239dbc1c 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -84,6 +84,7 @@ ResourceManagerType) from .sampler import SampleStateTensors from .scheduler import ScheduledRequests +from .trace_log_utils import log_mem_snapshot class ModelEngine(ABC): @@ -1033,6 +1034,7 @@ def warmup(self, resource_manager: ResourceManager) -> None: and self.guided_decoder is None and not isinstance(kv_cache_manager, MambaHybridCacheManager)) + log_mem_snapshot("warmup/before_warmup") self._run_attention_warmup(resource_manager, can_run_general_warmup) if can_run_general_warmup: @@ -1050,10 +1052,12 @@ def warmup(self, resource_manager: ResourceManager) -> None: # Memory pool will be warmed up later. gc.collect() torch.cuda.empty_cache() + # Autotuner warmup uses context-only requests. Helix CP # is decode-only and runs into issues with autotuner warmup. if not self.mapping.has_cp_helix(): self._run_autotuner_warmup(resource_manager) + log_mem_snapshot("warmup/after_autotuner") # Release the autotuner's exploration-mode intermediates. The # exploration leftovers are pure waste that hide tens of GiB from # non-torch allocators (cuBLAS handle workspace, UCX/NIXL, @@ -1062,12 +1066,14 @@ def warmup(self, resource_manager: ResourceManager) -> None: torch.cuda.empty_cache() with self.cuda_graph_runner.allow_capture(): self._run_cuda_graph_warmup(resource_manager) + log_mem_snapshot("warmup/after_cuda_graph_capture") if can_run_general_warmup: # Pre-populate the memory pool with max-shape allocations to reduce # fragmentation at runtime. warmup_requests_configs = self._get_max_shape_warmup_requests( resource_manager) self._general_warmup(resource_manager, warmup_requests_configs) + log_mem_snapshot("warmup/after_memory_pool_prepop") def _general_warmup(self, resource_manager: ResourceManager, warmup_requests_configs: List[Tuple[int, int]]): @@ -2636,9 +2642,14 @@ def _prepare_incremental_update_metadata( # Set iteration states - batch dictionary updates self.iter_states.update({ - 'num_ctx_requests': 0, - 'num_ctx_tokens': 0, - 'num_generation_tokens': num_generation_tokens + 'num_ctx_requests': + 0, + 'num_ctx_tokens': + 0, + 'num_generation_tokens': + num_generation_tokens, + 'cached_kv_tokens': + sum(num_cached_tokens_per_seq), }) return lora_params @@ -4009,6 +4020,8 @@ def previous_seq_slots_device(): self.iter_states['num_ctx_requests'] = num_ctx_requests self.iter_states['num_ctx_tokens'] = num_ctx_tokens self.iter_states['num_generation_tokens'] = num_generation_tokens + # Count the already-cached prefix for the sequences scheduled this iteration. + self.iter_states['cached_kv_tokens'] = sum(num_cached_tokens_per_seq) if not self.is_warmup: self.previous_request_ids = all_gen_request_ids diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 4257a6f61f59..352f41cdbe3a 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -89,6 +89,13 @@ if TYPE_CHECKING: from ray.actor import ActorHandle +_UNBOUNDED_STATS_MAX_LEN = -1 + + +def _stats_buffer_is_unbounded(max_stats_len: int) -> bool: + return max_stats_len == _UNBOUNDED_STATS_MAX_LEN + + # Environment variable to specify iteration ranges for profiling start/stop. # Format: "start1-stop1,start2-stop2,..." or single iterations "iter1,iter2,..." PROFILE_START_STOP_ENV_VAR_NAME = "TLLM_PROFILE_START_STOP" @@ -490,7 +497,7 @@ def __init__( self.max_draft_len = max_draft_len self.max_total_draft_tokens = max_total_draft_tokens self.llm_args = self.model_engine.llm_args - self.max_stats_len = max(self.llm_args.max_stats_len, 1) + self.max_stats_len = self.llm_args.max_stats_len self.max_num_tokens = self.llm_args.max_num_tokens self.print_log = self.llm_args.print_iter_log self.enable_iter_perf_stats = self.llm_args.enable_iter_perf_stats @@ -545,14 +552,10 @@ def __init__( # kv cache events self.kv_cache_manager = self.resource_manager.resource_managers.get( ResourceManagerType.KV_CACHE_MANAGER) - # V2 manager owns KV alloc + suspend during scheduling: it - # eagerly grows ctx/gen capacity in the schedule loop and calls - # suspend_request() when needed (offloads GPU pages while - # preserving the radix tree). The executor therefore does not - # need to call _terminate_requests (GPU resources are already - # freed by suspend) or _pause_requests (V2's prepare_context - # handles resume internally, so resetting to CONTEXT_INIT is - # unnecessary). Several revert/skip paths gate on this flag. + # V2 owns KV allocation, suspend, resume, and context finalization. + # The executor skips the V1 terminate/pause paths and finalizes V2 + # context resources before transfer or response handling can terminate + # a request. self._is_kv_manager_v2 = isinstance(self.kv_cache_manager, KVCacheManagerV2) self._prefetched_request_ids: set[int] = set() @@ -1032,11 +1035,13 @@ def _flush_iter_stats_synced(self): if not rank_dicts: return with self.stats_lock: + if not _stats_buffer_is_unbounded(self.max_stats_len): + cap = self.max_stats_len * tp_size + overflow = max(0, len(self.stats) + len(rank_dicts) - cap) + if overflow: + del self.stats[:overflow] for d in rank_dicts: self.stats.append(("per_rank_dict", d)) - cap = self.max_stats_len * tp_size - if len(self.stats) > cap: - del self.stats[:len(self.stats) - cap] # Performance metrics methods are in PerfMetricsManager (self.perf_manager) @@ -1977,7 +1982,8 @@ def _append_iter_stats(self, # [6] scheduler_mode: "overlap" | "non_overlap" # [7] gpu_forward_time_ms: Optional[float] with self.stats_lock: - if len(self.stats) > self.max_stats_len: + if (not _stats_buffer_is_unbounded(self.max_stats_len) + and len(self.stats) > self.max_stats_len): self.stats.pop(0) self.stats.append( (stats, req_stats, kv_iter_stats, attention_dp_rank, @@ -2580,6 +2586,11 @@ def _handle_executed_batch(self, executed_batch: Optional[BatchStatePP]): self._update_requests(executed_batch.sample_state) scheduled_requests = executed_batch.scheduled_requests + if self._is_kv_manager_v2: + # Finalize V2 context KV before disagg transfer/response + # handling can terminate the request. + self.kv_cache_manager.update_context_resources( + scheduled_requests) if self.kv_cache_transceiver: finished_ctx_reqs = scheduled_requests.context_requests_last_chunk self._send_kv_async(finished_ctx_reqs) @@ -2716,18 +2727,7 @@ def _revert_gen_alloc(self, scheduled_batch): self.kv_cache_manager.revert_allocate_generation(req) def _revert_ctx_alloc(self, dropped_context_requests): - """Revert KV cache capacity growth for ctx requests deferred by - delay batching. - - With KV cache manager V2 + scheduler V2, ctx KV cache is grown - during scheduling (``resize_context``). When delay batching - (``_balance_adp_requests`` for ADP, or ``_waiting_requests`` - for non-ADP batch waiting) defers ctx requests, the - freshly-allocated pages would otherwise sit idle until the - request is re-scheduled, blocking pool space — particularly - painful for long-context workloads where each deferred ctx can - hold GBs of KV. - """ + """Revert V2 context KV growth for requests deferred after scheduling.""" for req in dropped_context_requests: self.kv_cache_manager.revert_allocate_context(req) @@ -3336,6 +3336,11 @@ def _executor_loop(self): self._update_request_states(scheduled_batch) self._update_requests(sample_state, self.resource_manager) + if self._is_kv_manager_v2: + # Finalize V2 context KV before disagg transfer/response + # handling can terminate the request. + self.kv_cache_manager.update_context_resources( + scheduled_batch) self._send_kv_async(scheduled_batch.all_requests()) self._flush_pending_transfer_responses() @@ -3593,6 +3598,12 @@ def control_action(self, *, drain: bool = True): self.control_action_done.set() self.control_request_barrier.clear() + def _wait_for_model_engine_input_copy(self): + wait_for_input_copy = getattr(self.model_engine, "wait_for_input_copy", + None) + if wait_for_input_copy is not None: + wait_for_input_copy() + def _executor_loop_overlap(self): torch.cuda.set_device(self.device_id) # ensure the context is created, otherwise, some MPI calls will fail. @@ -3617,6 +3628,12 @@ def _executor_loop_overlap(self): self._handle_disagg_cache_errors_synced() + # Need to wait for the copy of previous iteration before + # modifying any host memory copied to GPU. Scheduler V2 + # modifies the host page table, so wait before scheduling. + # This wait is also needed for legacy scheduler, but it can + # be pushed later, e.g. before model_engine._prepare_inputs(). + self._wait_for_model_engine_input_copy() scheduled_batch, iter_stats = self._prepare_and_schedule_batch() if scheduled_batch is None: @@ -3802,11 +3819,29 @@ def _executor_loop_overlap(self): # causing _sample_async to fail when accessing context_chunk_size property. self._handle_guided_decoder_errors( scheduled_batch, guided_decoder_failed_requests) + # _update_request_states() can terminate attention-DP + # dummy requests, which frees V2 KV pages and overwrites + # host page-index entries with BAD_PAGE_INDEX. Wait until + # the current input preparation has consumed those buffers. + self._wait_for_model_engine_input_copy() self._update_request_states(scheduled_batch) + # Update context requests' KV cache so that sliding-window + # blocks freed by this chunk are visible to the next + # iteration's scheduler. + # Only applies to KV cache manager V2 + scheduler V2. + if (self._is_kv_manager_v2 + and scheduled_batch.context_requests): + self.kv_cache_manager.update_context_resources( + scheduled_batch) + if self.previous_batch is not None and should_process_previous_batch: self._commit_kv_cache_stats( self.previous_batch.scheduled_requests) + # _process_previous_batch may terminate requests or resize + # generation KV caches, both of which can mutate V2 page + # indices used by the current batch's input preparation. + self._wait_for_model_engine_input_copy() self._process_previous_batch() self.perf_manager.compute_batch_gpu_times( self.previous_batch.scheduled_requests.all_requests()) @@ -4051,7 +4086,9 @@ def _fetch_and_enqueue_requests(self, waiting_queue: WaitingQueue, new_requests.extend( self.executor_request_queue.get_from_request_queue(timeout)) - # Broadcast requests and handle Python objects + # Broadcast requests and handle Python objects. RequestBroadcaster probes + # the request count first and can skip the heavy payload broadcast on + # empty iterations. new_requests, py_request_objects = self.request_broadcaster.broadcast( new_requests) @@ -4157,7 +4194,8 @@ def _fetch_new_requests( # 6. Schedule requests across ranks (DP only) if self.enable_attention_dp: - if self.adp_router.needs_prefix_matches: + # Symmetric skip — after _pop_from_waiting_queue all ranks see identical new_requests. + if self.adp_router.needs_prefix_matches and new_requests: self.adp_router.gather_prefix_matches(new_requests) all_ranks_new_requests, self.expected_num_active_requests = \ @@ -4377,8 +4415,7 @@ def _schedule(self): scheduler_output = self.scheduler.schedule_request( self.active_requests, self.inflight_req_ids) - original_ctx_requests = scheduler_output.context_requests - scheduled_context_requests = original_ctx_requests + scheduled_context_requests = scheduler_output.context_requests if self.enable_attention_dp and self.attention_dp_enable_balance: scheduled_context_requests = self._balance_adp_requests( scheduler_output.context_requests, @@ -4389,6 +4426,8 @@ def _schedule(self): scheduler_output.context_requests) > 0 and len( scheduler_output.generation_requests) > 0 if should_check_waiting: + # With KV cache manager V2, scheduling has already grown context request KV cache capacity. Requests dropped + # for batch waiting still occupy KV cache and may reduce the batch size available for generation requests. scheduled_context_requests = self._waiting_requests( scheduler_output.context_requests, scheduler_output.generation_requests) @@ -4403,19 +4442,6 @@ def _schedule(self): scheduled_context_requests) num_fitting = len(scheduled_context_requests) - # V2 scheduler grew KV cache for ctx during scheduling; release - # those pages for any ctx that delay batching has dropped, so - # the wait window does not hold pool capacity hostage. V1 - # allocates after delay batching, so skip the dropped-set - # computation entirely on V1. - if (self._is_kv_manager_v2 and len(scheduled_context_requests) - < len(original_ctx_requests)): - kept = {r.py_request_id for r in scheduled_context_requests} - dropped = [ - r for r in original_ctx_requests if r.py_request_id not in kept - ] - self._revert_ctx_alloc(dropped) - scheduled_requests = ScheduledRequests() scheduled_requests.encoder_requests = scheduler_output.encoder_requests scheduled_requests.reset_context_requests(scheduled_context_requests) diff --git a/tensorrt_llm/_torch/pyexecutor/request_utils.py b/tensorrt_llm/_torch/pyexecutor/request_utils.py index ae1dfba65e6d..e7da86608153 100644 --- a/tensorrt_llm/_torch/pyexecutor/request_utils.py +++ b/tensorrt_llm/_torch/pyexecutor/request_utils.py @@ -512,6 +512,15 @@ def __init__(self, dist: Distributed, hang_detector: HangDetector): def broadcast(self, new_requests: List) -> Tuple[List, Optional[Tuple]]: """Broadcast requests and Python objects across ranks.""" + request_count = len(new_requests) if self.dist.rank == 0 else 0 + # Idle non-root ranks can wait here while rank 0 blocks in the + # pause-wrapped request queue fetch, so keep the probe pause-wrapped too. + with self.hang_detector.pause(): + request_count = self._broadcast_request_count(request_count) + + if request_count == 0: + return [], None + if self.dist.rank == 0: py_request_objects = self._collect_py_objects(new_requests) else: @@ -528,6 +537,30 @@ def broadcast(self, new_requests: List) -> Tuple[List, Optional[Tuple]]: return new_requests, py_request_objects + def _broadcast_request_count(self, request_count: int) -> int: + """Broadcast rank 0's request count using the same PP route as requests.""" + if self.dist.world_size == 1: + return request_count + + if not self.dist.has_pp: + return self.dist.broadcast(request_count, root=0) + + if self.dist.is_first_pp_rank: + with nvtx_range("tp_broadcast_request_count"): + request_count = self.dist.tp_cp_broadcast(request_count, root=0) + + tag = self.dist.pp_size + 1 # Avoid the heavy request payload tag. + + if not self.dist.is_first_pp_rank: + with nvtx_range("recv_request_count_from_prev_pp"): + request_count = self.dist.recv_object(self.dist.prev_pp_rank, tag) + + if not self.dist.is_last_pp_rank: + with nvtx_range("send_request_count_to_next_pp"): + self.dist.send_object(request_count, self.dist.next_pp_rank, tag) + + return request_count + def _collect_py_objects(self, new_requests: List) -> Tuple: """Collect Python-only objects from requests.""" py_logits_post_processors = collect_py_objects_from_requests( diff --git a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py index dbd2652058b9..2accbcc152b3 100644 --- a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py @@ -495,15 +495,9 @@ def _try_schedule_disagg_gen_init( Returns ``(action, tokens)``. *tokens* is 0 because disagg requests don't participate in the forward pass token budget. """ - if not self.kv_cache_manager.prepare_context(req): - logger.debug("prepare_context failed for disagg gen init request %s", req.py_request_id) - return ScheduleAction.SKIP, 0 - - if not self.kv_cache_manager.resize_context( - req, req.context_remaining_length + get_draft_token_length(req) - ): + if not self.kv_cache_manager.prepare_disagg_gen_init(req): + logger.debug("prepare_disagg_gen_init failed for request %s", req.py_request_id) return ScheduleAction.SKIP, 0 - return ScheduleAction.SCHEDULED, 0 def _try_schedule_context( diff --git a/tensorrt_llm/_torch/pyexecutor/trace_log_utils.py b/tensorrt_llm/_torch/pyexecutor/trace_log_utils.py new file mode 100644 index 000000000000..6a83ef4194d7 --- /dev/null +++ b/tensorrt_llm/_torch/pyexecutor/trace_log_utils.py @@ -0,0 +1,72 @@ +"""Gated trace/log utilities for pyexecutor. + +Leaf module — no other pyexecutor file is imported here, so any consumer +(``_util``, ``model_engine``, ``model_loader``, ``resource_manager``) +can import freely without creating circular dependencies. +""" + +import os + +import torch + +from tensorrt_llm.logger import logger + +_GIB = 1 << 30 + + +def log_mem_snapshot(tag: str) -> None: + """Log Torch alloc/reserved + alloc/reserved peak + free/total GPU memory. + + Gated by ``TLLM_LOG_MEM_PROFILE=1``; default OFF (zero overhead). + + Prints these fields: + + - ``torch_alloc`` = :func:`torch.cuda.memory_allocated` + - ``torch_reserved`` = :func:`torch.cuda.memory_reserved` + - ``torch_alloc_peak`` = :func:`torch.cuda.max_memory_allocated` + - ``torch_reserved_peak`` = :func:`torch.cuda.max_memory_reserved` + - ``free`` = ``cuMemGetInfo().free`` + - ``total`` = ``cuMemGetInfo().total`` + + Derived quantities the reader may need: + + - ``used = total - free`` — whole-process GPU consumption + - ``slack = reserved - alloc`` — Torch caching allocator free blocks + - ``non_torch = used - reserved`` — bytes outside Torch (KV pool C++ + cudaMalloc, NCCL buffers, cuBLAS workspace, CUDA driver context, + CUDA graph mempool, etc.) + """ + if os.environ.get("TLLM_LOG_MEM_PROFILE", "") != "1": + return + free, total = torch.cuda.mem_get_info() + alloc = torch.cuda.memory_allocated() + reserved = torch.cuda.memory_reserved() + alloc_peak = torch.cuda.max_memory_allocated() + reserved_peak = torch.cuda.max_memory_reserved() + logger.info( + f"[mem-profile/{tag}] " + f"torch_alloc={alloc / _GIB:.2f}GiB " + f"torch_reserved={reserved / _GIB:.2f}GiB " + f"torch_alloc_peak={alloc_peak / _GIB:.2f}GiB " + f"torch_reserved_peak={reserved_peak / _GIB:.2f}GiB " + f"free={free / _GIB:.2f}GiB total={total / _GIB:.2f}GiB" + ) + + +def log_tensor_size(tag: str, tensor: torch.Tensor, **extra) -> None: + """Log a single tensor's footprint (shape / dtype / bytes) at a tag. + + Gated by ``TLLM_LOG_MEM_PROFILE=1``; default OFF (zero overhead). + + Bytes = ``numel * element_size``. Any keyword arguments are appended + as ``key=value`` for caller-specific context (e.g. routing config). + """ + if os.environ.get("TLLM_LOG_MEM_PROFILE", "") != "1": + return + size_bytes = tensor.numel() * tensor.element_size() + extras = "".join(f" {k}={v}" for k, v in extra.items()) + logger.info( + f"[mem-profile/{tag}] " + f"shape={tuple(tensor.shape)} dtype={tensor.dtype} " + f"size={size_bytes / 1024 / 1024:.2f}MiB{extras}" + ) diff --git a/tensorrt_llm/executor/base_worker.py b/tensorrt_llm/executor/base_worker.py index b2263242a2e5..d3f324f755de 100644 --- a/tensorrt_llm/executor/base_worker.py +++ b/tensorrt_llm/executor/base_worker.py @@ -395,7 +395,18 @@ def _enqueue_request(self, else: lora_config = None - prompt_token_ids = list(request.prompt_token_ids) + # prompt_token_ids stays list[int] for all consumers. If an int32 buffer + # rode along on the wire (GenerationRequest._prompt_token_ids_i32), hand + # THAT to the C++ Request ctor (memcpy) instead of the list -- this avoids + # the O(ISL) list copy + element-wise nanobind cast on the GIL-held submit + # thread. No buffer (in-process, or prompt-adapter prepend) -> list path. + # If both forms exist, they are assumed to describe the same token ids; + # in-place mutations of request.prompt_token_ids must clear the i32 buffer. + i32_buf = getattr(request, "_prompt_token_ids_i32", None) + if i32_buf is not None and request.prompt_adapter_request is None: + prompt_token_ids = i32_buf + else: + prompt_token_ids = list(request.prompt_token_ids) prompt_tuning_config = None if request.prompt_adapter_request is not None: self._load_prompt_adapter(request.prompt_adapter_request) diff --git a/tensorrt_llm/executor/request.py b/tensorrt_llm/executor/request.py index a43e89c0e4f7..a003b387f45b 100644 --- a/tensorrt_llm/executor/request.py +++ b/tensorrt_llm/executor/request.py @@ -184,6 +184,79 @@ def set_id(self, id): self.id = id return self + # --- int32 token-id wire serialization + LAZY list materialization ---------- + # Pickling a flat list[int] on the hot proxy->worker RPC-submit path emits one + # PyLong frame per token (O(ISL)). Encode token-ids as int32 bytes for the wire. + # On decode we DO NOT eagerly rebuild the list: stash the int32 ndarray in + # `_prompt_token_ids_i32` (the C++ Request ctor memcpy's it -- see + # base_worker._enqueue_request) and leave the backing `_prompt_token_ids` None. + # The list is built lazily (and cached) by the `prompt_token_ids` property below + # only if a consumer actually reads it (e.g. prompt-logprobs, star-attention). + # Plain decode never reads it -> the O(ISL) `.tolist()` never runs. + # + # NOTE: implemented as a *property* (scoped to this one name), NOT a class-level + # __getattr__. A __getattr__ is invoked by the interpreter on EVERY missing + # attribute access on the object -- hasattr()/getattr(default) probes, copy/ + # pickle dunder lookups, duck-typing -- each paying a Python frame + an + # AttributeError raise. A property has identical lazy-materialize behavior with + # zero blast radius on any other attribute. + _I32 = "\x00i32be" + + @property + def prompt_token_ids(self): + ptids = self.__dict__.get("_prompt_token_ids") + if ptids is None: + buf = self.__dict__.get("_prompt_token_ids_i32") + if buf is not None: + ptids = buf.tolist() + self._prompt_token_ids = ptids # cache + # Callers that mutate this returned list in place must clear + # `_prompt_token_ids_i32`; base_worker assumes both forms stay synced. + return ptids + + @prompt_token_ids.setter + def prompt_token_ids(self, value): + self._prompt_token_ids = value + + @staticmethod + def _enc_tokens(v): + # flat list[int] -> (_I32, int32 bytes); leave None / list[list[int]] / + # ndarray untouched. + if type(v) is list and (len(v) == 0 or type(v[0]) is int): + return (GenerationRequest._I32, + np.asarray(v, dtype=np.int32).tobytes()) + return v + + def __getstate__(self): + state = self.__dict__.copy() + buf = state.pop("_prompt_token_ids_i32", None) + ptids = state.get("_prompt_token_ids") + if ptids is not None: + state["_prompt_token_ids"] = GenerationRequest._enc_tokens(ptids) + elif buf is not None: + # not yet materialized -> encode the buffer's bytes directly + state["_prompt_token_ids"] = (GenerationRequest._I32, buf.tobytes()) + if state.get("query_token_ids") is not None: + state["query_token_ids"] = GenerationRequest._enc_tokens( + state["query_token_ids"]) + return state + + def __setstate__(self, state): + buf = None + pt = state.get("_prompt_token_ids") + if type(pt) is tuple and len( + pt) == 2 and pt[0] == GenerationRequest._I32: + buf = np.frombuffer(pt[1], dtype=np.int32) + state["_prompt_token_ids"] = None # leave None -> lazy via property + qt = state.get("query_token_ids") + if type(qt) is tuple and len( + qt) == 2 and qt[0] == GenerationRequest._I32: + # query_token_ids is rare/small -> materialize to list eagerly + state["query_token_ids"] = np.frombuffer(qt[1], + dtype=np.int32).tolist() + self.__dict__.update(state) + self._prompt_token_ids_i32 = buf + class TruncateKVCacheRequest: diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 803714f0d46a..50dcae69f1cc 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -3367,6 +3367,13 @@ class KvCacheConfig(StrictBaseModel, PybindMirror): status="prototype", description="Whether to use the KV cache manager v2 (experimental).") + # This is a pure python field, not a pybind field. It is only for the Pytorch backend. + enable_swa_scratch_reuse: bool = Field( + default=False, + status="prototype", + description= + "Whether KV cache manager v2 uses SWA scratch reuse during prefill.") + kv_cache_event_hash_algo: Literal[ "auto", "v1_block_key", "v2_sha256", "v2_sha256_64"] = Field( default="auto", @@ -3407,6 +3414,34 @@ class KvCacheConfig(StrictBaseModel, PybindMirror): "Set to 0 to disable prefetch. Only effective with KV cache manager v2 and block reuse enabled." ) + # This is a pure python field, not a pybind field. It is only for the Pytorch backend. + pool_ratio: Optional[List[float]] = Field( + default=None, + min_length=1, + status="prototype", + description= + "Initial pool ratios for KV cache manager v2. When used by DeepSeek-V4, " + "values map to KVCacheManagerV2 pool_group_id order and must sum to 1.0. " + "When set, DeepSeek-V4 uses this directly and avg_seq_len does not take effect." + ) + + # This is a pure python field, not a pybind field. It is only for the Pytorch backend. + avg_seq_len: Optional[PositiveInt] = Field( + default=None, + status="prototype", + description= + "Average sequence length used by DeepSeek-V4 to build the KV cache manager v2 " + "typical step. If unset, max_seq_len is used. This does not take effect when " + "pool_ratio is set.") + + # This is a pure python field, not a pybind field. It is only for the Pytorch backend. + block_reuse_policy: Literal["all_reusable", "per_request"] = Field( + default="all_reusable", + status="prototype", + description="KV cache manager v2 block reuse policy. " + "With SWA scratch reuse and 'all_reusable', only non-scratch " + "blocks are saved for reuse.") + def _to_pybind(self): config = _KvCacheConfig( enable_block_reuse=self.enable_block_reuse, @@ -3508,6 +3543,19 @@ def validate_max_util_for_resume(cls, v: float): "kv_cache_config.max_util_for_resume must be between 0 and 1") return v + @field_validator('pool_ratio') + @classmethod + def validate_pool_ratio(cls, v: Optional[List[float]]): + if v is None: + return v + if any(r <= 0 for r in v): + raise ValueError( + "kv_cache_config.pool_ratio values must be positive") + if not math.isclose(sum(v), 1.0, rel_tol=0.0, abs_tol=1e-6): + raise ValueError( + "kv_cache_config.pool_ratio values must sum to 1.0") + return v + @PybindMirror.mirror_pybind_fields(_ExtendedRuntimePerfKnobConfig) class ExtendedRuntimePerfKnobConfig(StrictBaseModel, PybindMirror): @@ -3808,12 +3856,18 @@ class BaseLlmArgs(StrictBaseModel): iter_stats_max_iterations: Optional[int] = Field( default=None, - description="The maximum number of iterations for iter stats.", + ge=-1, + description= + "The maximum number of iterations for iter stats. Set to -1 to keep all iteration stats. " + "Set to 0 to disable iteration stats in the TensorRT executor.", status="prototype") request_stats_max_iterations: Optional[int] = Field( default=None, - description="The maximum number of iterations for request stats.", + ge=-1, + description= + "The maximum number of iterations for request stats. Set to -1 to keep all request stats. " + "Set to 0 to disable request stats.", status="prototype") # A handful of options from PretrainedConfig @@ -5014,8 +5068,19 @@ def validate_encoder_runtime_sizes(cls, v: Optional[int]) -> Optional[int]: max_stats_len: int = Field( default=1000, - description="The max number of performance statistic entries.", - status="prototype") + ge=-1, + description= + "The max number of performance statistic entries. Set to -1 to keep all entries. " + "Set to 0 to use a minimum buffer size of 1.", + status="prototype", + ) + + @field_validator('max_stats_len') + @classmethod + def normalize_max_stats_len(cls, v): + if v == -1: + return v + return max(v, 1) layer_wise_benchmarks_config: LayerwiseBenchmarksConfig = Field( default_factory=LayerwiseBenchmarksConfig, diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py index 8e118c11f013..28276acb310b 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py @@ -47,6 +47,7 @@ BeamIndex, KVCacheManager, PageIndexConverter, + PoolGroupPeakBlockStats, ScratchDesc, _KVCache, ) @@ -108,6 +109,7 @@ "AggregatedPageDesc", "BufferId", "PageIndexConverter", + "PoolGroupPeakBlockStats", "PageIndexMode", "ScratchDesc", "KVCacheIterationStatsDelta", diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi index ee4a4badc218..ee15c8575df2 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi @@ -86,6 +86,14 @@ class KVCacheIterationStatsDelta: iter_offload_bytes: int = 0 iter_intra_device_copy_blocks: int = 0 iter_intra_device_copy_bytes: int = 0 + iter_host_dropped_blocks: int = 0 + iter_host_dropped_bytes: int = 0 + +@dataclass(slots=True, frozen=True) +class PoolGroupPeakBlockStats: + available: int + unavailable: int + evictable: int # From _config.py DataRole = NewType("DataRole", str) @@ -171,6 +179,7 @@ class KVCacheManagerConfig: enable_partial_reuse: bool = True constraints: list[BatchDesc] = ... typical_step: BatchDesc | None = None + initial_pool_ratio: list[float] | None = None ssm_reuse_interval: int = 512 swa_scratch_reuse: SwaScratchReuseConfig | None = None enable_stats: bool = True @@ -449,6 +458,9 @@ class KVCacheManager: 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 get_and_reset_iteration_peak_block_stats( + self, cache_level: CacheLevel + ) -> Sequence[PoolGroupPeakBlockStats]: ... 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]: ... diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py index 9341141a139a..3bb9b0eb20a7 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py @@ -223,6 +223,12 @@ class KVCacheManagerConfig: layer groups. """ + initial_pool_ratio: list[float] | None = None + """ + User-provided initial memory partitioning between pool groups. When set, this + takes precedence over typical_step and constraints for initial sizing. + """ + ssm_reuse_interval: int = 512 """ Interval (in tokens) at which SSM state is snapshotted for prefix reuse. diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_core/__init__.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_core/__init__.py index dc9ddb706674..09157ba67d36 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_core/__init__.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_core/__init__.py @@ -15,7 +15,13 @@ from .._common import DEFAULT_BEAM_INDEX, BeamIndex from ._kv_cache import _KVCache -from ._kv_cache_manager import AggregatedPageDesc, KVCacheManager, PageIndexConverter, ScratchDesc +from ._kv_cache_manager import ( + AggregatedPageDesc, + KVCacheManager, + PageIndexConverter, + PoolGroupPeakBlockStats, + ScratchDesc, +) __all__ = [ "KVCacheManager", @@ -24,5 +30,6 @@ "DEFAULT_BEAM_INDEX", "AggregatedPageDesc", "PageIndexConverter", + "PoolGroupPeakBlockStats", "ScratchDesc", ] 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 d2c66568808a..f8063c961fd8 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 @@ -484,6 +484,32 @@ def _record_migrated_slots( if not stats.empty or not iteration_stats.empty: self.manager.commit_stats(stats, {life_cycle_key: iteration_stats}) + def _record_dropped_pages( + self, + pages: Sequence[Page], + cache_level: CacheLevel, + ) -> None: + """Record host-tier LRU drops (pages released without onboarding back to GPU). + + Mirrors _record_migrated_slots in structure: per-life-cycle attribution, + gated on _should_record_stats(), per-page bytes computed from slot_size. + cache_level is unused for now (we only have a 2-tier setup in practice; + all last-level drops are host drops) but kept in the signature for future + per-tier disambiguation. + """ + if not self._should_record_stats() or not pages: + return + 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)) + iteration_stats = KVCacheIterationStatsDelta() + iteration_stats.iter_host_dropped_blocks = 1 + iteration_stats.iter_host_dropped_bytes = page_size + self.manager.commit_stats(KVCacheStatsDelta(), {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: @@ -734,6 +760,7 @@ def resize(self, capacity: int | None, history_length: int | None = None) -> boo new_slots = storage.new_gpu_slots( make_typed(lambda lc: max(0, net_alloc_counts[lc]), num_life_cycles), self._record_migrated_slots, + self._record_dropped_pages, ) except OutOfPagesError: self._recover_excess_scratch_slots(excess_scratch_slots) @@ -996,7 +1023,9 @@ 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, self._record_migrated_slots) + tmp_slots = storage.new_gpu_slots( + num_slots, self._record_migrated_slots, self._record_dropped_pages + ) except OutOfPagesError: return False @@ -1029,7 +1058,9 @@ 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, self._record_migrated_slots) + locks = batched_lock_to_gpu( + self, tasks, self._record_migrated_slots, self._record_dropped_pages + ) except OutOfPagesError: for lc_idx, slot in typed_enumerate(deferred_slots): if slot is not None: @@ -1382,6 +1413,7 @@ def _commit_block(self, ordinal: BlockOrdinal, is_last: bool) -> None: self, [BatchedLockTarget(p, beam_idx, ordinal, lc) for lc, p in reuse_list], self._record_migrated_slots, + self._record_dropped_pages, ) for (lc, _), lock in zip(reuse_list, locks): beam_block[lc] = lock @@ -1472,6 +1504,7 @@ def _lock_held_blocks( for ordinal, beam_idx, lc, holder in backup_holders ], self._record_migrated_slots, + self._record_dropped_pages, ) for lock in locks: user = lock._user 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 d358d064c358..fb60f98bb94a 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 @@ -42,7 +42,7 @@ 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, StorageStatistics +from .._storage_manager import StorageManager from .._utils import ( HalfOpenRange, HomoTuple, @@ -183,12 +183,10 @@ def __call__( @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 PoolGroupPeakBlockStats: + available: int + unavailable: int + evictable: int class KVCacheManager: @@ -211,6 +209,7 @@ class KVCacheManager: "_stats_enabled", "_committed_stats", "_iteration_stats_by_life_cycle", + "_iteration_peak_num_blocks_by_cache_level", "_dirty_stats_kv_cache_ids", "_stats_excluded_kv_cache_ids", ) @@ -239,6 +238,9 @@ class KVCacheManager: _stats_enabled: bool _committed_stats: KVCacheStatsDelta _iteration_stats_by_life_cycle: dict[LifeCycleId, KVCacheIterationStatsDelta] + _iteration_peak_num_blocks_by_cache_level: TypedIndexList[ + CacheLevel, TypedIndexList[PoolGroupIndex, PoolGroupPeakBlockStats] + ] _dirty_stats_kv_cache_ids: set[int] _stats_excluded_kv_cache_ids: set[int] @@ -260,6 +262,7 @@ def __init__( config.swa_scratch_reuse, typical_batch=config.typical_step, constraints=config.constraints, + initial_pool_ratio=config.initial_pool_ratio, event_manager=event_manager, ) self._living_kv_caches = set[rawref.ref[_KVCache]]() @@ -277,6 +280,7 @@ def __init__( self._stats_enabled = config.enable_stats self._committed_stats = KVCacheStatsDelta() self._iteration_stats_by_life_cycle = {} + self._reset_iteration_peak_num_blocks() self._dirty_stats_kv_cache_ids = set() self._stats_excluded_kv_cache_ids = set() @@ -457,18 +461,54 @@ 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 _current_block_stats_by_cache_level( + self, + ) -> TypedIndexList[CacheLevel, TypedIndexList[PoolGroupIndex, PoolGroupPeakBlockStats]]: + def collect( + cache_level: CacheLevel, + ) -> TypedIndexList[PoolGroupIndex, PoolGroupPeakBlockStats]: + stats_by_pool_group = self._storage.get_statistics(cache_level) + return make_typed( + lambda pool_group_index: PoolGroupPeakBlockStats( + available=stats_by_pool_group[pool_group_index].available, + unavailable=stats_by_pool_group[pool_group_index].unavailable, + evictable=stats_by_pool_group[pool_group_index].evictable, + ), + self._storage.num_pool_groups, + ) + + return make_typed(collect, self._storage.num_cache_levels) + + def _reset_iteration_peak_num_blocks(self, cache_level: CacheLevel | None = None) -> None: + if cache_level is None: + self._iteration_peak_num_blocks_by_cache_level = ( + self._current_block_stats_by_cache_level() + ) + return + stats_by_pool_group = self._storage.get_statistics(cache_level) + self._iteration_peak_num_blocks_by_cache_level[cache_level] = make_typed( + lambda pool_group_index: PoolGroupPeakBlockStats( + available=stats_by_pool_group[pool_group_index].available, + unavailable=stats_by_pool_group[pool_group_index].unavailable, + evictable=stats_by_pool_group[pool_group_index].evictable, + ), + self._storage.num_pool_groups, ) + def _update_iteration_peak_num_blocks(self) -> None: + current = self._current_block_stats_by_cache_level() + for cache_level in typed_range(self._storage.num_cache_levels): + peak = self._iteration_peak_num_blocks_by_cache_level[cache_level] + current_level = current[cache_level] + for pool_group_index in typed_range(self._storage.num_pool_groups): + peak_stats = peak[pool_group_index] + current_stats = current_level[pool_group_index] + peak[pool_group_index] = PoolGroupPeakBlockStats( + available=max(peak_stats.available, current_stats.available), + unavailable=max(peak_stats.unavailable, current_stats.unavailable), + evictable=max(peak_stats.evictable, current_stats.evictable), + ) + def commit_stats( self, stats: KVCacheStatsDelta, @@ -476,6 +516,7 @@ def commit_stats( ) -> None: if not self._stats_enabled: return + self._update_iteration_peak_num_blocks() self._committed_stats.add(stats) if iteration_stats_by_life_cycle is None: return @@ -499,6 +540,27 @@ def get_and_reset_iteration_stats(self) -> dict[LifeCycleId, KVCacheIterationSta self._iteration_stats_by_life_cycle.clear() return stats + def get_and_reset_iteration_peak_block_stats( + self, cache_level: CacheLevel + ) -> TypedIndexList[PoolGroupIndex, PoolGroupPeakBlockStats]: + self._update_iteration_peak_num_blocks() + peak = make_typed( + lambda pool_group_index: PoolGroupPeakBlockStats( + available=self._iteration_peak_num_blocks_by_cache_level[cache_level][ + pool_group_index + ].available, + unavailable=self._iteration_peak_num_blocks_by_cache_level[cache_level][ + pool_group_index + ].unavailable, + evictable=self._iteration_peak_num_blocks_by_cache_level[cache_level][ + pool_group_index + ].evictable, + ), + self._storage.num_pool_groups, + ) + self._reset_iteration_peak_num_blocks(cache_level) + return peak + 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) diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_page.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_page.py index b276387e2385..4d6fbfd88ebb 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_page.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_page.py @@ -445,6 +445,7 @@ def batched_lock_to_gpu( tasks: Sequence[BatchedLockTarget], migration_recorder: Callable[[Sequence[Page], Sequence[Slot], CacheLevel, CacheLevel], None] | None = None, + drop_recorder: Callable[[Sequence[Page], 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 @@ -460,7 +461,7 @@ def batched_lock_to_gpu( requirements[lc2pg[t.life_cycle]] += 1 try: - storage.prepare_free_slots(GPU_LEVEL, requirements, migration_recorder) + storage.prepare_free_slots(GPU_LEVEL, requirements, migration_recorder, drop_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: diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_stats.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_stats.py index 73b104e7c02a..a02b55dbbf60 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_stats.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_stats.py @@ -64,6 +64,11 @@ class KVCacheIterationStatsDelta(_StatsDeltaMixin): iter_offload_bytes: int = 0 iter_intra_device_copy_blocks: int = 0 iter_intra_device_copy_bytes: int = 0 + # Host-tier pages released by LRU without ever being onboarded back to GPU + # in the lifetime since they were offloaded. Counted at the drop site in + # _storage_manager._prepare_free_slots when is_last_level(lvl). + iter_host_dropped_blocks: int = 0 + iter_host_dropped_bytes: int = 0 @property def iter_cache_hit_rate(self) -> float: 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 7a32c732c3bf..eb32717f8324 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_storage_manager.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_storage_manager.py @@ -170,6 +170,9 @@ def unavailable(self) -> int: MigrationRecorder = Callable[[Sequence[Page], Sequence[Slot], CacheLevel, CacheLevel], None] +# Invoked when pages at the last cache level are released to free their slots +# without being migrated to any further tier (i.e. dropped from the cache hierarchy). +DropRecorder = Callable[[Sequence[Page], CacheLevel], None] class StorageManager: @@ -208,6 +211,7 @@ def __init__( swa_scratch_reuse: SwaScratchReuseConfig | None, typical_batch: BatchDesc | None = None, constraints: list[BatchDesc] | None = None, + initial_pool_ratio: list[float] | None = None, event_manager: "KVCacheEventManager | None" = None, ) -> None: self.__rawref__ = rawref.NULL @@ -233,12 +237,24 @@ def __init__( gpu_quota = config.cache_tiers[GPU_LEVEL].quota gpu_granularity = CacheLevelManager.cache_tier_granularity(CacheTier.GPU_MEM, gpu_quota) + constraints_for_min_slots = [] if initial_pool_ratio is not None else constraints or [] self._min_slots = self._compute_min_slots_from_constraints( - constraints or [], tokens_per_block, swa_scratch_reuse + constraints_for_min_slots, tokens_per_block, swa_scratch_reuse ) - # Compute init_ratio from typical_batch, constraints, or fallback. - if typical_batch is not None: + # Compute init_ratio from explicit config, typical_batch, constraints, or fallback. + if initial_pool_ratio is not None: + if len(initial_pool_ratio) != self.num_pool_groups: + raise ValueError( + f"initial_pool_ratio length must match number of pool groups " + f"({self.num_pool_groups}), got {len(initial_pool_ratio)}" + ) + if any(r <= 0 for r in initial_pool_ratio): + raise ValueError("initial_pool_ratio values must be positive") + if not math.isclose(sum(initial_pool_ratio), 1.0, rel_tol=0.0, abs_tol=1e-6): + raise ValueError("initial_pool_ratio values must sum to 1.0") + init_ratio = cast(TypedIndexList[PoolGroupIndex, float], list(initial_pool_ratio)) + elif typical_batch is not None: init_ratio = self.ratio_from_batch( typical_batch, tokens_per_block, swa_scratch_reuse, gpu_granularity ) @@ -291,14 +307,16 @@ def new_gpu_slots( self, num_slots: TypedIndexList[LifeCycleId, int], migration_recorder: MigrationRecorder | None = None, + drop_recorder: DropRecorder | None = None, ) -> TypedIndexList[LifeCycleId, list[Slot]]: - return self.new_slots(GPU_LEVEL, num_slots, migration_recorder) + return self.new_slots(GPU_LEVEL, num_slots, migration_recorder, drop_recorder) def new_slots( self, level: CacheLevel, num_slots: TypedIndexList[LifeCycleId, int], migration_recorder: MigrationRecorder | None = None, + drop_recorder: DropRecorder | None = None, ) -> TypedIndexList[LifeCycleId, list[Slot]]: lc2pg = self._life_cycle_grouping pg_num_slots = filled_list(0, self.num_pool_groups) @@ -309,7 +327,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, migration_recorder) + self.prepare_free_slots(level, pg_num_slots, migration_recorder, drop_recorder) assert all( pg_num_slots[pg] <= storage.get_num_free_slots(pg) for pg in typed_range(self.num_pool_groups) @@ -334,12 +352,13 @@ def new_slots_for_pool_group( pg_idx: PoolGroupIndex, num_slots: int, migration_recorder: MigrationRecorder | None = None, + drop_recorder: DropRecorder | 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, migration_recorder) + self.prepare_free_slots(level, num_slots_list, migration_recorder, drop_recorder) assert num_slots <= storage.get_num_free_slots(pg_idx) try: return storage.allocate_multiple(pg_idx, num_slots) @@ -391,15 +410,19 @@ def prepare_free_slots( level: CacheLevel, requirements: TypedIndexList[PoolGroupIndex, int], migration_recorder: MigrationRecorder | None = None, + drop_recorder: DropRecorder | 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, migration_recorder) + self._prepare_free_slots(goals, level, fallen_pages, migration_recorder, drop_recorder) def force_evict( - self, level: CacheLevel, min_num_pages: TypedIndexList[PoolGroupIndex, int] + self, + level: CacheLevel, + min_num_pages: TypedIndexList[PoolGroupIndex, int], + drop_recorder: DropRecorder | None = None, ) -> None: # If we break inside this function with debugpy, pages in `evicted` won't be # released even after the function returns. This is a debugpy artifact. @@ -408,11 +431,18 @@ def force_evict( assert all(p.status == PageStatus.DROPPABLE for pages in evicted for p in pages), ( "Corrupted eviction controller" ) + if drop_recorder is not None: + for pg_idx in typed_range(self.num_pool_groups): + if evicted[pg_idx]: + drop_recorder(evicted[pg_idx], level) return next_lvl = CacheLevel(level + 1) goals = filled_array2d(self.num_cache_levels, self.num_pool_groups, 0) self._prepare_free_slots( - goals, next_lvl, cast(TypedIndexList[PoolGroupIndex, list[Page]], evicted) + goals, + next_lvl, + cast(TypedIndexList[PoolGroupIndex, list[Page]], evicted), + drop_recorder=drop_recorder, ) def _prepare_free_slots( @@ -421,6 +451,7 @@ def _prepare_free_slots( lvl_id: CacheLevel, fallen_pages: TypedIndexList[PoolGroupIndex, list[Page]], migration_recorder: MigrationRecorder | None = None, + drop_recorder: DropRecorder | None = None, ) -> None: assert NDEBUG or goals.rows == self.num_cache_levels and goals.cols == self.num_pool_groups assert NDEBUG or all( @@ -461,6 +492,10 @@ def _prepare_free_slots( assert NDEBUG or all(p.status == PageStatus.DROPPABLE for p in evicted[pg_idx]) if not NDEBUG: dbg_rawrefs = [rawref.ref(p) for p in evicted[pg_idx]] + # Record the drop event before releasing — these pages are leaving the + # cache hierarchy entirely without being onboarded back to GPU. + if drop_recorder is not None and num_evicted > 0: + drop_recorder(evicted[pg_idx], lvl_id) evicted[pg_idx].clear() if not NDEBUG: assert all(p() is None for p in dbg_rawrefs) # pyright: ignore @@ -497,6 +532,7 @@ def _prepare_free_slots( CacheLevel(lvl_id + 1), fallen_pages, migration_recorder, + drop_recorder, ) assert all(len(f) == 0 for f in fallen_pages) # migrate pages diff --git a/tensorrt_llm/serve/openai_server.py b/tensorrt_llm/serve/openai_server.py index f66e4a788c0e..729ade7e7f41 100644 --- a/tensorrt_llm/serve/openai_server.py +++ b/tensorrt_llm/serve/openai_server.py @@ -183,6 +183,15 @@ def _normalize_image_output(image) -> list: class OpenAIServer(_VideoRoutesMixin): + @staticmethod + def _iteration_stats_buffer_maxlen( + iter_stats_max_iterations: Optional[int]) -> Optional[int]: + if iter_stats_max_iterations is None or iter_stats_max_iterations == 0: + return 1000 + if iter_stats_max_iterations < 0: + return None + return iter_stats_max_iterations + def __init__( self, generator: Union[LLM, MultimodalEncoder, VisualGen], @@ -288,8 +297,14 @@ async def lifespan(app: FastAPI): # engine stats queue; /metrics reads from a tee buffer # bounded by iter_stats_max_iterations to avoid racing # the loop for the queue (nvbug 6102381). - max_buf = getattr(self.generator.args, - "iter_stats_max_iterations", 1000) or 1000 + # One shared buffer is sufficient while this collector task + # is the only consumer of the engine iteration-stats queue. + # Other consumers can read it through get_iteration_stats(), + # which clears the buffer. Adding another queue consumer + # requires revisiting the buffering and clearing ownership. + max_buf = self._iteration_stats_buffer_maxlen( + getattr(self.generator.args, + "iter_stats_max_iterations", 1000)) self._iteration_stats_buffer = deque(maxlen=max_buf) self._iteration_stats_collector_task = asyncio.create_task( self._iteration_stats_collector_loop()) diff --git a/tensorrt_llm/serve/perf_metrics.py b/tensorrt_llm/serve/perf_metrics.py index a8279e6cedc4..e2a4ff7607bb 100644 --- a/tensorrt_llm/serve/perf_metrics.py +++ b/tensorrt_llm/serve/perf_metrics.py @@ -1,4 +1,4 @@ -# Copyright (c) 2025, NVIDIA CORPORATION. +# Copyright (c) 2025-2026, NVIDIA CORPORATION. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. @@ -150,6 +150,7 @@ def __init__(self, max_requests: int): self._request_meteics = deque(maxlen=max_requests) self._server_metrics = defaultdict(dict) self._lock = asyncio.Lock() + self._collect_lock = asyncio.Lock() self._clients = [] self._metrics = { definition.name: instance_metric(definition) @@ -182,58 +183,60 @@ async def add_per_request_metrics( ) async def get_perf_metrics(self) -> List[Dict[str, Any]]: - perf_metrics = {} - for client in self._clients: - metrics_dict = await client.collect_metrics() - perf_metrics.update(metrics_dict) + async with self._collect_lock: + perf_metrics = {} + for client in self._clients: + metrics_dict = await client.collect_metrics() + perf_metrics.update(metrics_dict) + + return_metrics = [] + async with self._lock: + for server, metrics_data in perf_metrics.items(): + server_metrics = self._server_metrics[server] + # avoid metrics map inflation by limiting the number of requests to add + available_req_num = min( + max(0, self._max_requests - len(server_metrics)), + len(metrics_data), + ) + req_metrics_map = { + req_metrics["ctx_request_id"]: req_metrics + for req_metrics in metrics_data[:available_req_num] + if "ctx_request_id" in req_metrics + } + server_metrics.update(req_metrics_map) - return_metrics = [] - async with self._lock: - for server, metrics_data in perf_metrics.items(): - server_metrics = self._server_metrics[server] - # avoid metrics map inflation by limiting the number of requests to add - available_req_num = min( - max(0, self._max_requests - len(server_metrics)), len(metrics_data) - ) - req_metrics_map = { - req_metrics["ctx_request_id"]: req_metrics - for req_metrics in metrics_data[:available_req_num] - if "ctx_request_id" in req_metrics - } - server_metrics.update(req_metrics_map) - - remain_keys = [] - for ( - ctx_server, - gen_server, - ctx_request_id, - server_arrival_time, - server_first_token_time, - ) in self._request_meteics: - gen_perf_metrics = self._server_metrics[gen_server].pop(ctx_request_id, None) - if gen_perf_metrics is None: - # generation not finished - remain_keys.append( - ( - ctx_server, - gen_server, - ctx_request_id, - server_arrival_time, - server_first_token_time, + remain_keys = [] + for ( + ctx_server, + gen_server, + ctx_request_id, + server_arrival_time, + server_first_token_time, + ) in self._request_meteics: + gen_perf_metrics = self._server_metrics[gen_server].pop(ctx_request_id, None) + if gen_perf_metrics is None: + # generation not finished + remain_keys.append( + ( + ctx_server, + gen_server, + ctx_request_id, + server_arrival_time, + server_first_token_time, + ) ) + continue + ctx_perf_metrics = self._server_metrics[ctx_server].pop(ctx_request_id, None) + # TODO: strip the keys for less repeating and use table style response + return_metrics.append( + { + "ctx_server": ctx_server, + "gen_server": gen_server, + "disagg_server_arrival_time": server_arrival_time, + "disagg_server_first_token_time": server_first_token_time, + "ctx_perf_metrics": ctx_perf_metrics, + "gen_perf_metrics": gen_perf_metrics, + } ) - continue - ctx_perf_metrics = self._server_metrics[ctx_server].pop(ctx_request_id, None) - # TODO: strip the keys for less repeating and use table style response - return_metrics.append( - { - "ctx_server": ctx_server, - "gen_server": gen_server, - "disagg_server_arrival_time": server_arrival_time, - "disagg_server_first_token_time": server_first_token_time, - "ctx_perf_metrics": ctx_perf_metrics, - "gen_perf_metrics": gen_perf_metrics, - } - ) - self._request_meteics = deque(remain_keys, maxlen=self._max_requests) - return return_metrics + self._request_meteics = deque(remain_keys, maxlen=self._max_requests) + return return_metrics diff --git a/tests/integration/defs/accuracy/test_llm_api_pytorch.py b/tests/integration/defs/accuracy/test_llm_api_pytorch.py index 641ed8c11ca4..a20718b0cd8e 100644 --- a/tests/integration/defs/accuracy/test_llm_api_pytorch.py +++ b/tests/integration/defs/accuracy/test_llm_api_pytorch.py @@ -3449,6 +3449,9 @@ def _make_deepseekv4_eplb_config(model_path, layer_updates_per_iter, ep_size=8): layer_updates_per_iter=0) +DEEPSEEKV4_TEST_MAX_BATCH_SIZE = 128 + + def _run_deepseekv4_eplb(model_name, model_path, moe_backend, @@ -3470,6 +3473,7 @@ def _run_deepseekv4_eplb(model_name, moe_expert_parallel_size=tensor_parallel_size, kv_cache_config=kv_cache_config, enable_attention_dp=True, + max_batch_size=DEEPSEEKV4_TEST_MAX_BATCH_SIZE, max_seq_len=4096, **pytorch_config, speculative_config=mtp_config) as llm: @@ -3498,7 +3502,7 @@ def test_auto_dtype(self): moe_expert_parallel_size=4, moe_config=MoeConfig(backend="TRTLLM"), enable_attention_dp=True, - max_batch_size=4, + max_batch_size=DEEPSEEKV4_TEST_MAX_BATCH_SIZE, max_seq_len=4096, max_num_tokens=4096, kv_cache_config=kv_cache_config) as llm: @@ -3570,6 +3574,7 @@ def test_auto_dtype(self, moe_backend): moe_expert_parallel_size=4, moe_config=MoeConfig(backend=moe_backend), enable_attention_dp=True, + max_batch_size=DEEPSEEKV4_TEST_MAX_BATCH_SIZE, max_seq_len=4096, kv_cache_config=kv_cache_config) as llm: task = MMLU(self.MODEL_NAME) @@ -3586,11 +3591,12 @@ def test_fp8_chunked_prefill(self): tensor_parallel_size=4, moe_expert_parallel_size=4, moe_config=MoeConfig(backend="WIDEEP"), - cuda_graph_config=CudaGraphConfig(max_batch_size=16, - enable_padding=True), + cuda_graph_config=CudaGraphConfig( + max_batch_size=DEEPSEEKV4_TEST_MAX_BATCH_SIZE, + enable_padding=True), enable_attention_dp=True, enable_chunked_prefill=True, - max_batch_size=16, + max_batch_size=DEEPSEEKV4_TEST_MAX_BATCH_SIZE, max_num_tokens=128, max_seq_len=4096, kv_cache_config=kv_cache_config) as llm: diff --git a/tests/unittest/_torch/attention/sparse/deepseek_v4/test_compressor_module.py b/tests/unittest/_torch/attention/sparse/deepseek_v4/test_compressor_module.py index b9e47a021eb7..2ed72ee02cb5 100644 --- a/tests/unittest/_torch/attention/sparse/deepseek_v4/test_compressor_module.py +++ b/tests/unittest/_torch/attention/sparse/deepseek_v4/test_compressor_module.py @@ -36,7 +36,9 @@ KVCacheDtype, ) from tensorrt_llm._torch.attention_backend.sparse.deepseek_v4.deepseek_v4 import ( + DEEPSEEK_V4_SLIDING_ATTENTION, DeepseekV4AttentionType, + DeepseekV4Indexer, ) from tensorrt_llm._torch.modules.rotary_embedding import RopeParams from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest, LlmRequestState @@ -88,7 +90,9 @@ def __init__( num_ctx_tokens: int, num_tokens: int, kv_cache_manager: DeepseekV4CacheManager, - block_tables: dict, + sliding_block_tables: torch.Tensor, + compress_block_tables: Dict[int, torch.Tensor], + indexer_k_cache_block_offsets: torch.Tensor, cu_seq_lens: dict, cu_new_comp_kv: dict, compressed_position_ids: dict, @@ -106,7 +110,9 @@ def __init__( self.num_ctx_tokens = num_ctx_tokens self.num_tokens = num_tokens self.kv_cache_manager = kv_cache_manager - self.block_tables = block_tables + self.sliding_block_tables = sliding_block_tables + self.compress_block_tables = compress_block_tables + self.indexer_k_cache_block_offsets = indexer_k_cache_block_offsets self.cu_seq_lens_cuda = cu_seq_lens self.cu_new_comp_kv_cuda = cu_new_comp_kv self.compressed_position_ids_cuda = compressed_position_ids @@ -794,7 +800,7 @@ def _create_deepseek_v4_cache_manager(self, compress_ratio: int) -> DeepseekV4Ca dtype=cache_dtype, compressor_dtype=DataType.FLOAT, # State caches always use FP32 vocab_size=self.VOCAB_SIZE, - max_num_tokens=MAX_SEQ + MAX_BATCH, + max_num_tokens=MAX_SEQ * MAX_BATCH, sparse_attn_config=sparse_attn_config, ) @@ -1130,21 +1136,21 @@ def normalize_is_prefill( # Determine attention types based on is_indexer if self.is_indexer: compress_type = DeepseekV4AttentionType.INDEXER_COMPRESS - state_type = DeepseekV4AttentionType.INDEXER_COMPRESSOR_STATE + kv_type = DeepseekV4AttentionType.INDEXER_COMPRESSOR_KV score_type = DeepseekV4AttentionType.INDEXER_COMPRESSOR_SCORE else: compress_type = DeepseekV4AttentionType.COMPRESS - state_type = DeepseekV4AttentionType.COMPRESSOR_STATE + kv_type = DeepseekV4AttentionType.COMPRESSOR_KV score_type = DeepseekV4AttentionType.COMPRESSOR_SCORE - # Build block_tables dict keyed by DeepseekV4AttentionType using cache manager + # Build sliding block tables using cache manager indices. block_table_compress_list = [] block_table_kv_state_list = [] block_table_score_state_list = [] for b, req in enumerate(requests): block_table_compress_list.append(self._get_block_table_for_request(req, compress_type)) - block_table_kv_state_list.append(self._get_block_table_for_request(req, state_type)) + block_table_kv_state_list.append(self._get_block_table_for_request(req, kv_type)) block_table_score_state_list.append(self._get_block_table_for_request(req, score_type)) # Pad and stack block tables to handle variable-length block indices @@ -1170,11 +1176,45 @@ def normalize_is_prefill( # Update block_offsets for test compatibility self.block_offsets = block_table_compress - block_tables = { - (ratio, compress_type): block_table_compress, - (ratio, state_type): block_table_kv_state, - (ratio, score_type): block_table_score_state, + max_blocks = max(max_blocks_compress, max_blocks_state) + sliding_block_tables = torch.zeros( + 1, + len(DEEPSEEK_V4_SLIDING_ATTENTION), + bsz, + max_blocks, + dtype=torch.int32, + device=DEVICE, + ) + compress_block_tables = { + ratio: torch.zeros( + bsz, + max_blocks, + dtype=torch.int32, + device=DEVICE, + ) } + indexer_k_cache_block_offsets = torch.zeros( + bsz, + max_blocks, + dtype=torch.int32, + device=DEVICE, + ) + if self.is_indexer: + indexer_k_cache_block_offsets[:, :max_blocks_compress] = block_table_compress + else: + compress_block_tables[ratio][:, :max_blocks_compress] = block_table_compress + sliding_block_tables[ + 0, + kv_type.value, + :, + :max_blocks_state, + ] = block_table_kv_state + sliding_block_tables[ + 0, + score_type.value, + :, + :max_blocks_state, + ] = block_table_score_state # Both prefill and decode kernels use absolute token positions for the # state cache, so pass the absolute kv_lens directly. @@ -1206,7 +1246,9 @@ def normalize_is_prefill( num_ctx_tokens=num_ctx_tokens, num_tokens=num_ctx_tokens + num_gen_tokens, kv_cache_manager=self.cache_manager, - block_tables=block_tables, + sliding_block_tables=sliding_block_tables, + compress_block_tables=compress_block_tables, + indexer_k_cache_block_offsets=indexer_k_cache_block_offsets, cu_seq_lens=cu_seq_lens, cu_new_comp_kv=cu_new_comp_kv_dict, compressed_position_ids=compressed_position_ids_dict, @@ -1237,6 +1279,7 @@ def normalize_is_prefill( req.context_current_position = token_count # Call add_new_token for BOTH prefill and generation requests. req.add_new_token(token_count, 0) + self.cache_manager.update_context_resources(scheduled_batch) self.cache_manager.update_resources(scheduled_batch) # Compressor.forward() returns (kv_comp, scale) tuple. @@ -1596,6 +1639,7 @@ class _FakeCompressorCacheManager: def __init__(self, head_dim: int, tokens_per_block: int = 4): self.tokens_per_block = tokens_per_block self.compressed_block_sizes = {0: tokens_per_block} + self.layer_offsets = {0: 0} self._buffer = torch.empty(1, tokens_per_block * head_dim, device=DEVICE, dtype=DTYPE) def get_buffers(self, layer_idx, attn_type): @@ -1639,15 +1683,25 @@ def _create_small_compressor(kv_cache_dtype: str, is_indexer: bool) -> Compresso def _create_minimal_metadata(compressor: Compressor, total_compressed_tokens: int = 1): ratio = compressor.compress_ratio bsz = 1 - block_table = torch.zeros(bsz, 1, device=DEVICE, dtype=torch.int32) - block_tables = {(ratio, attn_type): block_table for attn_type in DeepseekV4AttentionType} + sliding_block_tables = torch.zeros( + 1, + len(DEEPSEEK_V4_SLIDING_ATTENTION), + bsz, + 1, + device=DEVICE, + dtype=torch.int32, + ) + compress_block_tables = {ratio: torch.zeros(bsz, 1, device=DEVICE, dtype=torch.int32)} + indexer_k_cache_block_offsets = torch.zeros(bsz, 1, device=DEVICE, dtype=torch.int32) metadata = DummyAttentionMetadata( num_contexts=1, num_generations=0, num_ctx_tokens=ratio, num_tokens=ratio, kv_cache_manager=_FakeCompressorCacheManager(compressor.head_dim), - block_tables=block_tables, + sliding_block_tables=sliding_block_tables, + compress_block_tables=compress_block_tables, + indexer_k_cache_block_offsets=indexer_k_cache_block_offsets, cu_seq_lens=torch.tensor([0, ratio], device=DEVICE, dtype=torch.int32), cu_new_comp_kv={ ratio: torch.tensor([0, total_compressed_tokens], device=DEVICE, dtype=torch.int32) @@ -1814,6 +1868,24 @@ def test_indexer_returns_fused_quant_outputs( assert torch.equal(scale_output, torch.full_like(scale_output, 0x7F)) +def test_deepseek_v4_indexer_keeps_shared_indexer_block_table(): + class _Metadata: + indexer_k_cache_block_offsets = torch.arange( + 4 * 5, dtype=torch.int32, device=DEVICE + ).reshape(4, 5) + + indexer = DeepseekV4Indexer.__new__(DeepseekV4Indexer) + indexer.layer_idx = 7 + metadata = _Metadata() + expected = metadata.indexer_k_cache_block_offsets + + indexer._update_k_cache(None, None, metadata) + + selected = metadata.indexer_k_cache_block_offsets + assert selected.data_ptr() == expected.data_ptr() + torch.testing.assert_close(selected, expected) + + # ============================================================================ # FP8 Blockwise Quantization Tests # ============================================================================ diff --git a/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_cache_manager.py b/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_cache_manager.py index 42a5b8b76a09..e4852b449cf0 100644 --- a/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_cache_manager.py +++ b/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_cache_manager.py @@ -22,15 +22,26 @@ from tensorrt_llm._torch.attention_backend.sparse.deepseek_v4 import DeepseekV4CacheManager from tensorrt_llm._torch.attention_backend.sparse.deepseek_v4.deepseek_v4 import ( + DEEPSEEK_V4_SLIDING_ATTENTION, DeepseekV4AttentionType, + compress_ratio_has_attention, ) -from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest +from tensorrt_llm._torch.disaggregation.native.peer import PeerRegistrar +from tensorrt_llm._torch.disaggregation.native.rank_info import RankInfo +from tensorrt_llm._torch.disaggregation.resource.kv_extractor import ( + KVRegionExtractorV1, + build_page_table_from_manager, +) +from tensorrt_llm._torch.disaggregation.resource.page import MapperKind +from tensorrt_llm._torch.pyexecutor._util import CacheCost +from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest, LlmRequestState from tensorrt_llm._torch.pyexecutor.scheduler import ScheduledRequests from tensorrt_llm._utils import binding_to_torch_dtype from tensorrt_llm.bindings import DataType, SamplingConfig from tensorrt_llm.bindings.internal.batch_manager import CacheType as CacheTypeCpp from tensorrt_llm.llmapi.llm_args import DeepSeekV4SparseAttentionConfig, KvCacheConfig from tensorrt_llm.mapping import Mapping +from tensorrt_llm.runtime.kv_cache_manager_v2 import GpuCacheTierConfig, PageIndexMode from tensorrt_llm.runtime.kv_cache_manager_v2._common import BAD_PAGE_INDEX _RequestCache = Dict[ @@ -45,6 +56,7 @@ class FakeModelConfig: index_head_dim=128, compress_ratios=[1, 4, 1, 128], indexer_k_dtype="fp8", + window_size=128, ) pretrained_config = SimpleNamespace( kv_lora_rank=512, @@ -58,10 +70,109 @@ def get_num_attention_layers(self) -> int: size_per_token = DeepseekV4CacheManager.get_cache_size_per_token( FakeModelConfig(), Mapping(world_size=1, rank=0, tp_size=1, pp_size=1), + tokens_per_block=128, is_disagg=True, ) - assert size_per_token > 0 + cost = CacheCost.from_raw(size_per_token) + assert cost.slope > 0 + assert cost.intercept == 0 + + +def test_quota_from_max_tokens_models_context_swa_scratch(): + manager = object.__new__(DeepseekV4CacheManager) + manager.pp_layers = [0, 1] + manager._compress_ratios = [4, 4] + manager.dtype = DataType.BF16 + manager.head_dim = 512 + 64 + manager.index_head_dim = 128 + manager._indexer_k_dtype = "fp8" + manager._swa_window_size = 128 + manager._max_draft_len = 0 + manager._max_num_tokens = 1024 + manager.tokens_per_block = 128 + manager.max_batch_size = 2 + + manager.enable_swa_scratch_reuse = True + small_quota = manager._get_quota_from_max_tokens(1) + large_quota = manager._get_quota_from_max_tokens(4096) + scratch_delta = large_quota - small_quota + + assert small_quota < large_quota + assert small_quota > manager._get_extra_quota_padding() + assert manager._get_max_tokens_from_quota(small_quota) == 1 + assert manager._get_max_tokens_from_quota(large_quota) == 4096 + + manager.enable_swa_scratch_reuse = False + no_scratch_small_quota = manager._get_quota_from_max_tokens(1) + no_scratch_large_quota = manager._get_quota_from_max_tokens(4096) + no_scratch_delta = no_scratch_large_quota - no_scratch_small_quota + + assert no_scratch_small_quota < no_scratch_large_quota + assert no_scratch_delta > scratch_delta + assert manager._get_max_tokens_from_quota(no_scratch_large_quota) == 4096 + + +def test_needed_resource_uses_context_swa_scratch_slope(): + manager = object.__new__(DeepseekV4CacheManager) + manager.pp_layers = [0, 1] + manager._compress_ratios = [4, 4] + manager.dtype = DataType.BF16 + manager.head_dim = 512 + 64 + manager.index_head_dim = 128 + manager._indexer_k_dtype = "fp8" + manager._swa_window_size = 128 + manager.tokens_per_block = 128 + manager.num_extra_kv_tokens = 0 + manager.enable_swa_scratch_reuse = True + + context_request = SimpleNamespace( + is_context_init_state=True, + is_generation_in_progress_state=False, + is_generation_to_complete_state=False, + is_disagg_generation_init_state=False, + state=LlmRequestState.CONTEXT_INIT, + prompt_len=10, + max_new_tokens=100, + ) + longer_context_request = SimpleNamespace( + is_context_init_state=True, + is_generation_in_progress_state=False, + is_generation_to_complete_state=False, + is_disagg_generation_init_state=False, + state=LlmRequestState.CONTEXT_INIT, + prompt_len=11, + max_new_tokens=100, + ) + generation_request = SimpleNamespace( + is_context_init_state=False, + is_generation_in_progress_state=True, + is_generation_to_complete_state=False, + is_disagg_generation_init_state=False, + state=LlmRequestState.GENERATION_IN_PROGRESS, + prompt_len=10, + max_new_tokens=20, + ) + longer_generation_request = SimpleNamespace( + is_context_init_state=False, + is_generation_in_progress_state=True, + is_generation_to_complete_state=False, + is_disagg_generation_init_state=False, + state=LlmRequestState.GENERATION_IN_PROGRESS, + prompt_len=10, + max_new_tokens=21, + ) + + non_sliding_attn_size_per_token = manager.get_cache_bytes_per_token() + context_bytes = manager.get_needed_resource_to_completion(context_request) + longer_context_bytes = manager.get_needed_resource_to_completion(longer_context_request) + generation_bytes = manager.get_needed_resource_to_completion(generation_request) + longer_generation_bytes = manager.get_needed_resource_to_completion(longer_generation_request) + + assert non_sliding_attn_size_per_token > 0 + assert context_bytes > context_request.prompt_len * non_sliding_attn_size_per_token + assert longer_context_bytes - context_bytes > non_sliding_attn_size_per_token + assert longer_generation_bytes - generation_bytes == non_sliding_attn_size_per_token def _view_fp8_as_uint8(buffer: torch.Tensor) -> torch.Tensor: @@ -71,6 +182,82 @@ def _view_fp8_as_uint8(buffer: torch.Tensor) -> torch.Tensor: return buffer +def _build_deepseek_v4_cache_config_for_test( + kv_cache_config: KvCacheConfig, + *, + max_batch_size: int = 4, + max_seq_len: int = 1024, + max_num_tokens: int | None = 2048, + max_draft_len: int = 0, +): + cache_manager = object.__new__(DeepseekV4CacheManager) + cache_manager.pp_layers = [0, 1, 2] + cache_manager._compress_ratios = [1, 4, 128] + cache_manager._swa_window_size = 128 + cache_manager._max_draft_len = max_draft_len + cache_manager._max_num_tokens = max_num_tokens + cache_manager.compressed_block_sizes = [128, 32, 1] + cache_manager.index_head_dim = 128 + cache_manager.head_dim = 512 + cache_manager.tokens_per_block = 128 + cache_manager.dtype = DataType.BF16 + cache_manager._indexer_k_dtype = "fp8" + cache_manager.max_batch_size = max_batch_size + cache_manager.max_seq_len = max_seq_len + cache_manager.enable_stats = False + cache_manager.enable_swa_scratch_reuse = False + cache_manager.num_extra_kv_tokens = 0 + + return cache_manager._build_cache_config( + kv_cache_config, + tokens_per_block=128, + vocab_size=129280, + cache_tiers=[GpuCacheTierConfig(quota=1 << 30)], + ) + + +def test_deepseek_v4_pool_ratio_overrides_typical_step_and_constraints(): + config = _build_deepseek_v4_cache_config_for_test( + KvCacheConfig(pool_ratio=[0.2, 0.3, 0.5], avg_seq_len=256) + ) + + assert config.initial_pool_ratio == [0.2, 0.3, 0.5] + assert config.typical_step is None + assert config.constraints == [] + + +def test_deepseek_v4_avg_seq_len_updates_typical_step(): + config = _build_deepseek_v4_cache_config_for_test( + KvCacheConfig(avg_seq_len=256), + max_batch_size=3, + max_seq_len=1024, + max_num_tokens=2048, + max_draft_len=2, + ) + + assert config.initial_pool_ratio is None + assert config.typical_step is not None + assert config.typical_step.kv_caches[0].capacity == 2048 + assert config.typical_step.kv_caches[0].history_length == 0 + assert [kv.capacity for kv in config.typical_step.kv_caches[1:]] == [256, 256] + assert [kv.history_length for kv in config.typical_step.kv_caches[1:]] == [253, 253] + assert config.constraints[0].kv_caches[0].capacity == 1024 + assert config.constraints[0].kv_caches[0].history_length == 1023 + + +def test_deepseek_v4_avg_seq_len_must_not_exceed_max_seq_len(): + with pytest.raises(ValueError, match="avg_seq_len"): + _build_deepseek_v4_cache_config_for_test( + KvCacheConfig(avg_seq_len=2048), + max_seq_len=1024, + ) + + +@pytest.fixture(params=[False, True], ids=["scratch_reuse_disabled", "scratch_reuse_enabled"]) +def scratch_reuse_enabled(request) -> bool: + return request.param + + @skip_pre_blackwell @pytest.mark.skip_less_device_memory(80000) class TestDeepseekV4CacheManager: @@ -90,6 +277,28 @@ class TestDeepseekV4CacheManager: # cache manager specific param tokens_per_block = 128 + @staticmethod + def _attention_op_block_offsets_shape( + cache_manager: DeepseekV4CacheManager, num_seqs: int + ) -> tuple[int, int, int, int]: + return ( + cache_manager.num_attention_op_pools, + num_seqs, + 2, + cache_manager.max_blocks_per_seq, + ) + + @staticmethod + def _sliding_block_tables_shape( + cache_manager: DeepseekV4CacheManager, num_seqs: int + ) -> tuple[int, int, int, int]: + return ( + cache_manager.num_local_layers, + len(DEEPSEEK_V4_SLIDING_ATTENTION), + num_seqs, + cache_manager.max_blocks_per_seq, + ) + def _is_compress_layer(self, compress_ratio: int) -> bool: """Check if a layer uses compression based on its compress ratio. @@ -137,9 +346,9 @@ def _get_window_size(self, compress_ratio: int, attn_type: DeepseekV4AttentionTy if attn_type == DeepseekV4AttentionType.SWA: return self.window_size elif attn_type in [ - DeepseekV4AttentionType.COMPRESSOR_STATE, + DeepseekV4AttentionType.COMPRESSOR_KV, DeepseekV4AttentionType.COMPRESSOR_SCORE, - DeepseekV4AttentionType.INDEXER_COMPRESSOR_STATE, + DeepseekV4AttentionType.INDEXER_COMPRESSOR_KV, DeepseekV4AttentionType.INDEXER_COMPRESSOR_SCORE, ]: return state_factor * compress_ratio @@ -158,14 +367,24 @@ def _create_deepseek_v4_cache_manager( dtype: DataType, compressor_dtype: DataType, max_input_len: Optional[int] = None, + is_draft: bool = False, + tp_size: int = 1, + enable_attention_dp: bool = False, + spec_config: object | None = None, + indexer_k_dtype: str | None = None, + enable_swa_scratch_reuse: bool = True, ) -> Tuple[DeepseekV4CacheManager, DeepSeekV4SparseAttentionConfig]: """Helper to create a DeepseekV4CacheManager for testing.""" # Create sparse attention config + config_kwargs = {} + if indexer_k_dtype is not None: + config_kwargs["indexer_k_dtype"] = indexer_k_dtype sparse_attn_config = DeepSeekV4SparseAttentionConfig( index_head_dim=self.index_head_dim, window_size=self.window_size, compress_ratios=compress_ratios, + **config_kwargs, ) # Create KV cache config @@ -175,10 +394,17 @@ def _create_deepseek_v4_cache_manager( enable_block_reuse=False, max_tokens=max_seq_len * max_batch_size, event_buffer_max_size=0, + enable_swa_scratch_reuse=enable_swa_scratch_reuse, ) # Create mapping (single GPU, no parallelism) - mapping = Mapping(world_size=1, rank=0, tp_size=1, pp_size=1) + mapping = Mapping( + world_size=tp_size, + rank=0, + tp_size=tp_size, + pp_size=1, + enable_attention_dp=enable_attention_dp, + ) # Create cache manager cache_manager = DeepseekV4CacheManager( @@ -197,6 +423,8 @@ def _create_deepseek_v4_cache_manager( vocab_size=self.vocab_size, max_num_tokens=max_batch_size * (max_input_len + 1), sparse_attn_config=sparse_attn_config, + is_draft=is_draft, + spec_config=spec_config, ) return cache_manager, sparse_attn_config @@ -271,7 +499,7 @@ def _create_random_cache( self._rand_tensor((seq_len // ratio, head_dim), dtype, device), None, ) - cache[layer, DeepseekV4AttentionType.COMPRESSOR_STATE] = ( + cache[layer, DeepseekV4AttentionType.COMPRESSOR_KV] = ( self._rand_tensor((seq_len, compressor_dim), compressor_dtype, device), None, ) @@ -297,7 +525,7 @@ def _create_random_cache( ) indexer_compressor_dim = 2 * indexer_dim if is_overlap else indexer_dim - cache[layer, DeepseekV4AttentionType.INDEXER_COMPRESSOR_STATE] = ( + cache[layer, DeepseekV4AttentionType.INDEXER_COMPRESSOR_KV] = ( self._rand_tensor((seq_len, indexer_compressor_dim), compressor_dtype, device), None, ) @@ -433,6 +661,123 @@ def _read_paged_cache( values = values[-window_size:] return values + def _get_page_indices( + self, + req: LlmRequest, + cache_manager: DeepseekV4CacheManager, + layer_idx: int, + attn_type: DeepseekV4AttentionType, + num_contexts: int = 1, + ) -> torch.Tensor: + if attn_type == DeepseekV4AttentionType.COMPRESS: + compress_block_tables = torch.empty( + 1, + cache_manager.max_blocks_per_seq, + dtype=torch.int32, + device="cpu", + ) + cache_manager.copy_batch_compress_block_tables( + compress_block_tables, + [req.py_request_id], + compress_ratio=cache_manager._compress_ratios[layer_idx], + beam_width=1, + num_contexts=num_contexts, + num_seqs=1, + ) + return compress_block_tables[0] + + if attn_type == DeepseekV4AttentionType.INDEXER_COMPRESS: + host_block_table = torch.empty( + 1, + cache_manager.max_blocks_per_seq, + dtype=torch.int32, + device="cpu", + ) + cache_manager.copy_batch_indexer_compress_block_tables( + host_block_table, + [req.py_request_id], + beam_width=1, + num_contexts=num_contexts, + num_seqs=1, + ) + return host_block_table[0] + + sliding_block_tables = torch.empty( + self._sliding_block_tables_shape(cache_manager, 1), + dtype=torch.int32, + device="cuda", + ) + with torch.cuda.stream(cache_manager._stream): + cache_manager.compute_sliding_block_tables( + [req.py_request_id], + num_contexts=num_contexts, + ) + cache_manager.copy_batch_sliding_block_tables( + sliding_block_tables, + [req.py_request_id], + num_contexts=num_contexts, + num_seqs=1, + ) + cache_manager._stream.synchronize() + return sliding_block_tables.cpu()[ + cache_manager.layer_offsets[layer_idx], + attn_type.value, + 0, + ] + + def _prepare_mixed_copy_batch( + self, + cache_manager: DeepseekV4CacheManager, + prompt_len: int, + ) -> tuple[list[LlmRequest], int]: + requests = [self._create_request(request_id, prompt_len) for request_id in range(3)] + for req in requests: + assert cache_manager.prepare_context(req) + assert cache_manager.resize_context(req, req.context_chunk_size) + + gen_req = requests[-1] + scheduled_batch = ScheduledRequests() + scheduled_batch.context_requests_last_chunk = [gen_req] + gen_req.context_current_position = prompt_len + gen_req.add_new_token(prompt_len, 0) + cache_manager.update_context_resources(scheduled_batch) + cache_manager.update_resources(scheduled_batch) + assert cache_manager.try_allocate_generation(gen_req) + return requests, len(requests) - 1 + + def _reference_copy_batch_page_indices( + self, + cache_manager: DeepseekV4CacheManager, + request_ids: list[int], + num_contexts: int, + layer_idx: int, + attn_type: DeepseekV4AttentionType, + page_index_mode: PageIndexMode, + ) -> torch.Tensor: + layer_id = cache_manager._layer_attn_to_layer_id[layer_idx, attn_type] + pool_id = cache_manager.layer_to_pool_mapping_dict[layer_id] + converter = cache_manager.impl.get_page_index_converter(layer_id, attn_type.role) + copy_idx = cache_manager.index_mapper.get_copy_index(request_ids, num_contexts, 1) + + expected = torch.full( + (len(request_ids), cache_manager.max_blocks_per_seq), + BAD_PAGE_INDEX, + dtype=torch.int32, + device="cpu", + ) + for row, request_id in enumerate(request_ids): + base_indices = cache_manager.host_kv_cache_block_offsets[ + pool_id, + int(copy_idx[row]), + 0, + ].tolist() + scratch = None + if row < num_contexts and page_index_mode == PageIndexMode.PER_LAYER: + scratch = cache_manager.kv_cache_map[request_id].get_scratch_desc(pool_id) + converted = converter(base_indices, page_index_mode, scratch) + expected[row, : len(converted)] = torch.tensor(converted, dtype=torch.int32) + return expected + def _write_request_prefill( self, req: LlmRequest, @@ -450,14 +795,12 @@ def _write_request_prefill( """ compress_ratios = cache_manager._compress_ratios for (layer_idx, attn_type), (values, scales) in cache_values.items(): - page_indices = cache_manager.get_batch_attn_offset( - [req.py_request_id], - beam_width=1, - num_contexts=1, - num_seqs=1, - attn_type=attn_type, - compress_ratio=compress_ratios[layer_idx], - ).squeeze(0) + page_indices = self._get_page_indices( + req, + cache_manager, + layer_idx, + attn_type, + ) if attn_type in [ DeepseekV4AttentionType.COMPRESS, @@ -505,14 +848,12 @@ def _write_request_decode( """ compress_ratios = cache_manager._compress_ratios for (layer_idx, attn_type), (values, scales) in cache_values.items(): - block_indices = cache_manager.get_batch_attn_offset( - [req.py_request_id], - beam_width=1, - num_contexts=1, - num_seqs=1, - attn_type=attn_type, - compress_ratio=compress_ratios[layer_idx], - ).squeeze(0) + block_indices = self._get_page_indices( + req, + cache_manager, + layer_idx, + attn_type, + ) compressed_token_idx = token_idx if attn_type in [ @@ -572,7 +913,7 @@ def _read_request( attn_types.extend( [ DeepseekV4AttentionType.COMPRESS, - DeepseekV4AttentionType.COMPRESSOR_STATE, + DeepseekV4AttentionType.COMPRESSOR_KV, DeepseekV4AttentionType.COMPRESSOR_SCORE, ] ) @@ -580,21 +921,19 @@ def _read_request( attn_types.extend( [ DeepseekV4AttentionType.INDEXER_COMPRESS, - DeepseekV4AttentionType.INDEXER_COMPRESSOR_STATE, + DeepseekV4AttentionType.INDEXER_COMPRESSOR_KV, DeepseekV4AttentionType.INDEXER_COMPRESSOR_SCORE, ] ) # read cache values for each attention type for attn_type in attn_types: - page_indices = cache_manager.get_batch_attn_offset( - [req.py_request_id], - beam_width=1, - num_contexts=1, - num_seqs=1, - attn_type=attn_type, - compress_ratio=ratio, - ).squeeze(0) + page_indices = self._get_page_indices( + req, + cache_manager, + layer, + attn_type, + ) if attn_type in [ DeepseekV4AttentionType.COMPRESS, DeepseekV4AttentionType.INDEXER_COMPRESS, @@ -732,6 +1071,7 @@ def test_write_read_cache( num_generation_steps: int, dtype: DataType, compressor_dtype: DataType, + scratch_reuse_enabled: bool, ): max_batch_size = len(prompt_lens) max_seq_len = max(prompt_lens) + num_generation_steps + 1 @@ -745,7 +1085,9 @@ def test_write_read_cache( dtype=dtype, compressor_dtype=compressor_dtype, max_input_len=max_input_len, + enable_swa_scratch_reuse=scratch_reuse_enabled, ) + assert cache_manager.enable_swa_scratch_reuse == scratch_reuse_enabled # Create requests and their cache values requests = list[LlmRequest]() @@ -784,6 +1126,7 @@ def test_write_read_cache( for req in requests: req.context_current_position = prompt_lens[req.py_request_id] req.add_new_token(prompt_lens[req.py_request_id], 0) + cache_manager.update_context_resources(scheduled_batch) cache_manager.update_resources(scheduled_batch) # Read context from cache and verify @@ -847,11 +1190,10 @@ def test_write_read_cache( @pytest.mark.parametrize( "dtype,compressor_dtype", [(DataType.BF16, DataType.FLOAT), (DataType.FP8, DataType.FLOAT)] ) - def test_kv_cache_pool_mapping( + def test_pool_mapping_tensors( self, compress_ratios: List[int], dtype: DataType, compressor_dtype: DataType ): # Create cache manager and sparse attention config - num_layers = len(compress_ratios) cache_manager, _ = self._create_deepseek_v4_cache_manager( tokens_per_block=self.tokens_per_block, max_batch_size=4, @@ -862,19 +1204,698 @@ def test_kv_cache_pool_mapping( ) try: - kv_cache_pool_mapping = cache_manager.kv_cache_pool_mapping - assert kv_cache_pool_mapping.shape == (num_layers, 2) + expected_mapping = torch.stack( + [ + torch.arange(cache_manager.num_local_layers, dtype=torch.int32), + torch.zeros(cache_manager.num_local_layers, dtype=torch.int32), + ], + dim=1, + ) + expected_pointers = torch.tensor( + [ + [ + cache_manager.impl.get_mem_pool_base_address( + cache_manager._layer_attn_to_layer_id[ + pp_layer, DeepseekV4AttentionType.SWA + ], + DeepseekV4AttentionType.SWA.role, + PageIndexMode.PER_LAYER, + ), + 0, + ] + for pp_layer in cache_manager.pp_layers + ], + dtype=torch.int64, + ) + + assert cache_manager.kv_cache_pool_mapping.shape == expected_mapping.shape + assert cache_manager.kv_cache_pool_mapping.dtype == expected_mapping.dtype + assert torch.equal(cache_manager.kv_cache_pool_mapping, expected_mapping) + assert cache_manager.kv_cache_pool_pointers.shape == expected_pointers.shape + assert cache_manager.kv_cache_pool_pointers.dtype == expected_pointers.dtype + assert torch.equal(cache_manager.kv_cache_pool_pointers, expected_pointers) + finally: + cache_manager.shutdown() + + def test_disagg_page_table_accepts_deepseek_v4_roles(self): + cache_manager, _ = self._create_deepseek_v4_cache_manager( + tokens_per_block=self.tokens_per_block, + max_batch_size=1, + max_seq_len=512, + compress_ratios=[1, 4, 128], + dtype=DataType.BF16, + compressor_dtype=DataType.FLOAT, + ) + + try: + page_table = build_page_table_from_manager(cache_manager) + saw_pool_view = False + for layer_group in page_table.layer_groups: + for pool_view in layer_group.pool_views: + saw_pool_view = True + # All DSv4 pools publish per-layer buffer_entries (INDEXED); + # the FLAT layout is reserved for the DSA (DeepSeek v3.2) + # indexer K cache pool. + assert pool_view.mapper_kind == MapperKind.INDEXED + # pool_role carries the manager-native role-name strings, + # which are all DSv4 attention-type names like + # "deepseek_v4_swa", "deepseek_v4_compress", etc. + assert pool_view.pool_role + for role_name in pool_view.pool_role: + assert role_name.startswith("deepseek_v4_"), role_name + assert saw_pool_view + finally: + cache_manager.shutdown() + + def test_disagg_pool_mapping_prefers_exact_pool_view_layer_set(self): + ctx_cache_manager, _ = self._create_deepseek_v4_cache_manager( + tokens_per_block=self.tokens_per_block, + max_batch_size=1, + max_seq_len=512, + compress_ratios=[1, 4, 128], + dtype=DataType.FP8, + compressor_dtype=DataType.FLOAT, + tp_size=2, + ) + gen_cache_manager, _ = self._create_deepseek_v4_cache_manager( + tokens_per_block=self.tokens_per_block, + max_batch_size=1, + max_seq_len=512, + compress_ratios=[1, 4, 128], + dtype=DataType.FP8, + compressor_dtype=DataType.FLOAT, + tp_size=2, + enable_attention_dp=True, + ) + + try: + ctx_rank_info = RankInfo.from_kv_cache_manager("ctx", ctx_cache_manager, 0) + gen_rank_info = RankInfo.from_kv_cache_manager("gen", gen_cache_manager, 0) + registrar = PeerRegistrar(gen_rank_info, KVRegionExtractorV1(gen_cache_manager)) + registrar.register("ctx", 0, ctx_rank_info) + + pool_mapping = registrar.get_pool_mapping(ctx_rank_info) + assert pool_mapping + for self_pool_key, peer_pool_key in pool_mapping.items(): + registrar.get_kv_map(ctx_rank_info, self_pool_key, peer_pool_key) + finally: + ctx_cache_manager.shutdown() + gen_cache_manager.shutdown() + + def test_dsv4_block_tables(self, scratch_reuse_enabled: bool): + prompt_len = self.tokens_per_block * 2 + cache_manager, _ = self._create_deepseek_v4_cache_manager( + tokens_per_block=self.tokens_per_block, + max_batch_size=1, + max_seq_len=prompt_len, + compress_ratios=[4, 4], + dtype=DataType.BF16, + compressor_dtype=DataType.FLOAT, + enable_swa_scratch_reuse=scratch_reuse_enabled, + ) + assert cache_manager.enable_swa_scratch_reuse == scratch_reuse_enabled + + req = self._create_request(0, prompt_len) + allocated = False + try: + assert cache_manager.prepare_context(req) + assert cache_manager.resize_context(req, req.context_chunk_size) + allocated = True + + attention_types = { + DeepseekV4AttentionType.COMPRESS, + DeepseekV4AttentionType.INDEXER_COMPRESS, + DeepseekV4AttentionType.COMPRESSOR_KV, + } + sliding_block_tables_cuda = torch.empty( + self._sliding_block_tables_shape(cache_manager, 1), + dtype=torch.int32, + device="cuda", + ) + compress_block_table = torch.empty( + 1, + cache_manager.max_blocks_per_seq, + dtype=torch.int32, + device="cpu", + ) + host_indexer_compress_block_table = torch.empty( + 1, + cache_manager.max_blocks_per_seq, + dtype=torch.int32, + device="cpu", + ) + with torch.cuda.stream(cache_manager._stream): + cache_manager.compute_sliding_block_tables( + [req.py_request_id], + num_contexts=1, + ) + cache_manager.copy_batch_sliding_block_tables( + sliding_block_tables_cuda, + [req.py_request_id], + num_contexts=1, + num_seqs=1, + ) + cache_manager.copy_batch_compress_block_tables( + compress_block_table, + [req.py_request_id], + compress_ratio=4, + beam_width=1, + num_contexts=1, + num_seqs=1, + ) + cache_manager.copy_batch_indexer_compress_block_tables( + host_indexer_compress_block_table, + [req.py_request_id], + beam_width=1, + num_contexts=1, + num_seqs=1, + ) + cache_manager._stream.synchronize() + sliding_block_tables = sliding_block_tables_cuda.cpu() + for attn_type in attention_types: + layer0_buffer = _view_fp8_as_uint8(cache_manager.get_buffers(0, attn_type)) + layer1_buffer = _view_fp8_as_uint8(cache_manager.get_buffers(1, attn_type)) + attn_len = prompt_len + if attn_type in [ + DeepseekV4AttentionType.COMPRESS, + DeepseekV4AttentionType.INDEXER_COMPRESS, + ]: + attn_len //= 4 + num_blocks = (attn_len + layer0_buffer.shape[1] - 1) // layer0_buffer.shape[1] + if attn_type == DeepseekV4AttentionType.COMPRESS: + layer0_offsets = compress_block_table[0, :num_blocks].tolist() + layer1_offsets = compress_block_table[0, :num_blocks].tolist() + elif attn_type == DeepseekV4AttentionType.INDEXER_COMPRESS: + layer0_offsets = host_indexer_compress_block_table[0, :num_blocks].tolist() + layer1_offsets = host_indexer_compress_block_table[0, :num_blocks].tolist() + else: + layer0_offsets = sliding_block_tables[ + cache_manager.layer_offsets[0], + attn_type.value, + 0, + :num_blocks, + ].tolist() + layer1_offsets = sliding_block_tables[ + cache_manager.layer_offsets[1], + attn_type.value, + 0, + :num_blocks, + ].tolist() + assert all(offset != BAD_PAGE_INDEX for offset in layer0_offsets) + assert all(offset != BAD_PAGE_INDEX for offset in layer1_offsets) + + layer0_values = torch.full( + (num_blocks, layer0_buffer.shape[1], layer0_buffer.shape[-1]), + 1, + dtype=layer0_buffer.dtype, + device=layer0_buffer.device, + ) + layer1_values = torch.full_like(layer0_values, 2) + + layer0_buffer[layer0_offsets] = layer0_values + layer1_buffer[layer1_offsets] = layer1_values + + torch.testing.assert_close(layer0_buffer[layer0_offsets], layer0_values) + torch.testing.assert_close(layer1_buffer[layer1_offsets], layer1_values) + finally: + if allocated: + cache_manager.free_resources(req) + cache_manager.shutdown() + + def test_copy_batch_block_offsets_matches_python_converter(self, scratch_reuse_enabled: bool): + prompt_len = self.tokens_per_block * 2 + 1 + cache_manager, _ = self._create_deepseek_v4_cache_manager( + tokens_per_block=self.tokens_per_block, + max_batch_size=3, + max_seq_len=1024, + compress_ratios=[1, 4, 128, 4], + dtype=DataType.BF16, + compressor_dtype=DataType.FLOAT, + enable_swa_scratch_reuse=scratch_reuse_enabled, + ) + assert cache_manager.enable_swa_scratch_reuse == scratch_reuse_enabled + + requests = [] + try: + requests, num_contexts = self._prepare_mixed_copy_batch(cache_manager, prompt_len) + request_ids = [req.py_request_id for req in requests] + actual = torch.empty( + self._attention_op_block_offsets_shape(cache_manager, len(request_ids)), + dtype=torch.int32, + device="cuda", + ) + sliding_block_tables = torch.empty( + self._sliding_block_tables_shape(cache_manager, len(request_ids)), + dtype=torch.int32, + device="cuda", + ) + + with torch.cuda.stream(cache_manager._stream): + cache_manager.compute_sliding_block_tables( + request_ids, + num_contexts=num_contexts, + ) + cache_manager.copy_batch_block_offsets( + actual, + request_ids, + beam_width=1, + num_contexts=num_contexts, + num_seqs=len(request_ids), + ) + cache_manager.copy_batch_sliding_block_tables( + sliding_block_tables, + request_ids, + num_contexts=num_contexts, + num_seqs=len(request_ids), + ) + + expected = torch.full( + ( + cache_manager.num_attention_op_pools, + len(request_ids), + cache_manager.max_blocks_per_seq, + ), + BAD_PAGE_INDEX, + dtype=torch.int32, + device="cpu", + ) + for local_layer_idx, pp_layer in enumerate(cache_manager.pp_layers): + ref = self._reference_copy_batch_page_indices( + cache_manager, + request_ids, + num_contexts, + pp_layer, + DeepseekV4AttentionType.SWA, + PageIndexMode.PER_LAYER, + ) + expected[local_layer_idx, : len(request_ids)] = ref + + # DSV4 AttentionOp only consumes the key table. + cache_manager._stream.synchronize() + actual_cpu = actual.cpu() + sliding_block_tables_cpu = sliding_block_tables.cpu() + torch.testing.assert_close( + actual_cpu[:, : len(request_ids), 0], + expected[:, : len(request_ids)], + ) + torch.testing.assert_close( + actual_cpu[:, : len(request_ids), 0], + sliding_block_tables_cpu[ + :, + DeepseekV4AttentionType.SWA.value, + : len(request_ids), + ], + ) + finally: + for req in requests: + cache_manager.free_resources(req) + cache_manager.shutdown() + + def test_copy_batch_sliding_block_tables_matches_python_converter( + self, scratch_reuse_enabled: bool + ): + prompt_len = self.tokens_per_block * 2 + 1 + cache_manager, _ = self._create_deepseek_v4_cache_manager( + tokens_per_block=self.tokens_per_block, + max_batch_size=3, + max_seq_len=1024, + compress_ratios=[1, 4, 128, 4], + dtype=DataType.BF16, + compressor_dtype=DataType.FLOAT, + enable_swa_scratch_reuse=scratch_reuse_enabled, + ) + assert cache_manager.enable_swa_scratch_reuse == scratch_reuse_enabled + + requests = [] + try: + requests, num_contexts = self._prepare_mixed_copy_batch(cache_manager, prompt_len) + request_ids = [req.py_request_id for req in requests] + actual = torch.empty( + self._sliding_block_tables_shape(cache_manager, len(request_ids)), + dtype=torch.int32, + device="cuda", + ) + + with torch.cuda.stream(cache_manager._stream): + cache_manager.compute_sliding_block_tables( + request_ids, + num_contexts=num_contexts, + ) + cache_manager.copy_batch_sliding_block_tables( + actual, + request_ids, + num_contexts=num_contexts, + num_seqs=len(request_ids), + ) + + expected = torch.full( + self._sliding_block_tables_shape(cache_manager, len(request_ids)), + BAD_PAGE_INDEX, + dtype=torch.int32, + device="cpu", + ) + for pp_layer in cache_manager.pp_layers: + compress_ratio = cache_manager._compress_ratios[pp_layer] + local_layer_idx = cache_manager.layer_offsets[pp_layer] + for attn_type in DEEPSEEK_V4_SLIDING_ATTENTION: + if not compress_ratio_has_attention(compress_ratio, attn_type): + continue + expected[local_layer_idx, attn_type.value, : len(request_ids)] = ( + self._reference_copy_batch_page_indices( + cache_manager, + request_ids, + num_contexts, + pp_layer, + attn_type, + PageIndexMode.PER_LAYER, + ) + ) + + cache_manager._stream.synchronize() + torch.testing.assert_close( + actual.cpu()[:, :, : len(request_ids)], + expected[:, :, : len(request_ids)], + ) + finally: + for req in requests: + cache_manager.free_resources(req) + cache_manager.shutdown() + + def test_copy_batch_compress_block_tables_matches_python_converter( + self, scratch_reuse_enabled: bool + ): + prompt_len = self.tokens_per_block * 2 + 1 + cache_manager, _ = self._create_deepseek_v4_cache_manager( + tokens_per_block=self.tokens_per_block, + max_batch_size=3, + max_seq_len=1024, + compress_ratios=[1, 4, 128, 4], + dtype=DataType.BF16, + compressor_dtype=DataType.FLOAT, + enable_swa_scratch_reuse=scratch_reuse_enabled, + ) + assert cache_manager.enable_swa_scratch_reuse == scratch_reuse_enabled + + requests = [] + try: + requests, num_contexts = self._prepare_mixed_copy_batch(cache_manager, prompt_len) + request_ids = [req.py_request_id for req in requests] + + for compress_ratio in (4, 128): + actual = torch.empty( + len(request_ids), + cache_manager.max_blocks_per_seq, + dtype=torch.int32, + device="cpu", + ) + + cache_manager.copy_batch_compress_block_tables( + actual, + request_ids, + compress_ratio=compress_ratio, + beam_width=1, + num_contexts=num_contexts, + num_seqs=len(request_ids), + ) + + pp_layer = next( + layer + for layer in cache_manager.pp_layers + if cache_manager._compress_ratios[layer] == compress_ratio + ) + expected = self._reference_copy_batch_page_indices( + cache_manager, + request_ids, + num_contexts, + pp_layer, + DeepseekV4AttentionType.COMPRESS, + PageIndexMode.SHARED, + ) + torch.testing.assert_close(actual, expected) + finally: + for req in requests: + cache_manager.free_resources(req) + cache_manager.shutdown() + + def test_copy_batch_indexer_compress_block_tables_matches_python_converter( + self, scratch_reuse_enabled: bool + ): + prompt_len = self.tokens_per_block * 2 + 1 + cache_manager, _ = self._create_deepseek_v4_cache_manager( + tokens_per_block=self.tokens_per_block, + max_batch_size=3, + max_seq_len=1024, + compress_ratios=[1, 4, 128, 4], + dtype=DataType.BF16, + compressor_dtype=DataType.FLOAT, + enable_swa_scratch_reuse=scratch_reuse_enabled, + ) + assert cache_manager.enable_swa_scratch_reuse == scratch_reuse_enabled + + requests = [] + try: + requests, num_contexts = self._prepare_mixed_copy_batch(cache_manager, prompt_len) + request_ids = [req.py_request_id for req in requests] + actual = torch.empty( + len(request_ids), + cache_manager.max_blocks_per_seq, + dtype=torch.int32, + device="cpu", + ) - assert torch.all(kv_cache_pool_mapping[:, 0] != -1), ( - "all layers should have swa attention pool" + cache_manager.copy_batch_indexer_compress_block_tables( + actual, + request_ids, + beam_width=1, + num_contexts=num_contexts, + num_seqs=len(request_ids), ) - assert torch.all(kv_cache_pool_mapping[:, 1] >= 0), ( - "buffer pointer offset should be non-negative" + + pp_layer = next( + layer + for layer in cache_manager.pp_layers + if compress_ratio_has_attention( + cache_manager._compress_ratios[layer], + DeepseekV4AttentionType.INDEXER_COMPRESS, + ) + ) + expected = self._reference_copy_batch_page_indices( + cache_manager, + request_ids, + num_contexts, + pp_layer, + DeepseekV4AttentionType.INDEXER_COMPRESS, + PageIndexMode.SHARED, ) - assert torch.all(kv_cache_pool_mapping[:, 0] == kv_cache_pool_mapping[0, 0]), ( - "all layers should have the same pool_id" + torch.testing.assert_close(actual, expected) + assert num_contexts == 2 + finally: + for req in requests: + cache_manager.free_resources(req) + cache_manager.shutdown() + + def test_swa_scratch_reuse_enabled_by_default_for_main_manager(self): + cache_manager, _ = self._create_deepseek_v4_cache_manager( + tokens_per_block=self.tokens_per_block, + max_batch_size=1, + max_seq_len=1024, + compress_ratios=[1], + dtype=DataType.BF16, + compressor_dtype=DataType.FLOAT, + ) + + try: + assert cache_manager.enable_swa_scratch_reuse + assert cache_manager.kv_cache_manager_py_config.enable_swa_scratch_reuse + assert cache_manager.kv_cache_manager_py_config.swa_scratch_reuse is not None + assert cache_manager.kv_cache_manager_py_config.swa_scratch_reuse.max_rewind_len == 0 + assert cache_manager.num_attention_op_pools == cache_manager.num_local_layers + finally: + cache_manager.shutdown() + + def test_swa_scratch_reuse_disabled_by_config_for_main_manager(self): + cache_manager, _ = self._create_deepseek_v4_cache_manager( + tokens_per_block=self.tokens_per_block, + max_batch_size=1, + max_seq_len=1024, + compress_ratios=[1], + dtype=DataType.BF16, + compressor_dtype=DataType.FLOAT, + enable_swa_scratch_reuse=False, + ) + + try: + assert not cache_manager.enable_swa_scratch_reuse + assert not cache_manager.kv_cache_manager_py_config.enable_swa_scratch_reuse + assert cache_manager.kv_cache_manager_py_config.swa_scratch_reuse is None + assert cache_manager.num_attention_op_pools == cache_manager.num_local_layers + finally: + cache_manager.shutdown() + + def test_swa_scratch_reuse_uses_extra_kv_tokens_for_rewind(self): + spec_config = SimpleNamespace( + max_draft_len=7, + max_total_draft_tokens=7, + spec_dec_mode=SimpleNamespace( + is_eagle3_one_model=lambda: False, + is_mtp_eagle_one_model=lambda: False, + is_mtp_one_model=lambda: False, + is_mtp_vanilla=lambda: False, + use_one_engine=lambda: True, + ), + ) + cache_manager, _ = self._create_deepseek_v4_cache_manager( + tokens_per_block=self.tokens_per_block, + max_batch_size=1, + max_seq_len=1024, + compress_ratios=[1], + dtype=DataType.BF16, + compressor_dtype=DataType.FLOAT, + spec_config=spec_config, + enable_swa_scratch_reuse=True, + ) + + try: + scratch_reuse = cache_manager.kv_cache_manager_py_config.swa_scratch_reuse + assert scratch_reuse is not None + assert scratch_reuse.max_rewind_len == spec_config.max_draft_len - 1 + finally: + cache_manager.shutdown() + + def test_draft_cache_manager_disables_swa_scratch_reuse(self): + cache_manager, _ = self._create_deepseek_v4_cache_manager( + tokens_per_block=self.tokens_per_block, + max_batch_size=1, + max_seq_len=1024, + compress_ratios=[1], + dtype=DataType.BF16, + compressor_dtype=DataType.FLOAT, + is_draft=True, + enable_swa_scratch_reuse=True, + ) + + try: + assert not cache_manager.enable_swa_scratch_reuse + assert not cache_manager.kv_cache_manager_py_config.enable_swa_scratch_reuse + assert cache_manager.kv_cache_manager_py_config.swa_scratch_reuse is None + assert cache_manager.num_attention_op_pools == cache_manager.num_local_layers + finally: + cache_manager.shutdown() + + def test_context_request_enable_scratch_reuse_until_generation(self): + cache_manager, _ = self._create_deepseek_v4_cache_manager( + tokens_per_block=self.tokens_per_block, + max_batch_size=1, + max_seq_len=1024, + compress_ratios=[1], + dtype=DataType.BF16, + compressor_dtype=DataType.FLOAT, + enable_swa_scratch_reuse=True, + ) + + req = self._create_request(request_id=0, prompt_len=self.tokens_per_block + 1) + allocated = False + try: + assert cache_manager.prepare_context(req) + assert cache_manager.resize_context(req, req.context_chunk_size) + allocated = True + + kv_cache = cache_manager.kv_cache_map[req.py_request_id] + assert kv_cache.enable_swa_scratch_reuse + + scheduled_batch = ScheduledRequests() + scheduled_batch.context_requests_last_chunk = [req] + req.context_current_position = req.prompt_len + req.add_new_token(req.prompt_len, 0) + cache_manager.update_context_resources(scheduled_batch) + assert not kv_cache.enable_swa_scratch_reuse + + assert cache_manager.try_allocate_generation(req) + assert not kv_cache.enable_swa_scratch_reuse + finally: + if allocated: + cache_manager.free_resources(req) + cache_manager.shutdown() + + def test_disagg_generation_init_disables_swa_scratch_reuse(self): + cache_manager, _ = self._create_deepseek_v4_cache_manager( + tokens_per_block=self.tokens_per_block, + max_batch_size=1, + max_seq_len=1024, + compress_ratios=[1], + dtype=DataType.BF16, + compressor_dtype=DataType.FLOAT, + enable_swa_scratch_reuse=True, + ) + + req = self._create_request(request_id=0, prompt_len=self.tokens_per_block + 1) + req.state = LlmRequestState.DISAGG_GENERATION_INIT + allocated = False + try: + assert cache_manager.prepare_disagg_gen_init(req) + allocated = True + + kv_cache = cache_manager.kv_cache_map[req.py_request_id] + assert not kv_cache.enable_swa_scratch_reuse + assert kv_cache.history_length == req.prompt_len + finally: + if allocated: + cache_manager.free_resources(req) + cache_manager.shutdown() + + def test_disagg_generation_init_rejects_context_apis(self): + cache_manager, _ = self._create_deepseek_v4_cache_manager( + tokens_per_block=self.tokens_per_block, + max_batch_size=1, + max_seq_len=1024, + compress_ratios=[1], + dtype=DataType.BF16, + compressor_dtype=DataType.FLOAT, + enable_swa_scratch_reuse=True, + ) + req = self._create_request(request_id=0, prompt_len=self.tokens_per_block + 1) + req.state = LlmRequestState.DISAGG_GENERATION_INIT + try: + with pytest.raises(AssertionError, match="prepare_disagg_gen_init"): + cache_manager.prepare_context(req) + with pytest.raises(AssertionError, match="prepare_disagg_gen_init"): + cache_manager.resize_context(req, req.context_remaining_length) + finally: + cache_manager.shutdown() + + def test_dummy_generation_requests_with_swa_scratch_reuse(self): + cache_manager, _ = self._create_deepseek_v4_cache_manager( + tokens_per_block=self.tokens_per_block, + max_batch_size=2, + max_seq_len=1024, + compress_ratios=[1], + dtype=DataType.BF16, + compressor_dtype=DataType.FLOAT, + enable_swa_scratch_reuse=True, + ) + + requests = [] + token_nums = [1, self.tokens_per_block * 4 + 1] + try: + requests = cache_manager.add_dummy_requests( + request_ids=[0, 1], + token_nums=token_nums, + is_gen=True, ) + assert requests is not None + assert len(requests) == len(token_nums) + + short_kv_cache = cache_manager.kv_cache_map[requests[0].py_request_id] + long_kv_cache = cache_manager.kv_cache_map[requests[1].py_request_id] + assert not short_kv_cache.enable_swa_scratch_reuse + assert not long_kv_cache.enable_swa_scratch_reuse + assert short_kv_cache.history_length == 0 + assert short_kv_cache.capacity == token_nums[0] + 1 + assert long_kv_cache.history_length == token_nums[1] - 1 + assert long_kv_cache.capacity == token_nums[1] + 1 finally: + for req in requests: + cache_manager.free_resources(req) cache_manager.shutdown() @pytest.mark.parametrize("compress_ratios", [[1, 4, 128]]) @@ -910,9 +1931,7 @@ def test_check_invalid_values_in_kv_cache( if invalid: # Inject invalid into a float buffer so NaN/Inf checks are supported. layer_idx = next(i for i, ratio in enumerate(compress_ratios) if ratio > 1) - buffer = cache_manager.get_buffers( - layer_idx, DeepseekV4AttentionType.COMPRESSOR_STATE - ) + buffer = cache_manager.get_buffers(layer_idx, DeepseekV4AttentionType.COMPRESSOR_KV) buffer[0, 0, 0] = torch.nan needs_invalid_cleanup = True diff --git a/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_indices_transform.py b/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_indices_transform.py index 9d619d2aa547..b60aff4c96a1 100644 --- a/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_indices_transform.py +++ b/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_indices_transform.py @@ -399,6 +399,7 @@ def _run_test(scenario: Scenario, context_lengths: List[int]): for req, ctx_len in zip(requests, context_lengths, strict=True): req.context_current_position = ctx_len req.add_new_token(ctx_len, 0) + cache_manager.update_context_resources(scheduled_batch) cache_manager.update_resources(scheduled_batch) cache_manager.shutdown() diff --git a/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_sparse_mla.py b/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_sparse_mla.py index f67ba16ef90f..e7df3560db53 100644 --- a/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_sparse_mla.py +++ b/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_sparse_mla.py @@ -1138,6 +1138,11 @@ def yarn_get_mscale(scale=1, mscale=1): ) print(" PASSED") + for req, ctx_len in zip(requests, context_lengths): + req.context_current_position = ctx_len + req.add_new_token(ctx_len, 0) + cache_manager.update_context_resources(scheduled_batch) + # 8. Generation steps for step in range(num_generation_steps): print(f"\n=== Generation step {step + 1} ===") @@ -1430,6 +1435,14 @@ def test_deepseek_v4_sparse_mla_mixed_batch(context_lengths: List[int]): topk_indices=prefill_topk, ) + gen_requests = requests[1:] + gen_prefill_batch = ScheduledRequests() + gen_prefill_batch.context_requests_last_chunk = gen_requests + for req, ctx_len in zip(gen_requests, gen_ctx_lengths): + req.context_current_position = ctx_len + req.add_new_token(ctx_len, 0) + cache_manager.update_context_resources(gen_prefill_batch) + # Pre-fill COMPRESS buffers for ratio > 1. compress_ref_data: Dict[int, List[torch.Tensor]] = {} for li in TEST_LAYERS: @@ -1445,7 +1458,6 @@ def test_deepseek_v4_sparse_mla_mixed_batch(context_lengths: List[int]): ) # 3. Allocate 1 gen step for gen requests. - gen_requests = requests[1:] _allocate_kv_cache_for_generation(cache_manager, gen_requests) gen_cached_lens = gen_ctx_lengths diff --git a/tests/unittest/_torch/executor/test_kv_cache_budget_split.py b/tests/unittest/_torch/executor/test_kv_cache_budget_split.py index 8aa9ad22196d..109460a5a57b 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_budget_split.py +++ b/tests/unittest/_torch/executor/test_kv_cache_budget_split.py @@ -46,7 +46,9 @@ def _make_creator( host_cache_size=host_cache_size, ) c._tokens_per_block = 64 + c._max_seq_len = 1024 c._max_batch_size = 1 + c._speculative_config = None c._mapping = Mock() c._model_engine = Mock() diff --git a/tests/unittest/_torch/executor/test_kv_cache_estimation.py b/tests/unittest/_torch/executor/test_kv_cache_estimation.py index 1b2ece9b4d1b..81ec3fba673a 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_estimation.py +++ b/tests/unittest/_torch/executor/test_kv_cache_estimation.py @@ -7,12 +7,14 @@ share, not all copies. """ +from types import SimpleNamespace from unittest.mock import Mock, patch import pytest from tensorrt_llm._torch.pyexecutor._util import CacheCost, KvCacheCreator from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import KVCacheManagerV2 +from tensorrt_llm.llmapi.llm_args import KvCacheConfig # --------------------------------------------------------------------------- # Helpers @@ -76,6 +78,7 @@ def _make_creator( # behavior; production sets this in __init__ via # _get_model_kv_cache_manager_cls(), which we bypass here. c._kv_cache_manager_cls = KVCacheManagerV2 + c._is_kv_cache_manager_v2 = True return c @@ -302,6 +305,93 @@ def test_pool_scaling_prevents_mmmu_pro_underestimation(): assert per_pool_tokens >= max_seq_len +def test_v2_cache_size_per_token_models_generation_swa_cost(): + class FakeModelConfig: + quant_config = None + pretrained_config = SimpleNamespace( + hidden_size=32, + num_attention_heads=4, + num_key_value_heads=2, + ) + + def get_num_attention_layers(self): + return 3 + + mapping = Mock(enable_attention_dp=False, tp_size=1) + mapping.pp_layers.return_value = [0, 1, 2] + + no_scratch_size_per_token = CacheCost.from_raw( + KVCacheManagerV2.get_cache_size_per_token( + FakeModelConfig(), + mapping, + tokens_per_block=64, + max_seq_len=4096, + max_batch_size=3, + kv_cache_config=KvCacheConfig(max_attention_window=[2048, 2048, 4096]), + ) + ) + scratch_size_per_token = CacheCost.from_raw( + KVCacheManagerV2.get_cache_size_per_token( + FakeModelConfig(), + mapping, + tokens_per_block=64, + max_seq_len=4096, + max_batch_size=3, + kv_cache_config=KvCacheConfig(max_attention_window=[2048, 2048, 4096]), + enable_swa_scratch_reuse=True, + ) + ) + + # Per layer: K+V * kv_heads * head_dim * bf16 bytes = 2 * 2 * 8 * 2. + expected = CacheCost(slope=64, intercept=3 * 2 * 2048 * 64) + assert no_scratch_size_per_token == expected + assert scratch_size_per_token == expected + + +def test_creator_uses_v2_affine_cache_cost(): + class FakeV2Manager(KVCacheManagerV2): + @staticmethod + def get_cache_size_per_token(model_config, mapping, **kwargs): + return 20, 21 + + creator = object.__new__(KvCacheCreator) + creator._mapping = Mock() + creator._tokens_per_block = 64 + creator._max_seq_len = 1024 + creator._max_batch_size = 3 + creator._kv_cache_config = KvCacheConfig() + creator._speculative_config = None + + cost = creator._per_manager_cache_cost(FakeV2Manager, Mock()) + + assert cost == CacheCost(slope=20, intercept=21) + + +def test_v2_quota_from_max_tokens_models_context_swa_scratch(): + manager = object.__new__(KVCacheManagerV2) + manager.num_local_layers = 3 + manager.pp_layers = [0, 1, 2] + manager.max_attention_window_vec = [128, 128, None] + manager.tokens_per_block = 64 + manager.max_batch_size = 4 + manager.max_num_tokens = 1000 + manager.get_layer_bytes_per_token = lambda local_layer_idx, data_role: [10, 10, 20][ + local_layer_idx + ] + + max_tokens = 1200 + + manager.enable_swa_scratch_reuse = False + no_scratch_quota = manager._get_quota_from_max_tokens(max_tokens) + assert no_scratch_quota == (max_tokens * 20 + manager.max_num_tokens * 20 + 4 * 2 * 128 * 10) + assert manager._get_max_tokens_from_quota(no_scratch_quota) == max_tokens + + manager.enable_swa_scratch_reuse = True + scratch_quota = manager._get_quota_from_max_tokens(max_tokens) + assert scratch_quota == (max_tokens * 20 + manager.max_num_tokens * 10 + 4 * 2 * 128 * 10) + assert manager._get_max_tokens_from_quota(scratch_quota) == max_tokens + + # --------------------------------------------------------------------------- # KVCacheManagerV2 clamp_max_seq_len_for_mem float-to-int cast regression # --------------------------------------------------------------------------- diff --git a/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py b/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py new file mode 100644 index 000000000000..17582156eded --- /dev/null +++ b/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py @@ -0,0 +1,39 @@ +from types import SimpleNamespace + +from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import KVCacheManagerV2 + + +class _FakeKVCache: + def __init__(self, num_committed_tokens: int): + self.num_committed_tokens = num_committed_tokens + self.committed_tokens = None + self.stopped_committing = False + + def commit(self, tokens): + self.committed_tokens = tokens + self.num_committed_tokens += len(tokens) + + def stop_committing(self): + self.stopped_committing = True + + +def test_try_commit_blocks_commits_uncommitted_tokens_and_stops_at_context_end(): + request = SimpleNamespace( + py_request_id=1, + is_dummy_request=False, + context_current_position=8, + context_remaining_length=0, + get_tokens=lambda beam_id: list(range(10)), + ) + kv_cache = _FakeKVCache(num_committed_tokens=4) + manager = object.__new__(KVCacheManagerV2) + manager.enable_block_reuse = True + manager.is_draft = False + manager.kv_cache_map = {request.py_request_id: kv_cache} + manager._augment_tokens_for_block_reuse = lambda tokens, request, start, end: tokens[start:end] + + manager.try_commit_blocks(request) + + assert kv_cache.committed_tokens == [4, 5, 6, 7] + assert kv_cache.num_committed_tokens == 8 + assert kv_cache.stopped_committing diff --git a/tests/unittest/_torch/executor/test_kv_cache_v2_scheduler.py b/tests/unittest/_torch/executor/test_kv_cache_v2_scheduler.py index d2b1875c7286..135740961927 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_v2_scheduler.py +++ b/tests/unittest/_torch/executor/test_kv_cache_v2_scheduler.py @@ -18,12 +18,10 @@ No GPU required. """ -from types import SimpleNamespace from unittest.mock import Mock, patch import pytest -from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import KVCacheManagerV2 from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequestState from tensorrt_llm.llmapi.llm_args import CapacitySchedulerPolicy @@ -106,13 +104,18 @@ def make_encoder_request(request_id, encoder_output_len, lora_task_id=None): def make_disagg_request( - request_id, context_remaining_length=100, lora_task_id=None, num_draft_tokens=0 + request_id, + context_remaining_length=100, + lora_task_id=None, + num_draft_tokens=0, + prompt_len=None, ): req = Mock() req.request_id = request_id req.py_request_id = request_id req.state_value = DISAGG_GEN_INIT req.context_remaining_length = context_remaining_length + req.prompt_len = prompt_len if prompt_len is not None else context_remaining_length req.is_context_init_state = False req.is_generation_in_progress_state = False req.is_first_context_chunk = True @@ -154,6 +157,7 @@ def make_kv_cache_manager( tokens_per_block=64, prepare_context_fn=None, resize_context_fn=None, + prepare_disagg_gen_init_fn=None, try_allocate_generation_fn=None, ): mgr = Mock() @@ -161,6 +165,7 @@ def make_kv_cache_manager( mgr.kv_cache_map = _KVCacheMap() mgr.prepare_context.side_effect = prepare_context_fn or (lambda req: True) mgr.resize_context.side_effect = resize_context_fn or (lambda req, n: True) + mgr.prepare_disagg_gen_init.side_effect = prepare_disagg_gen_init_fn or (lambda req: True) mgr.try_allocate_generation.side_effect = try_allocate_generation_fn or (lambda req: True) mgr.suspend_request.return_value = None mgr.is_request_active.side_effect = lambda req_id: mgr.kv_cache_map[req_id].is_active @@ -583,7 +588,7 @@ def test_ctx_resize_fail_then_gen_succeeds(self): def test_ctx_resize_fail_then_smaller_ctx(self): call_count = [0] - def resize_fn(req, n): + def resize_fn(req, n, history_length=None): call_count[0] += 1 return call_count[0] > 1 # first fails, second succeeds @@ -1207,7 +1212,7 @@ def test_disagg_not_blocked_when_batch_full(self): assert ids(out.generation_requests) == [0, 1] assert ids(out.fitting_disagg_gen_init_requests) == [2] - def test_disagg_gated_by_prepare_context(self): + def test_disagg_gated_by_prepare_disagg_gen_init(self): """Disagg is limited by KV cache / IndexMapper, not batch budget.""" call_count = [0] @@ -1215,7 +1220,7 @@ def prepare_fn(req): call_count[0] += 1 return call_count[0] <= 2 # first 2 succeed, rest fail - mgr = make_kv_cache_manager(prepare_context_fn=prepare_fn) + mgr = make_kv_cache_manager(prepare_disagg_gen_init_fn=prepare_fn) sched = make_scheduler(mgr, max_num_tokens=100) reqs = [make_disagg_request(i) for i in range(4)] out = sched.schedule_request(reqs, set()) @@ -1240,33 +1245,45 @@ def test_disagg_peft_shared_task_id(self): out = sched.schedule_request(reqs, set()) assert ids(out.fitting_disagg_gen_init_requests) == [0, 1] - def test_disagg_prepare_context_fails_skips(self): - """prepare_context failure skips the request, loop continues.""" + def test_disagg_prepare_disagg_gen_init_fails_skips(self): + """prepare_disagg_gen_init failure skips the request, loop continues.""" call_count = [0] def prepare_fn(req): call_count[0] += 1 return call_count[0] > 1 # first fails, second succeeds - mgr = make_kv_cache_manager(prepare_context_fn=prepare_fn) + mgr = make_kv_cache_manager(prepare_disagg_gen_init_fn=prepare_fn) sched = make_scheduler(mgr, max_num_tokens=100) reqs = [make_disagg_request(0), make_disagg_request(1)] out = sched.schedule_request(reqs, set()) assert ids(out.fitting_disagg_gen_init_requests) == [1] - def test_disagg_resize_context_fails_skips(self): - """resize_context failure skips the request, loop continues.""" - call_count = [0] + def test_disagg_uses_prepare_disagg_gen_init_api(self): + """Scheduler routes disagg through prepare_disagg_gen_init (not + prepare_context/resize_context). The single-call API encapsulates + the SWA history pre-declaration internally.""" + called = [] + + def prepare_disagg_fn(req): + called.append(req.py_request_id) + return True def resize_fn(req, n): - call_count[0] += 1 - return call_count[0] > 1 # first fails, second succeeds + raise AssertionError("resize_context must not be called for disagg") - mgr = make_kv_cache_manager(resize_context_fn=resize_fn) + def prepare_ctx_fn(req): + raise AssertionError("prepare_context must not be called for disagg") + + mgr = make_kv_cache_manager( + prepare_disagg_gen_init_fn=prepare_disagg_fn, + prepare_context_fn=prepare_ctx_fn, + resize_context_fn=resize_fn, + ) sched = make_scheduler(mgr, max_num_tokens=100) reqs = [make_disagg_request(0), make_disagg_request(1)] - out = sched.schedule_request(reqs, set()) - assert ids(out.fitting_disagg_gen_init_requests) == [1] + sched.schedule_request(reqs, set()) + assert called == [0, 1] def test_disagg_exceeds_max_batch_size(self): """Disagg count can exceed max_batch_size since it bypasses budget.""" @@ -1327,7 +1344,7 @@ def test_disagg_interleaved_gen_full_batch_continues(self): assert ids(out.context_requests) == [] def test_disagg_prepare_fail_continues_to_next_disagg(self): - """prepare_context failure for one disagg does not stop other disaggs.""" + """prepare_disagg_gen_init failure for one disagg does not stop other disaggs.""" call_count = [0] def selective_prepare(req): @@ -1335,7 +1352,7 @@ def selective_prepare(req): # Fail for request 0 and 2, succeed for 1 and 3 return req.request_id in (1, 3) - mgr = make_kv_cache_manager(prepare_context_fn=selective_prepare) + mgr = make_kv_cache_manager(prepare_disagg_gen_init_fn=selective_prepare) sched = make_scheduler(mgr, max_num_tokens=200) reqs = [make_disagg_request(i) for i in range(4)] out = sched.schedule_request(reqs, set()) @@ -1349,13 +1366,13 @@ def test_disagg_cross_iteration_slot_overflow(self): DISAGG_GENERATION_TRANS_IN_PROGRESS (value=9). Iter 2: TRANS_IN_PROGRESS (9) is invisible to both the disagg branch (checks ==8) and state gating ([10,14)), so budget resets to - 0. New disagg_gen_init passes budget → prepare_context called - again → IndexMapper would crash in production. + 0. New disagg_gen_init passes budget → prepare_disagg_gen_init + called again → IndexMapper would crash in production. - This test counts prepare_context calls across two iterations. With - scheduler_capacity=2 (simulating IndexMapper=3 slots, 1 dummy), a - correct implementation should cap total prepare_context calls to 2 - (the IndexMapper capacity). The bug allows 4. + This test counts prepare_disagg_gen_init calls across two iterations. + With scheduler_capacity=2 (simulating IndexMapper=3 slots, 1 dummy), a + correct implementation should cap total prepare_disagg_gen_init calls to + 2 (the IndexMapper capacity). The bug allows 4. """ prepare_count = [0] @@ -1363,7 +1380,7 @@ def counting_prepare(req): prepare_count[0] += 1 return True - mgr = make_kv_cache_manager(prepare_context_fn=counting_prepare) + mgr = make_kv_cache_manager(prepare_disagg_gen_init_fn=counting_prepare) sched = make_scheduler(mgr, max_num_tokens=200, scheduler_capacity=2) # Iteration 1: two disagg_gen_init requests @@ -1385,7 +1402,7 @@ def counting_prepare(req): out2 = sched.schedule_request(all_active, set()) # BUG: budget resets to 0, TRANS_IN_PROGRESS not counted, so - # both new disagg pass → prepare_context called 4 times total. + # both new disagg pass → prepare_disagg_gen_init called 4 times total. # In production, this would crash IndexMapper (only 3 slots = 2+1 dummy). assert ids(out2.fitting_disagg_gen_init_requests) == [2, 3] assert prepare_count[0] == 4 # <-- proves the overflow @@ -2461,48 +2478,3 @@ def track_resize(req, n): out = sched.schedule_request([req], set()) assert ids(out.context_requests) == [] assert resize_calls == [] # SKIP path: no commit to KV cache - - -# --------------------------------------------------------------------------- -# KVCacheManagerV2.trim_to_history (#14258): unbound on a fake self, all 5 branches. -# --------------------------------------------------------------------------- -class TestTrimToHistory: - @staticmethod - def _call(kv_cache, history_length, req_id=1): - kv_cache_map = {} if kv_cache is None else {req_id: kv_cache} - fake = SimpleNamespace(kv_cache_map=kv_cache_map) - req = SimpleNamespace(py_request_id=req_id) - return KVCacheManagerV2.trim_to_history(fake, req, history_length) - - def test_missing_cache_is_noop_true(self): - assert self._call(None, 50) is True - - def test_inactive_cache_is_noop_true(self): - kv = Mock(is_active=False) - assert self._call(kv, 50) is True - kv.resize.assert_not_called() - - def test_history_not_increasing_is_noop_true(self): - kv = Mock(is_active=True, history_length=50, capacity=100) - assert self._call(kv, 50) is True # 50 <= current 50 - kv.resize.assert_not_called() - - def test_resize_success_clamps_capacity_and_returns_true(self): - kv = Mock(is_active=True, history_length=10, capacity=8) - kv.resize.return_value = True - assert self._call(kv, 64) is True - # target_capacity = max(capacity=8, history=64) = 64 - kv.resize.assert_called_once_with(64, history_length=64) - - def test_resize_rejection_returns_false(self): - kv = Mock(is_active=True, history_length=10, capacity=100) - kv.resize.return_value = False - assert self._call(kv, 64) is False - kv.resize.assert_called_once_with(100, history_length=64) - - def test_resize_exception_degrades_to_false(self): - # Broad except: a non-ValueError (e.g. internal state assert) must - # degrade to False rather than propagate (do-not-narrow contract). - kv = Mock(is_active=True, history_length=10, capacity=100) - kv.resize.side_effect = RuntimeError("internal state assert") - assert self._call(kv, 64) is False diff --git a/tests/unittest/_torch/modeling/test_modeling_deepseekv4.py b/tests/unittest/_torch/modeling/test_modeling_deepseekv4.py index 13218f5c35a9..43ffcfbc4b53 100644 --- a/tests/unittest/_torch/modeling/test_modeling_deepseekv4.py +++ b/tests/unittest/_torch/modeling/test_modeling_deepseekv4.py @@ -148,13 +148,19 @@ def test_deepseek_v4_fused_hc_default_enabled(monkeypatch): assert _resolve_enable_fused_hc(config) is False -def test_deepseek_v4_model_defaults_keep_tokens_per_block(): +def test_deepseek_v4_model_defaults(): class LlmArgs: pass defaults = DeepseekV4ForCausalLM.get_model_defaults(LlmArgs()) - assert defaults == {"kv_cache_config": {"tokens_per_block": 128}} + assert defaults == { + "kv_cache_config": { + "tokens_per_block": 128, + "use_kv_cache_manager_v2": True, + "enable_swa_scratch_reuse": True, + } + } def test_deepseek_v4_weight_remap_for_mxfp4_routed_experts(): @@ -733,6 +739,7 @@ def test_deepseek_v4_sanity(): assert not model.model.layers[0].fusion_config.POST_MOE_FUSION context_sequence_length = [3, 2, 5] + num_contexts = len(context_sequence_length) sequence_length = context_sequence_length + [1, 1] # Total tokens = sum(sequence_length) = 3+2+5+1+1 = 12 @@ -742,7 +749,7 @@ def test_deepseek_v4_sanity(): past_seen_tokens = [0, 0, 0, 62, 75] request_ids = list(range(len(sequence_length))) token_nums = (torch.tensor(past_seen_tokens) + torch.tensor(sequence_length)).tolist() - prompt_lens = token_nums[:3] + past_seen_tokens[3:] + prompt_lens = token_nums[:num_contexts] + past_seen_tokens[num_contexts:] tokens_per_block = 128 # DeepSeek-V4 requirement max_new_tokens = 1024 required_blocks = sum( @@ -798,14 +805,27 @@ def test_deepseek_v4_sanity(): ) success = kv_cache_manager.prepare_context(req) assert success, f"Failed to prepare context for request {req_id}" - # Allocate enough capacity for context tokens plus generation headroom - success = kv_cache_manager.resize_context(req, token_nums[i] + max_new_tokens) + if i < num_contexts: + success = kv_cache_manager.resize_context(req, req.context_chunk_size) + else: + # Warm-cache setup for a generation request: simulate + # past_seen_tokens[i] worth of history without running forward. + # Reach into kv_cache.resize directly because resize_context no + # longer exposes a history_length override (production callers + # use prepare_disagg_gen_init or update_resources to advance it). + kv_cache = kv_cache_manager.kv_cache_map[req.py_request_id] + kv_cache.enable_swa_scratch_reuse = False + target = ( + req.context_current_position + token_nums[i] + kv_cache_manager.num_extra_kv_tokens + ) + capacity = max(kv_cache.capacity, target) + success = kv_cache.resize(capacity, past_seen_tokens[i]) assert success, f"Failed to resize context for request {req_id}" reqs.append(req) attn_metadata = DeepseekV4TrtllmAttentionMetadata( seq_lens=torch.tensor(sequence_length, dtype=torch.int32), - num_contexts=len(context_sequence_length), + num_contexts=num_contexts, max_num_requests=len(sequence_length), kv_cache_params=KVCacheParams( use_cache=True, @@ -833,7 +853,8 @@ def test_deepseek_v4_sanity(): extra_attrs["attention_metadata"] = weakref.ref(attn_metadata) with torch.inference_mode(), model_extra_attrs(extra_attrs): scheduled_batch = ScheduledRequests() - scheduled_batch.context_requests_last_chunk = reqs + scheduled_batch.context_requests_last_chunk = reqs[:num_contexts] + scheduled_batch.generation_requests = reqs[num_contexts:] kv_cache_manager.prepare_resources(scheduled_batch) attn_metadata.prepare() @@ -841,9 +862,11 @@ def test_deepseek_v4_sanity(): input_ids=input_ids, position_ids=position_ids, attn_metadata=attn_metadata ) - for req in reqs: + for req in reqs[:num_contexts]: req.context_current_position = seq_lens[req.py_request_id] + for req in reqs: req.add_new_token(seq_lens[req.py_request_id], 0) + kv_cache_manager.update_context_resources(scheduled_batch) kv_cache_manager.update_resources(scheduled_batch) assert len(past_seen_tokens) == logits.shape[0] @@ -852,6 +875,8 @@ def test_deepseek_v4_sanity(): seq_lens = [seq_len + 1 for seq_len in seq_lens] scheduled_batch = ScheduledRequests() scheduled_batch.generation_requests = reqs + for req in reqs: + assert kv_cache_manager.try_allocate_generation(req) kv_cache_manager.prepare_resources(scheduled_batch) attn_metadata.prepare() logits = model.forward( diff --git a/tests/unittest/_torch/speculative/test_eagle3.py b/tests/unittest/_torch/speculative/test_eagle3.py index 2dfb39269453..04a09ef65b39 100644 --- a/tests/unittest/_torch/speculative/test_eagle3.py +++ b/tests/unittest/_torch/speculative/test_eagle3.py @@ -56,6 +56,7 @@ def test_kv_lens_runtime_with_eagle3_one_model(): mock_kv_cache_manager = MagicMock() mock_kv_cache_manager.tokens_per_block = 32 mock_kv_cache_manager.num_pools = 1 + mock_kv_cache_manager.num_attention_op_pools = mock_kv_cache_manager.num_pools mock_kv_cache_manager.max_blocks_per_seq = 16 mock_kv_cache_manager.max_batch_size = num_seqs mock_kv_cache_manager.max_seq_len = 512 # Large enough to hold our test sequences diff --git a/tests/unittest/bindings/test_executor_bindings.py b/tests/unittest/bindings/test_executor_bindings.py index 3af86491610d..438101bba3da 100644 --- a/tests/unittest/bindings/test_executor_bindings.py +++ b/tests/unittest/bindings/test_executor_bindings.py @@ -1800,6 +1800,16 @@ def test_executor_config(): assert config.mm_embedding_offloading is False assert config.enable_trt_overlap is False + unbounded_stats_config = trtllm.ExecutorConfig( + iter_stats_max_iterations=-1, request_stats_max_iterations=-1) + assert unbounded_stats_config.iter_stats_max_iterations == -1 + assert unbounded_stats_config.request_stats_max_iterations == -1 + + with pytest.raises(Exception): + trtllm.ExecutorConfig(iter_stats_max_iterations=-2) + with pytest.raises(Exception): + trtllm.ExecutorConfig(request_stats_max_iterations=-2) + kwargs = { "max_beam_width": 2, diff --git a/tests/unittest/disaggregated/region/test_page.py b/tests/unittest/disaggregated/region/test_page.py index c11391b530a7..14042299fa3b 100644 --- a/tests/unittest/disaggregated/region/test_page.py +++ b/tests/unittest/disaggregated/region/test_page.py @@ -15,8 +15,8 @@ def _make_buffer_entries(): return np.array( [ - (0, 0, 0, 128), # local_layer_id=0, role=0(KEY), offset=0, size=128 - (0, 1, 128, 128), # local_layer_id=0, role=1(VALUE), offset=128, size=128 + (0, 0, 128), # local_layer_id=0, offset=0, size=128 + (0, 128, 128), # local_layer_id=0, offset=128, size=128 ], dtype=BUFFER_ENTRY_DTYPE, ) diff --git a/tests/unittest/disaggregated/region/test_region.py b/tests/unittest/disaggregated/region/test_region.py index dfe6955b82b5..91c9419077da 100644 --- a/tests/unittest/disaggregated/region/test_region.py +++ b/tests/unittest/disaggregated/region/test_region.py @@ -2,7 +2,6 @@ from tensorrt_llm._torch.disaggregation.base.region import ( DataLayout, - DataRole, IndexRange, KVRegionSpec, RegionSpec, @@ -39,14 +38,6 @@ def test_index_range_invalid_end_before_start(): IndexRange(start=10, end=5) -def test_data_role_flags(): - assert DataRole.KEY == 1 - assert DataRole.VALUE == 2 - combined = DataRole.KEY | DataRole.VALUE - assert DataRole.KEY in combined - assert DataRole.VALUE in combined - - def test_data_layout_flags(): assert DataLayout.HND == 1 assert DataLayout.NHD == 2 @@ -65,7 +56,6 @@ def test_region_spec_construction(): def test_kv_region_spec_defaults(): spec = KVRegionSpec() assert spec.layers is None - assert spec.role == DataRole.KEY | DataRole.VALUE assert spec.heads is None assert spec.tokens is None @@ -73,11 +63,9 @@ def test_kv_region_spec_defaults(): def test_kv_region_spec_with_all_axes(): spec = KVRegionSpec( layers=IndexRange(0, 15), - role=DataRole.KEY, heads=IndexRange(0, 7), tokens=IndexRange(0, 127), ) assert spec.layers == IndexRange(0, 15) - assert spec.role == DataRole.KEY assert spec.heads == IndexRange(0, 7) assert spec.tokens == IndexRange(0, 127) diff --git a/tests/unittest/disaggregated/test_cache_reuse_adapter.py b/tests/unittest/disaggregated/test_cache_reuse_adapter.py index c6f67328bc77..08d030a707b8 100644 --- a/tests/unittest/disaggregated/test_cache_reuse_adapter.py +++ b/tests/unittest/disaggregated/test_cache_reuse_adapter.py @@ -15,7 +15,7 @@ """Tests for CacheReuseAdapter, _create_kv_slice SWA trim, and Sender token-start derivation.""" from types import SimpleNamespace -from unittest.mock import MagicMock, Mock +from unittest.mock import MagicMock import numpy as np import pytest @@ -650,47 +650,6 @@ def test_chunked_slice_entirely_stale_for_swa(self): assert (src_start, dst_start) == (16, 16) -# --------------------------------------------------------------------------- -# KvCacheTransceiverV2._trim_kv_to_prompt_history: capability-gated, graceful-degrade. (#14258) -# --------------------------------------------------------------------------- -class TestTrimKvToPromptHistory: - @staticmethod - def _tc(kv_cache_manager): - tc = object.__new__(KvCacheTransceiverV2) - tc._kv_cache_manager = kv_cache_manager - return tc - - def test_noop_without_trim_capability(self): - # V1 / non-V2 managers lack trim_to_history -> getattr None -> no raise. - tc = self._tc(SimpleNamespace()) - assert ( - tc._trim_kv_to_prompt_history(SimpleNamespace(prompt_len=17, py_request_id=1)) is None - ) - - def test_noop_when_prompt_len_non_positive(self): - trim = Mock(return_value=True) - tc = self._tc(SimpleNamespace(trim_to_history=trim)) - for prompt_len in (0, None): - tc._trim_kv_to_prompt_history(SimpleNamespace(prompt_len=prompt_len, py_request_id=1)) - trim.assert_not_called() - - def test_trims_to_prompt_len_on_success(self): - trim = Mock(return_value=True) - tc = self._tc(SimpleNamespace(trim_to_history=trim)) - req = SimpleNamespace(prompt_len=17, py_request_id=1) - tc._trim_kv_to_prompt_history(req) - trim.assert_called_once_with(req, 17) - - def test_swallows_trim_failure(self): - # trim returns False (degraded) -> method still returns None and never - # raises, so the downstream TRANS_COMPLETE transition is not gated on it. - trim = Mock(return_value=False) - tc = self._tc(SimpleNamespace(trim_to_history=trim)) - req = SimpleNamespace(prompt_len=17, py_request_id=1) - assert tc._trim_kv_to_prompt_history(req) is None - trim.assert_called_once_with(req, 17) - - # --------------------------------------------------------------------------- # KvCacheTransceiverV2 context-manager (__enter__/__exit__) + shutdown idempotency. (#14137) # --------------------------------------------------------------------------- diff --git a/tests/unittest/disaggregated/test_cache_transceiver_single_process.py b/tests/unittest/disaggregated/test_cache_transceiver_single_process.py index 9a993a2f4d48..c5d4254c12f4 100644 --- a/tests/unittest/disaggregated/test_cache_transceiver_single_process.py +++ b/tests/unittest/disaggregated/test_cache_transceiver_single_process.py @@ -8,6 +8,7 @@ import os import threading import uuid +from types import SimpleNamespace # Exclude UCX IB transport (avoid NIXL setup hangs without IB) and gdr_copy # (avoid SIGSEGV at process exit from UCX rcache cleanup; gdr_copy disabled @@ -55,6 +56,49 @@ REQUEST_LENGTHS = [30, 60, 80] +def test_assert_disagg_history_declared_passes_when_contract_met(): + req = SimpleNamespace(py_request_id=7, prompt_len=17) + cache_manager = SimpleNamespace(get_history_length=lambda r: 17) + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._kv_cache_manager = cache_manager + + # Should not raise — history_length matches prompt_len exactly. + transceiver._assert_disagg_history_declared(req) + + # And also passes when history_length exceeds prompt_len (e.g., already advanced). + cache_manager.get_history_length = lambda r: 100 + transceiver._assert_disagg_history_declared(req) + + +def test_assert_disagg_history_declared_raises_on_contract_violation(): + req = SimpleNamespace(py_request_id=7, prompt_len=17) + cache_manager = SimpleNamespace(get_history_length=lambda r: 0) + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._kv_cache_manager = cache_manager + + with pytest.raises(RuntimeError, match="history_length=0 < prompt_len=17"): + transceiver._assert_disagg_history_declared(req) + + +def test_assert_disagg_history_declared_noops_without_capability(): + """V1 / non-V2 managers lack get_history_length; verify graceful skip.""" + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._kv_cache_manager = object() + + # Should not raise — no get_history_length attribute. + transceiver._assert_disagg_history_declared(SimpleNamespace(prompt_len=17)) + + +def test_assert_disagg_history_declared_noops_when_cache_released(): + """If the cache was released (e.g., cancelled mid-transfer), skip the check.""" + req = SimpleNamespace(py_request_id=7, prompt_len=17) + cache_manager = SimpleNamespace(get_history_length=lambda r: None) + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._kv_cache_manager = cache_manager + + transceiver._assert_disagg_history_declared(req) + + # --------------------------------------------------------------------------- # PP layer distribution helpers (mirrors C++ getLayerNumPPRank) # --------------------------------------------------------------------------- @@ -112,6 +156,10 @@ class KvCacheConfigV2: copy_on_partial_reuse: bool = False dtype: str = "auto" disk_prefetch_num_reqs: int = 4 + pool_ratio: Optional[List[float]] = None + avg_seq_len: Optional[int] = None + block_reuse_policy: str = "all_reusable" + enable_swa_scratch_reuse: bool = False max_util_for_resume: float = 0.95 @@ -439,13 +487,23 @@ def _init_pool_data(managers, tp, is_mla, use_v2, fill_random=True, seed_base=10 # --------------------------------------------------------------------------- # Add sequence to manager # --------------------------------------------------------------------------- -def _add_sequence(mgr, request_id: int, prompt_len: int, use_v2: bool): +def _add_sequence( + mgr, request_id: int, prompt_len: int, use_v2: bool, *, is_generation: bool = False +): """Add a sequence to the cache manager. Returns kv_cache for V2 (needed for cleanup).""" if use_v2: kv_cache = mgr._create_kv_cache(request_id, None, None) + if not mgr.enable_block_reuse: + kv_cache.stop_committing() success = kv_cache.resume(torch.cuda.current_stream().cuda_stream) assert success, f"Failed to resume kv_cache for request {request_id}" - kv_cache.resize(prompt_len) + if is_generation: + kv_cache.enable_swa_scratch_reuse = False + kv_cache.resize(prompt_len, history_length=prompt_len) + else: + kv_cache.resize(prompt_len) + kv_cache.resize(None, history_length=prompt_len) + kv_cache.enable_swa_scratch_reuse = False return kv_cache else: # V1: create a dummy LlmRequest for add_sequence_batch @@ -872,7 +930,13 @@ def run_transfer_test( # 5. Add sequences and gen receive first for rank in range(gen_world): for req_idx, req in gen_handle_map[rank]: - kv = _add_sequence(gen_managers[rank], req.py_request_id, req.prompt_len, use_v2) + kv = _add_sequence( + gen_managers[rank], + req.py_request_id, + req.prompt_len, + use_v2, + is_generation=True, + ) if kv is not None: gen_kv_caches[rank].append(kv) gen_tcs[rank].request_and_receive_async(req) diff --git a/tests/unittest/disaggregated/test_deepseek_v4_kv_transfer.py b/tests/unittest/disaggregated/test_deepseek_v4_kv_transfer.py new file mode 100644 index 000000000000..0fbe520ee067 --- /dev/null +++ b/tests/unittest/disaggregated/test_deepseek_v4_kv_transfer.py @@ -0,0 +1,963 @@ +"""Test KV Transfer for DeepseekV4CacheManager with KvCacheTransceiverV2. + +Uses threading + ThreadSafeDistributed to create KvCacheTransceiverV2 instances +in a single process. Validates all DeepseekV4AttentionType cache transfers across +different TP/PP/DP configurations. +""" + +import os +import threading +import uuid +from typing import Dict, List, Optional, Tuple + +import pytest +import torch + +import tensorrt_llm +import tensorrt_llm.bindings +import tensorrt_llm.tensorrt_llm_transfer_agent_binding # noqa: F401 +from tensorrt_llm import DisaggregatedParams, Mapping, SamplingParams +from tensorrt_llm._torch.attention_backend.sparse.deepseek_v4 import DeepseekV4CacheManager +from tensorrt_llm._torch.attention_backend.sparse.deepseek_v4.deepseek_v4 import ( + DEEPSEEK_V4_OVERLAP_COMPRESSOR_RATIO, + DEEPSEEK_V4_SPARSE_RATIO, + DeepseekV4AttentionType, +) +from tensorrt_llm._torch.disaggregation.resource.kv_extractor import KVRegionExtractorV1 +from tensorrt_llm._torch.disaggregation.resource.utils import get_physical_pool, get_pool_bytes +from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 +from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest, LlmRequestType +from tensorrt_llm._torch.pyexecutor.scheduler import ScheduledRequests +from tensorrt_llm._utils import TensorWrapper, convert_to_torch_tensor +from tensorrt_llm.bindings import DataType +from tensorrt_llm.bindings.internal.batch_manager import CacheType as CacheTypeCpp +from tensorrt_llm.llmapi.llm_args import ( + CacheTransceiverConfig, + DeepSeekV4SparseAttentionConfig, + KvCacheConfig, +) + +# Reduce NIXL threads for unit test: default 8 threads per agent causes heavy +# contention when creating multiple agents on a single GPU in the same process. +os.environ.setdefault("TRTLLM_NIXL_NUM_THREADS", "0") + + +# --------------------------------------------------------------------------- +# Constants matching DeepseekV4CacheManager defaults +# --------------------------------------------------------------------------- +HEAD_DIM = 256 # Reduced from 512: test validates transfer, not attention correctness +INDEX_HEAD_DIM = 128 +WINDOW_SIZE = 128 +TOKENS_PER_BLOCK = 128 +MAX_SEQ_LEN = 512 +MAX_BATCH_SIZE = 16 +VOCAB_SIZE = 129280 +NUM_KV_HEADS = 1 +INDEXER_QUANT_BLOCK_SIZE = 128 + + +# DeepSeek-V4 specific ratios (mirrors module constants) +SPARSE_RATIO = DEEPSEEK_V4_SPARSE_RATIO +OVERLAP_COMPRESSOR_RATIO = DEEPSEEK_V4_OVERLAP_COMPRESSOR_RATIO + + +# --------------------------------------------------------------------------- +# ThreadSafeDistributed: threading.Barrier-based Distributed mock +# --------------------------------------------------------------------------- +class ThreadSafeDistributed: + """Distributed mock using threading.Barrier for single-process multi-rank testing. + + Provides the same interface as TorchDistributedWrapper from test_py_cache_transceiver_mp.py + but uses Barrier + Lock + shared dict instead of torch.distributed. + """ + + def __init__( + self, + local_rank: int, + world_size: int, + tp_size: int, + pp_size: int, + tp_rank: int, + pp_rank: int, + shared: dict, + ): + self.rank = local_rank + self._world_size = world_size + self._tp_size = tp_size + self._pp_size = pp_size + self._tp_rank = tp_rank + self._pp_rank = pp_rank + self._s = shared + self._bcast_idx = 0 + self._ag_idx = 0 + self._pp_ag_idx = 0 + self._tp_ag_idx = 0 + + @property + def tp_size(self): + return self._tp_size + + @property + def pp_size(self): + return self._pp_size + + @property + def world_size(self): + return self._world_size + + def broadcast(self, obj, root=0): + idx = self._bcast_idx + self._bcast_idx += 1 + key = f"bcast_{idx}" + if self.rank == root: + self._s[key] = obj + self._s["barrier"].wait() + result = self._s[key] + self._s["barrier"].wait() + return result + + def allgather(self, obj): + idx = self._ag_idx + self._ag_idx += 1 + key = f"ag_{idx}" + with self._s["lock"]: + if key not in self._s: + self._s[key] = [None] * self._world_size + self._s[key][self.rank] = obj + self._s["barrier"].wait() + result = list(self._s[key]) + self._s["barrier"].wait() + return result + + def pp_allgather(self, obj): + idx = self._pp_ag_idx + self._pp_ag_idx += 1 + key = f"pp_ag_{idx}_tp{self._tp_rank}" + with self._s["lock"]: + if key not in self._s: + self._s[key] = [None] * self._pp_size + self._s[key][self._pp_rank] = obj + self._s["barrier"].wait() + result = list(self._s[key]) + self._s["barrier"].wait() + return result + + def tp_allgather(self, obj): + idx = self._tp_ag_idx + self._tp_ag_idx += 1 + key = f"tp_ag_{idx}_pp{self._pp_rank}" + with self._s["lock"]: + if key not in self._s: + self._s[key] = [None] * self._tp_size + self._s[key][self._tp_rank] = obj + self._s["barrier"].wait() + result = list(self._s[key]) + self._s["barrier"].wait() + return result + + +# --------------------------------------------------------------------------- +# Threading helpers +# --------------------------------------------------------------------------- +def run_concurrent(items, fn): + """Run fn(item) for each item concurrently in threads and propagate errors.""" + errors = [None] * len(items) + results = [None] * len(items) + + def _worker(idx, item): + try: + results[idx] = fn(item) + except Exception as e: + errors[idx] = e + + threads = [threading.Thread(target=_worker, args=(i, item)) for i, item in enumerate(items)] + for t in threads: + t.start() + for t in threads: + t.join() + for i, err in enumerate(errors): + if err is not None: + raise err + return results + + +def _create_transceiver_in_thread(rank, mapping, cache_manager, dist_mock, config, results, errors): + """Thread target: create one KvCacheTransceiverV2.""" + try: + tc = KvCacheTransceiverV2( + mapping=mapping, + dist=dist_mock, + kv_cache_manager=cache_manager, + cache_transceiver_config=config, + ) + results[rank] = tc + except Exception as e: + errors[rank] = e + + +def create_instance_transceivers( + tp: int, pp: int, enable_dp: bool, cache_managers: List, config: CacheTransceiverConfig +) -> List[KvCacheTransceiverV2]: + """Create KvCacheTransceiverV2 for all ranks via threaded init.""" + world_size = tp * pp + shared = {"barrier": threading.Barrier(world_size), "lock": threading.Lock()} + results = [None] * world_size + errors = [None] * world_size + threads = [] + + for rank in range(world_size): + pp_rank = rank // tp + tp_rank = rank % tp + mapping = Mapping( + world_size=world_size, + rank=rank, + tp_size=tp, + pp_size=pp, + enable_attention_dp=enable_dp, + ) + dist_mock = ThreadSafeDistributed(rank, world_size, tp, pp, tp_rank, pp_rank, shared) + t = threading.Thread( + target=_create_transceiver_in_thread, + args=( + rank, + mapping, + cache_managers[rank], + dist_mock, + config, + results, + errors, + ), + ) + threads.append(t) + + for t in threads: + t.start() + for t in threads: + t.join() + + for rank, err in enumerate(errors): + if err is not None: + raise err + + return results + + +# --------------------------------------------------------------------------- +# DeepseekV4CacheManager creation helpers +# --------------------------------------------------------------------------- +def _create_deepseek_v4_manager( + mapping: Mapping, + compress_ratios: List[int], + dtype: DataType = DataType.BF16, + compressor_dtype: DataType = DataType.FLOAT, +) -> DeepseekV4CacheManager: + """Create a DeepseekV4CacheManager for the given mapping.""" + sparse_attn_config = DeepSeekV4SparseAttentionConfig( + index_head_dim=INDEX_HEAD_DIM, + window_size=WINDOW_SIZE, + compress_ratios=compress_ratios, + ) + max_num_tokens = MAX_SEQ_LEN * MAX_BATCH_SIZE + kv_cache_config = KvCacheConfig( + enable_block_reuse=False, + max_tokens=max_num_tokens, + event_buffer_max_size=0, + ) + return DeepseekV4CacheManager( + kv_cache_config=kv_cache_config, + kv_cache_type=CacheTypeCpp.SELFKONLY, + num_layers=len(compress_ratios), + num_kv_heads=NUM_KV_HEADS, + head_dim=HEAD_DIM, + tokens_per_block=TOKENS_PER_BLOCK, + max_seq_len=MAX_SEQ_LEN, + max_batch_size=MAX_BATCH_SIZE, + mapping=mapping, + dtype=dtype, + compressor_dtype=compressor_dtype, + vocab_size=VOCAB_SIZE, + max_num_tokens=max_num_tokens, + sparse_attn_config=sparse_attn_config, + ) + + +def _create_managers_for_instance( + tp: int, + pp: int, + enable_dp: bool, + compress_ratios: List[int], +) -> List[DeepseekV4CacheManager]: + """Create DeepseekV4CacheManagers for all ranks in an instance.""" + world_size = tp * pp + managers = [] + for rank in range(world_size): + mapping = Mapping( + world_size=world_size, + rank=rank, + tp_size=tp, + pp_size=pp, + enable_attention_dp=enable_dp, + ) + managers.append(_create_deepseek_v4_manager(mapping, compress_ratios)) + return managers + + +def _init_pool_data(managers: List, tp: int, seed_base: int = 0, fill_random: bool = True): + """Initialize pool data for all managers. + + Uses half-precision view of pool memory for initialization. + Since pools have mixed dtypes (BF16, FP32, FP8), we use HALF (2 bytes) as a + common element type that evenly divides all pool sizes. + + Args: + managers: List of DeepseekV4CacheManagers. + tp: TP size. + seed_base: Base seed offset (different for ctx vs gen to avoid collisions). + fill_random: If True fill with random data (ctx), if False fill with zeros (gen). + """ + for rank, mgr in enumerate(managers): + pp_rank = rank // tp + page_table = KVRegionExtractorV1(mgr).page_table + + # Collect unique pools (deduplicate by base_address) + unique_pools: Dict[int, int] = {} + for lg_idx, lg in enumerate(page_table.layer_groups): + for pv in lg.pool_views: + pool = get_physical_pool(page_table, lg_idx, pv.pool_idx) + key = pool.base_address + if key not in unique_pools or get_pool_bytes(pool) > unique_pools[key]: + unique_pools[key] = get_pool_bytes(pool) + + # Use HALF (2 bytes) as element type - divides all pool sizes evenly + elem_bytes = 2 # sizeof(half) + for pool_base_ptr, pool_size in unique_pools.items(): + pool_size_elements = pool_size // elem_bytes + pool_tensor = convert_to_torch_tensor( + TensorWrapper(pool_base_ptr, DataType.HALF, [pool_size_elements]) + ) + if fill_random: + # Same seed for same pp_rank across TP ranks (kv_heads=1) + seed = seed_base + pp_rank + generator = torch.Generator(device=pool_tensor.device).manual_seed(seed) + random_values = torch.rand( + pool_tensor.shape, + dtype=pool_tensor.dtype, + device=pool_tensor.device, + generator=generator, + ) + pool_tensor.copy_(random_values) + else: + pool_tensor.zero_() + + +# --------------------------------------------------------------------------- +# Attention type helpers +# --------------------------------------------------------------------------- +def _get_attn_types_for_layer( + layer_idx: int, compress_ratios: List[int] +) -> List[DeepseekV4AttentionType]: + """Get the attention types for a given layer based on its compress ratio.""" + ratio = compress_ratios[layer_idx] + is_compress = ratio > 1 + is_sparse = ratio == SPARSE_RATIO + types = [DeepseekV4AttentionType.SWA] + if is_compress: + types.extend( + [ + DeepseekV4AttentionType.COMPRESS, + DeepseekV4AttentionType.COMPRESSOR_KV, + DeepseekV4AttentionType.COMPRESSOR_SCORE, + ] + ) + if is_sparse: + types.extend( + [ + DeepseekV4AttentionType.INDEXER_COMPRESS, + DeepseekV4AttentionType.INDEXER_COMPRESSOR_KV, + DeepseekV4AttentionType.INDEXER_COMPRESSOR_SCORE, + ] + ) + return types + + +# Mirror transceiver.py windowed-block trim: only in-window blocks are transferred. +_WINDOWED_ATTN_TYPES = { + DeepseekV4AttentionType.SWA, + DeepseekV4AttentionType.COMPRESSOR_KV, + DeepseekV4AttentionType.COMPRESSOR_SCORE, + DeepseekV4AttentionType.INDEXER_COMPRESSOR_KV, + DeepseekV4AttentionType.INDEXER_COMPRESSOR_SCORE, +} + + +def _expected_valid_blocks( + attn_type: DeepseekV4AttentionType, compress_ratio: int, prompt_len: int +) -> Optional[int]: + if attn_type == DeepseekV4AttentionType.SWA: + window = WINDOW_SIZE + elif attn_type in _WINDOWED_ATTN_TYPES: + state_factor = 2 if compress_ratio == OVERLAP_COMPRESSOR_RATIO else 1 + window = state_factor * compress_ratio + else: + return None + total = (prompt_len + TOKENS_PER_BLOCK - 1) // TOKENS_PER_BLOCK + stale = max(0, (prompt_len + 1 - window) // TOKENS_PER_BLOCK) + return total - stale + + +def _split_blockwise_buffer( + buffer: torch.Tensor, + index_head_dim: int = INDEX_HEAD_DIM, + quant_block_size: int = INDEXER_QUANT_BLOCK_SIZE, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Split a blockwise FP8 quantized buffer into value and scale buffers. + + Args: + buffer: shape [num_blocks, tokens_per_block, bytes_per_token] + + Returns: + (values_buffer, scales_buffer) where values are uint8 and scales are float32 + """ + num_blocks, tokens_per_block, bytes_per_token = buffer.shape + bytes_per_block = bytes_per_token * tokens_per_block + + # Value buffer + value_shape = (num_blocks, tokens_per_block, index_head_dim) + value_stride = (bytes_per_block, index_head_dim, 1) + value_buffer = buffer.as_strided(value_shape, value_stride, 0).view(torch.uint8) + + # Scale buffer + scale_dim = index_head_dim // quant_block_size + scale_bytes = scale_dim * 4 # float32 = 4 bytes + scale_shape = (num_blocks, tokens_per_block, scale_bytes) + scale_stride = (bytes_per_block, scale_bytes, 1) + scale_offset = index_head_dim * tokens_per_block + scale_buffer = buffer.as_strided(scale_shape, scale_stride, scale_offset).view(torch.float32) + + return value_buffer, scale_buffer + + +# --------------------------------------------------------------------------- +# Verification +# --------------------------------------------------------------------------- +def _read_cache_data( + mgr: DeepseekV4CacheManager, + layer_idx: int, + attn_type: DeepseekV4AttentionType, + request_id: int, +) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: + """Read cache data for a layer/attn_type from a DeepseekV4CacheManager. + + Returns all non-BAD block data. For INDEXER_COMPRESS returns (values, scales). + """ + buffer = mgr.get_buffers(layer_idx, attn_type) + # get_cache_indices may contain BAD_PAGE_INDEX (-1) for evicted blocks; filter them out. + indices = [i for i in mgr.get_cache_indices(request_id, layer_idx, attn_type) if i >= 0] + + if not indices: + return torch.tensor([]), None + + if attn_type == DeepseekV4AttentionType.INDEXER_COMPRESS: + values_buf, scales_buf = _split_blockwise_buffer(buffer) + return values_buf[indices], scales_buf[indices] + + return buffer[indices], None + + +def _find_ctx_rank_for_layer( + layer_idx: int, + ctx_managers: List[DeepseekV4CacheManager], + ctx_tp: int, + ctx_enable_dp: bool, + req_idx: int, +) -> int: + """Find the ctx rank that owns a given model layer. + + Returns a rank with the correct PP rank and correct TP rank for DP. + """ + for rank, mgr in enumerate(ctx_managers): + tp_rank = rank % ctx_tp + if ctx_enable_dp: + if req_idx % ctx_tp != tp_rank: + continue + elif tp_rank != 0: + # Without DP, all TP ranks have same data; use tp_rank=0 + continue + if layer_idx in mgr.pp_layers: + return rank + raise ValueError(f"No ctx rank found for layer {layer_idx}") + + +def verify_all_requests( + request_lengths: List[int], + compress_ratios: List[int], + ctx_managers: List[DeepseekV4CacheManager], + gen_managers: List[DeepseekV4CacheManager], + ctx_tp: int, + ctx_pp: int, + gen_tp: int, + gen_pp: int, + ctx_enable_dp: bool, + gen_enable_dp: bool, + ctx_request_ids: List[int], + gen_request_ids: List[int], +): + """Verify transferred cache data for all requests across all gen ranks.""" + gen_world = gen_tp * gen_pp + + for req_idx, req_len in enumerate(request_lengths): + ctx_rid = ctx_request_ids[req_idx] + gen_rid = gen_request_ids[req_idx] + + for gen_rank in range(gen_world): + tp_rank = gen_rank % gen_tp + # Skip ranks that didn't handle this request (DP mode) + if gen_enable_dp and req_idx % gen_tp != tp_rank: + continue + + gen_mgr = gen_managers[gen_rank] + gen_pp_layers = gen_mgr.pp_layers + + for layer_idx in gen_pp_layers: + # Find the ctx rank owning this layer + ctx_rank = _find_ctx_rank_for_layer( + layer_idx, + ctx_managers, + ctx_tp, + ctx_enable_dp, + req_idx, + ) + ctx_mgr = ctx_managers[ctx_rank] + + for attn_type in _get_attn_types_for_layer(layer_idx, compress_ratios): + ctx_data, ctx_scales = _read_cache_data(ctx_mgr, layer_idx, attn_type, ctx_rid) + gen_data, gen_scales = _read_cache_data(gen_mgr, layer_idx, attn_type, gen_rid) + + expected_valid = _expected_valid_blocks( + attn_type, compress_ratios[layer_idx], req_len + ) + if expected_valid is not None: + if expected_valid <= 0: + ctx_data = ctx_data[:0] + gen_data = gen_data[:0] + else: + ctx_data = ctx_data[-expected_valid:] + gen_data = gen_data[-expected_valid:] + + assert ctx_data.shape == gen_data.shape, ( + f"Shape mismatch at req={req_idx} layer={layer_idx} " + f"attn={attn_type.name}: ctx={ctx_data.shape} gen={gen_data.shape}" + ) + + torch.testing.assert_close( + gen_data, + ctx_data, + rtol=0, + atol=0, + msg=lambda m: ( + f"Data mismatch at req={req_idx} layer={layer_idx} " + f"attn={attn_type.name} gen_rank={gen_rank}: {m}" + ), + ) + + if ctx_scales is not None: + assert gen_scales is not None, ( + f"Expected scales at req={req_idx} layer={layer_idx} attn={attn_type.name}" + ) + torch.testing.assert_close( + gen_scales, + ctx_scales, + rtol=0, + atol=0, + msg=lambda m: ( + f"Scale mismatch at req={req_idx} layer={layer_idx} " + f"attn={attn_type.name} gen_rank={gen_rank}: {m}" + ), + ) + else: + assert gen_scales is None + + +# --------------------------------------------------------------------------- +# Main test function +# --------------------------------------------------------------------------- +def _get_ctx_info_endpoint(tc: KvCacheTransceiverV2) -> Optional[str]: + """Extract the context_info_endpoint from a transceiver's disaggregated params.""" + endpoints = tc.get_disaggregated_params().get("ctx_info_endpoint") or [] + return endpoints[0] if endpoints else None + + +def run_deepseek_v4_transfer_test( + ctx_tp: int, + ctx_pp: int, + gen_tp: int, + gen_pp: int, + ctx_enable_dp: bool, + gen_enable_dp: bool, + compress_ratios: List[int], + update_before_transfer: bool = True, +): + """Run a full DeepSeek-V4 KV transfer test.""" + ctx_world = ctx_tp * ctx_pp + gen_world = gen_tp * gen_pp + + # Mix of block-aligned and non-aligned lengths for boundary testing. + # TOKENS_PER_BLOCK=128: 65=half+1, 256=2x exact, 129=1x+1, 383=3x-1 + request_lengths = [65, 256, 129, 383] + + # ===== 1. Create DeepseekV4CacheManagers ===== + ctx_managers = _create_managers_for_instance(ctx_tp, ctx_pp, ctx_enable_dp, compress_ratios) + gen_managers = _create_managers_for_instance(gen_tp, gen_pp, gen_enable_dp, compress_ratios) + + # ===== 2. Initialize data ===== + # ctx: random data, seed=pp_rank (same across TP, different across PP) + _init_pool_data(ctx_managers, ctx_tp, seed_base=1000, fill_random=True) + # gen: zeros + _init_pool_data(gen_managers, gen_tp, fill_random=False) + + # ===== 3. Create KvCacheTransceiverV2 instances (threaded init) ===== + config = CacheTransceiverConfig( + backend="NIXL", + transceiver_runtime="PYTHON", + max_tokens_in_buffer=512, + ) + ctx_tcs = create_instance_transceivers(ctx_tp, ctx_pp, ctx_enable_dp, ctx_managers, config) + gen_tcs = create_instance_transceivers(gen_tp, gen_pp, gen_enable_dp, gen_managers, config) + + try: + ctx_info_endpoint = _get_ctx_info_endpoint(ctx_tcs[0]) + + # ===== 4. Create requests and determine handle map ===== + # handle_map: rank -> [(req_idx, ctx_request, gen_request)] + ctx_handle_map: Dict[int, List] = {r: [] for r in range(ctx_world)} + gen_handle_map: Dict[int, List] = {r: [] for r in range(gen_world)} + ctx_request_ids: List[int] = [] + gen_request_ids: List[int] = [] + + sampling_params = SamplingParams() + + for req_idx, req_len in enumerate(request_lengths): + unique_rid = uuid.uuid4().int & 0x7FFFFFFFFFFFFFFF + ctx_rid = req_idx * 2 + gen_rid = req_idx * 2 + 1 + ctx_request_ids.append(ctx_rid) + gen_request_ids.append(gen_rid) + + ctx_dp_rank = req_idx % ctx_tp if ctx_enable_dp else 0 + + ctx_request = LlmRequest( + request_id=ctx_rid, + max_new_tokens=1, + input_tokens=list(range(req_len)), + sampling_config=tensorrt_llm.bindings.SamplingConfig( + sampling_params._get_sampling_config() + ), + is_streaming=False, + llm_request_type=LlmRequestType.LLMREQUEST_TYPE_CONTEXT_ONLY, + ) + ctx_request.py_disaggregated_params = DisaggregatedParams(disagg_request_id=unique_rid) + + gen_request = LlmRequest( + request_id=gen_rid, + max_new_tokens=1, + input_tokens=list(range(req_len)), + sampling_config=tensorrt_llm.bindings.SamplingConfig( + sampling_params._get_sampling_config() + ), + is_streaming=False, + llm_request_type=LlmRequestType.LLMREQUEST_TYPE_GENERATION_ONLY, + ) + gen_request.py_disaggregated_params = DisaggregatedParams( + ctx_request_id=ctx_rid, + ctx_dp_rank=ctx_dp_rank, + ctx_info_endpoint=ctx_info_endpoint, + disagg_request_id=unique_rid, + ) + + for rank in range(ctx_world): + tp_rank = rank % ctx_tp + should_handle = (not ctx_enable_dp) or (req_idx % ctx_tp == tp_rank) + if should_handle: + ctx_handle_map[rank].append((req_idx, ctx_request)) + + for rank in range(gen_world): + tp_rank = rank % gen_tp + should_handle = (not gen_enable_dp) or (req_idx % gen_tp == tp_rank) + if should_handle: + gen_handle_map[rank].append((req_idx, gen_request)) + + # ===== 5. Allocate KV cache for all ranks ===== + # prepare_resources is a no-op for non-draft KVCacheManagerV2. + # All ranks must allocate BEFORE mutating shared request objects + # (add_new_token changes is_first_context_chunk). + # + # Gen ranks take the disagg-gen-init path: prepare_disagg_gen_init + # sizes the cache for the full prompt and pre-declares + # history_length=prompt_len, matching what the V2 scheduler's + # _try_schedule_disagg_gen_init does in production so the + # transceiver's TRANS_COMPLETE contract check is satisfied. + # Ctx ranks take the regular prefill path (prepare_context + + # resize_context). + gen_batches: Dict[int, ScheduledRequests] = {} + for rank in range(gen_world): + reqs = [req for _, req in gen_handle_map[rank]] + if reqs: + batch = ScheduledRequests() + batch.context_requests_last_chunk = reqs + for req in reqs: + gen_managers[rank].prepare_disagg_gen_init(req) + gen_batches[rank] = batch + + ctx_batches: Dict[int, ScheduledRequests] = {} + for rank in range(ctx_world): + reqs = [req for _, req in ctx_handle_map[rank]] + if reqs: + batch = ScheduledRequests() + batch.context_requests_last_chunk = reqs + for req in reqs: + ctx_managers[rank].prepare_context(req) + ctx_managers[rank].resize_context(req, req.context_chunk_size) + ctx_batches[rank] = batch + + # ===== 5.5. context_current_position + add_new_token ===== + # Set position on each unique request once (needed for transfer metadata). + seen: set = set() + for rank in range(ctx_world): + for _, req in ctx_handle_map[rank]: + if req.py_request_id not in seen: + req.context_current_position = req.prompt_len + req.add_new_token(req.prompt_len, 0) + seen.add(req.py_request_id) + + seen = set() + for rank in range(gen_world): + for _, req in gen_handle_map[rank]: + if req.py_request_id not in seen: + req.context_current_position = req.prompt_len + req.add_new_token(req.prompt_len, 0) + seen.add(req.py_request_id) + + # ===== 5.6. update_resources BEFORE transfer (mode: update_before) ===== + if update_before_transfer: + for rank, batch in ctx_batches.items(): + ctx_managers[rank].update_resources(batch) + for rank, batch in gen_batches.items(): + gen_managers[rank].update_resources(batch) + + # ===== 6. gen receive + ctx send ===== + for rank in range(gen_world): + for _, req in gen_handle_map[rank]: + gen_tcs[rank].request_and_receive_async(req) + for rank in range(ctx_world): + for _, req in ctx_handle_map[rank]: + ctx_tcs[rank].respond_and_send_async(req) + + # ===== 7. Wait for completion (threaded, dist calls inside) ===== + run_concurrent( + ctx_tcs, lambda tc: tc.check_context_transfer_status(None, mark_complete=True) + ) + run_concurrent(gen_tcs, lambda tc: tc.check_gen_transfer_status(None)) + + # ===== 7.5. update_resources AFTER transfer (mode: update_after) ===== + if not update_before_transfer: + for rank, batch in ctx_batches.items(): + ctx_managers[rank].update_resources(batch) + for rank, batch in gen_batches.items(): + gen_managers[rank].update_resources(batch) + + # ===== 8. Verify ===== + verify_all_requests( + request_lengths=request_lengths, + compress_ratios=compress_ratios, + ctx_managers=ctx_managers, + gen_managers=gen_managers, + ctx_tp=ctx_tp, + ctx_pp=ctx_pp, + gen_tp=gen_tp, + gen_pp=gen_pp, + ctx_enable_dp=ctx_enable_dp, + gen_enable_dp=gen_enable_dp, + ctx_request_ids=ctx_request_ids, + gen_request_ids=gen_request_ids, + ) + + finally: + for tc in ctx_tcs + gen_tcs: + try: + tc.shutdown() + except Exception: + pass + for mgr in ctx_managers + gen_managers: + try: + mgr.shutdown() + except Exception: + pass + + +# --------------------------------------------------------------------------- +# Test configurations +# --------------------------------------------------------------------------- +TEST_CONFIGS = [ + # (ctx_tp, ctx_pp, gen_tp, gen_pp, ctx_enable_dp, gen_enable_dp, test_id) + # Basic + (1, 1, 1, 1, False, False, "tp1_pp1"), + (2, 1, 2, 1, False, False, "tp2_pp1"), + (1, 2, 1, 2, False, False, "pp2_symmetric"), + (1, 2, 1, 1, False, False, "pp2_to_pp1"), + (2, 2, 2, 2, False, False, "tp2_pp2"), + # DP + (2, 1, 2, 1, True, True, "tp2_dp_both"), + (2, 1, 1, 2, True, False, "ctx_dp_gen_pp2"), + (2, 2, 2, 2, True, True, "tp2_pp2_dp_both"), +] + + +# @pytest.mark.threadleak(enabled=False) +@pytest.mark.timeout(180) +@pytest.mark.parametrize( + "ctx_tp,ctx_pp,gen_tp,gen_pp,ctx_enable_dp,gen_enable_dp", + [(c[0], c[1], c[2], c[3], c[4], c[5]) for c in TEST_CONFIGS], + ids=[c[6] for c in TEST_CONFIGS], +) +@pytest.mark.parametrize( + "compress_ratios", + [[1, 4, 128], [128, 1, 4, 128]], + ids=["cr_1_4_128", "cr_128_1_4_128"], +) +@pytest.mark.parametrize( + "update_before_transfer", + [True, False], + ids=["update_before", "update_after"], +) +def test_deepseek_v4_kv_transfer( + ctx_tp, + ctx_pp, + gen_tp, + gen_pp, + ctx_enable_dp, + gen_enable_dp, + compress_ratios, + update_before_transfer, +): + """Test KvCacheTransceiverV2 with DeepseekV4CacheManager.""" + mode = "update_before" if update_before_transfer else "update_after" + print( + f"\nRunning DeepSeek-V4 transfer test [{mode}]: " + f"ctx_tp={ctx_tp} ctx_pp={ctx_pp} gen_tp={gen_tp} gen_pp={gen_pp} " + f"ctx_dp={ctx_enable_dp} gen_dp={gen_enable_dp} " + f"compress_ratios={compress_ratios}" + ) + + run_deepseek_v4_transfer_test( + ctx_tp=ctx_tp, + ctx_pp=ctx_pp, + gen_tp=gen_tp, + gen_pp=gen_pp, + ctx_enable_dp=ctx_enable_dp, + gen_enable_dp=gen_enable_dp, + compress_ratios=compress_ratios, + update_before_transfer=update_before_transfer, + ) + + print("PASSED") + + +# --------------------------------------------------------------------------- +# PP layer distribution helpers (mirrors C++ getLayerNumPPRank) +# --------------------------------------------------------------------------- +def _get_layers_per_pp(num_layers: int, pp_size: int) -> List[int]: + """Return a list of layer counts per PP rank. + + When num_layers is not evenly divisible by pp_size, the first + (num_layers % pp_size) ranks get one extra layer. + Matches Mapping.pp_layers / torch.tensor_split behaviour. + """ + base = num_layers // pp_size + extra = num_layers % pp_size + return [base + (1 if r < extra else 0) for r in range(pp_size)] + + +# --------------------------------------------------------------------------- +# Uneven PP layer test configurations +# --------------------------------------------------------------------------- +# compress_ratios templates for different num_layers (covering ratios 1, 4, 128) +UNEVEN_PP_COMPRESS_RATIOS = { + 5: [1, 4, 128, 1, 4], + 7: [1, 4, 128, 1, 4, 128, 1], +} + +UNEVEN_PP_CONFIGS = [ + # (ctx_tp, ctx_pp, gen_tp, gen_pp, ctx_dp, gen_dp, num_layers, test_id) + # 5 layers, pp=2 → [3, 2] + (1, 2, 1, 2, False, False, 5, "5L_tp1_pp2"), + (2, 2, 2, 2, False, False, 5, "5L_tp2_pp2"), + # 5 layers, pp=3 → [2, 2, 1] + (1, 3, 1, 3, False, False, 5, "5L_tp1_pp3"), + # 7 layers, pp=2 → [4, 3] + (1, 2, 1, 2, False, False, 7, "7L_tp1_pp2"), + (2, 2, 2, 2, False, False, 7, "7L_tp2_pp2"), + # 7 layers, pp=3 → [3, 2, 2] + (1, 3, 1, 3, False, False, 7, "7L_tp1_pp3"), + # 7 layers, pp=4 → [2, 2, 2, 1] + (1, 4, 1, 4, False, False, 7, "7L_tp1_pp4"), + # Asymmetric TP/PP with uneven layers + (2, 1, 1, 2, False, False, 5, "5L_tp2_to_pp2"), + (1, 2, 2, 1, False, False, 5, "5L_pp2_to_tp2"), + (4, 1, 1, 4, False, False, 7, "7L_tp4_to_pp4"), + (1, 4, 4, 1, False, False, 7, "7L_pp4_to_tp4"), + (2, 2, 1, 4, False, False, 7, "7L_tp2pp2_to_pp4"), + # Uneven layers + DP + (2, 2, 2, 2, True, True, 5, "5L_tp2_pp2_dp_both"), +] + + +@pytest.mark.timeout(180) +@pytest.mark.parametrize( + "ctx_tp,ctx_pp,gen_tp,gen_pp,ctx_enable_dp,gen_enable_dp,num_layers", + [(c[0], c[1], c[2], c[3], c[4], c[5], c[6]) for c in UNEVEN_PP_CONFIGS], + ids=[c[7] for c in UNEVEN_PP_CONFIGS], +) +@pytest.mark.parametrize( + "update_before_transfer", + [True, False], + ids=["update_before", "update_after"], +) +def test_deepseek_v4_kv_transfer_uneven_pp( + ctx_tp, + ctx_pp, + gen_tp, + gen_pp, + ctx_enable_dp, + gen_enable_dp, + num_layers, + update_before_transfer, +): + """Test KvCacheTransceiverV2 with DeepseekV4CacheManager and uneven layers-per-PP-rank. + + When num_layers is not evenly divisible by pp_size, the first + (num_layers % pp_size) PP ranks get one extra layer. + """ + compress_ratios = UNEVEN_PP_COMPRESS_RATIOS[num_layers] + mode = "update_before" if update_before_transfer else "update_after" + + print( + f"\nRunning DeepSeek-V4 uneven PP transfer test [{mode}]: " + f"ctx_tp={ctx_tp} ctx_pp={ctx_pp} gen_tp={gen_tp} gen_pp={gen_pp} " + f"ctx_dp={ctx_enable_dp} gen_dp={gen_enable_dp} " + f"num_layers={num_layers} compress_ratios={compress_ratios} " + f"layers_per_pp(ctx)={_get_layers_per_pp(num_layers, ctx_pp)} " + f"layers_per_pp(gen)={_get_layers_per_pp(num_layers, gen_pp)}" + ) + + run_deepseek_v4_transfer_test( + ctx_tp=ctx_tp, + ctx_pp=ctx_pp, + gen_tp=gen_tp, + gen_pp=gen_pp, + ctx_enable_dp=ctx_enable_dp, + gen_enable_dp=gen_enable_dp, + compress_ratios=compress_ratios, + update_before_transfer=update_before_transfer, + ) + + print("PASSED") diff --git a/tests/unittest/disaggregated/test_extractor.py b/tests/unittest/disaggregated/test_extractor.py index 57b38ce317ea..97d4503a6ee0 100644 --- a/tests/unittest/disaggregated/test_extractor.py +++ b/tests/unittest/disaggregated/test_extractor.py @@ -1,20 +1,18 @@ import numpy as np import pytest -from tensorrt_llm._torch.disaggregation.base.region import DataRole, MemRegionGroup, SpecRegion +from tensorrt_llm._torch.disaggregation.base.region import MemRegionGroup, SpecRegion from tensorrt_llm._torch.disaggregation.resource.kv_extractor import ( KVRegionExtractorV1, build_page_table, ) +from tensorrt_llm._torch.disaggregation.resource.page import MapperKind from tensorrt_llm._torch.disaggregation.resource.utils import ( - PoolRole, - get_device_pointer, get_global_layer_ids, get_layer_to_layer_group, get_num_layer_groups, get_num_layers, get_physical_pool, - get_pool_role, get_unique_layers, ) from tensorrt_llm._torch.pyexecutor.resource_manager import ( @@ -154,26 +152,19 @@ def test_build_page_table(): assert pool.base_address > 0 assert pool.num_slots == 50 # 3200 tokens / 64 tokens_per_block assert len(pv.buffer_entries) > 0 - assert get_pool_role(pv, kv_factor=2) == PoolRole.KV_CACHE + assert pv.pool_role == frozenset({"key", "value"}) + assert pv.mapper_kind == MapperKind.INDEXED assert len(get_global_layer_ids(lg)) > 0 local_layer_id = list(get_unique_layers(pv))[0] - ptr_key = get_device_pointer( - page_table, - lg_idx=0, - pool_view=pv, - slot_id=0, - local_layer_id=local_layer_id, - role=int(DataRole.KEY), - ) - ptr_value = get_device_pointer( - page_table, - lg_idx=0, - pool_view=pv, - slot_id=0, - local_layer_id=local_layer_id, - role=int(DataRole.VALUE), + layer_entries = sorted( + (e for e in pv.buffer_entries if int(e["local_layer_id"]) == local_layer_id), + key=lambda e: int(e["offset"]), ) + assert len(layer_entries) >= 2 # KV layout: K + V buffer per layer + pool = get_physical_pool(page_table, 0, pv.pool_idx) + ptr_key = int(pool.base_address) + int(layer_entries[0]["offset"]) + ptr_value = int(pool.base_address) + int(layer_entries[1]["offset"]) assert ptr_key > 0 assert ptr_value > ptr_key @@ -199,7 +190,6 @@ def test_build_page_table(): def test_layer_group_meta_serialization(): import numpy as np - from tensorrt_llm._torch.disaggregation.base.region import DataRole from tensorrt_llm._torch.disaggregation.resource.page import ( BUFFER_ENTRY_DTYPE, AttentionLayerGroup, @@ -211,7 +201,7 @@ def test_layer_group_meta_serialization(): ) entries = np.array( - [(0, int(DataRole.KEY), 0, 256), (0, int(DataRole.VALUE), 256, 256)], + [(0, 0, 256), (0, 256, 256)], dtype=BUFFER_ENTRY_DTYPE, ) kv_pool = PhysicalPool(base_address=1000, slot_bytes=512, num_slots=10) @@ -278,7 +268,6 @@ def test_mamba_layer_group_serialization(): def test_mixed_page_table_serialization(): import numpy as np - from tensorrt_llm._torch.disaggregation.base.region import DataRole from tensorrt_llm._torch.disaggregation.resource.page import ( BUFFER_ENTRY_DTYPE, AttentionLayerGroup, @@ -292,7 +281,7 @@ def test_mixed_page_table_serialization(): # Attention layer group entries = np.array( - [(0, int(DataRole.KEY), 0, 256), (0, int(DataRole.VALUE), 256, 256)], + [(0, 0, 256), (0, 256, 256)], dtype=BUFFER_ENTRY_DTYPE, ) attn_lg = AttentionLayerGroup( diff --git a/tests/unittest/disaggregated/test_kv_transfer.py b/tests/unittest/disaggregated/test_kv_transfer.py index 130d074dd9fe..2f4d697f6364 100644 --- a/tests/unittest/disaggregated/test_kv_transfer.py +++ b/tests/unittest/disaggregated/test_kv_transfer.py @@ -70,6 +70,10 @@ class KvCacheConfigV2: copy_on_partial_reuse: bool = False dtype: str = "auto" disk_prefetch_num_reqs: int = 4 + pool_ratio: Optional[List[float]] = None + avg_seq_len: Optional[int] = None + block_reuse_policy: str = "all_reusable" + enable_swa_scratch_reuse: bool = False # V2 specific field max_util_for_resume: float = 0.95 diff --git a/tests/unittest/disaggregated/test_peer.py b/tests/unittest/disaggregated/test_peer.py index da595cfd093c..d2329d74e1aa 100644 --- a/tests/unittest/disaggregated/test_peer.py +++ b/tests/unittest/disaggregated/test_peer.py @@ -1,7 +1,6 @@ import numpy as np import pytest -from tensorrt_llm._torch.disaggregation.base.region import DataRole from tensorrt_llm._torch.disaggregation.native.mixers.attention.peer import ( HeadMatchMapper, HeadMismatchMapper, @@ -32,13 +31,13 @@ def make_page_table(pool_ptrs=None, block_bytes=None, global_layer_ids=None): if global_layer_ids is None: global_layer_ids = [0, 1] - # Build buffer entries: KEY + VALUE per local layer + # Build buffer entries: K + V per local layer buffer_size = 256 # bytes per buffer entry (arbitrary for tests) entries = [] for i in range(len(global_layer_ids)): base_offset = i * buffer_size * 2 - entries.append((i, int(DataRole.KEY), base_offset, buffer_size)) - entries.append((i, int(DataRole.VALUE), base_offset + buffer_size, buffer_size)) + entries.append((i, base_offset, buffer_size)) + entries.append((i, base_offset + buffer_size, buffer_size)) buffer_entries = np.array(entries, dtype=BUFFER_ENTRY_DTYPE) local_layers = [ diff --git a/tests/unittest/disaggregated/test_pool_matching.py b/tests/unittest/disaggregated/test_pool_matching.py new file mode 100644 index 000000000000..884852c3a539 --- /dev/null +++ b/tests/unittest/disaggregated/test_pool_matching.py @@ -0,0 +1,350 @@ +"""Golden tests for ``PeerRegistrar.get_pool_mapping``. + +Locks the current (pre-refactor) behavior of disagg pool matching across the +representative scenarios that the refactor must keep working: + + * single KV pool, full layer overlap (basic MHA) + * MLA (kv_factor=1, KEY-only pool) + * KV + block-scale pools coexisting in one LG + * KV + FLAT pools coexisting in one LG (FLAT has empty buffer_entries) + * PP partial layer overlap (Step-1 LG match by global_layer_id) + * Two pools with same role but different layer sets within a peer LG + (DSv4 virtual-layer scenario, exercises ``best_overlap``) + +The refactor (PoolRole -> PoolMatchKey + TransferLayout) must keep these +results stable. If a test needs updating, the change should be deliberate and +called out in the refactor commit message. +""" + +import numpy as np +import pytest + +from tensorrt_llm._torch.disaggregation.native.mixers.attention.spec import AttentionInfo +from tensorrt_llm._torch.disaggregation.native.peer import PeerRegistrar +from tensorrt_llm._torch.disaggregation.native.rank_info import RankInfo +from tensorrt_llm._torch.disaggregation.resource.kv_extractor import KVRegionExtractorV1 +from tensorrt_llm._torch.disaggregation.resource.page import ( + BUFFER_ENTRY_DTYPE, + AttentionLayerGroup, + KVCachePageTable, + LocalLayer, + MapperKind, + PhysicalPool, + PhysicalPoolGroup, + PoolView, +) + +# --------------------------------------------------------------------------- +# Builders +# --------------------------------------------------------------------------- + + +def _entries(layer_role_pairs, per_buf_size=128): + """Build a structured numpy array of buffer entries. + + ``layer_role_pairs`` is a list of (local_layer_id, role_label_str). The + role label only feeds ``pool_role`` — buffer_entries itself only carries + (lid, offset, size). + """ + rows = [ + (lid, i * per_buf_size, per_buf_size) for i, (lid, _role) in enumerate(layer_role_pairs) + ] + if not rows: + return np.array([], dtype=BUFFER_ENTRY_DTYPE) + return np.array(rows, dtype=BUFFER_ENTRY_DTYPE) + + +def _pool_view(pool_idx, layer_role_pairs, *, pool_role=None, mapper_kind=MapperKind.INDEXED): + """A PoolView with buffer entries. + + ``pool_role`` defaults to the set of distinct role labels appearing in + ``layer_role_pairs`` (matching the V1 builder convention). + """ + if pool_role is None: + pool_role = frozenset(role for _, role in layer_role_pairs) + return PoolView( + pool_idx=pool_idx, + buffer_entries=_entries(layer_role_pairs), + pool_role=pool_role, + mapper_kind=mapper_kind, + ) + + +def _empty_pool_view(pool_idx, *, pool_role=frozenset({"indexer_k"})): + """Pool view with no buffer entries — FLAT pool convention.""" + return PoolView( + pool_idx=pool_idx, + buffer_entries=np.array([], dtype=BUFFER_ENTRY_DTYPE), + pool_role=pool_role, + mapper_kind=MapperKind.FLAT, + ) + + +def _attn_lg(pool_group_idx, local_global_pairs, pool_views, sliding_window=None): + return AttentionLayerGroup( + pool_group_idx=pool_group_idx, + kv_head_num_per_rank=2, + sliding_window_size=sliding_window, + local_layers=[ + LocalLayer(local_layer_id=lid, global_layer_id=gid) for lid, gid in local_global_pairs + ], + pool_views=pool_views, + ) + + +def _page_table(layer_groups, pool_specs=None): + """Build a KVCachePageTable. + + pool_specs: ``{pool_group_idx: [(slot_bytes, num_slots, base_addr), ...]}``. + Defaults to one 1024-byte/64-slot pool per group. + """ + num_pgs = max((lg.pool_group_idx for lg in layer_groups), default=-1) + 1 + if pool_specs is None: + pool_specs = {pg: [(1024, 64, 0x1000 * (pg + 1))] for pg in range(num_pgs)} + + pool_groups = [ + PhysicalPoolGroup( + pools=[ + PhysicalPool(base_address=base, slot_bytes=sb, num_slots=ns) + for sb, ns, base in pool_specs.get(pg, [(1024, 64, 0x1000 * (pg + 1))]) + ] + ) + for pg in range(num_pgs) + ] + + return KVCachePageTable( + tokens_per_block=16, + layer_groups=layer_groups, + pool_groups=pool_groups, + ) + + +def _rank_info( + name="self", + rank=0, + layer_num_per_pp=None, + page_table=None, + is_mla=False, + kv_heads=2, +): + if layer_num_per_pp is None: + layer_num_per_pp = [2] + return RankInfo( + instance_name=name, + instance_rank=rank, + tp_size=1, + tp_rank=0, + pp_size=len(layer_num_per_pp), + pp_rank=0, + dp_size=1, + dp_rank=0, + cp_size=1, + cp_rank=0, + device_id=0, + layer_num_per_pp=layer_num_per_pp, + server_endpoint="", + self_endpoint="", + transfer_engine_info=b"", + attention=AttentionInfo( + kv_heads_per_rank=kv_heads, + tokens_per_block=16, + dims_per_head=8, + element_bytes=2, + enable_attention_dp=False, + is_mla=is_mla, + ), + aux_meta=None, + page_table=page_table, + sender_endpoints=[], + ) + + +def _registrar(self_pt, **kwargs): + self_ri = _rank_info(name="self", page_table=self_pt, **kwargs) + return PeerRegistrar(self_ri, KVRegionExtractorV1(self_pt)) + + +def _kv_pool_view(pool_idx, local_layer_ids): + """Convenience: KV pool view (key+value for each given local layer).""" + return _pool_view( + pool_idx, + [(lid, role) for lid in local_layer_ids for role in ("key", "value")], + ) + + +def _key_only_pool_view(pool_idx, local_layer_ids): + """Convenience: KEY-only pool view (MLA).""" + return _pool_view(pool_idx, [(lid, "key") for lid in local_layer_ids]) + + +# --------------------------------------------------------------------------- +# Test cases +# --------------------------------------------------------------------------- + + +def test_kv_only_full_overlap(): + """Single LG, single KV pool on each side. Identity mapping.""" + self_lg = _attn_lg(0, [(0, 10), (1, 11)], [_kv_pool_view(0, [0, 1])]) + peer_lg = _attn_lg(0, [(0, 10), (1, 11)], [_kv_pool_view(0, [0, 1])]) + + reg = _registrar(_page_table([self_lg])) + peer_ri = _rank_info(name="peer", rank=1, page_table=_page_table([peer_lg])) + + mapping = reg.get_pool_mapping(peer_ri) + assert mapping == {(0, 0): (0, 0)} + + +def test_mla_kv_only_full_overlap(): + """MLA: KEY-only pool, kv_factor=1.""" + self_lg = _attn_lg(0, [(0, 0), (1, 1)], [_key_only_pool_view(0, [0, 1])]) + peer_lg = _attn_lg(0, [(0, 0), (1, 1)], [_key_only_pool_view(0, [0, 1])]) + + reg = _registrar(_page_table([self_lg]), is_mla=True) + peer_ri = _rank_info(name="peer", rank=1, page_table=_page_table([peer_lg]), is_mla=True) + + mapping = reg.get_pool_mapping(peer_ri) + assert mapping == {(0, 0): (0, 0)} + + +def test_kv_and_block_scale_in_same_lg(): + """KV pool (idx=0) + block-scale pool (idx=1) match by role within the LG.""" + + def _bq_pool(pool_idx, lids): + return _pool_view( + pool_idx, + [(lid, role) for lid in lids for role in ("key_block_scale", "value_block_scale")], + ) + + self_lg = _attn_lg(0, [(0, 100), (1, 101)], [_kv_pool_view(0, [0, 1]), _bq_pool(1, [0, 1])]) + peer_lg = _attn_lg(0, [(0, 100), (1, 101)], [_kv_pool_view(0, [0, 1]), _bq_pool(1, [0, 1])]) + + self_pt = _page_table([self_lg], pool_specs={0: [(1024, 64, 0x1000), (256, 64, 0x2000)]}) + peer_pt = _page_table([peer_lg], pool_specs={0: [(1024, 64, 0x3000), (256, 64, 0x4000)]}) + + reg = _registrar(self_pt) + peer_ri = _rank_info(name="peer", rank=1, page_table=peer_pt) + + mapping = reg.get_pool_mapping(peer_ri) + assert mapping == {(0, 0): (0, 0), (0, 1): (0, 1)} + + +def test_kv_and_indexer_in_same_lg(): + """KV pool + FLAT pool. FLAT has empty buffer_entries; matches by role.""" + self_lg = _attn_lg( + 0, + [(0, 0), (1, 1)], + [_kv_pool_view(0, [0, 1]), _empty_pool_view(1)], + ) + peer_lg = _attn_lg( + 0, + [(0, 0), (1, 1)], + [_kv_pool_view(0, [0, 1]), _empty_pool_view(1)], + ) + + self_pt = _page_table([self_lg], pool_specs={0: [(1024, 64, 0x1000), (512, 64, 0x2000)]}) + peer_pt = _page_table([peer_lg], pool_specs={0: [(1024, 64, 0x3000), (512, 64, 0x4000)]}) + + reg = _registrar(self_pt) + peer_ri = _rank_info(name="peer", rank=1, page_table=peer_pt) + + mapping = reg.get_pool_mapping(peer_ri) + assert mapping == {(0, 0): (0, 0), (0, 1): (0, 1)} + + +def test_pp_partial_layer_overlap(): + """Self covers global layers {10,11}, peer covers {11,12}. Match via overlap on layer 11.""" + self_lg = _attn_lg(0, [(0, 10), (1, 11)], [_kv_pool_view(0, [0, 1])]) + peer_lg = _attn_lg(0, [(0, 11), (1, 12)], [_kv_pool_view(0, [0, 1])]) + + reg = _registrar(_page_table([self_lg]), layer_num_per_pp=[1, 1]) + peer_ri = _rank_info( + name="peer", rank=1, page_table=_page_table([peer_lg]), layer_num_per_pp=[1, 1] + ) + + mapping = reg.get_pool_mapping(peer_ri) + assert mapping == {(0, 0): (0, 0)} + + +def test_two_pools_distinct_roles_in_same_lg(): + """Two pools with distinct pool_role in one LG match by role, not by layer overlap. + + Within one LG, pool_role normally identifies a pool uniquely (builder + invariant). Even when the pools' layer sets are different sizes / not + equal, the matching is purely role-based: each self pool finds the peer + pool with the same pool_role. + """ + self_kv = _kv_pool_view(0, [0, 1]) + self_indexer = _empty_pool_view(1) + self_lg = _attn_lg(0, [(0, 10), (1, 11)], [self_kv, self_indexer]) + + peer_kv = _kv_pool_view(0, [0, 1]) + peer_indexer = _empty_pool_view(1) + peer_lg = _attn_lg(0, [(0, 10), (1, 11)], [peer_kv, peer_indexer]) + + self_pt = _page_table([self_lg], pool_specs={0: [(1024, 64, 0x1000), (512, 64, 0x2000)]}) + peer_pt = _page_table([peer_lg], pool_specs={0: [(1024, 64, 0x3000), (512, 64, 0x4000)]}) + + reg = _registrar(self_pt) + peer_ri = _rank_info(name="peer", rank=1, page_table=peer_pt) + + mapping = reg.get_pool_mapping(peer_ri) + assert mapping == {(0, 0): (0, 0), (0, 1): (0, 1)} + + +def test_same_role_pools_disambiguated_by_layer_overlap(): + """Defensive: same pool_role + different layer sets is resolved by overlap. + + If a peer LG holds two pools with the same pool_role but different layer + sets, matching picks the one whose layer set actually overlaps self. A + peer pool with zero overlap is never matched even if the role matches. + """ + self_lg = _attn_lg(0, [(0, 10), (1, 11)], [_kv_pool_view(0, [0, 1])]) + peer_lg = _attn_lg( + 0, + [(0, 10), (1, 11), (2, 12)], + [ + # First peer pool covers layers {12} only — same role, zero overlap with self. + _kv_pool_view(0, [2]), + # Second peer pool covers layers {10, 11} — same role, full overlap. + _kv_pool_view(1, [0, 1]), + ], + ) + + self_pt = _page_table([self_lg]) + peer_pt = _page_table([peer_lg], pool_specs={0: [(1024, 64, 0x3000), (1024, 64, 0x4000)]}) + + reg = _registrar(self_pt) + peer_ri = _rank_info(name="peer", rank=1, page_table=peer_pt, layer_num_per_pp=[3]) + + mapping = reg.get_pool_mapping(peer_ri) + assert mapping == {(0, 0): (0, 1)} + + +@pytest.mark.parametrize( + ("self_mapper_kind", "peer_mapper_kind"), + [ + (MapperKind.FLAT, MapperKind.INDEXED), + (MapperKind.INDEXED, MapperKind.FLAT), + ], +) +def test_mixed_mapper_kinds_are_rejected(self_mapper_kind, peer_mapper_kind): + """Pool layouts must use the same mapper kind on both peers.""" + + def _indexer_view(mapper_kind): + if mapper_kind == MapperKind.FLAT: + return _empty_pool_view(0) + return _pool_view( + 0, + [(0, "indexer_k"), (1, "indexer_k")], + pool_role=frozenset({"indexer_k"}), + ) + + self_lg = _attn_lg(0, [(0, 10), (1, 11)], [_indexer_view(self_mapper_kind)]) + peer_lg = _attn_lg(0, [(0, 10), (1, 11)], [_indexer_view(peer_mapper_kind)]) + reg = _registrar(_page_table([self_lg])) + peer_ri = _rank_info(name="peer", rank=1, page_table=_page_table([peer_lg])) + + with pytest.raises(ValueError, match="incompatible mapper kinds"): + reg.get_pool_mapping(peer_ri) + with pytest.raises(ValueError, match="incompatible mapper kinds"): + reg.get_kv_map(peer_ri, (0, 0), (0, 0)) diff --git a/tests/unittest/executor/test_stats_serializer.py b/tests/unittest/executor/test_stats_serializer.py index 07402ae0cde4..bb061e4ec80d 100644 --- a/tests/unittest/executor/test_stats_serializer.py +++ b/tests/unittest/executor/test_stats_serializer.py @@ -49,6 +49,17 @@ def _make_mock_kv_iter_stats( window_size=16, primary_used=10, primary_max=20, + primary_evictable=0, + primary_peak_free=None, + primary_peak_used=None, + primary_peak_evictable=None, + secondary_max=None, + secondary_free=0, + secondary_used=0, + secondary_evictable=0, + secondary_peak_free=None, + secondary_peak_used=None, + secondary_peak_evictable=None, reused=5, full_reused=4, partial_reused=1, @@ -56,13 +67,36 @@ def _make_mock_kv_iter_stats( gen_alloc=2, ): """Create a mock KvCacheIterationStats nanobind object.""" + primary_free = primary_max - primary_used + if primary_peak_free is None: + primary_peak_free = primary_free + if primary_peak_used is None: + primary_peak_used = primary_used + if primary_peak_evictable is None: + primary_peak_evictable = primary_evictable + if secondary_max is None: + secondary_max = secondary_free + secondary_used + if secondary_peak_free is None: + secondary_peak_free = secondary_free + if secondary_peak_used is None: + secondary_peak_used = secondary_used + if secondary_peak_evictable is None: + secondary_peak_evictable = secondary_evictable s = SimpleNamespace( primary_max_num_blocks=primary_max, - primary_free_num_blocks=primary_max - primary_used, + primary_free_num_blocks=primary_free, primary_used_num_blocks=primary_used, - secondary_max_num_blocks=0, - secondary_free_num_blocks=0, - secondary_used_num_blocks=0, + primary_evictable_num_blocks=primary_evictable, + primary_peak_free_num_blocks=primary_peak_free, + primary_peak_used_num_blocks=primary_peak_used, + primary_peak_evictable_num_blocks=primary_peak_evictable, + secondary_max_num_blocks=secondary_max, + secondary_free_num_blocks=secondary_free, + secondary_used_num_blocks=secondary_used, + secondary_evictable_num_blocks=secondary_evictable, + secondary_peak_free_num_blocks=secondary_peak_free, + secondary_peak_used_num_blocks=secondary_peak_used, + secondary_peak_evictable_num_blocks=secondary_peak_evictable, iter_alloc_total_blocks=reused + missed, iter_alloc_new_blocks=missed, iter_reused_blocks=reused, @@ -77,10 +111,42 @@ def _make_mock_kv_iter_stats( iter_offload_bytes=0, iter_intra_device_copy_blocks=2, iter_intra_device_copy_bytes=8192, + iter_host_dropped_blocks=0, + iter_host_dropped_bytes=0, ) return {window_size: s} +class _FakeStorageStatistics(SimpleNamespace): + @property + def unavailable(self): + return self.total - self.available + + +class _FakePeakStorage: + num_pool_groups = 2 + num_cache_levels = 2 + + def __init__(self): + self._levels = [] + self.primary_stats = [ + _FakeStorageStatistics(total=10, available=8, evictable=1), + _FakeStorageStatistics(total=10, available=9, evictable=0), + ] + self.secondary_stats = [ + _FakeStorageStatistics(total=5, available=4, evictable=1), + _FakeStorageStatistics(total=5, available=5, evictable=0), + ] + + def get_statistics(self, level): + if int(level) == 0: + return self.primary_stats + return self.secondary_stats + + def destroy(self): + pass + + class TestStatsSerializer: def test_serializer_without_kv_iter_stats(self): """Legacy 2-tuple and 3-tuple with None should produce same output.""" @@ -123,6 +189,12 @@ def test_serializer_with_kv_iter_stats(self): assert ws_stats["primaryMaxNumBlocks"] == 20 assert ws_stats["primaryUsedNumBlocks"] == 10 assert ws_stats["primaryFreeNumBlocks"] == 10 + assert ws_stats["primaryPeakFreeNumBlocks"] == 10 + assert ws_stats["primaryPeakUsedNumBlocks"] == 10 + assert ws_stats["primaryPeakEvictableNumBlocks"] == 0 + assert ws_stats["secondaryPeakFreeNumBlocks"] == 0 + assert ws_stats["secondaryPeakUsedNumBlocks"] == 0 + assert ws_stats["secondaryPeakEvictableNumBlocks"] == 0 assert ws_stats["iterReusedBlocks"] == 5 assert ws_stats["iterFullReusedBlocks"] == 4 assert ws_stats["iterPartialReusedBlocks"] == 1 @@ -250,6 +322,17 @@ def test_serializer_with_v2_pool_group_stats(self): window_size=16, primary_used=10, primary_max=20, + primary_evictable=2, + primary_peak_free=12, + primary_peak_used=15, + primary_peak_evictable=4, + secondary_max=8, + secondary_free=5, + secondary_used=3, + secondary_evictable=1, + secondary_peak_free=6, + secondary_peak_used=4, + secondary_peak_evictable=2, reused=5, full_reused=4, partial_reused=1, @@ -260,6 +343,17 @@ def test_serializer_with_v2_pool_group_stats(self): window_size=16, primary_used=10, primary_max=20, + primary_evictable=2, + primary_peak_free=12, + primary_peak_used=15, + primary_peak_evictable=4, + secondary_max=8, + secondary_free=5, + secondary_used=3, + secondary_evictable=1, + secondary_peak_free=6, + secondary_peak_used=4, + secondary_peak_evictable=2, reused=0, full_reused=0, partial_reused=0, @@ -301,11 +395,23 @@ def test_serializer_with_v2_pool_group_stats(self): d = json.loads(result) assert d["kvCacheIterationStats"]["16"]["iterReusedBlocks"] == 5 + assert d["kvCacheIterationStats"]["16"]["primaryPeakFreeNumBlocks"] == 12 + assert d["kvCacheIterationStats"]["16"]["primaryPeakUsedNumBlocks"] == 15 + assert d["kvCacheIterationStats"]["16"]["primaryPeakEvictableNumBlocks"] == 4 + assert d["kvCacheIterationStats"]["16"]["secondaryPeakFreeNumBlocks"] == 6 + assert d["kvCacheIterationStats"]["16"]["secondaryPeakUsedNumBlocks"] == 4 + assert d["kvCacheIterationStats"]["16"]["secondaryPeakEvictableNumBlocks"] == 2 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["primaryPeakFreeNumBlocks"] == 12 + assert pool_group["primaryPeakUsedNumBlocks"] == 15 + assert pool_group["primaryPeakEvictableNumBlocks"] == 4 + assert pool_group["secondaryPeakFreeNumBlocks"] == 6 + assert pool_group["secondaryPeakUsedNumBlocks"] == 4 + assert pool_group["secondaryPeakEvictableNumBlocks"] == 2 assert pool_group["iterGenAllocBlocks"] == 2 assert "iterReusedBlocks" not in pool_group assert "iterMissedBlocks" not in pool_group @@ -319,3 +425,46 @@ def test_serializer_with_v2_pool_group_stats(self): assert life_cycle["iterReusedBlocks"] == 5 assert life_cycle["iterMissedBlocks"] == 3 assert "iterGenAllocBlocks" not in life_cycle + + def test_v2_peak_block_stats_reset_tracks_interval_peak(self): + """Peak block stats should cover the interval since the previous reset.""" + from tensorrt_llm.runtime.kv_cache_manager_v2._common import GPU_LEVEL, CacheLevel + from tensorrt_llm.runtime.kv_cache_manager_v2._core._kv_cache_manager import KVCacheManager + + storage = _FakePeakStorage() + manager = object.__new__(KVCacheManager) + manager._storage = storage + manager._radix_tree = SimpleNamespace(clear=lambda: []) + manager._reset_iteration_peak_num_blocks() + + # Some gauges rise above the reset baseline, then fall before drain. + storage.primary_stats[0].available = 5 # primary used = 5 + storage.primary_stats[0].evictable = 3 + storage.primary_stats[1].available = 6 # primary used = 4 + storage.primary_stats[1].evictable = 4 + storage.secondary_stats[0].available = 2 # secondary used = 3 + storage.secondary_stats[0].evictable = 2 + manager._update_iteration_peak_num_blocks() + storage.primary_stats[0].available = 7 # primary used = 3 + storage.primary_stats[0].evictable = 1 + storage.secondary_stats[0].available = 4 # secondary used = 1 + storage.secondary_stats[0].evictable = 1 + + primary_peak = manager.get_and_reset_iteration_peak_block_stats(GPU_LEVEL) + secondary_peak = manager.get_and_reset_iteration_peak_block_stats(CacheLevel(1)) + assert [stats.available for stats in primary_peak] == [8, 9] + assert [stats.unavailable for stats in primary_peak] == [5, 4] + assert [stats.evictable for stats in primary_peak] == [3, 4] + assert [stats.available for stats in secondary_peak] == [4, 5] + assert [stats.unavailable for stats in secondary_peak] == [3, 0] + assert [stats.evictable for stats in secondary_peak] == [2, 0] + + # The next interval starts from current usage, not zero. + primary_peak = manager.get_and_reset_iteration_peak_block_stats(GPU_LEVEL) + secondary_peak = manager.get_and_reset_iteration_peak_block_stats(CacheLevel(1)) + assert [stats.available for stats in primary_peak] == [7, 6] + assert [stats.unavailable for stats in primary_peak] == [3, 4] + assert [stats.evictable for stats in primary_peak] == [1, 4] + assert [stats.available for stats in secondary_peak] == [4, 5] + assert [stats.unavailable for stats in secondary_peak] == [1, 0] + assert [stats.evictable for stats in secondary_peak] == [1, 0] diff --git a/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_manager_v2.py b/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_manager_v2.py index 583fa5da72c9..9fc966dcdb2d 100755 --- a/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_manager_v2.py +++ b/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_manager_v2.py @@ -65,8 +65,10 @@ ) from kv_cache_manager_v2._copy_engine import CopyTask, batched_copy from kv_cache_manager_v2._exceptions import OutOfPagesError - from kv_cache_manager_v2._life_cycle_registry import SsmLifeCycle + from kv_cache_manager_v2._life_cycle_registry import LifeCycleRegistry, SsmLifeCycle + from kv_cache_manager_v2._storage._config import create_storage_config from kv_cache_manager_v2._storage._core import CacheLevelStorage, SlotAllocator + from kv_cache_manager_v2._storage_manager import StorageManager from kv_cache_manager_v2._utils import ( CachedCudaStream, HalfOpenRange, @@ -121,11 +123,16 @@ ) from tensorrt_llm.runtime.kv_cache_manager_v2._copy_engine import CopyTask, batched_copy from tensorrt_llm.runtime.kv_cache_manager_v2._exceptions import OutOfPagesError - from tensorrt_llm.runtime.kv_cache_manager_v2._life_cycle_registry import SsmLifeCycle + from tensorrt_llm.runtime.kv_cache_manager_v2._life_cycle_registry import ( + LifeCycleRegistry, + SsmLifeCycle, + ) + from tensorrt_llm.runtime.kv_cache_manager_v2._storage._config import create_storage_config from tensorrt_llm.runtime.kv_cache_manager_v2._storage._core import ( CacheLevelStorage, SlotAllocator, ) + from tensorrt_llm.runtime.kv_cache_manager_v2._storage_manager import StorageManager from tensorrt_llm.runtime.kv_cache_manager_v2._utils import ( CachedCudaStream, HalfOpenRange, @@ -885,6 +892,11 @@ def __getitem__(self, key: tuple[int, int]) -> Node: def __iter__(self) -> Iterator[Node]: return iter(self._nodes) + def shutdown(self) -> None: + for node in reversed(self._nodes): + node.kv_cache.close() + node.manager.shutdown() + def __init__( self, full_config: KVCacheManagerConfig, num_heads: int, tp_size: int, pp_size: int ): @@ -928,6 +940,12 @@ def setUp(self) -> None: def tearDown(self) -> None: gc.enable() + if hasattr(self, "decode"): + self.decode.shutdown() + del self.decode + if hasattr(self, "prefill"): + self.prefill.shutdown() + del self.prefill def next_token(self) -> TokenIdExt: token_id = next(self._token_id_gen) @@ -1708,6 +1726,7 @@ def _make_config( num_windowed_layers: int = 1, num_full_layers: int = 1, enable_swa_scratch_reuse: bool = False, + initial_pool_ratio: list[float] | None = None, ) -> KVCacheManagerConfig: """Create a config with two pool groups (windowed vs non-windowed). @@ -1748,6 +1767,7 @@ def _make_config( layers=layers, typical_step=typical_step, constraints=constraints or [], + initial_pool_ratio=initial_pool_ratio, swa_scratch_reuse=(SwaScratchReuseConfig() if enable_swa_scratch_reuse else None), ) @@ -1805,6 +1825,38 @@ def test_constraints_floor_typical_step(self): mgr_unconstrained.shutdown() mgr_constrained.shutdown() + def test_initial_pool_ratio_overrides_typical_step_and_constraints(self): + """Explicit initial_pool_ratio takes precedence over inferred sizing inputs.""" + typical = BatchDesc(kv_caches=[KVCacheDesc(capacity=4096, history_length=4000)] * 32) + constraint = BatchDesc(kv_caches=[KVCacheDesc(capacity=256, history_length=128)] * 256) + cfg = self._make_config( + typical_step=typical, + constraints=[constraint], + initial_pool_ratio=[0.8, 0.2], + ) + manager = KVCacheManager(cfg) + ratio = manager._current_gpu_ratio + + self.assertGreater(ratio[0], ratio[1]) + self.assertAlmostEqual(sum(ratio), 1.0, places=6) + manager.shutdown() + + def test_initial_pool_ratio_length_must_match_pool_groups(self): + cfg = self._make_config(initial_pool_ratio=[1.0]) + life_cycles = LifeCycleRegistry(cfg) + storage_config = create_storage_config(cfg) + + with self.assertRaisesRegex(ValueError, "initial_pool_ratio length"): + StorageManager( + life_cycles, + storage_config, + cfg.tokens_per_block, + cfg.swa_scratch_reuse, + typical_batch=cfg.typical_step, + constraints=cfg.constraints, + initial_pool_ratio=cfg.initial_pool_ratio, + ) + @parameterized.expand([(0,), (64,), (50,), (256,)]) def test_constraint_guarantees_batch_can_run(self, system_prompt_length: int): """Quota is tight; without constraint clamping the batch would fail. 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 index aa20cb4f56e7..d9e26516d04f 100644 --- 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 @@ -19,7 +19,7 @@ import torch from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import KVCacheManagerV2 -from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest +from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest, LlmRequestState from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager as KVCacheManagerV1 from tensorrt_llm._torch.pyexecutor.scheduler import ScheduledRequests from tensorrt_llm.bindings import DataType, SamplingConfig @@ -51,10 +51,13 @@ class _StatsRequest: is_dummy_request: bool = False is_attention_dp_dummy: bool = False is_cuda_graph_dummy: bool = False + is_disagg_generation_init_state: bool = False is_disagg_generation_transmission_complete: bool = False + is_finished_due_to_cancellation: bool = False context_phase_params: None = None py_draft_tokens: list[int] = field(default_factory=list) draft_tokens: list[int] = field(default_factory=list) + state: LlmRequestState = LlmRequestState.GENERATION_IN_PROGRESS context_current_position: int = 0 context_chunk_size: int = 0 prepopulated_prompt: tuple[int, int] | None = None @@ -111,6 +114,7 @@ def _create_manager( num_layers: int = 1, max_attention_window: list[int] | None = None, enable_block_reuse: bool = True, + block_reuse_policy: str = "all_reusable", enable_stats: bool = True, ) -> KVCacheManagerV2: return KVCacheManagerV2( @@ -120,6 +124,7 @@ def _create_manager( max_gpu_total_bytes=gpu_bytes, max_util_for_resume=1.0, max_attention_window=max_attention_window, + block_reuse_policy=block_reuse_policy, ), CacheType.SELF, num_layers=num_layers, @@ -403,6 +408,72 @@ def test_reverted_context_allocation_does_not_report_pending_stats(resource_guar ] +def test_reverted_disagg_gen_init_allocation_does_not_report_pending_stats( + resource_guard, +) -> None: + request = _StatsRequest( + 1, + list(range(8)), + context_remaining_length=8, + is_disagg_generation_init_state=True, + ) + manager = resource_guard(_create_manager(gpu_bytes=8 << 20), request) + + assert manager.prepare_disagg_gen_init(request) + manager.revert_allocate_context(request) + manager.commit_scheduled_kv_cache_stats(_context_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]) + assert request.kv_cache_perf_metric_calls == [] + kv_stats = manager.get_kv_cache_stats() + assert kv_stats.alloc_total_blocks == 0 + assert kv_stats.alloc_new_blocks == 0 + assert kv_stats.missed_blocks == 0 + + +def test_waited_context_allocation_reports_pending_stats_when_scheduled(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) + + 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), + ] + kv_stats = manager.get_kv_cache_stats() + assert kv_stats.alloc_total_blocks == 2 + assert kv_stats.alloc_new_blocks == 2 + assert kv_stats.missed_blocks == 2 + + 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) + 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) @@ -439,6 +510,50 @@ def test_chunked_context_reports_generation_alloc_only_in_generation(resource_gu ] +def test_all_reusable_policy_commits_each_context_chunk(resource_guard) -> None: + request = _StatsRequest(1, list(range(8)), context_remaining_length=8) + manager = resource_guard( + _create_manager(gpu_bytes=8 << 20, block_reuse_policy="all_reusable"), 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)) + + kv_cache = manager.kv_cache_map[request.py_request_id] + assert kv_cache.num_committed_tokens == 4 + assert kv_cache.history_length == 4 + + +def test_per_request_policy_delays_commit_until_last_context_chunk(resource_guard) -> None: + request = _StatsRequest(1, list(range(8)), context_remaining_length=8) + manager = resource_guard( + _create_manager(gpu_bytes=8 << 20, block_reuse_policy="per_request"), 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)) + + kv_cache = manager.kv_cache_map[request.py_request_id] + assert kv_cache.num_committed_tokens == 0 + assert kv_cache.history_length == 4 + + request.is_first_context_chunk = False + assert manager.prepare_context(request) + assert manager.resize_context(request, num_tokens=4) + request.context_current_position = 8 + request.context_remaining_length = 0 + manager.update_context_resources(_context_batch(request)) + + assert kv_cache.num_committed_tokens == 8 + assert kv_cache.history_length == 8 + + 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))) diff --git a/tests/unittest/llmapi/apps/test_disagg_perf_metrics_collector.py b/tests/unittest/llmapi/apps/test_disagg_perf_metrics_collector.py new file mode 100644 index 000000000000..70ab2432d98f --- /dev/null +++ b/tests/unittest/llmapi/apps/test_disagg_perf_metrics_collector.py @@ -0,0 +1,68 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import asyncio + +import pytest + +from tensorrt_llm.serve import perf_metrics + + +class _DummyMetric: + def inc(self): + pass + + def observe(self, _value): + pass + + +class _BlockingClient: + def __init__(self): + self.entered = asyncio.Event() + self.release = asyncio.Event() + self.active_collectors = 0 + self.max_active_collectors = 0 + self.calls = 0 + + async def collect_metrics(self): + self.calls += 1 + self.active_collectors += 1 + self.max_active_collectors = max(self.max_active_collectors, self.active_collectors) + self.entered.set() + await self.release.wait() + self.active_collectors -= 1 + return {} + + +@pytest.mark.asyncio +async def test_disagg_perf_metrics_collection_is_serialized(monkeypatch): + monkeypatch.setattr(perf_metrics, "instance_metric", lambda _definition: _DummyMetric()) + collector = perf_metrics.DisaggPerfMetricsCollector(max_requests=8) + client = _BlockingClient() + collector.add_client(client) + + first_task = asyncio.create_task(collector.get_perf_metrics()) + await client.entered.wait() + + second_task = asyncio.create_task(collector.get_perf_metrics()) + await asyncio.sleep(0) + + assert client.calls == 1 + assert client.max_active_collectors == 1 + + client.release.set() + assert await first_task == [] + assert await second_task == [] + assert client.calls == 2 + assert client.max_active_collectors == 1 diff --git a/tests/unittest/llmapi/apps/test_openai_server_iteration_stats.py b/tests/unittest/llmapi/apps/test_openai_server_iteration_stats.py new file mode 100644 index 000000000000..dd24a9b3df15 --- /dev/null +++ b/tests/unittest/llmapi/apps/test_openai_server_iteration_stats.py @@ -0,0 +1,127 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import json +from collections import deque +from types import SimpleNamespace + +import pytest + +from tensorrt_llm.serve.openai_server import OpenAIServer + + +class _FakeGenerator: + def __init__( + self, stat_batches: list[list[dict]], iter_stats_max_iterations: int | None = None + ): + self.args = SimpleNamespace(iter_stats_max_iterations=iter_stats_max_iterations) + self._stat_batches = deque(stat_batches) + self.stats_timeouts = [] + + def get_stats_async(self, timeout: float | None): + self.stats_timeouts.append(timeout) + + async def _stats_iter(): + if self._stat_batches: + for stat in self._stat_batches.popleft(): + yield stat + + return _stats_iter() + + +class _FakeMetricsCollector: + def __init__(self): + self.logged_stats = [] + + def log_iteration_stats(self, iteration_stats: dict) -> None: + self.logged_stats.append(iteration_stats) + + +def _make_server( + stat_batches: list[list[dict]], + *, + with_stats_buffer: bool = True, + is_visual_gen: bool = False, + iter_stats_max_iterations: int | None = None, +) -> OpenAIServer: + server = object.__new__(OpenAIServer) + server.generator = _FakeGenerator(stat_batches, iter_stats_max_iterations) + server.metrics_collector = _FakeMetricsCollector() + server._is_visual_gen = is_visual_gen + max_buffer_size = OpenAIServer._iteration_stats_buffer_maxlen( + server.generator.args.iter_stats_max_iterations + ) + server._iteration_stats_buffer = deque(maxlen=max_buffer_size) if with_stats_buffer else None + return server + + +def _response_content(response): + return json.loads(response.body.decode("utf-8")) + + +@pytest.mark.asyncio +async def test_metrics_endpoint_reads_background_buffer(): + stats = [{"iter": 1}, {"iter": 2}] + server = _make_server([]) + server._iteration_stats_buffer.extend(stats) + + response = await server.get_iteration_stats() + + assert _response_content(response) == stats + assert server.generator.stats_timeouts == [] + + response = await server.get_iteration_stats() + assert _response_content(response) == [] + + +@pytest.mark.asyncio +async def test_metrics_endpoint_reads_unbounded_background_buffer(): + stats = [{"iter": idx} for idx in range(3)] + server = _make_server([], iter_stats_max_iterations=-1) + assert server._iteration_stats_buffer.maxlen is None + server._iteration_stats_buffer.extend(stats) + + response = await server.get_iteration_stats() + + assert _response_content(response) == stats + + +@pytest.mark.asyncio +async def test_metrics_endpoint_reads_queue_without_background_buffer(): + stats = [{"iter": 5}, {"iter": 6}] + server = _make_server([stats], with_stats_buffer=False) + + response = await server.get_iteration_stats() + + assert _response_content(response) == stats + assert server.generator.stats_timeouts == [2] + assert server.metrics_collector.logged_stats == [] + + +@pytest.mark.asyncio +async def test_metrics_endpoint_reads_visual_gen_stats(): + stats = [ + { + "iter": 7, + "numQueuedRequests": 2, + "numActiveRequests": 1, + } + ] + server = _make_server([stats], is_visual_gen=True) + + response = await server.get_iteration_stats() + + assert _response_content(response) == stats + assert server.generator.stats_timeouts == [None] + assert server.metrics_collector.logged_stats == [] diff --git a/tests/unittest/llmapi/test_llm_args.py b/tests/unittest/llmapi/test_llm_args.py index 69bab9c13e00..6509cb4048e7 100644 --- a/tests/unittest/llmapi/test_llm_args.py +++ b/tests/unittest/llmapi/test_llm_args.py @@ -408,6 +408,8 @@ def get_model_defaults(cls, llm_args): def test_KvCacheConfig_declaration(): assert KvCacheConfig().kv_cache_event_hash_algo == "auto" + assert KvCacheConfig().block_reuse_policy == "all_reusable" + assert KvCacheConfig().enable_swa_scratch_reuse is False config = KvCacheConfig(enable_block_reuse=True, max_tokens=1024, @@ -420,8 +422,12 @@ def test_KvCacheConfig_declaration(): secondary_offload_min_priority=1, event_buffer_max_size=0, kv_cache_event_hash_algo="v2_sha256_64", + enable_swa_scratch_reuse=True, enable_partial_reuse=True, copy_on_partial_reuse=True, + pool_ratio=[0.25, 0.75], + avg_seq_len=2048, + block_reuse_policy="per_request", attention_dp_events_gather_period_ms=10) pybind_config = config._to_pybind() @@ -432,10 +438,19 @@ def test_KvCacheConfig_declaration(): assert pybind_config.host_cache_size == 1024 assert config.disk_cache_size == 2048 assert config.disk_cache_path == "/tmp" + assert config.enable_swa_scratch_reuse is True + assert KvCacheConfig().enable_swa_scratch_reuse is False assert pybind_config.cross_kv_cache_fraction == 0.5 assert pybind_config.secondary_offload_min_priority == 1 assert pybind_config.event_buffer_max_size == 0 assert config.kv_cache_event_hash_algo == "v2_sha256_64" + assert config.pool_ratio == [0.25, 0.75] + assert config.avg_seq_len == 2048 + assert config.block_reuse_policy == "per_request" + assert not hasattr(pybind_config, "pool_ratio") + assert not hasattr(pybind_config, "avg_seq_len") + assert not hasattr(pybind_config, "block_reuse_policy") + assert not hasattr(pybind_config, "enable_swa_scratch_reuse") assert KvCacheConfig( kv_cache_event_hash_algo="auto").kv_cache_event_hash_algo == "auto" assert KvCacheConfig(kv_cache_event_hash_algo="v1_block_key" @@ -443,6 +458,8 @@ def test_KvCacheConfig_declaration(): assert pybind_config.enable_partial_reuse == True assert pybind_config.copy_on_partial_reuse == True assert pybind_config.attention_dp_events_gather_period_ms == 10 + with pytest.raises(ValidationError): + KvCacheConfig(block_reuse_policy="invalid") def test_KvCacheConfig_disk_cache_validation(tmp_path): @@ -535,6 +552,25 @@ def test_torch_llm_args_with_encoder_cuda_graph_buckets_yaml(self): assert encoder_graph["vision"].buckets == [(1035, 1), (2069, 2)] +@pytest.mark.parametrize("kwargs", [ + { + "pool_ratio": [] + }, + { + "pool_ratio": [0.0, 1.0] + }, + { + "pool_ratio": [0.25, 0.25] + }, + { + "avg_seq_len": 0 + }, +]) +def test_KvCacheConfig_pool_ratio_avg_seq_len_validation(kwargs): + with pytest.raises(ValidationError): + KvCacheConfig(**kwargs) + + def test_CapacitySchedulerPolicy(): val = CapacitySchedulerPolicy.MAX_UTILIZATION assert PybindMirror.maybe_to_pybind(