diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/CMakeLists.txt b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/CMakeLists.txt index a2a737db9821..eb0bcb8fa8c1 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/CMakeLists.txt +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/CMakeLists.txt @@ -56,7 +56,9 @@ set(KV_CACHE_MANAGER_V2_SRCS kv_cache_manager_v2/blockRadixTree.cpp kv_cache_manager_v2/batchedPageCopy.cu kv_cache_manager_v2/coldPageCodec.cpp + kv_cache_manager_v2/eventData.cpp kv_cache_manager_v2/eventManager.cpp + kv_cache_manager_v2/streamingEventSink.cpp kv_cache_manager_v2/page.cpp kv_cache_manager_v2/storageManager.cpp kv_cache_manager_v2/introspection.cpp diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/eventData.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/eventData.cpp new file mode 100644 index 000000000000..ec95aa30ada6 --- /dev/null +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/eventData.cpp @@ -0,0 +1,90 @@ +/* + * 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. + */ + +#include "kv_cache_manager_v2/eventData.h" + +#include "kv_cache_manager_v2/blockRadixTree.h" + +#include +#include + +namespace tensorrt_llm::batch_manager::kv_cache_manager_v2 +{ + +std::string digestToHex(Digest const& digest) +{ + constexpr char kHex[] = "0123456789abcdef"; + std::string result(digest.size() * 2, '\0'); + auto out = result.begin(); + for (auto byte : digest) + { + auto const value = std::to_integer(byte); + *out++ = kHex[value >> 4U]; + *out++ = kHex[value & 0x0FU]; + } + return result; +} + +DecodedEventBlock decodeEventBlock(Block const& block, std::optional mmTokenIdOffset) +{ + DecodedEventBlock result; + result.tokenIds.reserve(block.tokens.size()); + + Digest const* itemDigest = nullptr; + if (mmTokenIdOffset.has_value() && block.prev != nullptr && block.prev->type() == NodeBase::Type::kBLOCK) + { + itemDigest = static_cast(block.prev)->getLastTokenDigest().get(); + } + bool inMmRun = false; + for (auto const& token : block.tokens) + { + if (token.isDigest()) + { + result.tokenIds.emplace_back(std::in_place_index<1>, digestToHex(token.digest())); + if (mmTokenIdOffset.has_value()) + { + itemDigest = &token.digest(); + result.mmKeys.push_back( + {std::string(reinterpret_cast(itemDigest->data()), itemDigest->size()), 0, + std::nullopt, false}); + inMmRun = true; + } + continue; + } + + auto const tokenId = token.tokenId(); + result.tokenIds.emplace_back(std::in_place_index<0>, tokenId); + if (itemDigest != nullptr && tokenId > *mmTokenIdOffset) + { + if (!inMmRun) + { + result.mmKeys.push_back( + {std::string(reinterpret_cast(itemDigest->data()), itemDigest->size()), + tokenId - *mmTokenIdOffset, std::nullopt, false}); + } + inMmRun = true; + } + else + { + // Text separates runs of the same item, so retain its digest for later continuations. + inMmRun = false; + } + } + return result; +} + +} // namespace tensorrt_llm::batch_manager::kv_cache_manager_v2 diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/eventData.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/eventData.h new file mode 100644 index 000000000000..e52017ba3927 --- /dev/null +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/eventData.h @@ -0,0 +1,61 @@ +/* + * 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. + */ + +#pragma once + +#include "kv_cache_manager_v2/common.h" + +#include +#include +#include +#include +#include + +namespace tensorrt_llm::batch_manager::kv_cache_manager_v2 +{ + +struct Block; + +using EventTokenId = std::variant; + +struct MmKey +{ + std::string hash; + int startOffset = 0; + std::optional uuid; + bool hasUuidField = false; + + bool operator==(MmKey const& other) const + { + return hash == other.hash && startOffset == other.startOffset && uuid == other.uuid + && hasUuidField == other.hasUuidField; + } +}; + +struct DecodedEventBlock +{ + std::vector tokenIds; + std::vector mmKeys; +}; + +[[nodiscard]] std::string digestToHex(Digest const& digest); + +//! Decode the V2 digest-first multimodal representation used by KV event consumers. +//! When mmTokenIdOffset is absent, tokens are still preserved but no MM segments are derived. +[[nodiscard]] DecodedEventBlock decodeEventBlock(Block const& block, std::optional mmTokenIdOffset); + +} // namespace tensorrt_llm::batch_manager::kv_cache_manager_v2 diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/eventManager.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/eventManager.cpp index 8ca96b4edde8..bd7067080b1d 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/eventManager.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/eventManager.cpp @@ -485,20 +485,6 @@ int EventManager::getWindowSize(EventLayerGroupId layerGroupId) const return windowSize == mWindowSizeByLayerGroup.end() ? mWindowSize : windowSize->second; } -std::string EventManager::digestToHex(Digest const& digest) -{ - constexpr char kHex[] = "0123456789abcdef"; - std::string result; - result.resize(digest.size() * 2); - for (size_t i = 0; i < digest.size(); ++i) - { - auto const value = std::to_integer(digest[i]); - result[2 * i] = kHex[value >> 4U]; - result[2 * i + 1] = kHex[value & 0x0FU]; - } - return result; -} - uint64_t EventManager::truncateDigestToInt64(Digest const& digest) { uint64_t result = 0; @@ -566,54 +552,15 @@ std::optional EventManager::storedBlockFromBlock( return std::nullopt; } - std::vector mmKeys; - Digest const* itemDigest = nullptr; - if (mMmTokenIdOffset.has_value() && block.prev != nullptr && block.prev->type() == NodeBase::Type::kBLOCK) - { - itemDigest = static_cast(block.prev)->getLastTokenDigest().get(); - } - bool inMmRun = false; + auto decoded = decodeEventBlock(block, mMmTokenIdOffset); std::vector tokens; - tokens.reserve(block.tokens.size()); - for (auto const& token : block.tokens) + tokens.reserve(decoded.tokenIds.size()); + for (auto& tokenId : decoded.tokenIds) { - if (!token.isDigest()) - { - UniqueToken uniqueToken; - uniqueToken.tokenId = EventTokenId{std::in_place_index<0>, token.tokenId()}; - tokens.push_back(std::move(uniqueToken)); - if (itemDigest != nullptr && token.tokenId() > *mMmTokenIdOffset) - { - if (!inMmRun) - { - mmKeys.push_back( - {std::string(reinterpret_cast(itemDigest->data()), itemDigest->size()), - token.tokenId() - *mMmTokenIdOffset, std::nullopt, false}); - } - inMmRun = true; - } - else - { - // Text separates runs of the same item, so retain its digest for later continuations. - inMmRun = false; - } - } - else - { - UniqueToken uniqueToken; - uniqueToken.tokenId = EventTokenId{std::in_place_index<1>, digestToHex(token.digest())}; - tokens.push_back(std::move(uniqueToken)); - if (mMmTokenIdOffset.has_value()) - { - itemDigest = &token.digest(); - mmKeys.push_back({std::string(reinterpret_cast(itemDigest->data()), itemDigest->size()), 0, - std::nullopt, false}); - inMmRun = true; - } - } + tokens.push_back(UniqueToken{std::move(tokenId)}); } return KVCacheStoredBlockData{ - hashFromBlock(block), std::move(tokens), cacheLevel.value(), priority, std::move(mmKeys), std::nullopt}; + hashFromBlock(block), std::move(tokens), cacheLevel.value(), priority, std::move(decoded.mmKeys), std::nullopt}; } uint64_t EventManager::hashV1BlockKey(std::vector const& tokens, uint64_t parentHash, diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/eventManager.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/eventManager.h index c1ca4ae83450..351120c5c47b 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/eventManager.h +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/eventManager.h @@ -18,6 +18,7 @@ #pragma once #include "kv_cache_manager_v2/common.h" +#include "kv_cache_manager_v2/eventData.h" #include "kv_cache_manager_v2/eventSink.h" #include @@ -40,7 +41,6 @@ namespace tensorrt_llm::batch_manager::kv_cache_manager_v2 { using EventBlockHash = std::variant; -using EventTokenId = std::variant; using EventLayerGroupId = std::optional; struct UniqueToken @@ -64,20 +64,6 @@ struct KVCacheCreatedData } }; -struct MmKey -{ - std::string hash; - int startOffset = 0; - std::optional uuid; - bool hasUuidField = false; - - bool operator==(MmKey const& other) const - { - return hash == other.hash && startOffset == other.startOffset && uuid == other.uuid - && hasUuidField == other.hasUuidField; - } -}; - struct KVCacheStoredBlockData { EventBlockHash blockHash; @@ -220,7 +206,6 @@ class EventManager final : public EventSink using V1RootAttrs = std::pair, std::optional>; static std::pair parseHashAlgorithm(std::string const& hashAlgo); - static std::string digestToHex(Digest const& digest); static uint64_t truncateDigestToInt64(Digest const& digest); static std::vector trimEvents(std::vector events, int maxKvEventEntries); diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/streamingEventSink.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/streamingEventSink.cpp new file mode 100644 index 000000000000..81278c7f3d54 --- /dev/null +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/streamingEventSink.cpp @@ -0,0 +1,258 @@ +/* + * 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. + */ + +#include "kv_cache_manager_v2/streamingEventSink.h" + +#include "kv_cache_manager_v2/blockRadixTree.h" +#include "kv_cache_manager_v2/page.h" +#include "tensorrt_llm/common/logger.h" + +#include +#include +#include +#include +#include +#include + +namespace tensorrt_llm::batch_manager::kv_cache_manager_v2 +{ + +StreamingEventSink::StreamingEventSink(int maxEntries, std::optional mmTokenIdOffset) + : mMaxEntries(maxEntries) + , mMmTokenIdOffset(mmTokenIdOffset) +{ + if (mMaxEntries <= 0) + { + throw std::invalid_argument("maxEntries must be positive"); + } + if (mMmTokenIdOffset.has_value() && *mMmTokenIdOffset < 0) + { + throw std::invalid_argument("mmTokenIdOffset must be non-negative"); + } +} + +void StreamingEventSink::setTargetLifeCycle(LifeCycleId lifeCycle) +{ + if (lifeCycle.value() < 0) + { + throw std::invalid_argument("lifeCycle must be non-negative"); + } + std::lock_guard lock(mMutex); + mTargetLifeCycle = lifeCycle; +} + +std::vector StreamingEventSink::drainIterationEvents() +{ + std::lock_guard lock(mMutex); + auto events = std::move(mPendingEvents); + mPendingEvents.clear(); + mPendingEntries = 0; + return events; +} + +StreamingEventStats StreamingEventSink::getStats() const +{ + std::lock_guard lock(mMutex); + return mStats; +} + +void StreamingEventSink::addStoredBlock(Block const& block) +{ + std::lock_guard lock(mMutex); + addStoredBlockUnlocked(block); +} + +void StreamingEventSink::addStoredLifeCycle(Block const& block, LifeCycleId lifeCycle) +{ + std::lock_guard lock(mMutex); + if (!mTargetLifeCycle.has_value()) + { + return; + } + if (lifeCycle != *mTargetLifeCycle) + { + ++mStats.nonTargetLifeCyclesIgnored; + return; + } + addStoredBlockUnlocked(block); +} + +void StreamingEventSink::addRemovedBlock(Digest const& blockKey) +{ + std::lock_guard lock(mMutex); + addRemovedBlockUnlocked(blockKey); +} + +void StreamingEventSink::addRemovedLifeCycle(Digest const& blockKey, LifeCycleId lifeCycle) +{ + std::lock_guard lock(mMutex); + if (!mTargetLifeCycle.has_value()) + { + return; + } + if (lifeCycle != *mTargetLifeCycle) + { + ++mStats.nonTargetLifeCyclesIgnored; + return; + } + addRemovedBlockUnlocked(blockKey); +} + +void StreamingEventSink::addCacheLevelUpdated(Digest const&, CacheLevel, CacheLevel, LifeCycleId) +{ + // The streaming protocol currently tracks radix-tree residency, not cache-tier movement. +} + +void StreamingEventSink::addStoredBlockUnlocked(Block const& block) +{ + if (!mTargetLifeCycle.has_value() || *mTargetLifeCycle >= block.storage.size()) + { + return; + } + auto const* page = block.getPage(*mTargetLifeCycle); + if (page == nullptr) + { + return; + } + if (!block.isFull() || page->numTokensInBlock < static_cast(block.tokens.size())) + { + ++mStats.partialBlocksSuppressed; + return; + } + if (mStoredBlocks.count(block.key) != 0) + { + return; + } + if (block.prev == nullptr) + { + throw std::logic_error("Cannot publish an orphan KV cache block"); + } + + int64_t const blockHash = wireHash(block.key); + std::optional parentHash; + std::optional loraId; + if (block.prev->type() == NodeBase::Type::kBLOCK) + { + auto const& parent = *static_cast(block.prev); + auto const storedParent = mStoredBlocks.find(parent.key); + if (storedParent == mStoredBlocks.end()) + { + recordDroppedEventUnlocked("parent block has not been published"); + return; + } + parentHash = storedParent->second.blockHash; + // Inherit scope metadata from the published parent without walking the ancestor chain. + loraId = storedParent->second.loraId; + } + else + { + loraId = static_cast(block.prev)->reuseScope.loraId; + } + + auto decoded = decodeEventBlock(block, mMmTokenIdOffset); + if (!reserveEntryUnlocked()) + { + return; + } + + mStoredBlocks.emplace(block.key, StoredBlock{blockHash, loraId}); + if (!mPendingEvents.empty()) + { + auto* stored = std::get_if(&mPendingEvents.back()); + if (stored != nullptr && !stored->blockHashes.empty() && parentHash.has_value() + && stored->blockHashes.back() == *parentHash && stored->loraId == loraId) + { + stored->blockHashes.push_back(blockHash); + stored->tokenIds.insert(stored->tokenIds.end(), std::make_move_iterator(decoded.tokenIds.begin()), + std::make_move_iterator(decoded.tokenIds.end())); + stored->mmKeys.push_back(std::move(decoded.mmKeys)); + ++mStats.storedBlocks; + return; + } + } + std::vector> mmKeys; + mmKeys.push_back(std::move(decoded.mmKeys)); + mPendingEvents.emplace_back( + StreamingBlockStoredData{{blockHash}, parentHash, std::move(decoded.tokenIds), std::move(mmKeys), loraId}); + ++mStats.storedBlocks; +} + +void StreamingEventSink::addRemovedBlockUnlocked(Digest const& blockKey) +{ + auto const stored = mStoredBlocks.find(blockKey); + if (stored == mStoredBlocks.end()) + { + return; + } + int64_t const blockHash = stored->second.blockHash; + mStoredBlocks.erase(stored); + addRemovedHashUnlocked(blockHash); +} + +void StreamingEventSink::addRemovedHashUnlocked(int64_t blockHash) +{ + if (!mPendingEvents.empty()) + { + auto* removed = std::get_if(&mPendingEvents.back()); + if (removed != nullptr) + { + removed->blockHashes.push_back(blockHash); + ++mStats.removedBlocks; + return; + } + } + mPendingEvents.emplace_back(StreamingBlockRemovedData{{blockHash}}); + ++mStats.removedBlocks; +} + +bool StreamingEventSink::reserveEntryUnlocked() +{ + if (mPendingEntries < mMaxEntries) + { + ++mPendingEntries; + return true; + } + recordDroppedEventUnlocked("per-iteration safety cap was exceeded"); + return false; +} + +void StreamingEventSink::recordDroppedEventUnlocked(char const* reason) +{ + ++mStats.droppedEvents; + int64_t const dropped = mStats.droppedEvents; + if (dropped == 1 || (dropped & (dropped - 1)) == 0) + { + TLLM_LOG_WARNING("Dropping streaming KV store event because %s; dropped_events=%" PRId64, reason, dropped); + } +} + +int64_t StreamingEventSink::wireHash(Digest const& digest) +{ + uint64_t value = 0; + for (size_t i = 0; i < sizeof(value); ++i) + { + value = (value << 8U) | std::to_integer(digest[i]); + } + uint64_t constexpr kSignedMax = static_cast(std::numeric_limits::max()); + if (value <= kSignedMax) + { + return static_cast(value); + } + return -static_cast(~value) - 1; +} + +} // namespace tensorrt_llm::batch_manager::kv_cache_manager_v2 diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/streamingEventSink.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/streamingEventSink.h new file mode 100644 index 000000000000..89e519ce7bd5 --- /dev/null +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/streamingEventSink.h @@ -0,0 +1,108 @@ +/* + * 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. + */ + +#pragma once + +#include "kv_cache_manager_v2/eventData.h" +#include "kv_cache_manager_v2/eventSink.h" + +#include +#include +#include +#include +#include +#include + +namespace tensorrt_llm::batch_manager::kv_cache_manager_v2 +{ + +//! Semantic data for one wire-level BlockStored event. +struct StreamingBlockStoredData +{ + std::vector blockHashes; + std::optional parentBlockHash; + std::vector tokenIds; + //! One entry per blockHash. Empty block entries represent text-only blocks. + std::vector> mmKeys; + std::optional loraId; +}; + +//! Semantic data for one wire-level BlockRemoved event. +struct StreamingBlockRemovedData +{ + std::vector blockHashes; +}; + +using StreamingEventData = std::variant; + +//! Counters accumulated over the lifetime of a streaming event sink. +struct StreamingEventStats +{ + int64_t storedBlocks = 0; + int64_t removedBlocks = 0; + int64_t partialBlocksSuppressed = 0; + int64_t nonTargetLifeCyclesIgnored = 0; + int64_t droppedEvents = 0; +}; + +//! Captures streaming KV-cache lifecycle events without depending on Python or a transport. +class StreamingEventSink final : public EventSink +{ +public: + StreamingEventSink(int maxEntries, std::optional mmTokenIdOffset = std::nullopt); + + bool needsTokenDigestContext() const override + { + return mMmTokenIdOffset.has_value(); + } + + void setTargetLifeCycle(LifeCycleId lifeCycle); + [[nodiscard]] std::vector drainIterationEvents(); + [[nodiscard]] StreamingEventStats getStats() const; + + void addStoredBlock(Block const& block) override; + void addStoredLifeCycle(Block const& block, LifeCycleId lifeCycle) override; + void addRemovedBlock(Digest const& blockKey) override; + void addRemovedLifeCycle(Digest const& blockKey, LifeCycleId lifeCycle) override; + void addCacheLevelUpdated( + Digest const& blockKey, CacheLevel oldLevel, CacheLevel newLevel, LifeCycleId lifeCycle) override; + +private: + struct StoredBlock + { + int64_t blockHash; + std::optional loraId; + }; + + void addStoredBlockUnlocked(Block const& block); + void addRemovedBlockUnlocked(Digest const& blockKey); + void addRemovedHashUnlocked(int64_t blockHash); + void recordDroppedEventUnlocked(char const* reason); + [[nodiscard]] bool reserveEntryUnlocked(); + [[nodiscard]] static int64_t wireHash(Digest const& digest); + + int mMaxEntries; + std::optional mMmTokenIdOffset; + std::optional mTargetLifeCycle; + int mPendingEntries = 0; + std::unordered_map mStoredBlocks; + std::vector mPendingEvents; + StreamingEventStats mStats; + mutable std::mutex mMutex; +}; + +} // namespace tensorrt_llm::batch_manager::kv_cache_manager_v2 diff --git a/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp b/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp index 3988b84bbc3d..82811a8c30dd 100644 --- a/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp +++ b/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp @@ -32,6 +32,7 @@ #include "kv_cache_manager_v2/stats.h" #include "kv_cache_manager_v2/storage/config.h" #include "kv_cache_manager_v2/storage/core.h" +#include "kv_cache_manager_v2/streamingEventSink.h" #include "kv_cache_manager_v2/utils/optionalGilRelease.h" #include @@ -461,10 +462,10 @@ static std::vector castMmKeys(nb::handle values) return result; } -static nb::list castMmKeys(kv::KVCacheStoredBlockData const& data) +static nb::list castMmKeys(std::vector const& mmKeys) { nb::list result; - for (auto const& mmKey : data.mmKeys) + for (auto const& mmKey : mmKeys) { auto hash = nb::bytes(mmKey.hash.data(), mmKey.hash.size()); if (mmKey.hasUuidField) @@ -479,6 +480,21 @@ static nb::list castMmKeys(kv::KVCacheStoredBlockData const& data) return result; } +static nb::list castMmKeys(kv::KVCacheStoredBlockData const& data) +{ + return castMmKeys(data.mmKeys); +} + +static nb::list castStreamingMmKeys(kv::StreamingBlockStoredData const& data) +{ + nb::list result; + for (auto const& mmKeys : data.mmKeys) + { + result.append(castMmKeys(mmKeys)); + } + return result; +} + static nb::object castEventData(kv::KVCacheEventData const& data) { return std::visit([](auto const& concreteData) { return nb::cast(concreteData); }, data); @@ -1050,7 +1066,38 @@ void KvCacheManagerV2Bindings::initBindings(nb::module_& m) self.attentionDpRank, self.layerGroupId)); }); - nb::class_(m, "KVCacheEventManager") + nb::class_(m, "KVCacheEventSink"); + + nb::class_(m, "StreamingBlockStoredData") + .def_ro("block_hashes", &kv::StreamingBlockStoredData::blockHashes) + .def_ro("parent_block_hash", &kv::StreamingBlockStoredData::parentBlockHash) + .def_ro("token_ids", &kv::StreamingBlockStoredData::tokenIds) + .def_ro("lora_id", &kv::StreamingBlockStoredData::loraId) + .def_prop_ro("mm_keys", [](kv::StreamingBlockStoredData const& self) { return castStreamingMmKeys(self); }); + + nb::class_(m, "StreamingBlockRemovedData") + .def_ro("block_hashes", &kv::StreamingBlockRemovedData::blockHashes); + + nb::class_(m, "StreamingEventStats") + .def_ro("stored_blocks", &kv::StreamingEventStats::storedBlocks) + .def_ro("removed_blocks", &kv::StreamingEventStats::removedBlocks) + .def_ro("partial_blocks_suppressed", &kv::StreamingEventStats::partialBlocksSuppressed) + .def_ro("non_target_life_cycles_ignored", &kv::StreamingEventStats::nonTargetLifeCyclesIgnored) + .def_ro("dropped_events", &kv::StreamingEventStats::droppedEvents); + + nb::class_(m, "StreamingEventSink") + .def(nb::init>(), nb::arg("max_entries") = 50'000, + nb::arg("mm_token_id_offset") = std::nullopt) + .def( + "set_target_life_cycle", + [](kv::StreamingEventSink& self, int lifeCycleId) + { self.setTargetLifeCycle(kv::LifeCycleId{lifeCycleId}); }, + nb::arg("life_cycle_id"), nb::call_guard()) + .def("drain_iteration_events", &kv::StreamingEventSink::drainIterationEvents, + nb::call_guard()) + .def_prop_ro("stats", &kv::StreamingEventSink::getStats, nb::call_guard()); + + nb::class_(m, "KVCacheEventManager") .def( "__init__", [](kv::EventManager* self, int maxKvEventEntries, int windowSize, std::optional attentionDpRank, @@ -2048,6 +2095,26 @@ void KvCacheManagerV2Bindings::initBindings(nb::module_& m) [](kv::EventManager& eventManager, EventManagerTestBlock const& block, int lifeCycleId) { eventManager.addStoredLifeCycle(*block.block, kv::LifeCycleId{lifeCycleId}); }, nb::arg("event_manager"), nb::arg("block"), nb::arg("life_cycle_id"), nb::call_guard()); + mIntrospection.def( + "streaming_event_sink_add_stored_block", + [](kv::StreamingEventSink& eventSink, EventManagerTestBlock const& block) + { eventSink.addStoredBlock(*block.block); }, + nb::arg("event_sink"), nb::arg("block"), nb::call_guard()); + mIntrospection.def( + "streaming_event_sink_add_stored_life_cycle", + [](kv::StreamingEventSink& eventSink, EventManagerTestBlock const& block, int lifeCycleId) + { eventSink.addStoredLifeCycle(*block.block, kv::LifeCycleId{lifeCycleId}); }, + nb::arg("event_sink"), nb::arg("block"), nb::arg("life_cycle_id"), nb::call_guard()); + mIntrospection.def( + "streaming_event_sink_add_removed_block", + [](kv::StreamingEventSink& eventSink, EventManagerTestBlock const& block) + { eventSink.addRemovedBlock(block.block->key); }, + nb::arg("event_sink"), nb::arg("block"), nb::call_guard()); + mIntrospection.def( + "streaming_event_sink_add_removed_life_cycle", + [](kv::StreamingEventSink& eventSink, EventManagerTestBlock const& block, int lifeCycleId) + { eventSink.addRemovedLifeCycle(block.block->key, kv::LifeCycleId{lifeCycleId}); }, + nb::arg("event_sink"), nb::arg("block"), nb::arg("life_cycle_id"), nb::call_guard()); mIntrospection.def( "active_page_stats", [](kv::KvCache const& kvCache) @@ -2275,7 +2342,7 @@ void KvCacheManagerV2Bindings::initBindings(nb::module_& m) std::shared_ptr eventSink; if (!eventManager.is_none()) { - eventSink = nb::cast>(eventManager); + eventSink = nb::cast>(eventManager); } std::unique_ptr codec; diff --git a/docs/source/features/kvcache.md b/docs/source/features/kvcache.md index 65465c09e08d..fdff367cc0c7 100644 --- a/docs/source/features/kvcache.md +++ b/docs/source/features/kvcache.md @@ -309,29 +309,26 @@ shows whether the switch actually did anything. ### KV Cache Events -KV cache events report block **stored**, **removed**, **created** and **updated** operations -so an external KV-cache-aware router (for example NVIDIA Dynamo) can route a request to the -engine that already holds its prefix. Two delivery paths are available. +KV cache events let an external KV-cache-aware router (for example NVIDIA Dynamo) route a +request to the engine that already holds its prefix. The buffered path reports block **stored**, +**removed**, **created** and **updated** operations. The streaming path exposes the narrower +router-facing contract described below. #### Buffered path (default) Set ```event_buffer_max_size``` to a positive integer and ```enable_block_reuse``` to True. Events are buffered per rank, gathered onto rank 0 under attention data parallelism, and -pulled per iteration through `LLM.get_kv_cache_events()` / `LLM.get_kv_cache_events_async()`, -or over the `/kv_cache_events` endpoint of `trtllm-serve`. +exposed through `LLM.get_kv_cache_events()` / `LLM.get_kv_cache_events_async()`, or over the +`/kv_cache_events` endpoint of `trtllm-serve`. -#### Streaming path (unsupported) +#### Streaming path (prototype) -```{note} -The streaming path has no implementation: `kv_cache_config.kv_events_config` is rejected -at startup. Use the buffered path via `kv_cache_config.event_buffer_max_size` instead. The -wire format and endpoint convention below describe the contract a future native event sink -must satisfy. -``` - -Configured with ```kv_cache_config.kv_events_config```. Each rank encodes its own events and -publishes them directly over a ZeroMQ `PUB` socket from a background thread, so there is no -rank-0 gather and no per-iteration pull. +Configured with ```kv_cache_config.kv_events_config```. Streaming is intended to reduce event +publishing overhead under attention data parallelism: each emitting rank publishes its own +events over a ZeroMQ `PUB` socket, removing the rank-0 gather and the consumer pull through the +LLM API. Each emitting rank still drains its local native event source, converts the events to +wire structs and enqueues one batch at the iteration boundary. A background thread performs +msgpack encoding and socket I/O. ```python from tensorrt_llm.llmapi import KvCacheConfig, KVEventsConfig @@ -346,10 +343,22 @@ kv_cache_config = KvCacheConfig( ) ``` -**Constraints.** Enabling the streaming path raises at startup. A Python event sink cannot -serve it, because the KV cache manager V2 radix tree calls its sink natively rather than -through Python; re-enabling it needs a native sink. Pipeline parallelism and context -parallelism are rejected independently. +**Limitations.** Pipeline and context parallelism are unsupported. Buffered polling returns +no events while streaming is enabled. Draft models and KV-cache-size estimation do not publish +events. Use buffered mode for cache-tier, priority, and other lifecycle updates. + +Streaming exposes only the full-block residency information external routers need, using the +attention lifecycle with the largest window. Conversion to `BlockStored`/`BlockRemoved` is +centralized at the once-per-iteration publisher boundary, keeping internal lifecycle details +contained and avoiding Python callbacks from native cache operations. + +**Multimodal payloads.** `token_ids` can contain integers and hexadecimal digest strings; +integer-only consumers are incompatible. Each `mm_keys[i]` describes the multimodal segments +in `block_hashes[i]`, using the `hash` and `start_offset` fields defined above. This payload +support does not establish end-to-end multimodal-aware routing compatibility. The final +identity and normalization contract will be revisited separately after +[Dynamo #15095](https://github.com/ai-dynamo/dynamo/pull/15095) and +[TensorRT-LLM #19529](https://github.com/NVIDIA/TensorRT-LLM/pull/19529) merge. **Endpoint convention.** Every attention-DP rank binds `base_port + rank` using its **global** rank, so `N` ranks occupy `[base_port, base_port + N - 1]` cluster-wide and @@ -366,15 +375,17 @@ have no port, each rank appends a `_dp` suffix instead. **Wire format.** Each batch is sent as three ZeroMQ frames: the subscription ```topic```, an 8-byte big-endian sequence number, and a msgpack payload `[timestamp, [events], data_parallel_rank]`. Each event is a map tagged with a `type` key — -`BlockStored`, `BlockRemoved` or `AllBlocksCleared` — carrying int64 block hashes derived -from the V2 radix block keys. This is the format documented for custom router backends; it -differs from vLLM's positional-array encoding of the individual events, though the batch -envelope is positional in both. +`BlockStored` or `BlockRemoved` — carrying int64 block hashes derived from the V2 radix block +keys. This is the format documented for custom router backends; it differs from vLLM's +positional-array encoding of the individual events, though the batch envelope is positional in +both. -**Delivery guarantees.** Delivery is best effort, but loss is observable. Every accepted +**Delivery guarantees.** Delivery is best effort, and batch loss is observable. Every accepted batch reserves a sequence number up front, so a batch dropped by a full publisher queue (```max_queue_size```) or by a failed send leaves a hole in the sequence. Subscribers must -treat any gap as lost KV-cache state and resynchronize rather than assuming continuity. +treat any gap as lost KV-cache state and resynchronize rather than assuming continuity. A +producer-side safety-cap drop is reported in the producer's logs and counters but does not +create a sequence gap. **Replay.** If ```replay_endpoint``` is set, the publisher also binds a `ROUTER` socket. A subscriber sends an empty delimiter frame plus an 8-byte big-endian start sequence, and diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py index 41471d4c573b..12496c28768b 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py @@ -82,7 +82,6 @@ KVCacheEventManager, KVCacheIterationStatsDelta, LayerId, - LifeCycleId, OutOfPagesError, PageIndexMode, PlannedDropHandle, @@ -1326,8 +1325,8 @@ def __init__( ) if streaming_events_enabled: assert kv_events_config is not None - # Rejects unsupported parallelism and streaming itself, before any socket is - # bound and before any claim is made about which event path is in use. + # Reject unsupported parallelism and colliding publish/replay port ranges + # before any socket is bound. validate_streaming_support( kv_events_config, pp_size=mapping.pp_size, @@ -1352,6 +1351,7 @@ def __init__( data_parallel_rank=event_rank, block_size=self.tokens_per_block, max_window_size=event_window_size, + mm_token_id_offset=vocab_size, ) elif self.event_buffer_max_size > 0: if mapping.enable_attention_dp: @@ -1566,10 +1566,15 @@ def create_cold_page_codec(cache_config: object) -> Optional[object]: ) candidate: Optional[KVCacheManagerPy] = None + event_sink = ( + self.event_manager.event_sink + if isinstance(self.event_manager, StreamingKVCacheEventManager) + else self.event_manager + ) if not has_host_cache_tier: candidate = KVCacheManagerPy( config, - event_manager=self.event_manager, + event_manager=event_sink, cold_page_codec=create_cold_page_codec(config), ) else: @@ -1578,7 +1583,7 @@ def create_cold_page_codec(cache_config: object) -> Optional[object]: try: candidate = KVCacheManagerPy( config, - event_manager=self.event_manager, + event_manager=event_sink, cold_page_codec=create_cold_page_codec(config), ) except Exception as error: @@ -1618,7 +1623,7 @@ def create_cold_page_codec(cache_config: object) -> Optional[object]: ) candidate = KVCacheManagerPy( config, - event_manager=self.event_manager, + event_manager=event_sink, cold_page_codec=create_cold_page_codec(config), ) except Exception as error: @@ -2479,18 +2484,17 @@ def _get_event_window_sizes_by_layer_group( # tying with the attention life cycle and being selected as the event target. # The buffered manager keeps every layer group, so its windows are unchanged. - def get_event_window_size(layer_id: int) -> int: - layer_config = self.kv_cache_manager_py_config.layers[layer_id] + def get_event_window_size(layer_config: object) -> int: window_size = getattr(layer_config, "sliding_window_size", None) return self.max_seq_len if window_size is None else int(window_size) window_sizes: Dict[int, int] = {} for layer_group_id, layer_ids in enumerate(self.impl.layer_grouping): - if attention_only: - life_cycle = self.impl._life_cycles.get_life_cycle(LifeCycleId(layer_group_id)) - if not isinstance(life_cycle, AttnLifeCycle): - continue - window_sizes[int(layer_group_id)] = get_event_window_size(int(layer_ids[0])) + # Native bindings expose grouping, not the private Python lifecycle registry. + layer_config = self.kv_cache_manager_py_config.layers[int(layer_ids[0])] + if attention_only and not isinstance(layer_config, AttentionLayerConfig): + continue + window_sizes[int(layer_group_id)] = get_event_window_size(layer_config) return window_sizes def _format_kv_cache_pool_lifecycle_entry(self, layer_id: LayerId, role: DataRole) -> str: diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_events.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_events.py index 7d109712671f..86a13e09a7ad 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_events.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_events.py @@ -39,11 +39,13 @@ from tensorrt_llm.llmapi.llm_args import KVEventsConfig from tensorrt_llm.logger import logger -from tensorrt_llm.runtime.kv_cache_manager_v2 import KVCacheEvent, KVCacheEventDiff +from tensorrt_llm.runtime import kv_cache_manager_v2 as kv_cache_manager_v2_runtime +from tensorrt_llm.runtime.kv_cache_manager_v2 import KVCacheEvent # Subscribers decode block hashes as 64-bit ints, so a bytes value would fail the # decode for the entire batch. ExternalBlockHash = int +EventTokenId = int | str class EventBatch( @@ -68,12 +70,24 @@ class KVCacheWireEvent( """Base class for KV cache event wire messages.""" +class MultimodalKey( + msgspec.Struct, + omit_defaults=True, # type: ignore[call-arg] + gc=False, # type: ignore[call-arg] + tag="mm_key", +): + """One continuous multimodal segment within a stored block.""" + + hash: str + start_offset: int + + class BlockStored(KVCacheWireEvent): """A sequence of full KV cache blocks was stored.""" block_hashes: list[ExternalBlockHash] parent_block_hash: ExternalBlockHash | None - token_ids: list[int] + token_ids: list[EventTokenId] block_size: int lora_id: int | None medium: str | None @@ -83,6 +97,8 @@ class BlockStored(KVCacheWireEvent): kv_cache_spec_kind: str | None = None kv_cache_spec_sliding_window: int | None = None locality: str | None = None + # Aligned one-for-one with block_hashes when multimodal decoding is enabled. + mm_keys: list[list[MultimodalKey]] | None = None class BlockRemoved(KVCacheWireEvent): @@ -406,24 +422,12 @@ def validate_streaming_support( Split out of ``KVCacheManagerV2.__init__`` so the preconditions are testable without building a manager, which needs a GPU. - Streaming is currently unsupported, so this always raises; ``config``, - ``ranks_per_host`` and ``data_parallel_size`` are retained for the caller's - signature and are consumed again once a native event sink exists. """ - del config, ranks_per_host, data_parallel_size if pp_size > 1: raise ValueError("Streaming KV events do not support pipeline parallelism") if cp_size > 1: raise ValueError("Streaming KV events do not support context parallelism") - # StreamingKVCacheEventManager is a duck-typed Python event sink. It cannot satisfy the - # nanobind constructor's nb::cast>, and the C++ radix - # tree calls its sink natively rather than through Python, so there is no live path that - # can produce streaming events. Fail with an actionable message pointing at the buffered - # alternative instead of an opaque TypeError from the cast. - raise ValueError( - "Streaming KV events (kv_cache_config.kv_events_config) are not supported. Use the " - "buffered path via kv_cache_config.event_buffer_max_size instead." - ) + validate_endpoint_ranges(config, ranks_per_host, data_parallel_size) def validate_endpoint_ranges( @@ -494,19 +498,67 @@ def create_event_publisher(config: KVEventsConfig, data_parallel_rank: int) -> E raise ValueError(f"Unsupported KV event publisher: {config.publisher!r}") -class StreamingKVCacheEventManager: - """Event-sink hook interface for out-of-band KV cache event publishing. +class _StreamingEventSource: + """Own the native sink and translate its semantic DTOs to wire structs.""" - The interface is duck typed rather than derived from ``KVCacheEventManager``: a sink - fully replaces event production (reusing the radix block hashes) and shares none of - the base manager's state. + def __init__( + self, + *, + block_size: int, + max_entries: int, + mm_token_id_offset: int | None, + ) -> None: + if mm_token_id_offset is not None and mm_token_id_offset < 0: + raise ValueError("mm_token_id_offset must be non-negative") + self._block_size = block_size + self._include_mm_keys = mm_token_id_offset is not None + self._event_sink = kv_cache_manager_v2_runtime.StreamingEventSink( + max_entries=max_entries, + mm_token_id_offset=mm_token_id_offset, + ) + + @property + def event_sink(self) -> kv_cache_manager_v2_runtime.StreamingEventSink: + return self._event_sink + + def set_target_life_cycle(self, life_cycle_id: int) -> None: + self._event_sink.set_target_life_cycle(life_cycle_id) + + def drain_events(self) -> list[BlockStored | BlockRemoved | AllBlocksCleared]: + result: list[BlockStored | BlockRemoved | AllBlocksCleared] = [] + for event in self._event_sink.drain_iteration_events(): + if isinstance(event, kv_cache_manager_v2_runtime.StreamingBlockStoredData): + result.append( + BlockStored( + block_hashes=list(event.block_hashes), + parent_block_hash=event.parent_block_hash, + token_ids=list(event.token_ids), + block_size=self._block_size, + lora_id=event.lora_id, + medium="GPU", + lora_name=None, + mm_keys=( + [ + [ + MultimodalKey(hash=bytes(key[0]).hex(), start_offset=key[1]) + for key in block_keys + ] + for block_keys in event.mm_keys + ] + if self._include_mm_keys + else None + ), + ) + ) + elif isinstance(event, kv_cache_manager_v2_runtime.StreamingBlockRemovedData): + result.append(BlockRemoved(block_hashes=list(event.block_hashes), medium="GPU")) + else: + raise TypeError(f"Unsupported native streaming KV event: {type(event)!r}") + return result - No implementation is currently wired up. The C++ radix tree invokes its sink natively - rather than through Python, so a Python object cannot receive these callbacks; the - signatures below record the contract a native sink must satisfy. Constructing this - class raises -- ``validate_streaming_support`` rejects the configuration earlier, so - reaching here means that check was bypassed. - """ + +class StreamingKVCacheEventManager: + """Publish iteration batches captured by the native streaming event sink.""" def __init__( self, @@ -516,49 +568,104 @@ def __init__( block_size: int, max_window_size: int, max_entries: int = 50_000, + mm_token_id_offset: int | None = None, ) -> None: - raise NotImplementedError( - "Streaming KV events are not supported. Use the buffered path via " - "kv_cache_config.event_buffer_max_size instead." + self._rank = data_parallel_rank + self._publisher = create_event_publisher(config, data_parallel_rank) + self._max_window_size = max_window_size + self._event_source = _StreamingEventSource( + block_size=block_size, + max_entries=max_entries, + mm_token_id_offset=mm_token_id_offset, ) + self._closed = False + self.enqueued_batches = 0 + self.enqueued_events = 0 + self.dropped_batches = 0 + + @property + def stored_blocks(self) -> int: + return self.event_sink.stats.stored_blocks - def needs_token_digest_context(self) -> bool: - # Streaming events do not emit multimodal keys. - return False + @property + def removed_blocks(self) -> int: + return self.event_sink.stats.removed_blocks + + @property + def partial_blocks_suppressed(self) -> int: + return self.event_sink.stats.partial_blocks_suppressed + + @property + def non_target_life_cycles_ignored(self) -> int: + return self.event_sink.stats.non_target_life_cycles_ignored + + @property + def dropped_events(self) -> int: + return self.event_sink.stats.dropped_events def start(self) -> None: """Bind the publisher's sockets and start its background thread.""" + self._publisher.start() def set_layer_group_window_sizes(self, window_sizes: dict[int, int]) -> None: - """Select the attention life cycle whose blocks are published.""" + target_ids = [ + int(life_cycle_id) + for life_cycle_id, window_size in window_sizes.items() + if int(window_size) == self._max_window_size + ] + if not target_ids and window_sizes: + largest_window = max(window_sizes.values()) + target_ids = [ + int(life_cycle_id) + for life_cycle_id, window_size in window_sizes.items() + if window_size == largest_window + ] + if not target_ids: + raise ValueError("Streaming KV events require an attention KV cache life cycle") + target_life_cycle_id = min(target_ids) + self._event_source.set_target_life_cycle(target_life_cycle_id) + logger.info( + "Streaming KV event fast path selected " + f"lifecycle_id={target_life_cycle_id} " + f"window_size={self._max_window_size}" + ) def add_created_event( self, num_blocks_per_cache_level: Any, layer_group_ids: Any = None, - ) -> None: ... - - def add_stored_event(self, *args: Any, **kwargs: Any) -> None: ... - - def add_stored_block_event_from_block(self, block: Any) -> None: ... - - def add_stored_life_cycle_event_from_block(self, block: Any, life_cycle_id: int) -> None: ... - - def add_removed_event(self, block_hashes: Any) -> None: ... - - def add_removed_life_cycle_event(self, block_hash: bytes, life_cycle_id: int) -> None: ... - - def add_updated_event( - self, - block_hash: Any, - *, - cache_level: KVCacheEventDiff | None = None, - priority: KVCacheEventDiff | None = None, - layer_group_id: int | None = None, - ) -> None: ... + ) -> None: + """Accept the common manager notification; streaming publishes no creation event.""" + return def flush_iteration_events(self) -> None: - """Publish the events accumulated during this iteration.""" + if self._closed: + return + events = self._event_source.drain_events() + if not events: + return + batch = KVEventBatch( + ts=time.time(), + events=events, + data_parallel_rank=self._rank, + ) + try: + if self._publisher.publish(batch): + self.enqueued_batches += 1 + self.enqueued_events += len(events) + else: + self.dropped_batches += 1 + except Exception: + self.dropped_batches += 1 + logger.error( + f"Dropping streaming KV event iteration batch on rank={self._rank}\n" + f"{traceback.format_exc()}" + ) + + @property + def event_sink(self) -> kv_cache_manager_v2_runtime.StreamingEventSink: + """Return the native sink installed in KVCacheManager.""" + return self._event_source.event_sink def get_latest_events(self, timeout_ms: float | None = None) -> list[KVCacheEvent]: # Streaming publishing pushes events out-of-band, so the pull API has @@ -567,4 +674,20 @@ def get_latest_events(self, timeout_ms: float | None = None) -> list[KVCacheEven return [] def shutdown(self) -> None: - """Flush pending events and stop the publisher.""" + if self._closed: + return + self.flush_iteration_events() + self._closed = True + self._publisher.shutdown() + stats = self.event_sink.stats + logger.info( + "Streaming KV event fast path " + f"rank={self._rank} " + f"stored_blocks={stats.stored_blocks} " + f"removed_blocks={stats.removed_blocks} " + f"partial_blocks_suppressed={stats.partial_blocks_suppressed} " + f"non_target_life_cycles_ignored={stats.non_target_life_cycles_ignored} " + f"dropped_events={stats.dropped_events} " + f"enqueued_batches={self.enqueued_batches} " + f"dropped_batches={self.dropped_batches}" + ) diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py index 921959815b2e..698a4408e723 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py @@ -93,6 +93,10 @@ class _BatchDescFieldSpec: KVCacheIterationStatsDelta = _cpp.KVCacheIterationStatsDelta KVCacheManager = _cpp.KVCacheManager KVCacheManagerConfig = _cpp.KVCacheManagerConfig +StreamingBlockRemovedData = _cpp.StreamingBlockRemovedData +StreamingBlockStoredData = _cpp.StreamingBlockStoredData +StreamingEventSink = _cpp.StreamingEventSink +StreamingEventStats = _cpp.StreamingEventStats IKvCacheColdPageCodec = _cpp.IKvCacheColdPageCodec create_default_kv_cache_cold_page_codec = _cpp.create_default_kv_cache_cold_page_codec # The C++ KVCacheManagerConfig binding replaces the Python @dataclass, but @@ -256,6 +260,10 @@ def typed_range(*args: int) -> range: "SlotDesc", "SlotDescVariant", "SsmLayerConfig", + "StreamingBlockRemovedData", + "StreamingBlockStoredData", + "StreamingEventSink", + "StreamingEventStats", "SwaScratchReuseConfig", "TokenId", "TokenIdExt", diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi index d50d3d894174..32ef372bcc75 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi @@ -332,6 +332,47 @@ class KVCacheEvent: attention_dp_rank: int | None = None layer_group_id: int | None = None +class StreamingBlockStoredData: + @property + def lora_id(self) -> int | None: ... + @property + def block_hashes(self) -> list[int]: ... + @property + def parent_block_hash(self) -> int | None: ... + @property + def token_ids(self) -> list[EventTokenId]: ... + @property + def mm_keys(self) -> list[list[MmKey]]: ... + +class StreamingBlockRemovedData: + @property + def block_hashes(self) -> list[int]: ... + +class StreamingEventStats: + @property + def stored_blocks(self) -> int: ... + @property + def removed_blocks(self) -> int: ... + @property + def partial_blocks_suppressed(self) -> int: ... + @property + def non_target_life_cycles_ignored(self) -> int: ... + @property + def dropped_events(self) -> int: ... + +class StreamingEventSink: + def __init__( + self, + max_entries: int = ..., + mm_token_id_offset: int | None = None, + ) -> None: ... + def set_target_life_cycle(self, life_cycle_id: int) -> None: ... + def drain_iteration_events( + self, + ) -> list[StreamingBlockStoredData | StreamingBlockRemovedData]: ... + @property + def stats(self) -> StreamingEventStats: ... + class KVCacheEventManager: def __init__( self, @@ -594,7 +635,7 @@ class KVCacheManager: def __init__( self, config: KVCacheManagerConfig, - event_manager: KVCacheEventManager | None = None, + event_manager: KVCacheEventManager | StreamingEventSink | None = None, cold_page_codec: IKvCacheColdPageCodec | None = None, ) -> None: ... def __del__(self) -> None: ... diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_introspection.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_introspection.py index e4d7cb7080e1..891d1d5df3a3 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_introspection.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_introspection.py @@ -106,6 +106,30 @@ def event_manager_add_stored_life_cycle(event_manager: Any, block: Any, life_cyc _cpp().event_manager_add_stored_life_cycle(event_manager, block, life_cycle_id) +def streaming_event_sink_add_stored_block(event_sink: Any, block: Any) -> None: + """Feed a real test block to the native streaming sink.""" + _cpp().streaming_event_sink_add_stored_block(event_sink, block) + + +def streaming_event_sink_add_stored_life_cycle( + event_sink: Any, block: Any, life_cycle_id: int +) -> None: + """Feed one lifecycle of a real test block to the native streaming sink.""" + _cpp().streaming_event_sink_add_stored_life_cycle(event_sink, block, life_cycle_id) + + +def streaming_event_sink_add_removed_block(event_sink: Any, block: Any) -> None: + """Remove a real test block from the native streaming sink.""" + _cpp().streaming_event_sink_add_removed_block(event_sink, block) + + +def streaming_event_sink_add_removed_life_cycle( + event_sink: Any, block: Any, life_cycle_id: int +) -> None: + """Remove one lifecycle of a real test block from the native streaming sink.""" + _cpp().streaming_event_sink_add_removed_life_cycle(event_sink, block, life_cycle_id) + + def active_page_stats(kv_cache: Any) -> tuple[list[int], list[int]]: """Return active pages and unscheduled evictable active pages by cache level.""" counts, unscheduled_evictable = _cpp().active_page_stats(kv_cache) diff --git a/tests/unittest/_torch/executor/kv_cache/test_kvcm2_integration.py b/tests/unittest/_torch/executor/kv_cache/test_kvcm2_integration.py index 98aed949cc4b..2f1d9f40cee3 100644 --- a/tests/unittest/_torch/executor/kv_cache/test_kvcm2_integration.py +++ b/tests/unittest/_torch/executor/kv_cache/test_kvcm2_integration.py @@ -595,6 +595,42 @@ def test_zero_size_filter_rejects_empty_local_cache() -> None: manager._remove_zero_size_buffers(config) +@pytest.mark.parametrize("sliding_window_size", [None, 8]) +@pytest.mark.parametrize("attention_first", [False, True]) +def test_event_window_sizes_filter_attention_without_backend_internals( + sliding_window_size: int | None, attention_first: bool +) -> None: + manager = object.__new__(KVCacheManagerV2) + manager.max_seq_len = MAX_SEQ_LEN + manager.kv_cache_manager_py_config = SimpleNamespace( + layers=[ + AttentionLayerConfig( + layer_id=LayerId(0), + buffers=[BufferConfig(role=Role.KEY, size=128)], + sliding_window_size=sliding_window_size, + ), + SsmLayerConfig( + layer_id=LayerId(1), + buffers=[BufferConfig(role=DataRole("ssm_state"), size=128)], + ), + ] + ) + # Match the C++ binding surface: layer_grouping is public, while the + # Python implementation's private _life_cycles registry is absent. + manager.impl = SimpleNamespace(layer_grouping=[[0], [1]] if attention_first else [[1], [0]]) + + attention_group_id = 0 if attention_first else 1 + window_size = MAX_SEQ_LEN if sliding_window_size is None else sliding_window_size + # SSM must be excluded even when its window ties with full attention's window. + assert manager._get_event_window_sizes_by_layer_group(attention_only=True) == { + attention_group_id: window_size + } + assert manager._get_event_window_sizes_by_layer_group() == { + attention_group_id: window_size, + 1 - attention_group_id: MAX_SEQ_LEN, + } + + def test_draft_token_relocation_uses_local_cache_layout(monkeypatch: pytest.MonkeyPatch) -> None: request = SimpleNamespace( state=LlmRequestState.GENERATION_IN_PROGRESS, diff --git a/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_event_manager.py b/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_event_manager.py index f427394b56f2..3f6f0ad92332 100644 --- a/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_event_manager.py +++ b/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_event_manager.py @@ -22,9 +22,15 @@ from importlib.util import find_spec from typing import TYPE_CHECKING, cast +import msgspec import pytest +from tensorrt_llm._torch.pyexecutor.kv_cache_events import ( + KVEventBatch, + StreamingKVCacheEventManager, +) from tensorrt_llm._utils import KVCacheEventSerializer +from tensorrt_llm.llmapi.llm_args import KVEventsConfig from tensorrt_llm.runtime.kv_cache_hash import ( KV_CACHE_HASH_ALGO_V1, KV_CACHE_HASH_ALGO_V2_SHA256_64, @@ -164,6 +170,22 @@ def _add_stored_life_cycle(event_manager, block, life_cycle_id): _introspection.event_manager_add_stored_life_cycle(event_manager, block, life_cycle_id) +def _add_streaming_stored_block(event_sink, block): + _introspection.streaming_event_sink_add_stored_block(event_sink, block) + + +def _add_streaming_stored_life_cycle(event_sink, block, life_cycle_id): + _introspection.streaming_event_sink_add_stored_life_cycle(event_sink, block, life_cycle_id) + + +def _add_streaming_removed_block(event_sink, block): + _introspection.streaming_event_sink_add_removed_block(event_sink, block) + + +def _add_streaming_removed_life_cycle(event_sink, block, life_cycle_id): + _introspection.streaming_event_sink_add_removed_life_cycle(event_sink, block, life_cycle_id) + + def _token_ids(start, end): return [TokenId(token_id) for token_id in range(start, end)] @@ -245,6 +267,264 @@ def test_event_manager_queue_and_stored_coalescing(): ] +@pytest.mark.parametrize("lora_id", [None, 0, 7, 2**64 - 1]) +def test_native_streaming_sink_to_python_wire_structs(real_block_factory, lora_id): + manager = StreamingKVCacheEventManager( + KVEventsConfig(enable_kv_cache_events=True, publisher="null"), + data_parallel_rank=0, + block_size=2, + max_window_size=128, + max_entries=8, + ) + event_sink = manager.event_sink + manager.start() + try: + manager.set_layer_group_window_sizes({0: 128, 1: 64}) + published = [] + manager._publisher.publish = lambda batch: published.append(batch) or True + make_block = real_block_factory(event_sink, num_life_cycles=2, tokens_per_block=2) + + first = make_block(_token_ids(1, 3), [2, 2], reuse_scope=ReuseScope(lora_id=lora_id)) + partial = make_block(_token_ids(5, 7), [1, 2], parent=first) + second = make_block(_token_ids(3, 5), [2, 2], parent=first) + + _add_streaming_stored_block(event_sink, first) + _add_streaming_stored_block(event_sink, partial) + _add_streaming_stored_life_cycle(event_sink, second, 1) + _add_streaming_stored_life_cycle(event_sink, second, 0) + # Native capture statistics are available before draining or publishing. + assert manager.stored_blocks == 2 + assert manager.partial_blocks_suppressed == 1 + assert manager.non_target_life_cycles_ignored == 1 + assert manager.dropped_events == 0 + assert published == [] + manager.flush_iteration_events() + + first_hash = int.from_bytes(_block_key(first)[:8], byteorder="big", signed=True) + second_hash = int.from_bytes(_block_key(second)[:8], byteorder="big", signed=True) + assert len(published) == 1 + stored = published[0].events + assert len(stored) == 1 + assert stored[0].block_hashes == [first_hash, second_hash] + assert stored[0].parent_block_hash is None + assert stored[0].token_ids == [1, 2, 3, 4] + assert stored[0].lora_id == lora_id + assert stored[0].lora_name is None + wire_batch = msgspec.msgpack.decode(msgspec.msgpack.encode(published[0]), type=KVEventBatch) + assert wire_batch.events[0].lora_id == lora_id + + # Descendants retain the scope after their parent's event has been drained. + third = make_block(_token_ids(7, 9), [2, 2], parent=second) + _add_streaming_stored_block(event_sink, third) + manager.flush_iteration_events() + assert published[-1].events[0].parent_block_hash == second_hash + assert published[-1].events[0].lora_id == lora_id + + _add_streaming_removed_life_cycle(event_sink, second, 1) + _add_streaming_removed_block(event_sink, first) + _add_streaming_removed_life_cycle(event_sink, second, 0) + assert manager.removed_blocks == 2 + assert manager.non_target_life_cycles_ignored == 2 + manager.flush_iteration_events() + + assert len(published) == 3 + removed = published[2].events + assert len(removed) == 1 + assert removed[0].block_hashes == [first_hash, second_hash] + assert manager.stored_blocks == 3 + assert manager.removed_blocks == 2 + assert manager.partial_blocks_suppressed == 1 + assert manager.non_target_life_cycles_ignored == 2 + assert manager.dropped_events == 0 + _add_streaming_stored_block(event_sink, first) + manager.flush_iteration_events() + assert published[-1].events[0].block_hashes == [first_hash] + assert published[-1].events[0].lora_id == lora_id + finally: + manager.shutdown() + + +def test_native_streaming_sink_separates_lora_scopes(real_block_factory): + manager = StreamingKVCacheEventManager( + KVEventsConfig(enable_kv_cache_events=True, publisher="null"), + data_parallel_rank=0, + block_size=2, + max_window_size=128, + ) + manager.start() + try: + manager.set_layer_group_window_sizes({0: 128}) + published = [] + manager._publisher.publish = lambda batch: published.append(batch) or True + make_block = real_block_factory(manager.event_sink, tokens_per_block=2) + for lora_id in (None, 0, 7, 8): + first = make_block(_token_ids(1, 3), [2], reuse_scope=ReuseScope(lora_id=lora_id)) + second = make_block(_token_ids(3, 5), [2], parent=first) + _add_streaming_stored_block(manager.event_sink, first) + _add_streaming_stored_block(manager.event_sink, second) + manager.flush_iteration_events() + stored = published[0].events + assert [event.lora_id for event in stored] == [None, 0, 7, 8] + assert all(event.token_ids == [1, 2, 3, 4] for event in stored) + assert all(len(event.block_hashes) == 2 for event in stored) + assert len({block_hash for event in stored for block_hash in event.block_hashes}) == 8 + finally: + manager.shutdown() + + +def test_native_streaming_sink_preserves_multimodal_event_data(real_block_factory): + manager = StreamingKVCacheEventManager( + KVEventsConfig(enable_kv_cache_events=True, publisher="null"), + data_parallel_rank=0, + block_size=4, + max_window_size=128, + max_entries=8, + mm_token_id_offset=1000, + ) + event_sink = manager.event_sink + manager.start() + try: + manager.set_layer_group_window_sizes({0: 128}) + published = [] + manager._publisher.publish = lambda batch: published.append(batch) or True + make_block = real_block_factory(event_sink, tokens_per_block=4) + digest_a = bytes(range(32)) + digest_b = bytes(reversed(range(32))) + first = make_block([1, digest_a, 1001, 1002], [4]) + gap = make_block([2, 3, 4, 5], [4], parent=first) + continued = make_block([1003, 7, 1004, 1005], [4], parent=gap) + last = make_block([1006, digest_b, 1001, 9], [4], parent=continued) + + for block in (first, gap, continued, last): + _add_streaming_stored_block(event_sink, block) + manager.flush_iteration_events() + + assert len(published) == 1 + assert len(published[0].events) == 1 + stored = published[0].events[0] + assert stored.token_ids == [ + 1, + digest_a.hex(), + 1001, + 1002, + 2, + 3, + 4, + 5, + 1003, + 7, + 1004, + 1005, + 1006, + digest_b.hex(), + 1001, + 9, + ] + assert [ + [(key.hash, key.start_offset) for key in block_keys] for block_keys in stored.mm_keys + ] == [ + [(digest_a.hex(), 0)], + [], + [(digest_a.hex(), 3), (digest_a.hex(), 4)], + [(digest_a.hex(), 6), (digest_b.hex(), 0)], + ] + assert manager.stored_blocks == 4 + finally: + manager.shutdown() + + +@pytest.mark.parametrize("overflow", [False, True]) +def test_native_streaming_removals_are_never_dropped_by_the_entry_cap(real_block_factory, overflow): + manager = StreamingKVCacheEventManager( + KVEventsConfig(enable_kv_cache_events=True, publisher="null"), + data_parallel_rank=0, + block_size=2, + max_window_size=128, + max_entries=2, + ) + manager.start() + try: + manager.set_layer_group_window_sizes({0: 128}) + published = [] + manager._publisher.publish = lambda batch: published.append(batch) or True + event_sink = manager.event_sink + make_block = real_block_factory(event_sink, tokens_per_block=2) + first = make_block(_token_ids(1, 3), [2]) + second = make_block(_token_ids(3, 5), [2], parent=first) + _add_streaming_stored_block(event_sink, first) + _add_streaming_stored_block(event_sink, second) + if overflow: + third = make_block(_token_ids(5, 7), [2], parent=second) + _add_streaming_stored_block(event_sink, third) + + # Both removal entry points must bypass the saturated store cap. + _add_streaming_removed_block(event_sink, first) + _add_streaming_removed_life_cycle(event_sink, second, 0) + manager.flush_iteration_events() + + assert manager.stored_blocks == 2 + assert manager.removed_blocks == 2 + assert manager.dropped_events == int(overflow) + assert len(published) == 1 + decoded = msgspec.msgpack.decode(msgspec.msgpack.encode(published[0])) + assert [event["type"] for event in decoded[1]] == ["BlockStored", "BlockRemoved"] + expected_hashes = [ + int.from_bytes(_block_key(block)[:8], byteorder="big", signed=True) + for block in (first, second) + ] + assert decoded[1][0]["block_hashes"] == expected_hashes + assert decoded[1][1]["block_hashes"] == expected_hashes + finally: + manager.shutdown() + + +def test_native_streaming_sink_drops_descendants_of_unpublished_parent(real_block_factory): + manager = StreamingKVCacheEventManager( + KVEventsConfig(enable_kv_cache_events=True, publisher="null"), + data_parallel_rank=0, + block_size=2, + max_window_size=128, + max_entries=1, + ) + event_sink = manager.event_sink + manager.start() + try: + manager.set_layer_group_window_sizes({0: 128}) + published = [] + manager._publisher.publish = lambda batch: published.append(batch) or True + make_block = real_block_factory(event_sink, tokens_per_block=2) + first = make_block(_token_ids(1, 3), [2]) + dropped_parent = make_block(_token_ids(3, 5), [2], parent=first) + child = make_block(_token_ids(5, 7), [2], parent=dropped_parent) + + _add_streaming_stored_block(event_sink, first) + _add_streaming_stored_block(event_sink, dropped_parent) + manager.flush_iteration_events() + _add_streaming_stored_block(event_sink, child) + manager.flush_iteration_events() + + assert len(published) == 1 + assert published[0].events[0].token_ids == [1, 2] + assert manager.stored_blocks == 1 + assert manager.dropped_events == 2 + finally: + manager.shutdown() + + +def test_native_streaming_sink_rejects_negative_life_cycle_id(): + manager = StreamingKVCacheEventManager( + KVEventsConfig(enable_kv_cache_events=True, publisher="null"), + data_parallel_rank=0, + block_size=2, + max_window_size=128, + ) + try: + with pytest.raises(ValueError, match="lifeCycle must be non-negative"): + manager.event_sink.set_target_life_cycle(-1) + finally: + manager.shutdown() + + def test_event_manager_attention_dp_gather_callback(): gathered_events = [] diff --git a/tests/unittest/kv_cache_manager_v2_tests/test_streaming_kv_events.py b/tests/unittest/kv_cache_manager_v2_tests/test_streaming_kv_events.py index 6866f780864c..b1dab0736a11 100644 --- a/tests/unittest/kv_cache_manager_v2_tests/test_streaming_kv_events.py +++ b/tests/unittest/kv_cache_manager_v2_tests/test_streaming_kv_events.py @@ -14,12 +14,8 @@ # limitations under the License. import socket -from types import SimpleNamespace -from typing import Callable -import msgspec import pytest -import zmq from tensorrt_llm._torch.pyexecutor.kv_cache_events import ( KVEventBatch, @@ -30,23 +26,6 @@ ) from tensorrt_llm.llmapi.llm_args import KVEventsConfig -# Streaming KV cache events have no implementation: StreamingKVCacheEventManager -# construction and validate_streaming_support() both raise. The supported route for KV -# cache events is the buffered one, via kv_cache_config.event_buffer_max_size. -streaming_unsupported = pytest.mark.skip( - reason="streaming KV cache events have no implementation; " - "use the buffered path via kv_cache_config.event_buffer_max_size" -) - -_ZMQ_SETUP_ATTEMPTS = 4 -_RECEIVE_TIMEOUT_MS = 2_000 -_SUBSCRIBE_ATTEMPTS = 50 -_PROBE_TIMEOUT_MS = 100 - - -class _NotReceived(Exception): - """No batch arrived: the subscription had not propagated yet.""" - def _unused_tcp_port() -> int: with socket.socket() as sock: @@ -54,219 +33,6 @@ def _unused_tcp_port() -> int: return int(sock.getsockname()[1]) -def _await_subscription(publisher: ZmqEventPublisher, subscriber: zmq.Socket) -> int: - """Publish probe batches until the subscriber's subscription is live. - - A PUB socket silently drops everything published before a subscriber's - subscription has propagated, and that window is not bounded by any delay the test - can pick -- so synchronise on an actual received message instead of sleeping. - Returns the number of probes published, which is the sequence number the next real - batch will carry. - """ - for probes in range(1, _SUBSCRIBE_ATTEMPTS + 1): - publisher.publish(KVEventBatch(ts=0.0, events=[])) - if subscriber.poll(_PROBE_TIMEOUT_MS): - while subscriber.poll(0): - subscriber.recv_multipart() - return probes - raise _NotReceived("subscription never propagated") - - -def _run_on_fresh_port(scenario: Callable[[int], None]) -> None: - """Retry `scenario(port)` on a fresh port if its sockets could not come up. - - `_unused_tcp_port()` releases its port before the publisher binds it, so another - process can take it in between. Assertion failures inside `scenario` are not - retried. - """ - for _ in range(_ZMQ_SETUP_ATTEMPTS): - try: - scenario(_unused_tcp_port()) - return - except _NotReceived: - pass - except zmq.ZMQError as exc: - if exc.errno != zmq.EADDRINUSE: - raise - pytest.fail(f"ZeroMQ setup failed after {_ZMQ_SETUP_ATTEMPTS} attempts") - - -@streaming_unsupported -def test_streaming_fast_path_publishes_only_full_max_window_blocks() -> None: - """Protect radix hash reuse, filtering, wire format, and shutdown.""" - topic = "kv-events" - context = zmq.Context.instance() - - # Wire hashes come from truncate_sha256_hash_to_int64 = the FIRST 8 bytes - # of the radix key, so put the distinguishing bytes -- including the high - # bit that exercises the signed-wraparound branch -- at the front. - first_hash = b"\x80\x00\x00\x00\x00\x00\x00\x01" + b"\x11" * 24 - partial_hash = b"\x22" * 32 - second_hash = b"\x00\x00\x00\x00\x00\x00\x00\x02" + b"\x33" * 24 - first_wire_hash = int.from_bytes(first_hash[:8], "big") - 2**64 - second_wire_hash = int.from_bytes(second_hash[:8], "big") - - # A fresh manager per attempt restarts sequence numbers at 0 and clears the - # stored-block dedup state, so a retry replays the scenario exactly. - def scenario(port: int) -> None: - bind_endpoint = f"tcp://*:{port}" - subscriber = context.socket(zmq.SUB) - subscriber.setsockopt_string(zmq.SUBSCRIBE, topic) - subscriber.connect(f"tcp://127.0.0.1:{port}") - manager = None - try: - manager = StreamingKVCacheEventManager( - KVEventsConfig( - enable_kv_cache_events=True, - publisher="zmq", - endpoint=bind_endpoint, - topic=topic, - max_queue_size=8, - ), - data_parallel_rank=0, - block_size=4, - max_window_size=128, - ) - manager.start() - base_seq = _await_subscription(manager._publisher, subscriber) - manager.set_layer_group_window_sizes({0: 128, 1: 64}) - - root = SimpleNamespace(ordinal=-1) - - def block(key: bytes, tokens: list[int], prev: object) -> SimpleNamespace: - max_window_page = SimpleNamespace(num_tokens_in_block=len(tokens)) - smaller_window_page = SimpleNamespace(num_tokens_in_block=len(tokens)) - return SimpleNamespace( - key=key, - tokens=tokens, - prev=prev, - ordinal=getattr(prev, "ordinal", -1) + 1, - storage=[lambda: max_window_page, lambda: smaller_window_page], - ) - - first = block(first_hash, [1, 2, 3, 4], root) - partial = block(partial_hash, [5, 6], first) - second = block(second_hash, [5, 6, 7, 8], first) - - manager.add_stored_block_event_from_block(first) - manager.add_stored_block_event_from_block(partial) - manager.add_stored_life_cycle_event_from_block(second, 1) - manager.add_stored_life_cycle_event_from_block(second, 0) - manager.flush_iteration_events() - manager.add_removed_event([first_hash, partial_hash, second_hash]) - manager.flush_iteration_events() - - frames = [] - for _ in range(2): - if not subscriber.poll(_RECEIVE_TIMEOUT_MS): - raise _NotReceived(port) - frames.append(subscriber.recv_multipart()) - - assert [frame[0] for frame in frames] == [topic.encode(), topic.encode()] - # Sequence numbers stay dense across the probes and the real batches. - assert [int.from_bytes(frame[1], "big") for frame in frames] == [ - base_seq, - base_seq + 1, - ] - stored_batch = msgspec.msgpack.decode(frames[0][2]) - removed_batch = msgspec.msgpack.decode(frames[1][2]) - assert stored_batch[2] == 0 - assert stored_batch[1] == [ - { - "type": "BlockStored", - "block_hashes": [first_wire_hash, second_wire_hash], - "parent_block_hash": None, - "token_ids": [1, 2, 3, 4, 5, 6, 7, 8], - "block_size": 4, - "lora_id": None, - "medium": "GPU", - "lora_name": None, - } - ] - assert removed_batch[1] == [ - { - "type": "BlockRemoved", - "block_hashes": [first_wire_hash, second_wire_hash], - "medium": "GPU", - } - ] - assert manager.stored_blocks == 2 - assert manager.removed_blocks == 2 - assert manager.partial_blocks_suppressed == 1 - assert manager.non_target_life_cycles_ignored == 1 - assert manager.dropped_events == 0 - - # Streaming publishing pushes events out-of-band, so the buffered pull - # API must degrade to an empty result rather than raising. - assert manager.get_latest_events() == [] - - # shutdown() must be idempotent and must release the bound port. - manager.shutdown() - manager.shutdown() - replacement = context.socket(zmq.PUB) - replacement.bind(bind_endpoint) - replacement.close(linger=0) - finally: - if manager is not None: - manager.shutdown() - subscriber.close(linger=0) - - _run_on_fresh_port(scenario) - - -@streaming_unsupported -def test_streaming_removals_are_never_dropped_by_the_entry_cap() -> None: - """Removals must survive the per-iteration cap or the consumer desyncs.""" - manager = StreamingKVCacheEventManager( - KVEventsConfig(enable_kv_cache_events=True, publisher="null"), - data_parallel_rank=0, - block_size=2, - max_window_size=128, - max_entries=2, - ) - manager.start() - try: - manager.set_layer_group_window_sizes({0: 128}) - - # Capture what actually reaches the publisher so the test proves the - # removals are emitted on flush, not merely queued in _pending_events. - published: list[object] = [] - manager._publisher.publish = lambda batch: published.append(batch) or True - - root = SimpleNamespace(ordinal=-1) - - def block(key: bytes, tokens: list[int], prev: object) -> SimpleNamespace: - page = SimpleNamespace(num_tokens_in_block=len(tokens)) - return SimpleNamespace( - key=key, - tokens=tokens, - prev=prev, - ordinal=getattr(prev, "ordinal", -1) + 1, - storage=[lambda: page], - ) - - first = block(b"\x01" * 32, [1, 2], root) - second = block(b"\x02" * 32, [3, 4], first) - manager.add_stored_block_event_from_block(first) - manager.add_stored_block_event_from_block(second) - - # Both stores fill the entry cap (max_entries=2); the removals must still - # be emitted rather than dropped, or the consumer treats the blocks as - # resident forever. - manager.add_removed_event([b"\x01" * 32, b"\x02" * 32]) - manager.flush_iteration_events() - - assert manager.removed_blocks == 2 - assert len(published) == 1 - # Round-trip through msgpack to prove the removals reach the wire as a - # BlockRemoved batch carrying both hashes. - decoded = msgspec.msgpack.decode(msgspec.msgpack.encode(published[0])) - removed = [event for event in decoded[1] if event.get("type") == "BlockRemoved"] - assert sum(len(event["block_hashes"]) for event in removed) == 2 - finally: - manager.shutdown() - - def test_dropped_batches_leave_a_sequence_gap() -> None: """A batch lost to a full queue must be observable as a missing sequence number.""" # Left unstarted on purpose: publish() only touches the queue, so the drop path is @@ -293,7 +59,6 @@ def test_dropped_batches_leave_a_sequence_gap() -> None: publisher.shutdown() -@streaming_unsupported def test_construction_binds_nothing_until_start() -> None: """A constructed-but-unstarted publisher must hold no socket and no thread.""" port = _unused_tcp_port() @@ -318,7 +83,6 @@ def test_construction_binds_nothing_until_start() -> None: manager.shutdown() -@streaming_unsupported def test_shutdown_without_start_is_safe() -> None: """Tearing down a manager that never started must not raise.""" manager = StreamingKVCacheEventManager( @@ -332,17 +96,15 @@ def test_shutdown_without_start_is_safe() -> None: manager.shutdown() -def test_validate_streaming_support_rejects_unsupported_setups() -> None: +def test_validate_streaming_support_rejects_unsupported_parallelism() -> None: config = KVEventsConfig(enable_kv_cache_events=True, endpoint="tcp://*:5557") supported = dict(pp_size=1, cp_size=1, ranks_per_host=1, data_parallel_size=1) + validate_streaming_support(config, **supported) with pytest.raises(ValueError, match="pipeline parallelism"): validate_streaming_support(config, **{**supported, "pp_size": 2}) with pytest.raises(ValueError, match="context parallelism"): validate_streaming_support(config, **{**supported, "cp_size": 2}) - # Streaming has no implementation, so even an otherwise supported setup is rejected. - with pytest.raises(ValueError, match="event_buffer_max_size"): - validate_streaming_support(config, **supported) @pytest.mark.parametrize( @@ -354,12 +116,7 @@ def test_validate_streaming_support_rejects_unsupported_setups() -> None: ("tcp://*:5557", "tcp://*:5657", 2, False), # Replay below the publish base overlaps just the same. ("tcp://*:5558", "tcp://*:5557", 2, True), - ( - "tcp://*:5557", - "tcp://*:5559", - 2, - False, - ), + ("tcp://*:5557", "tcp://*:5559", 2, False), # No replay endpoint means no second range to collide with. ("tcp://*:5557", None, 8, False), # 16 attention-DP ranks over 2 nodes collide only within a node, so spacing @@ -379,80 +136,9 @@ def test_validate_endpoint_ranges(endpoint, replay_endpoint, ranks_per_host, ove validate_endpoint_ranges(config, ranks_per_host, ranks_per_host) -@streaming_unsupported -def test_partial_target_page_coverage_is_suppressed_until_fully_covered() -> None: - """A page adopted from a shorter sibling must not be published as a full block.""" - manager = StreamingKVCacheEventManager( - KVEventsConfig(enable_kv_cache_events=True, publisher="null"), - data_parallel_rank=0, - block_size=4, - max_window_size=128, - ) - manager.start() - try: - manager.set_layer_group_window_sizes({0: 128}) - published: list[object] = [] - manager._publisher.publish = lambda batch: published.append(batch) or True - - root = SimpleNamespace(ordinal=-1) - # The block holds 4 tokens but its target page only covers 2 of them. - page = SimpleNamespace(num_tokens_in_block=2) - block = SimpleNamespace( - key=b"\x01" * 32, - tokens=[1, 2, 3, 4], - prev=root, - ordinal=0, - storage=[lambda: page], - ) - - manager.add_stored_block_event_from_block(block) - manager.flush_iteration_events() - assert manager.stored_blocks == 0 - assert manager.partial_blocks_suppressed == 1 - assert published == [] - - # Once the page covers the whole block, the same block is published. - page.num_tokens_in_block = 4 - manager.add_stored_life_cycle_event_from_block(block, 0) - manager.flush_iteration_events() - assert manager.stored_blocks == 1 - assert len(published) == 1 - decoded = msgspec.msgpack.decode(msgspec.msgpack.encode(published[0])) - stored = [event for event in decoded[1] if event["type"] == "BlockStored"] - assert sum(len(event["block_hashes"]) for event in stored) == 1 - finally: - manager.shutdown() - - -@streaming_unsupported -def test_life_cycle_hooks_ignore_none_ids() -> None: - """A None life-cycle id must not reach int() before the target is configured.""" - manager = StreamingKVCacheEventManager( - KVEventsConfig(enable_kv_cache_events=True, publisher="null"), - data_parallel_rank=0, - block_size=4, - max_window_size=128, - ) - manager.start() - try: - # Before set_layer_group_window_sizes(), and with a None id, both hooks are - # no-ops rather than raising TypeError. - manager.add_stored_life_cycle_event_from_block(object(), None) - manager.add_removed_life_cycle_event(b"\x01" * 32, None) - manager.set_layer_group_window_sizes({0: 128}) - manager.add_stored_life_cycle_event_from_block(object(), None) - manager.add_removed_life_cycle_event(b"\x01" * 32, None) - assert manager.stored_blocks == 0 - assert manager.removed_blocks == 0 - finally: - manager.shutdown() - - @pytest.mark.parametrize( "endpoint,replay_endpoint,dp_size,overflows", [ - # rank 1 would bind 65536, which offset_endpoint_port rejects on that rank - # alone -- before rank 0 reaches the collective. ("tcp://*:65535", None, 2, True), ("tcp://*:65535", None, 1, False), ("tcp://*:65534", None, 2, False), @@ -477,11 +163,10 @@ def test_validate_endpoint_ranges_rejects_port_overflow( validate_endpoint_ranges(config, 1, dp_size) -@streaming_unsupported def test_validate_streaming_support_rejects_overflowing_port_span() -> None: """Every rank must reject the span, so none reaches the following collective.""" config = KVEventsConfig(enable_kv_cache_events=True, endpoint="tcp://*:65535") - supported = dict(pp_size=1, cp_size=1, ranks_per_host=1, backend="python") + supported = dict(pp_size=1, cp_size=1, ranks_per_host=1) # One rank fits; two do not, and rank 0 must refuse it just as rank 1 would. validate_streaming_support(config, **supported, data_parallel_size=1) with pytest.raises(ValueError, match="above the maximum port 65535"):