From 1c9731d2741b45019339747f15594bd83b2df5ec Mon Sep 17 00:00:00 2001 From: Guan Luo Date: Wed, 23 Sep 2026 16:31:50 -0700 Subject: [PATCH 1/3] [None][fix] preserve multimodal identity in V2 KV events Signed-off-by: Guan Luo --- .../kv_cache_manager_v2/blockRadixTree.cpp | 10 +- .../kv_cache_manager_v2/blockRadixTree.h | 12 +-- .../kv_cache_manager_v2/eventManager.cpp | 21 ++-- .../kv_cache_manager_v2/eventManager.h | 11 +- .../kv_cache_manager_v2/tokenIdExt.cpp | 73 +++++++++---- .../kv_cache_manager_v2/tokenIdExt.h | 29 ++++- .../nanobind/batch_manager/bindings.cpp | 10 ++ .../batch_manager/kvCacheManagerV2.cpp | 100 +++++++++++++++--- .../kvCacheManagerV2DigestPoolTest.cpp | 39 ++++++- docs/source/features/kvcache.md | 8 +- .../kv_cache/kv_cache_manager_v2.py | 28 ++++- .../_torch/pyexecutor/kv_cache_events.py | 12 ++- tensorrt_llm/_utils.py | 20 ++-- tensorrt_llm/inputs/data.py | 14 +-- tensorrt_llm/inputs/multimodal.py | 6 +- tensorrt_llm/inputs/registry.py | 5 +- .../runtime/kv_cache_manager_v2/__init__.py | 5 +- .../runtime/kv_cache_manager_v2/__init__.pyi | 15 ++- .../kv_cache_manager_v2/_block_radix_tree.py | 65 +++++++++--- .../runtime/kv_cache_manager_v2/_common.py | 28 ++++- .../kv_cache_manager_v2/_core/_kv_cache.py | 3 +- .../kv_cache_manager_v2/_event_manager.py | 28 +++-- .../test_kv_cache_v2_multimodal_runs.py | 58 +++++++++- .../multimodal/test_mm_encoder_standalone.py | 22 +++- .../kv_cache_manager_v2_tests/kernels.py | 17 ++- .../test_kv_cache_event_manager.py | 82 +++++++++++--- .../test_streaming_kv_events.py | 36 +++++++ .../llmapi/test_llm_kv_cache_events.py | 20 ++++ 28 files changed, 631 insertions(+), 146 deletions(-) diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/blockRadixTree.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/blockRadixTree.cpp index d38b7ec2d7ef..21caeba2a180 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/blockRadixTree.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/blockRadixTree.cpp @@ -193,8 +193,8 @@ BlockKey Hasher::digest() const // genMultimodalCacheKeyTokens // --------------------------------------------------------------------------- -std::vector genMultimodalCacheKeyTokens( - int idOffset, std::vector const& multiModalDataDigest, int numTokens, int tokenOffset) +std::vector genMultimodalCacheKeyTokens(int idOffset, std::vector const& multiModalDataDigest, + int numTokens, int tokenOffset, std::optional uuid) { TLLM_CHECK(numTokens > 0); TLLM_CHECK(tokenOffset >= 0); @@ -207,7 +207,7 @@ std::vector genMultimodalCacheKeyTokens( { Digest digest; std::memcpy(digest.data(), multiModalDataDigest.data(), kDIGEST_LEN); - result.emplace_back(digest); + result.emplace_back(digest, std::move(uuid)); } else { @@ -330,13 +330,13 @@ Block::Block(BlockKey k, std::vector toks, NodeBase* prevNode) { if (prevNode->type() == Type::kBLOCK) { - mLastTokenDigest = static_cast(prevNode)->getLastTokenDigest(); + mLastMmItemContext = static_cast(prevNode)->getLastMmItemContext(); } for (auto iter = tokens.rbegin(); iter != tokens.rend(); ++iter) { if (iter->isDigest()) { - mLastTokenDigest = std::make_shared(iter->digest()); + mLastMmItemContext = iter->sharedMmItemContext(); break; } } diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/blockRadixTree.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/blockRadixTree.h index bdae500113db..bf38058818e7 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/blockRadixTree.h +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/blockRadixTree.h @@ -160,8 +160,8 @@ inline auto sequenceToBlockchainKeys( } // Generate multi-modal token IDs (mirrors gen_multimodal_cache_key_tokens in Python). -std::vector genMultimodalCacheKeyTokens( - int idOffset, std::vector const& multiModalDataDigest, int numTokens, int tokenOffset = 0); +std::vector genMultimodalCacheKeyTokens(int idOffset, std::vector const& multiModalDataDigest, + int numTokens, int tokenOffset = 0, std::optional uuid = std::nullopt); // --------------------------------------------------------------------------- // NodeBase — common base for RootBlock and Block (nodes in the radix tree). @@ -273,10 +273,10 @@ struct Block : NodeBase, EnableSharedFromThis return storage.size(); } - //! Latest digest token in this block's prefix, when requested by the event sink. - std::shared_ptr const& getLastTokenDigest() const noexcept + //! Latest multimodal item context in this block's prefix, when requested by the event sink. + std::shared_ptr const& getLastMmItemContext() const noexcept { - return mLastTokenDigest; + return mLastMmItemContext; } bool isFull() const noexcept @@ -344,7 +344,7 @@ struct Block : NodeBase, EnableSharedFromThis BlockOrdinal mOrdinal; // Share an immutable value through descendants without retaining any ancestor block. // Unlike prev, this context remains valid while the block is detached from the tree. - std::shared_ptr mLastTokenDigest; + std::shared_ptr mLastMmItemContext; }; // --------------------------------------------------------------------------- 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..cd93d13f60dd 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 @@ -567,10 +567,10 @@ std::optional EventManager::storedBlockFromBlock( } std::vector mmKeys; - Digest const* itemDigest = nullptr; + MmItemContext const* itemContext = nullptr; if (mMmTokenIdOffset.has_value() && block.prev != nullptr && block.prev->type() == NodeBase::Type::kBLOCK) { - itemDigest = static_cast(block.prev)->getLastTokenDigest().get(); + itemContext = static_cast(block.prev)->getLastMmItemContext().get(); } bool inMmRun = false; std::vector tokens; @@ -582,13 +582,14 @@ std::optional EventManager::storedBlockFromBlock( 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 (itemContext != nullptr && token.tokenId() > *mMmTokenIdOffset) { if (!inMmRun) { - mmKeys.push_back( - {std::string(reinterpret_cast(itemDigest->data()), itemDigest->size()), - token.tokenId() - *mMmTokenIdOffset, std::nullopt, false}); + mmKeys.push_back({std::string(reinterpret_cast(itemContext->digest.data()), + itemContext->digest.size()), + token.tokenId() - *mMmTokenIdOffset, itemContext->uuid, + itemContext->uuid.has_value() ? MmKeyUuidMode::kAdditive : MmKeyUuidMode::kNone}); } inMmRun = true; } @@ -605,9 +606,11 @@ std::optional EventManager::storedBlockFromBlock( 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}); + itemContext = &token.mmItemContext(); + mmKeys.push_back( + {std::string(reinterpret_cast(itemContext->digest.data()), itemContext->digest.size()), + 0, itemContext->uuid, + itemContext->uuid.has_value() ? MmKeyUuidMode::kAdditive : MmKeyUuidMode::kNone}); inMmRun = true; } } 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..8edf442c0533 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 @@ -64,17 +64,24 @@ struct KVCacheCreatedData } }; +enum class MmKeyUuidMode : uint8_t +{ + kNone, + kReplacesHash, + kAdditive, +}; + struct MmKey { std::string hash; int startOffset = 0; std::optional uuid; - bool hasUuidField = false; + MmKeyUuidMode uuidMode = MmKeyUuidMode::kNone; bool operator==(MmKey const& other) const { return hash == other.hash && startOffset == other.startOffset && uuid == other.uuid - && hasUuidField == other.hasUuidField; + && uuidMode == other.uuidMode; } }; diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/tokenIdExt.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/tokenIdExt.cpp index b6ab10dab190..4b257ba9d362 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/tokenIdExt.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/tokenIdExt.cpp @@ -24,14 +24,16 @@ #include #include #include +#include #include +#include namespace tensorrt_llm::batch_manager::kv_cache_manager_v2 { // --------------------------------------------------------------------------- -// DigestPool — process-global, address-stable store of 32-byte multi-modal -// Digests, referenced by a 31-bit slot index packed into a TokenIdExt. It is a +// DigestPool — process-global, address-stable store of immutable multi-modal +// item contexts, referenced by a 31-bit slot index packed into a TokenIdExt. It is a // pure implementation detail of TokenIdExt, so it lives here (anonymous // namespace) rather than in the header. // @@ -74,28 +76,35 @@ class DigestPool DigestPool(DigestPool&&) = delete; DigestPool& operator=(DigestPool&&) = delete; - // Store a copy of `digest` in the lowest free slot; return its index. - uint32_t alloc(Digest const& digest) + // Store `context` in the lowest free slot; return its index. + uint32_t alloc(MmItemContext context) { std::lock_guard const lock(mMutex); - return allocLocked(digest); + return allocLocked(std::make_shared(std::move(context))); } - // Duplicate the digest at slot `idx` into a fresh slot; return the new index. + // Share the context at slot `idx` through a fresh slot; return the new index. uint32_t duplicate(uint32_t idx) { std::lock_guard const lock(mMutex); - return allocLocked(mStore[idx]); // Safe for deque + return allocLocked(liveSlotLocked(idx)); // Share immutable context, including long UUID strings. } - // The digest at slot `idx`. The reference stays valid after the lock is + // The context at slot `idx`. The reference stays valid after the lock is // released and across later shrinks (which only pop free tail slots). Uses // at() so a bad index (e.g. the sentinel of a moved-from handle) throws - // rather than reading out of bounds; digests are rare so the check is cheap. - [[nodiscard]] Digest const& get(uint32_t idx) const + // rather than reading out of bounds; contexts are rare so the check is cheap. + [[nodiscard]] MmItemContext const& get(uint32_t idx) const { std::lock_guard const lock(mMutex); - return mStore.at(idx); + return *liveSlotLocked(idx); + } + + // Shared ownership of the context at slot `idx`. + [[nodiscard]] std::shared_ptr getShared(uint32_t idx) const + { + std::lock_guard const lock(mMutex); + return liveSlotLocked(idx); } // Clear slot `idx` and reclaim trailing free slots. @@ -106,6 +115,7 @@ class DigestPool return; // the default / moved-from sentinel index — nothing to free } std::lock_guard const lock(mMutex); + mStore[idx].reset(); mInUse.clear(idx); if (idx < mMinFreeHint) { @@ -123,6 +133,14 @@ class DigestPool private: DigestPool() = default; + // Precondition: caller holds mMutex. Validate that `idx` names a live slot. + [[nodiscard]] std::shared_ptr const& liveSlotLocked(uint32_t idx) const + { + auto const& slot = mStore.at(idx); + TLLM_CHECK_WITH_INFO(slot != nullptr, "DigestPool slot %u is not live", idx); + return slot; + } + // Slot count. Also checks the occupancy bitset stays sized to the store. // Precondition: caller holds mMutex and mStore/mInUse are in sync (i.e. not // called between growing/shrinking one and resizing the other). @@ -132,9 +150,9 @@ class DigestPool return mStore.size(); } - // Precondition: caller holds mMutex. Store `digest` in the lowest free slot + // Precondition: caller holds mMutex. Store `context` in the lowest free slot // (front-packing), growing the deque only when no free slot exists. - uint32_t allocLocked(Digest const& digest) + uint32_t allocLocked(std::shared_ptr context) { size_t const cap = capacity(); size_t idx = mMinFreeHint; @@ -147,12 +165,12 @@ class DigestPool // No free slot below the high-water mark — grow by one. Indices stay // strictly below kValueMask, which is reserved as the bad-handle sentinel. TLLM_CHECK_WITH_INFO(cap < TokenIdExt::kValueMask, "DigestPool exhausted the 31-bit index space"); - mStore.push_back(digest); + mStore.push_back(std::move(context)); mInUse.resize(mStore.size()); // re-sync the bitset; new bit is clear } else { - mStore[idx] = digest; + mStore[idx] = std::move(context); } mInUse.set(idx); mMinFreeHint = idx + 1; // everything below is now occupied @@ -191,9 +209,9 @@ class DigestPool static constexpr size_t kSlackLow = 64; mutable std::mutex mMutex; - std::deque mStore; // slot storage; mStore.size() == the slot count (== bitset capacity) - DynamicBitset mInUse{0}; // bit i set == slot i occupied - size_t mMinFreeHint{0}; // lower bound on the lowest free slot index + std::deque> mStore; // address-stable slot storage + DynamicBitset mInUse{0}; // bit i set == slot i occupied + size_t mMinFreeHint{0}; // lower bound on the lowest free slot index }; } // namespace @@ -202,8 +220,13 @@ class DigestPool // TokenIdExt — RAII members that touch the pool (construct/copy=alloc, dtor=free). // --------------------------------------------------------------------------- -TokenIdExt::TokenIdExt(Digest const& digestValue) - : mBits(DigestPool::instance().alloc(digestValue) | kTagMask) +TokenIdExt::TokenIdExt(Digest const& digestValue, std::optional uuid) + : TokenIdExt(MmItemContext{digestValue, std::move(uuid)}) +{ +} + +TokenIdExt::TokenIdExt(MmItemContext context) + : mBits(DigestPool::instance().alloc(std::move(context)) | kTagMask) { } @@ -218,10 +241,20 @@ uint32_t TokenIdExt::duplicateSlot(uint32_t index) } Digest const& TokenIdExt::digest() const +{ + return mmItemContext().digest; +} + +MmItemContext const& TokenIdExt::mmItemContext() const { return DigestPool::instance().get(digestIndex()); } +std::shared_ptr TokenIdExt::sharedMmItemContext() const +{ + return DigestPool::instance().getShared(digestIndex()); +} + namespace detail { diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/tokenIdExt.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/tokenIdExt.h index c0248361dd9e..614d435a7dbd 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/tokenIdExt.h +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/tokenIdExt.h @@ -23,6 +23,9 @@ #include #include #include +#include +#include +#include #include #include @@ -51,6 +54,19 @@ struct alignas(kDIGEST_LEN) Digest : std::array } }; +//! Multimodal item metadata carried by a digest token. UUID does not +//! participate in equality or radix-tree hashing. +struct MmItemContext +{ + Digest digest; + std::optional uuid; + + bool operator==(MmItemContext const& other) const noexcept + { + return digest == other.digest; + } +}; + // --------------------------------------------------------------------------- // TokenIdExt — 4-byte self-describing token handle (RAII value type). // @@ -58,7 +74,7 @@ struct alignas(kDIGEST_LEN) Digest : std::array // - tag 0: normal token id (stored verbatim). An all-normal array is a // contiguous little-endian int32 array, hashed in one CSHA256::Write(N*4). // - tag 1: multi-modal digest; low bits index a slot in an internal pool that -// holds the 32-byte Digest. +// holds immutable digest and optional UUID context. // // A digest handle owns its pool slot: construct from a Digest to allocate, copy // to clone into a fresh slot, and destroy to free. Normal handles own nothing. @@ -89,8 +105,9 @@ class TokenIdExt TLLM_CHECK_DEBUG(id >= 0 && id <= static_cast(kMaxValue)); } - // Multi-modal digest (tag 1): copies `digest` into a fresh pool slot. - explicit TokenIdExt(Digest const& digest); + // Multi-modal digest (tag 1): stores immutable item context in a fresh pool slot. + explicit TokenIdExt(Digest const& digest, std::optional uuid = std::nullopt); + explicit TokenIdExt(MmItemContext context); ~TokenIdExt() { @@ -155,6 +172,12 @@ class TokenIdExt return static_cast(mBits); } + // The pooled multimodal item context. Precondition: isDigest(). + [[nodiscard]] MmItemContext const& mmItemContext() const; + + // Shared ownership of the pooled multimodal item context. Precondition: isDigest(). + [[nodiscard]] std::shared_ptr sharedMmItemContext() const; + // The pooled 32-byte digest. Precondition: isDigest(). [[nodiscard]] Digest const& digest() const; diff --git a/cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp b/cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp index 7209556bb6db..7c817cd93c85 100644 --- a/cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp @@ -259,6 +259,16 @@ void initBindings(nb::module_& m) } return hashes; }) + .def_prop_ro("multimodal_uuids", + [](GenLlmReq& self) + { + std::optional>> uuids = std::nullopt; + if (self.getMultimodalUuids()) + { + uuids = *self.getMultimodalUuids().value(); + } + return uuids; + }) .def_prop_ro("multimodal_positions", [](GenLlmReq& self) { diff --git a/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp b/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp index f45743b79290..ae4393d95017 100644 --- a/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp +++ b/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp @@ -231,7 +231,12 @@ static std::pair, bool> castTokenIterable(nb::handle bool knownNoDigest = true; for (auto item : nb::cast(tokens)) { - if (nb::isinstance(item)) + if (nb::isinstance(item)) + { + vec.emplace_back(nb::cast(item)); + knownNoDigest = false; + } + else if (nb::isinstance(item)) { auto b = nb::cast(item); if (nb::len(b) != kv::kDIGEST_LEN) @@ -372,8 +377,15 @@ static nb::list tokenList(std::vector const& tokens) } else { - auto const& digest = tok.digest(); - result.append(nb::bytes(reinterpret_cast(digest.data()), digest.size())); + auto const& context = tok.mmItemContext(); + if (context.uuid.has_value()) + { + result.append(nb::cast(context)); + } + else + { + result.append(nb::bytes(reinterpret_cast(context.digest.data()), context.digest.size())); + } } } return result; @@ -440,12 +452,12 @@ static std::vector castMmKeys(nb::handle values) for (nb::handle value : nb::cast(values)) { nb::tuple tuple = nb::cast(value); - if (tuple.size() != 2 && tuple.size() != 3) + if (tuple.size() != 2 && tuple.size() != 3 && tuple.size() != 4) { - throw std::invalid_argument("mm_key must have two or three entries"); + throw std::invalid_argument("mm_key must have two, three, or four entries"); } std::optional uuid; - if (tuple.size() == 3 && !tuple[2].is_none()) + if (tuple.size() >= 3 && !tuple[2].is_none()) { uuid = nb::cast(tuple[2]); } @@ -454,8 +466,14 @@ static std::vector castMmKeys(nb::handle values) throw std::invalid_argument("mm_key hash must be bytes"); } auto hash = nb::cast(tuple[0]); + kv::MmKeyUuidMode uuidMode = kv::MmKeyUuidMode::kNone; + if (tuple.size() >= 3) + { + uuidMode = tuple.size() == 4 && nb::cast(tuple[3]) ? kv::MmKeyUuidMode::kAdditive + : kv::MmKeyUuidMode::kReplacesHash; + } result.push_back(kv::MmKey{std::string(hash.c_str(), static_cast(nb::len(hash))), - nb::cast(tuple[1]), std::move(uuid), tuple.size() == 3}); + nb::cast(tuple[1]), std::move(uuid), uuidMode}); } return result; } @@ -466,13 +484,15 @@ static nb::list castMmKeys(kv::KVCacheStoredBlockData const& data) for (auto const& mmKey : data.mmKeys) { auto hash = nb::bytes(mmKey.hash.data(), mmKey.hash.size()); - if (mmKey.hasUuidField) + switch (mmKey.uuidMode) { + case kv::MmKeyUuidMode::kAdditive: + result.append(nb::make_tuple(std::move(hash), mmKey.startOffset, mmKey.uuid, true)); + break; + case kv::MmKeyUuidMode::kReplacesHash: result.append(nb::make_tuple(std::move(hash), mmKey.startOffset, mmKey.uuid)); - } - else - { - result.append(nb::make_tuple(std::move(hash), mmKey.startOffset)); + break; + case kv::MmKeyUuidMode::kNone: result.append(nb::make_tuple(std::move(hash), mmKey.startOffset)); break; } } return result; @@ -884,6 +904,53 @@ void KvCacheManagerV2Bindings::initBindings(nb::module_& m) .value("SUSPENDED", kv::KvCache::Status::SUSPENDED) .value("CLOSED", kv::KvCache::Status::CLOSED); + nb::class_(m, "MmItemContext") + .def( + "__init__", + [](kv::MmItemContext* self, nb::bytes digestBytes, std::optional uuid) + { + if (nb::len(digestBytes) != kv::kDIGEST_LEN) + { + throw std::invalid_argument("digest must have length kDIGEST_LEN"); + } + kv::Digest digest; + std::memcpy(digest.data(), digestBytes.c_str(), kv::kDIGEST_LEN); + new (self) kv::MmItemContext{digest, std::move(uuid)}; + }, + nb::arg("digest"), nb::arg("uuid") = std::nullopt) + .def_prop_ro("digest", + [](kv::MmItemContext const& self) + { return nb::bytes(reinterpret_cast(self.digest.data()), self.digest.size()); }) + .def_ro("uuid", &kv::MmItemContext::uuid) + .def( + "__eq__", + [](kv::MmItemContext const& self, nb::handle other) + { + if (nb::isinstance(other)) + { + return self == nb::cast(other); + } + if (nb::isinstance(other) && nb::len(other) == kv::kDIGEST_LEN) + { + auto const bytes = nb::cast(other); + return std::memcmp(self.digest.data(), bytes.c_str(), kv::kDIGEST_LEN) == 0; + } + return false; + }, + nb::arg("other")) + .def("__hash__", + [](kv::MmItemContext const& self) + { + auto digest = nb::bytes(reinterpret_cast(self.digest.data()), self.digest.size()); + return PyObject_Hash(digest.ptr()); + }) + .def("__reduce__", + [](kv::MmItemContext const& self) + { + auto digest = nb::bytes(reinterpret_cast(self.digest.data()), self.digest.size()); + return nb::make_tuple(nb::type(), nb::make_tuple(std::move(digest), self.uuid)); + }); + // ---- KV cache events ---------------------------------------------------- nb::class_(m, "UniqueToken") .def(nb::init(), nb::arg("token_id"), nb::arg("token_extra_id") = 0) @@ -2578,7 +2645,8 @@ void KvCacheManagerV2Bindings::initBindings(nb::module_& m) m.def( "gen_multimodal_cache_key_tokens", - [](int idOffset, nb::bytes multiModalDataDigest, int numTokens, int tokenOffset) + [](int idOffset, nb::bytes multiModalDataDigest, int numTokens, int tokenOffset, + std::optional uuid) { auto const digestSize = nb::len(multiModalDataDigest); if (digestSize != kv::kDIGEST_LEN) @@ -2587,9 +2655,11 @@ void KvCacheManagerV2Bindings::initBindings(nb::module_& m) } auto const* first = reinterpret_cast(multiModalDataDigest.c_str()); std::vector digest(first, first + digestSize); - return tokenList(kv::genMultimodalCacheKeyTokens(idOffset, digest, numTokens, tokenOffset)); + return tokenList( + kv::genMultimodalCacheKeyTokens(idOffset, digest, numTokens, tokenOffset, std::move(uuid))); }, - nb::arg("id_offset"), nb::arg("multi_modal_data_digest"), nb::arg("num_tokens"), nb::arg("token_offset") = 0); + nb::arg("id_offset"), nb::arg("multi_modal_data_digest"), nb::arg("num_tokens"), nb::arg("token_offset") = 0, + nb::arg("uuid") = std::nullopt); // Lazy iterator yielding (token_block, key) pairs; hashes one block per __next__. nb::class_(m, "_BlockchainKeyIterator") .def("__iter__", [](nb::handle self) { return self; }) diff --git a/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2DigestPoolTest.cpp b/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2DigestPoolTest.cpp index 77f03b74e490..44092372d99a 100644 --- a/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2DigestPoolTest.cpp +++ b/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2DigestPoolTest.cpp @@ -27,6 +27,7 @@ #include #include #include +#include #include #include @@ -78,12 +79,14 @@ TEST(DigestPoolTest, DigestValueEqualityAcrossDistinctSlots) size_t const baseline = detail::digestPoolLiveCount(); Digest const bytes = makeDigest(std::byte{0x42}); - TokenIdExt const tokA(bytes); - TokenIdExt const tokB(bytes); + TokenIdExt const tokA(bytes, "frontend-a"); + TokenIdExt const tokB(bytes, "frontend-b"); EXPECT_EQ(detail::digestPoolLiveCount(), baseline + 2); // two distinct slots ASSERT_TRUE(tokA.isDigest()); EXPECT_EQ(tokA, tokB); // by-value (pooled) equality - EXPECT_NE(tokA, TokenIdExt(TokenId{5})); // digest != normal + EXPECT_EQ(tokA.mmItemContext().uuid, "frontend-a"); + EXPECT_EQ(tokB.mmItemContext().uuid, "frontend-b"); + EXPECT_NE(tokA, TokenIdExt(TokenId{5})); // digest != normal // Hashing the two distinct-slot digests yields the same contribution. Hasher hashA; @@ -98,13 +101,16 @@ TEST(DigestPoolTest, CopyDigestTokenClonesSlot) { size_t const baseline = detail::digestPoolLiveCount(); Digest const bytes = makeDigest(std::byte{0x5A}); + std::string const uuid(4096, 'u'); { - TokenIdExt const original(bytes); + TokenIdExt const original(bytes, uuid); EXPECT_EQ(detail::digestPoolLiveCount(), baseline + 1); TokenIdExt const copy = original; // clone → second slot EXPECT_EQ(detail::digestPoolLiveCount(), baseline + 2); EXPECT_EQ(original, copy); EXPECT_EQ(copy.digest(), bytes); + EXPECT_EQ(copy.mmItemContext().uuid, uuid); + EXPECT_EQ(&original.mmItemContext(), ©.mmItemContext()); } EXPECT_EQ(detail::digestPoolLiveCount(), baseline); // both slots freed } @@ -116,13 +122,14 @@ TEST(DigestPoolTest, MoveTransfersSlotWithoutCloning) Digest const bytes = makeDigest(std::byte{0x33}); Digest const other = makeDigest(std::byte{0x77}); { - TokenIdExt source(bytes); + TokenIdExt source(bytes, "routing-identity"); EXPECT_EQ(detail::digestPoolLiveCount(), baseline + 1); TokenIdExt moved(std::move(source)); EXPECT_EQ(detail::digestPoolLiveCount(), baseline + 1); EXPECT_EQ(source.raw(), TokenIdExt::kBadToken); EXPECT_EQ(moved.digest(), bytes); + EXPECT_EQ(moved.mmItemContext().uuid, "routing-identity"); EXPECT_THROW((void) (source == moved), std::out_of_range); TokenIdExt target(other); @@ -248,4 +255,26 @@ TEST(DigestPoolTest, MixedBlockHashesDeterministically) hb.update(b.data(), b.size()); EXPECT_EQ(ha.digest(), hb.digest()); } + +TEST(DigestPoolTest, UuidMetadataDoesNotChangeTokenIdentityOrBlockHash) +{ + Digest const digest = makeDigest(std::byte{0x2A}); + std::vector digestBytes(kDIGEST_LEN); + std::memcpy(digestBytes.data(), digest.data(), kDIGEST_LEN); + auto const withoutUuid = genMultimodalCacheKeyTokens(1000, digestBytes, 3); + auto const withUuid + = genMultimodalCacheKeyTokens(1000, digestBytes, 3, /*tokenOffset=*/0, "frontend-routing-identity"); + + ASSERT_EQ(withUuid.size(), withoutUuid.size()); + EXPECT_EQ(withUuid, withoutUuid); + ASSERT_TRUE(withUuid.front().isDigest()); + EXPECT_EQ(withUuid.front().digest(), digest); + EXPECT_EQ(withUuid.front().mmItemContext().uuid, "frontend-routing-identity"); + + Hasher withoutUuidHash; + withoutUuidHash.update(withoutUuid.data(), withoutUuid.size()); + Hasher withUuidHash; + withUuidHash.update(withUuid.data(), withUuid.size()); + EXPECT_EQ(withUuidHash.digest(), withoutUuidHash.digest()); +} } // namespace diff --git a/docs/source/features/kvcache.md b/docs/source/features/kvcache.md index 7ccea7387c3a..8edc9a71ea8e 100644 --- a/docs/source/features/kvcache.md +++ b/docs/source/features/kvcache.md @@ -172,7 +172,7 @@ This isolation is enforced entirely by the block-key hash: the salt is mixed int When working with multimodal models (e.g., vision-language models), the KV cache system needs to identify which cached blocks correspond to which multimodal inputs (images, videos, etc.). By default, the system uses content-based hashing to generate unique identifiers for each multimodal input. However, this approach has limitations for cache management across sessions, as the same content must be re-processed to generate the same hash. -You can provide custom UUID strings for your multimodal data using the `multi_modal_uuids` parameter when creating requests. Both cache managers compute the item digest from **both** the UUID and content together for correctness. V1 returns the original UUID in the KV cache event's `mm_keys[].hash` field when one is supplied. V2 returns the item digest as a hexadecimal string, including for items with UUIDs. +You can provide custom UUID strings for your multimodal data using the `multi_modal_uuids` parameter when creating requests. Both cache managers compute the item digest from **both** the UUID and content together for correctness. V1 returns the original UUID in the KV cache event's `mm_keys[].hash` field when one is supplied. V2 always returns the derived digest in `mm_keys[].hash` and adds the original identity in the optional `mm_keys[].uuid` field. **Usage Example:** @@ -191,16 +191,16 @@ prompt = TextPrompt( - **Cache Correctness**: When a UUID is provided, the cache key is computed from both the UUID and content together using `BLAKE3(UUID || Content)`. This ensures different content always produces different cache entries, even with the same UUID. - **User Isolation**: Same content with different UUIDs produces different cache entries, enabling per-user or per-session cache isolation. -- **Stable Event Identifiers**: `get_kv_cache_events()` returns the original UUID for V1, or the item digest for V2. V2 consumers can use the same digest that appears in its cache-key token sequence. +- **Stable Event Identifiers**: `get_kv_cache_events()` returns the original UUID in V1's `hash` field. V2 returns the item digest in `hash` and the original UUID in the optional `uuid` field. - **Partial UUID Support**: You can provide UUIDs for some items and use `None` for others to fall back to content-only hashing. - **Cross-Modality Support**: Different modalities (images, videos) can each have their own UUIDs. **UUID Format:** - Can be any string (e.g., "image-123", "user-session-img-a", database keys) -- Original UUID strings are preserved in request metadata and returned in V1 KV cache events +- Original UUID strings are preserved in request metadata and returned in KV cache events -V2 derives `mm_keys` directly from the cached token sequence. Each entry identifies a continuous multimodal segment within that block: `hash` is the item's digest, and `start_offset` is the segment's first token offset within the item. An item spanning multiple blocks retains the same digest with increasing offsets. Text may separate segments of the same item. Items are processed in prompt order; one item cannot resume after another item has started. The item digest is distinct from `block_hash`, which also depends on the preceding token sequence. +V2 derives `mm_keys` directly from the cached token sequence. Each entry identifies a continuous multimodal segment within that block: `hash` is the item's digest, optional `uuid` is the externally supplied identity, and `start_offset` is the segment's first token offset within the item. An item spanning multiple blocks retains the same digest and UUID with increasing offsets. Text may separate segments of the same item. Items are processed in prompt order; one item cannot resume after another item has started. The item digest is distinct from `block_hash`, which also depends on the preceding token sequence. Clients that require the externally supplied identity must read `uuid`; truncating or otherwise deriving an identity from `hash` does not recover the UUID. ### Enable Offloading to Host Memory 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 2b47f3f5982d..e6c8caac1feb 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 @@ -826,6 +826,7 @@ def _augment_tokens_with_mm_run_metadata( vocab_size: int, result: list[TokenIdExt], multimodal_hashes: Sequence[Sequence[int]], + multimodal_uuids: Sequence[str | None] | None, metadata: _MmRunMetadata, chunk_start: int, chunk_end: int, @@ -860,6 +861,7 @@ def _augment_tokens_with_mm_run_metadata( current_item_idx: Optional[int] = None digest = b"" + uuid: str | None = None for item_idx, chunk_result_offset, item_token_offset, length in zip( overlap_run_item_indices.tolist(), chunk_result_offsets.tolist(), @@ -870,13 +872,22 @@ def _augment_tokens_with_mm_run_metadata( if item_idx != current_item_idx: current_item_idx = item_idx digest = _hash_to_digest(multimodal_hashes[item_idx]) + uuid = ( + multimodal_uuids[item_idx] + if multimodal_uuids is not None and item_idx < len(multimodal_uuids) + else None + ) # Feed the coarse item property (content digest) and granular run # properties (item-local offset and span length) into the key # generator, so cache keys reflect the actual multimodal tokens being # rewritten. result[chunk_result_offset : chunk_result_offset + length] = ( gen_multimodal_cache_key_tokens( - vocab_size, digest, length, token_offset=item_token_offset + vocab_size, + digest, + length, + token_offset=item_token_offset, + uuid=uuid, ) ) @@ -887,6 +898,7 @@ def _augment_tokens_with_contiguous_mm_metadata( vocab_size: int, result: list[TokenIdExt], multimodal_hashes: Sequence[Sequence[int]], + multimodal_uuids: Sequence[str | None] | None, multimodal_positions: Sequence[int] | torch.Tensor, multimodal_lengths: Sequence[int] | torch.Tensor, chunk_start: int, @@ -914,6 +926,11 @@ def _augment_tokens_with_contiguous_mm_metadata( _hash_to_digest(multimodal_hashes[item_idx]), overlap_length, token_offset=source_offset, + uuid=( + multimodal_uuids[item_idx] + if multimodal_uuids is not None and item_idx < len(multimodal_uuids) + else None + ), ) return result @@ -4217,13 +4234,20 @@ def _augment_tokens_for_block_reuse( run_metadata = _resolve_multimodal_run_metadata(req) if run_metadata is not None: return _augment_tokens_with_mm_run_metadata( - self.vocab_size, result, req.multimodal_hashes, run_metadata, chunk_start, chunk_end + self.vocab_size, + result, + req.multimodal_hashes, + req.multimodal_uuids, + run_metadata, + chunk_start, + chunk_end, ) return _augment_tokens_with_contiguous_mm_metadata( self.vocab_size, result, req.multimodal_hashes, + req.multimodal_uuids, req.multimodal_positions, req.multimodal_lengths, chunk_start, diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_events.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_events.py index f2d4dab1f7b7..b2f5fce4b722 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_events.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_events.py @@ -40,6 +40,7 @@ from tensorrt_llm.llmapi.llm_args import KVEventsConfig from tensorrt_llm.logger import logger from tensorrt_llm.runtime.kv_cache_hash import truncate_sha256_hash_to_int64 +from tensorrt_llm.runtime.kv_cache_manager_v2 import MmItemContext from tensorrt_llm.runtime.kv_cache_manager_v2._event_manager import KVCacheEvent, KVCacheEventDiff # Subscribers decode block hashes as 64-bit ints, so a bytes value would fail the @@ -506,11 +507,12 @@ def _kv_event_wire_hash_from_radix_key(block_key: bytes) -> int: class _MultimodalBlockError(ValueError): - """A block token is a multimodal cache-key digest (bytes), not a wire int. + """A block token is a multimodal cache-key digest, not a wire int. - ``gen_multimodal_cache_key_tokens`` stores the per-item digest as ``bytes``, - which has no integer wire representation. Such blocks are skipped - quietly rather than routed through the malformed-data traceback path. + ``gen_multimodal_cache_key_tokens`` stores the per-item digest as ``bytes`` + or an ``MmItemContext`` carrying its UUID. Neither has an integer wire + representation. Such blocks are skipped quietly rather than routed through + the malformed-data traceback path. """ @@ -678,7 +680,7 @@ def _add_full_block(self, block: Any) -> None: def _token_ids(tokens: Any) -> list[int]: token_ids: list[int] = [] for token in tokens: - if type(token) is bytes: + if type(token) is bytes or isinstance(token, MmItemContext): # Multimodal cache-key digest; not representable as a wire int. raise _MultimodalBlockError if type(token) is not int: diff --git a/tensorrt_llm/_utils.py b/tensorrt_llm/_utils.py index 372f4e1ee70e..c6cde3e04344 100644 --- a/tensorrt_llm/_utils.py +++ b/tensorrt_llm/_utils.py @@ -1183,9 +1183,13 @@ def _unique_tokens_to_json(data): @staticmethod def _mm_key_to_json(data): - # MmKey is a tuple of (hash_bytes, start_offset, uuid) - # where uuid is optional (None if content-hashed) - if len(data) == 3: + # V2 uses a four-element internal form to mark UUID as additive: + # (digest_bytes, start_offset, uuid, preserve_digest_hash). V1 retains + # the legacy three-element form where UUID replaces the hash field. + preserve_digest_hash = len(data) == 4 and data[3] + if len(data) == 4: + hash_array, start_offset, uuid, _ = data + elif len(data) == 3: hash_array, start_offset, uuid = data else: # Backward compatibility: old format (hash_array, start_offset) @@ -1195,14 +1199,14 @@ def _mm_key_to_json(data): # Convert array to hex string hash_hex = ''.join(f'{b:02x}' for b in hash_array) - # Use UUID from C++ if available, otherwise use hash_hex - hash_or_uuid = uuid if uuid is not None else hash_hex - - return { + result = { "type": "mm_key", - "hash": hash_or_uuid, + "hash": hash_hex if preserve_digest_hash or uuid is None else uuid, "start_offset": start_offset } + if preserve_digest_hash and uuid is not None: + result["uuid"] = uuid + return result @staticmethod def _mm_keys_to_json(data): diff --git a/tensorrt_llm/inputs/data.py b/tensorrt_llm/inputs/data.py index fe16de2b941b..acd9eb48491f 100644 --- a/tensorrt_llm/inputs/data.py +++ b/tensorrt_llm/inputs/data.py @@ -21,9 +21,10 @@ class TextPrompt(TypedDict): """ Optional user-provided UUIDs for multimodal items. Structure mirrors multi_modal_data: {"image": ["uuid1", None, "uuid3"]}. - When a UUID is provided for an item, it will be returned in KV cache events - instead of the computed content hash. Use None to fall back to content - hashing for specific items. + When a UUID is provided for an item, it is preserved in KV cache events. + V1 uses it as the legacy hash value; V2 returns it in the optional uuid + field while retaining the derived digest as hash. Use None for specific + items that need content-only hashing. """ mm_processor_kwargs: NotRequired[Dict[str, Any]] @@ -49,9 +50,10 @@ class TokensPrompt(TypedDict): """ Optional user-provided UUIDs for multimodal items. Structure mirrors multi_modal_data: {"image": ["uuid1", None, "uuid3"]}. - When a UUID is provided for an item, it will be returned in KV cache events - instead of the computed content hash. Use None to fall back to content - hashing for specific items. + When a UUID is provided for an item, it is preserved in KV cache events. + V1 uses it as the legacy hash value; V2 returns it in the optional uuid + field while retaining the derived digest as hash. Use None for specific + items that need content-only hashing. """ mm_processor_kwargs: NotRequired[Dict[str, Any]] diff --git a/tensorrt_llm/inputs/multimodal.py b/tensorrt_llm/inputs/multimodal.py index 072a51c86b11..7d297e72bc6a 100644 --- a/tensorrt_llm/inputs/multimodal.py +++ b/tensorrt_llm/inputs/multimodal.py @@ -216,9 +216,9 @@ class MultimodalInput: multimodal_uuids: Optional[List[Optional[str]]] = None """Optional user-provided UUIDs for multimodal data items. - When provided, these UUIDs will be returned in KV cache events instead of the - computed hash hex string. This enables deterministic cache identification across - sessions using user-defined stable identifiers. + When provided, these UUIDs are preserved in KV cache events. V1 reports the UUID + through the legacy ``hash`` field; V2 keeps the derived digest in ``hash`` and + reports the UUID through the optional ``uuid`` field. Each element can be: - A string UUID: Used as the cache identifier (returned in events) diff --git a/tensorrt_llm/inputs/registry.py b/tensorrt_llm/inputs/registry.py index 985075db25ff..0ef2353c22ab 100644 --- a/tensorrt_llm/inputs/registry.py +++ b/tensorrt_llm/inputs/registry.py @@ -1342,8 +1342,9 @@ def multimodal_hashing_process( positions, and masks on top. Supports optional user-provided UUIDs via 'multi_modal_uuids' in inputs. - When a UUID is provided for a multimodal item, it will be used as the - cache identifier and returned in KV cache events instead of the content hash. + When a UUID is provided for a multimodal item, it contributes to the + derived cache digest and is preserved as the external identity in KV + cache events. """ assert 'multi_modal_data' in inputs, "multi_modal_data must be provided for hashing support." mm_data = inputs['multi_modal_data'] diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py index 5a5eee4ecda0..ce106d6edcd9 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py @@ -43,6 +43,7 @@ CudaStream, LayerId, MemAddress, + MmItemContext, PageIndexMode, PageStatus, Priority, @@ -264,6 +265,7 @@ class _KVCacheManagerConfigFieldSpec: ReusedBlocksByLevel = _cpp.ReusedBlocksByLevel SwaScratchReuseConfig = getattr(_cpp, "SwaScratchReuseConfig", None) UniqueToken = _cpp.UniqueToken + MmItemContext = _cpp.MmItemContext BeamIndex = int CacheLevel = int @@ -280,7 +282,7 @@ class _KVCacheManagerConfigFieldSpec: Priority = int SlidingWindowSize = Optional[int] TokenId = int - TokenIdExt = Union[int, bytes] + TokenIdExt = Union[int, bytes, MmItemContext] BAD_PAGE_INDEX = -1 DEFAULT_BEAM_INDEX = 0 @@ -362,6 +364,7 @@ def typed_range(*args: int) -> range: "LayerId", "LifeCycleId", "MemAddress", + "MmItemContext", "NDEBUG", "OutOfPagesError", "PageIndexConverter", diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi index 91ccd9142528..f763585185a5 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi @@ -64,7 +64,15 @@ LifeCycleId = NewType("LifeCycleId", int) LayerGroupId: TypeAlias = LifeCycleId CacheLevel = NewType("CacheLevel", int) TokenId = NewType("TokenId", int) -TokenIdExt = Union[TokenId, bytes] + +class MmItemContext: + def __init__(self, digest: bytes, uuid: str | None = None) -> None: ... + @property + def digest(self) -> bytes: ... + @property + def uuid(self) -> str | None: ... + +TokenIdExt = Union[TokenId, bytes, MmItemContext] class PlannedDropHandle: def drop(self) -> None: ... @@ -245,7 +253,9 @@ EventBlockHash: TypeAlias = int | str BlockHashLike: TypeAlias = bytes | EventBlockHash BlockHashesLike: TypeAlias = BlockHashLike | Iterable[BlockHashLike] EventTokenId: TypeAlias = int | str -MmKey: TypeAlias = tuple[bytes, int] | tuple[bytes, int, str | None] +MmKey: TypeAlias = ( + tuple[bytes, int] | tuple[bytes, int, str | None] | tuple[bytes, int, str | None, bool] +) AttentionDpGatherFn: TypeAlias = Callable[[list["KVCacheEvent"]], list[list["KVCacheEvent"]]] @dataclass(slots=True, frozen=True) @@ -340,6 +350,7 @@ def gen_multimodal_cache_key_tokens( multi_modal_data_digest: bytes, num_tokens: int, token_offset: int = 0, + uuid: str | None = None, ) -> list[TokenIdExt]: ... def sequence_to_blockchain_keys( tokens_per_block: int, diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_block_radix_tree.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_block_radix_tree.py index d1e3ec662cc7..2db82a6051bb 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_block_radix_tree.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_block_radix_tree.py @@ -20,7 +20,7 @@ from typing import TYPE_CHECKING, Iterable, Iterator, NamedTuple, Sequence, TypeVar, cast from . import rawref -from ._common import NDEBUG, BlockOrdinal, PageStatus, TokenId, TokenIdExt +from ._common import NDEBUG, BlockOrdinal, MmItemContext, PageStatus, TokenId, TokenIdExt from ._life_cycle_registry import AttnLifeCycle, LifeCycle, LifeCycleId, LifeCycleRegistry from ._utils import TypedIndexList, filled_list, map_optional, typed_range, unwrap_rawref @@ -41,7 +41,11 @@ # id_offset is usually vocab_size. Backend-neutral (depends only on _common); the # C++ backend exposes a native gen_multimodal_cache_key_tokens via nanobind instead. def gen_multimodal_cache_key_tokens( - id_offset: int, multi_modal_data_digest: bytes, num_tokens: int, token_offset: int = 0 + id_offset: int, + multi_modal_data_digest: bytes, + num_tokens: int, + token_offset: int = 0, + uuid: str | None = None, ) -> list[TokenIdExt]: """Create synthetic tokens used only when building multimodal KV-cache keys. @@ -55,6 +59,7 @@ def gen_multimodal_cache_key_tokens( num_tokens: Number of synthetic tokens to generate. Must be positive. token_offset: Item-local index of the first generated token. Must be non-negative; only offset 0 carries the digest. + uuid: Optional external routing identity to retain with the digest token. Returns: The generated tokens, digest first when ``token_offset`` is 0. @@ -70,7 +75,9 @@ def gen_multimodal_cache_key_tokens( if token_offset < 0: raise ValueError("token_offset must be non-negative") return [ - multi_modal_data_digest if token_offset + i == 0 else TokenId(id_offset + token_offset + i) + (multi_modal_data_digest if uuid is None else MmItemContext(multi_modal_data_digest, uuid)) + if token_offset + i == 0 + else TokenId(id_offset + token_offset + i) for i in range(num_tokens) ] @@ -79,8 +86,9 @@ class Hasher: """Incremental SHA-256 hasher used to derive block keys for the radix tree. Accepts ints (encoded as 4 little-endian bytes each, matching the C++ backend's - 4-byte ``TokenIdExt`` layout), raw ``bytes`` (multimodal content digests and - reuse-scope fields), or a sequence mixing the two. Both backends must produce + 4-byte ``TokenIdExt`` layout), raw ``bytes`` or ``MmItemContext`` values + (multimodal content digests and reuse-scope fields), or a sequence mixing them. + UUID metadata in ``MmItemContext`` is deliberately ignored. Both backends must produce identical digests for the same logical input, so the encoding is part of the on-disk/cross-process contract and cannot change unilaterally. @@ -102,18 +110,23 @@ class Hasher: __slots__ = "_hasher" _hasher: "hashlib._Hash" - def __init__(self, data: int | bytes | Sequence[int | bytes] | None = None) -> None: + def __init__( + self, + data: int | bytes | MmItemContext | Sequence[int | bytes | MmItemContext] | None = None, + ) -> None: self._hasher = hashlib.sha256() if data is not None: self.update(data) - def update(self, data: int | bytes | Sequence[int | bytes]) -> "Hasher": + def update( + self, data: int | bytes | MmItemContext | Sequence[int | bytes | MmItemContext] + ) -> "Hasher": """Fold ``data`` into the running digest. Args: - data: An int token id (0 <= id < 2**31), raw ``bytes``, or a sequence of - either. An all-int sequence takes a single-call fast path; a sequence - containing ``bytes`` (multimodal blocks) falls back to per-item hashing. + data: An int token id (0 <= id < 2**31), raw ``bytes``, an + ``MmItemContext``, or a sequence of these. An all-int sequence takes a + single-call fast path; a multimodal block falls back to per-item hashing. Returns: This ``Hasher``, to allow chaining. @@ -124,6 +137,8 @@ def update(self, data: int | bytes | Sequence[int | bytes]) -> "Hasher": self._hasher.update(data.to_bytes(4, "little")) elif type(data) is bytes: self._hasher.update(data) + elif isinstance(data, MmItemContext): + self._hasher.update(data.digest) else: # Hash the whole token block in one C call instead of one per token. # array("I", data).tobytes() packs each int as 4 native-endian bytes @@ -140,8 +155,14 @@ def update(self, data: int | bytes | Sequence[int | bytes]) -> "Hasher": NDEBUG or (type(item) is int and (0 <= item < (1 << 31))) or type(item) is bytes + or isinstance(item, MmItemContext) ) - self._hasher.update(item.to_bytes(4, "little") if (type(item) is int) else item) # type: ignore + if type(item) is int: + self._hasher.update(item.to_bytes(4, "little")) + elif isinstance(item, MmItemContext): + self._hasher.update(item.digest) + else: + self._hasher.update(item) # type: ignore return self @property @@ -406,7 +427,7 @@ class Block: "_needs_token_digest_context", "_prev", "key", - "last_token_digest", + "last_mm_item_context", "next", "ordinal", "storage", @@ -414,7 +435,7 @@ class Block: ) key: BlockKey tokens: Sequence[TokenIdExt] - last_token_digest: bytes | None + last_mm_item_context: MmItemContext | None ordinal: BlockOrdinal _needs_token_digest_context: bool _prev: rawref.ref["Block | RootBlock"] @@ -438,7 +459,7 @@ def __init__(self, tokens: Sequence[TokenIdExt], prev: "Block | RootBlock") -> N self.storage = filled_list(None, prev.num_life_cycles) self.__rawref__ = rawref.NULL self._needs_token_digest_context = prev._needs_token_digest_context - self.last_token_digest = None + self.last_mm_item_context = None # a Block is useless if all its tokens are covered by a sibling block. Raise UselessBlockError if so. if self.key in prev.next: raise UselessBlockError(prev.next[self.key]) @@ -451,10 +472,15 @@ def __init__(self, tokens: Sequence[TokenIdExt], prev: "Block | RootBlock") -> N if self._needs_token_digest_context: # Share the last digest through text-only descendants, including ancestors # without committable pages that never publish a stored event themselves. - self.last_token_digest = prev.last_token_digest if isinstance(prev, Block) else None + self.last_mm_item_context = ( + prev.last_mm_item_context if isinstance(prev, Block) else None + ) for token in reversed(tokens): + if isinstance(token, MmItemContext): + self.last_mm_item_context = token + break if isinstance(token, bytes): - self.last_token_digest = token + self.last_mm_item_context = MmItemContext(token) break # A later turn may extend a partial endpoint to this longer block, replacing the # partial sibling. That turn may not have a committable SWA page for this block: @@ -484,6 +510,13 @@ def __init__(self, tokens: Sequence[TokenIdExt], prev: "Block | RootBlock") -> N assert b.is_orphan # _KVCache may still hold it. # prev.next keeps a strong ref to this _Block, so no need to remove self from prev.next in __del__(). + @property + def last_token_digest(self) -> bytes | None: + """Digest-only compatibility view of the inherited multimodal context.""" + if self.last_mm_item_context is None: + return None + return self.last_mm_item_context.digest + def page_coverage(self, lc_idx: LifeCycleId) -> int: """Return the page's recorded token count, or zero if the slot is empty. diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_common.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_common.py index 5782e59df03b..6ed689b66204 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_common.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_common.py @@ -54,12 +54,36 @@ class PageIndexMode(enum.IntEnum): # Normal token id that falls in the tokenizer vocabulary. TokenId = NewType("TokenId", int) + +@dataclass(slots=True, frozen=True, eq=False) +class MmItemContext: + """Multimodal cache identity carried by a digest token. + + ``digest`` alone defines token identity and radix-tree keys. ``uuid`` is + request metadata retained only for KV-cache event routing. + """ + + digest: bytes + uuid: str | None = None + + def __eq__(self, other: object) -> bool: + if isinstance(other, MmItemContext): + return self.digest == other.digest + if isinstance(other, bytes): + return self.digest == other + return False + + def __hash__(self) -> int: + return hash(self.digest) + + # For multi-modal tokens, we can handle it in either of the following ways: # 1. Hash combine image digest and local_token_id, then use digest for every multi-modal token. # 2. Use digest only for the first multi-modal token, and use int(vocab_size + local_token_id) for the rest. -# 3. Hash the multi-modal token embedding data and use the digest as TokenIdExt for every multi-modal token. +# 3. Hash the multi-modal token embedding data and use the digest context as +# TokenIdExt for every multi-modal token. # If we do this, we can't skip the encoder. -TokenIdExt = TokenId | bytes +TokenIdExt = TokenId | bytes | MmItemContext BlockOrdinal = NewType("BlockOrdinal", int) 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 9eb3d20b15f0..1d69749fe73a 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 @@ -35,6 +35,7 @@ BlockOrdinalT, CacheLevel, CudaStream, + MmItemContext, PageIndex, PageIndexMode, PageStatus, @@ -799,7 +800,7 @@ def text_only(self, text_only: bool) -> None: "configured text_only=True" ) # Claiming text-only is a fast-path claim; verify committed tokens are digest-free. - if text_only and any(isinstance(t, bytes) for t in self._committed_tokens): + if text_only and any(isinstance(t, (bytes, MmItemContext)) for t in self._committed_tokens): raise ValueError( "Cannot set text_only=True: this sequence has already committed digest tokens" ) diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_event_manager.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_event_manager.py index ff3e55f7b6ed..e28ed29af17e 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_event_manager.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_event_manager.py @@ -61,14 +61,14 @@ truncate_sha256_hash_to_int64, ) -from ._common import GPU_LEVEL, PRIORITY_DEFAULT, CacheLevel, Priority, TokenIdExt +from ._common import GPU_LEVEL, PRIORITY_DEFAULT, CacheLevel, MmItemContext, Priority, TokenIdExt EventBlockHash = int | str BlockHashLike = bytes | EventBlockHash BlockHashesLike = BlockHashLike | Iterable[BlockHashLike] LayerGroupId = int | None EventTokenId = int | str -MmKey = tuple[bytes, int] | tuple[bytes, int, str | None] +MmKey = tuple[bytes, int] | tuple[bytes, int, str | None] | tuple[bytes, int, str | None, bool] AttentionDpGatherFn = Callable[[list["KVCacheEvent"]], list[list["KVCacheEvent"]]] @@ -539,6 +539,8 @@ def _drop_hash_cache(self, block_hash: bytes) -> None: @staticmethod def _normalize_token(token: TokenIdExt) -> UniqueToken: + if isinstance(token, MmItemContext): + return UniqueToken(token.digest.hex()) if isinstance(token, bytes): return UniqueToken(token.hex()) return UniqueToken(int(token)) @@ -550,17 +552,27 @@ def _mm_keys_from_radix_block(self, block: Any) -> list[MmKey]: id_offset = self._mm_token_id_offset assert id_offset is not None parent = block.prev - digest = parent.last_token_digest if parent.ordinal >= 0 else None + context = parent.last_mm_item_context if parent.ordinal >= 0 else None in_mm_run = False mm_keys: list[MmKey] = [] + + def make_mm_key(item_context: MmItemContext, start_offset: int) -> MmKey: + if item_context.uuid is None: + return (item_context.digest, start_offset) + return (item_context.digest, start_offset, item_context.uuid, True) + for token in block.tokens: - if isinstance(token, bytes): - digest = token - mm_keys.append((digest, 0)) + if isinstance(token, MmItemContext): + context = token + mm_keys.append(make_mm_key(context, 0)) + in_mm_run = True + elif isinstance(token, bytes): + context = MmItemContext(token) + mm_keys.append(make_mm_key(context, 0)) in_mm_run = True - elif token > id_offset and digest is not None: + elif token > id_offset and context is not None: if not in_mm_run: - mm_keys.append((digest, int(token) - id_offset)) + mm_keys.append(make_mm_key(context, int(token) - id_offset)) in_mm_run = True else: # Text may separate runs of the same item, so retain its digest. diff --git a/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_multimodal_runs.py b/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_multimodal_runs.py index e60a759b662b..03f72602005b 100644 --- a/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_multimodal_runs.py +++ b/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_multimodal_runs.py @@ -12,7 +12,9 @@ from tensorrt_llm._torch.pyexecutor.kv_cache.mamba_cache_manager import MambaHybridCacheManagerV2 from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest, SamplingConfig from tensorrt_llm.runtime.kv_cache_manager_v2._block_radix_tree import ( + ReuseScope, gen_multimodal_cache_key_tokens, + sequence_to_blockchain_keys, ) pytestmark = pytest.mark.cpu_only @@ -50,6 +52,26 @@ def test_gen_multimodal_cache_key_tokens_uses_token_offset(): ] +def test_multimodal_uuid_metadata_does_not_change_digest_token_identity_or_block_keys(): + vocab_size = 1000 + digest = b"".join(v.to_bytes(4, "big", signed=True) for v in _HASH_INTS) + uuid = "frontend-routing-identity" + without_uuid = gen_multimodal_cache_key_tokens(vocab_size, digest, 3) + with_uuid = gen_multimodal_cache_key_tokens(vocab_size, digest, 3, uuid=uuid) + with_other_uuid = gen_multimodal_cache_key_tokens( + vocab_size, digest, 3, uuid="other-routing-identity" + ) + + assert with_uuid == with_other_uuid == without_uuid + assert with_uuid[0].digest == digest + assert with_uuid[0].uuid == uuid + assert hash(with_uuid[0]) == hash(with_other_uuid[0]) == hash(without_uuid[0]) + assert len({with_uuid[0], with_other_uuid[0], without_uuid[0]}) == 1 + assert [key for _, key in sequence_to_blockchain_keys(2, ReuseScope(), with_uuid)] == [ + key for _, key in sequence_to_blockchain_keys(2, ReuseScope(), without_uuid) + ] + + def test_augment_tokens_for_block_reuse_uses_exact_multimodal_runs(): vocab_size = 1000 digest = b"".join(v.to_bytes(4, "big", signed=True) for v in _HASH_INTS) @@ -110,8 +132,10 @@ def test_augment_tokens_for_block_reuse_uses_two_item_exact_multimodal_runs(): def test_augment_tokens_for_block_reuse_skips_out_of_slice_runs(monkeypatch): calls = [] - def fake_gen_multimodal_cache_key_tokens(vocab_size, digest, num_tokens, token_offset=0): - calls.append((vocab_size, digest, num_tokens, token_offset)) + def fake_gen_multimodal_cache_key_tokens( + vocab_size, digest, num_tokens, token_offset=0, uuid=None + ): + calls.append((vocab_size, digest, num_tokens, token_offset, uuid)) return [digest, *range(vocab_size + 1, vocab_size + num_tokens)] monkeypatch.setattr( @@ -165,7 +189,7 @@ def test_hash_to_digest_rejects_malformed_hashes(): resource_manager._hash_to_digest([*_HASH_INTS[:-1], "8"]) -def test_augment_tokens_for_block_reuse_preserves_supplied_item_digest(): +def test_augment_tokens_for_block_reuse_preserves_supplied_item_digest_and_uuid(): vocab_size = 1000 tokens = list(range(8)) manager = _make_manager(vocab_size) @@ -190,8 +214,33 @@ def augmented_tokens(uuid): content_digest = resource_manager._hash_to_digest(_HASH_INTS) assert no_uuid[2:5] == gen_multimodal_cache_key_tokens(vocab_size, content_digest, 3) # Input preprocessing owns item identity, including any UUID contribution. - # The cache must use that digest without applying another UUID hash. + # The cache must use that digest without applying another UUID hash, while + # retaining the external identity independently of token equality. assert uuid_a == no_uuid == uuid_b + assert uuid_a[2].digest == content_digest + assert uuid_a[2].uuid == "image-a" + assert uuid_b[2].digest == content_digest + assert uuid_b[2].uuid == "image-b" + + +def test_augment_tokens_for_block_reuse_handles_partial_uuid_list(): + vocab_size = 1000 + tokens = list(range(10)) + manager = _make_manager(vocab_size) + req = _make_request( + tokens, + multimodal_hashes=[_HASH_INTS, _OTHER_HASH_INTS], + multimodal_positions=[1, 6], + multimodal_lengths=[2, 2], + multimodal_uuids=["first-item"], + multimodal_item_run_cu_offsets=None, + multimodal_run_positions=None, + multimodal_run_lengths=None, + ) + + augmented = KVCacheManagerV2._augment_tokens_for_block_reuse(manager, tokens, req) + assert augmented[1].uuid == "first-item" + assert isinstance(augmented[6], bytes) def test_augment_tokens_for_block_reuse_canonicalizes_adjacent_runs(): @@ -249,6 +298,7 @@ def commit(chunk): py_request_id=req.py_request_id, is_dummy_request=False, multimodal_hashes=req.multimodal_hashes, + multimodal_uuids=req.multimodal_uuids, multimodal_positions=req.multimodal_positions, multimodal_lengths=req.multimodal_lengths, multimodal_item_run_cu_offsets=None, diff --git a/tests/unittest/_torch/multimodal/test_mm_encoder_standalone.py b/tests/unittest/_torch/multimodal/test_mm_encoder_standalone.py index d7b61c119c12..bd305913447c 100644 --- a/tests/unittest/_torch/multimodal/test_mm_encoder_standalone.py +++ b/tests/unittest/_torch/multimodal/test_mm_encoder_standalone.py @@ -315,7 +315,7 @@ def _expected_mm_event_hashes(inp: TextPrompt, @pytest.mark.parametrize("use_kv_cache_manager_v2", [False, True]) @pytest.mark.parametrize("use_uuids", [False, True]) def test_kv_event_mm_keys_with_uuid(use_uuids, use_kv_cache_manager_v2): - """V2 emits item digests; V1 preserves the optional UUID event label.""" + """V2 emits digest plus UUID; V1 preserves the UUID as its hash label.""" encoder_model_dir = _QWEN_3_VL_DIR max_tokens = 16 @@ -372,6 +372,10 @@ def test_kv_event_mm_keys_with_uuid(use_uuids, use_kv_cache_manager_v2): assert len(mm_keys_found) > 0, "Expected mm_keys in stored events" assert {mm_key["hash"] for mm_key in mm_keys_found} == expected_hashes + expected_uuids = {test_uuid + } if use_uuids and use_kv_cache_manager_v2 else set() + assert {mm_key["uuid"] + for mm_key in mm_keys_found if "uuid" in mm_key} == expected_uuids def _load_inputs_with_uuids(llm: LLM, prompts, media, uuids): @@ -448,17 +452,23 @@ def test_kv_event_mm_keys_with_partial_uuids(uuids, use_kv_cache_manager_v2): # Collect all unique mm_key hashes from stored events mm_key_hashes = set() + mm_key_uuids = set() for event in events: if event and event.get("data", {}).get("type") == "stored": for block in event["data"].get("blocks", []): if block.get("mm_keys"): for mm_key in block["mm_keys"]: mm_key_hashes.add(mm_key["hash"]) + if "uuid" in mm_key: + mm_key_uuids.add(mm_key["uuid"]) # Verify we got mm_keys assert len(mm_key_hashes) > 0, "Expected mm_keys in stored events" assert mm_key_hashes == expected_hashes + assert mm_key_uuids == ({uuid + for uuid in uuids if uuid is not None} + if use_kv_cache_manager_v2 else set()) @pytest.mark.parametrize("use_kv_cache_manager_v2", [False, True]) @@ -514,22 +524,26 @@ def test_kv_event_mm_keys_with_uuid_multiple_prompts(use_kv_cache_manager_v2): # Collect all unique mm_key hashes from stored events mm_key_hashes = set() + mm_key_uuids = set() for event in events: if event and event.get("data", {}).get("type") == "stored": for block in event["data"].get("blocks", []): if block.get("mm_keys"): for mm_key in block["mm_keys"]: mm_key_hashes.add(mm_key["hash"]) + if "uuid" in mm_key: + mm_key_uuids.add(mm_key["uuid"]) # Verify we got mm_keys assert len(mm_key_hashes) > 0, "Expected mm_keys in stored events" assert mm_key_hashes == expected_hashes + assert mm_key_uuids == (set(uuids) if use_kv_cache_manager_v2 else set()) @pytest.mark.parametrize("use_kv_cache_manager_v2", [False, True]) def test_kv_event_mm_keys_with_very_long_uuid(use_kv_cache_manager_v2): - """Long UUIDs remain intact in V1 labels and contribute to V2 digests.""" + """Long UUIDs remain intact in V1 labels and V2 UUID fields.""" encoder_model_dir = _QWEN_3_VL_DIR max_tokens = 16 @@ -593,17 +607,21 @@ def test_kv_event_mm_keys_with_very_long_uuid(use_kv_cache_manager_v2): # Collect all unique mm_key hashes from stored events mm_key_hashes = set() + mm_key_uuids = set() for event in events: if event and event.get("data", {}).get("type") == "stored": for block in event["data"].get("blocks", []): if block.get("mm_keys"): for mm_key in block["mm_keys"]: mm_key_hashes.add(mm_key["hash"]) + if "uuid" in mm_key: + mm_key_uuids.add(mm_key["uuid"]) # Verify we got mm_keys assert len(mm_key_hashes) > 0, "Expected mm_keys in stored events" assert mm_key_hashes == expected_hashes + assert mm_key_uuids == (set(uuids) if use_kv_cache_manager_v2 else set()) @pytest.fixture(scope="module", diff --git a/tests/unittest/kv_cache_manager_v2_tests/kernels.py b/tests/unittest/kv_cache_manager_v2_tests/kernels.py index d5955026a52f..f30a6c79b170 100644 --- a/tests/unittest/kv_cache_manager_v2_tests/kernels.py +++ b/tests/unittest/kv_cache_manager_v2_tests/kernels.py @@ -28,13 +28,20 @@ from cuda.core.experimental._module import ObjectCode if not TYPE_CHECKING and find_spec("kv_cache_manager_v2") is not None: - from kv_cache_manager_v2._common import CudaStream, LayerId, MemAddress, TokenIdExt + from kv_cache_manager_v2._common import ( + CudaStream, + LayerId, + MemAddress, + MmItemContext, + TokenIdExt, + ) from kv_cache_manager_v2._utils import _unwrap, div_up, exact_div else: from tensorrt_llm.runtime.kv_cache_manager_v2._common import ( CudaStream, LayerId, MemAddress, + MmItemContext, TokenIdExt, ) from tensorrt_llm.runtime.kv_cache_manager_v2._utils import _unwrap, div_up, exact_div @@ -224,7 +231,13 @@ def _make_tokens(tokens: Sequence[TokenIdExt], max_tokens: int) -> ctypes.Struct return Tokens( tokens=(ctypes.c_uint32 * max_tokens)( *[ - t if isinstance(t, int) else int.from_bytes(t[:4], "little", signed=False) + t + if isinstance(t, int) + else int.from_bytes( + t.digest[:4] if isinstance(t, MmItemContext) else t[:4], + "little", + signed=False, + ) for t in padded ] ) 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 16854c737b9e..08ac5bdf3f4b 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 @@ -41,6 +41,7 @@ ) from tensorrt_llm.runtime.kv_cache_manager_v2 import KVCacheStoredData as NativeKVCacheStoredData from tensorrt_llm.runtime.kv_cache_manager_v2 import KVCacheUpdatedData as NativeKVCacheUpdatedData +from tensorrt_llm.runtime.kv_cache_manager_v2 import MmItemContext from tensorrt_llm.runtime.kv_cache_manager_v2 import UniqueToken as NativeUniqueToken from tensorrt_llm.runtime.kv_cache_manager_v2._event_manager import ( KVCacheCreatedData, @@ -63,6 +64,10 @@ RootBlock, detach_next, ) + from kv_cache_manager_v2._common import MmItemContext as PythonMmItemContext + from kv_cache_manager_v2._event_manager import ( + KVCacheEventManager as PythonBackendKVCacheEventManager, + ) from kv_cache_manager_v2._life_cycle_registry import LifeCycleRegistry from kv_cache_manager_v2._utils import CachedCudaStream, init_cuda_once, temporary_sys_path else: @@ -80,6 +85,12 @@ RootBlock, detach_next, ) + from tensorrt_llm.runtime.kv_cache_manager_v2._common import ( + MmItemContext as PythonMmItemContext, + ) + from tensorrt_llm.runtime.kv_cache_manager_v2._event_manager import ( + KVCacheEventManager as PythonBackendKVCacheEventManager, + ) from tensorrt_llm.runtime.kv_cache_manager_v2._life_cycle_registry import LifeCycleRegistry from tensorrt_llm.runtime.kv_cache_manager_v2._utils import ( CachedCudaStream, @@ -94,6 +105,10 @@ _USING_CPP_BACKEND = os.environ.get("TLLM_KV_CACHE_MANAGER_V2_BACKEND", "cpp").lower() != "python" +_BackendMmItemContext = MmItemContext if _USING_CPP_BACKEND else PythonMmItemContext +_BackendKVCacheEventManager = ( + NativeKVCacheEventManager if _USING_CPP_BACKEND else PythonBackendKVCacheEventManager +) with temporary_sys_path(os.path.dirname(os.path.abspath(__file__))): @@ -304,7 +319,11 @@ def test_native_event_data_value_semantics_match_python_reference(): [native_token], cache_level=1, priority=35, - mm_keys=[(b"short-mm-key", 3), (b"another-key", 5, "uuid")], + mm_keys=[ + (b"short-mm-key", 3), + (b"another-key", 5, "uuid"), + (b"digest-key", 7, "routing-identity", True), + ], cache_salt="salt", ) native_diff = NativeKVCacheEventDiff(0, 1) @@ -333,7 +352,11 @@ def test_native_event_data_value_semantics_match_python_reference(): [python_token], cache_level=1, priority=35, - mm_keys=[(b"short-mm-key", 3), (b"another-key", 5, "uuid")], + mm_keys=[ + (b"short-mm-key", 3), + (b"another-key", 5, "uuid"), + (b"digest-key", 7, "routing-identity", True), + ], cache_salt="salt", ) python_diff = KVCacheEventDiff(0, 1) @@ -1048,13 +1071,16 @@ def test_v2_kv_cache_event_manager_uses_stored_registry_for_removed_event( def test_v2_kv_cache_event_manager_derives_mm_keys_across_blocks( real_block_factory, ancestor_coverage ): - event_manager = NativeKVCacheEventManager( + event_manager = _BackendKVCacheEventManager( max_kv_event_entries=8, window_size=128, mm_token_id_offset=1000 ) make_block = real_block_factory(event_manager) digest_a = bytes(range(32)) digest_b = bytes(reversed(range(32))) - first = make_block([1, digest_a, 1001, 1002], [ancestor_coverage]) + uuid_a = "frontend-h-a" + first = make_block( + [1, _BackendMmItemContext(digest_a, uuid_a), 1001, 1002], [ancestor_coverage] + ) gap = make_block([2, 3, 4, 5], [ancestor_coverage], parent=first) continued = make_block([1003, 7, 1004, 1005], [4], parent=gap) last = make_block([1006, digest_b, 1001, 9], [4], parent=continued) @@ -1071,29 +1097,51 @@ def test_v2_kv_cache_event_manager_derives_mm_keys_across_blocks( } expected = { _block_key(continued).hex(): [ - {"type": "mm_key", "hash": digest_a.hex(), "start_offset": 3}, - {"type": "mm_key", "hash": digest_a.hex(), "start_offset": 4}, + { + "type": "mm_key", + "hash": digest_a.hex(), + "uuid": uuid_a, + "start_offset": 3, + }, + { + "type": "mm_key", + "hash": digest_a.hex(), + "uuid": uuid_a, + "start_offset": 4, + }, ], _block_key(last).hex(): [ - {"type": "mm_key", "hash": digest_a.hex(), "start_offset": 6}, + { + "type": "mm_key", + "hash": digest_a.hex(), + "uuid": uuid_a, + "start_offset": 6, + }, {"type": "mm_key", "hash": digest_b.hex(), "start_offset": 0}, ], } if ancestor_coverage == 4: expected[_block_key(first).hex()] = [ - {"type": "mm_key", "hash": digest_a.hex(), "start_offset": 0} + { + "type": "mm_key", + "hash": digest_a.hex(), + "uuid": uuid_a, + "start_offset": 0, + } ] expected[_block_key(gap).hex()] = [] assert mm_keys_by_hash == expected def test_v2_kv_cache_event_manager_preserves_mm_keys_after_life_cycle_removal(real_block_factory): - event_manager = NativeKVCacheEventManager( + event_manager = _BackendKVCacheEventManager( max_kv_event_entries=8, window_size=128, mm_token_id_offset=1000 ) make_block = real_block_factory(event_manager, num_life_cycles=2) mm_hash = bytes(range(32)) - block = make_block([mm_hash, 1001], [2, 2]) + uuid = "frontend-routing-identity-" + "x" * 1024 + assert uuid != mm_hash.hex()[:16] + block = make_block([_BackendMmItemContext(mm_hash, uuid), 1001], [2, 2]) block_key = _block_key(block) _add_stored_block(event_manager, block) @@ -1103,6 +1151,7 @@ def test_v2_kv_cache_event_manager_preserves_mm_keys_after_life_cycle_removal(re { "type": "mm_key", "hash": mm_hash.hex(), + "uuid": uuid, "start_offset": 0, } ] @@ -1153,18 +1202,25 @@ def test_python_v2_mm_digest_context_survives_detached_ancestor(): tree = BlockRadixTree(life_cycles, tokens_per_block=4, event_manager=event_manager) root = tree.add_or_get_existing(ReuseScope()) digest = bytes(range(32)) - first = Block([0, digest, 1001, 1002], root) + uuid = "routing-identity" + first = Block([0, PythonMmItemContext(digest, uuid), 1001, 1002], root) gap = Block([1, 2, 3, 4], first) continued = Block([1003, 5, 1004, 1005], gap) parent_ref = gap._prev try: assert detach_next(first, gap.key) is gap assert gap.last_token_digest == digest - assert event_manager._mm_keys_from_radix_block(continued) == [(digest, 3), (digest, 4)] + assert event_manager._mm_keys_from_radix_block(continued) == [ + (digest, 3, uuid, True), + (digest, 4, uuid, True), + ] gap._prev = parent_ref first.next[gap.key] = gap - assert event_manager._mm_keys_from_radix_block(continued) == [(digest, 3), (digest, 4)] + assert event_manager._mm_keys_from_radix_block(continued) == [ + (digest, 3, uuid, True), + (digest, 4, uuid, True), + ] finally: gap._prev = parent_ref first.next[gap.key] = gap 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 c9c25731e3b6..1d0bbf65bce2 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 @@ -13,6 +13,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import logging import socket from types import SimpleNamespace from typing import Callable @@ -29,6 +30,7 @@ validate_streaming_support, ) from tensorrt_llm.llmapi.llm_args import KVEventsConfig +from tensorrt_llm.runtime.kv_cache_manager_v2 import MmItemContext from tensorrt_llm.runtime.kv_cache_manager_v2._block_radix_tree import ( Block, BlockRadixTree, @@ -139,6 +141,40 @@ def test_streaming_sink_supports_real_radix_blocks(monkeypatch: pytest.MonkeyPat manager.shutdown() +def test_streaming_sink_suppresses_multimodal_context_blocks( + caplog: pytest.LogCaptureFixture, +) -> None: + """Multimodal digest tokens are expected omissions, not malformed events.""" + manager = StreamingKVCacheEventManager( + KVEventsConfig(enable_kv_cache_events=True, publisher="null"), + data_parallel_rank=0, + block_size=4, + max_window_size=128, + ) + published: list[KVEventBatch] = [] + manager._publisher.publish = lambda batch: published.append(batch) or True + root = SimpleNamespace(ordinal=-1) + block = SimpleNamespace( + key=b"\x01" * 32, + tokens=[1, MmItemContext(bytes(range(32)), "routing-identity"), 1001, 1002], + prev=root, + ) + + try: + with caplog.at_level(logging.ERROR, logger="tensorrt_llm"): + manager._add_full_block(block) + manager.flush_iteration_events() + + assert manager.multimodal_blocks_suppressed == 1 + assert manager.dropped_events == 0 + assert manager.stored_blocks == 0 + assert manager._pending_events == [] + assert published == [] + assert not [record for record in caplog.records if record.levelno >= logging.ERROR] + finally: + manager.shutdown() + + def test_streaming_fast_path_publishes_only_full_max_window_blocks() -> None: """Protect radix hash reuse, filtering, wire format, and shutdown.""" topic = "kv-events" diff --git a/tests/unittest/llmapi/test_llm_kv_cache_events.py b/tests/unittest/llmapi/test_llm_kv_cache_events.py index d2447c1aea2f..4b7332f86464 100644 --- a/tests/unittest/llmapi/test_llm_kv_cache_events.py +++ b/tests/unittest/llmapi/test_llm_kv_cache_events.py @@ -246,6 +246,26 @@ def test_mm_key_with_uuid(): mock_mm_key_old_format) assert result_old_format["hash"] == expected_hash + # V2 marks UUID as additive: hash remains the digest and UUID is emitted + # separately for routing identity. + mock_mm_key_v2 = (mock_hash, mock_offset, test_uuid, True) + result_v2 = KVCacheEventSerializer._mm_key_to_json(mock_mm_key_v2) + assert result_v2 == { + "type": "mm_key", + "hash": expected_hash, + "uuid": test_uuid, + "start_offset": 42, + } + + mock_mm_key_v2_without_uuid = (mock_hash, mock_offset, None, True) + result_v2_without_uuid = KVCacheEventSerializer._mm_key_to_json( + mock_mm_key_v2_without_uuid) + assert result_v2_without_uuid == { + "type": "mm_key", + "hash": expected_hash, + "start_offset": 42, + } + def test_apply_mm_hashes_with_uuids(): """Test apply_mm_hashes with user-provided UUIDs.""" From 9c9a5487fcb413d9a5207a134e56defe2f41eaf9 Mon Sep 17 00:00:00 2001 From: Guan Luo Date: Wed, 30 Sep 2026 15:15:31 -0700 Subject: [PATCH 2/3] test: cover multimodal UUIDs in exact runs and chunked commits Signed-off-by: Guan Luo --- .../test_kv_cache_v2_multimodal_runs.py | 35 +++++++++++++++++-- 1 file changed, 33 insertions(+), 2 deletions(-) diff --git a/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_multimodal_runs.py b/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_multimodal_runs.py index a331b6232c00..b6f0f82e6112 100644 --- a/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_multimodal_runs.py +++ b/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_multimodal_runs.py @@ -2,7 +2,7 @@ # Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. from types import SimpleNamespace -from unittest.mock import Mock +from unittest.mock import Mock, call import pytest import torch @@ -98,18 +98,22 @@ def test_augment_tokens_for_block_reuse_uses_exact_multimodal_runs(): assert sliced == [tokens[7], mm_tokens[2], mm_tokens[3]] -def test_augment_tokens_for_block_reuse_uses_two_item_exact_multimodal_runs(): +def test_augment_tokens_for_block_reuse_uses_two_item_exact_multimodal_runs(monkeypatch): + """Keep each item's UUID and token offsets across separated and sliced runs.""" vocab_size = 1000 digest_a = b"".join(v.to_bytes(4, "big", signed=True) for v in _HASH_INTS) digest_b = b"".join(v.to_bytes(4, "big", signed=True) for v in _OTHER_HASH_INTS) mm_tokens_a = gen_multimodal_cache_key_tokens(vocab_size, digest_a, 3) mm_tokens_b = gen_multimodal_cache_key_tokens(vocab_size, digest_b, 3) + generator = Mock(wraps=gen_multimodal_cache_key_tokens) + monkeypatch.setattr(resource_manager, "gen_multimodal_cache_key_tokens", generator) tokens = list(range(16)) manager = _make_manager(vocab_size) req = _make_request( tokens, multimodal_hashes=[_HASH_INTS, _OTHER_HASH_INTS], + multimodal_uuids=["item-a", "item-b"], multimodal_positions=[1, 9], multimodal_lengths=[3, 3], multimodal_item_run_cu_offsets=[0, 2, 4], @@ -124,9 +128,33 @@ def test_augment_tokens_for_block_reuse_uses_two_item_exact_multimodal_runs(): assert augmented[12:14] == mm_tokens_b[1:3] assert augmented[3:6] == tokens[3:6] assert augmented[10:12] == tokens[10:12] + assert augmented[1].uuid == "item-a" + assert augmented[9].uuid == "item-b" + assert generator.call_args_list == [ + call(vocab_size, digest_a, 2, token_offset=0, uuid="item-a"), + call(vocab_size, digest_a, 1, token_offset=2, uuid="item-a"), + call(vocab_size, digest_b, 1, token_offset=0, uuid="item-b"), + call(vocab_size, digest_b, 2, token_offset=1, uuid="item-b"), + ] + generator.reset_mock() sliced = KVCacheManagerV2._augment_tokens_for_block_reuse(manager, tokens, req, start=8, end=13) assert sliced == [tokens[8], mm_tokens_b[0], tokens[10], tokens[11], mm_tokens_b[1]] + assert sliced[1].uuid == "item-b" + assert generator.call_args_list == [ + call(vocab_size, digest_b, 1, token_offset=0, uuid="item-b"), + call(vocab_size, digest_b, 1, token_offset=1, uuid="item-b"), + ] + + generator.reset_mock() + continuation = KVCacheManagerV2._augment_tokens_for_block_reuse( + manager, tokens, req, start=2, end=7 + ) + assert continuation == [mm_tokens_a[1], *tokens[3:6], mm_tokens_a[2]] + assert generator.call_args_list == [ + call(vocab_size, digest_a, 1, token_offset=1, uuid="item-a"), + call(vocab_size, digest_a, 1, token_offset=2, uuid="item-a"), + ] def test_augment_tokens_for_block_reuse_skips_out_of_slice_runs(monkeypatch): @@ -265,10 +293,12 @@ def test_augment_tokens_for_block_reuse_canonicalizes_adjacent_runs(): @pytest.mark.parametrize("hybrid", [False, True]) def test_multimodal_events_keep_chunked_commit_incremental(hybrid): + """Preserve UUID context while committing each prefill chunk only once.""" tokens = list(range(12)) req = _make_request( tokens, multimodal_hashes=[_HASH_INTS], + multimodal_uuids=["chunked-item"], multimodal_positions=[1], multimodal_lengths=[8], multimodal_item_run_cu_offsets=None, @@ -320,6 +350,7 @@ def commit(chunk): digest = resource_manager._hash_to_digest(_HASH_INTS) assert calls == [[0, digest, 1001, 1002], [1003, 1004, 1005, 1006], [1007, 9, 10, 11]] + assert calls[0][1].uuid == "chunked-item" assert [call.kwargs for call in manager._augment_tokens_for_block_reuse.call_args_list] == [ {"start": 0, "end": 4}, {"start": 4, "end": 8}, From cb4141ad6c10118d9ba51fb1d9ac40ba773cfaf7 Mon Sep 17 00:00:00 2001 From: Guan Luo Date: Thu, 1 Oct 2026 00:03:14 -0700 Subject: [PATCH 3/3] test: align prefix probe request stub with multimodal UUIDs Signed-off-by: Guan Luo --- .../test_kv_cache_v2_first_new_block_probe.py | 15 +++++++++++++-- 1 file changed, 13 insertions(+), 2 deletions(-) diff --git a/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_first_new_block_probe.py b/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_first_new_block_probe.py index d55fc3a957d9..fcff6280ec49 100644 --- a/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_first_new_block_probe.py +++ b/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_first_new_block_probe.py @@ -106,6 +106,7 @@ def __init__( self.total_input_len_cp = len(self._tokens) self.context_current_position = 0 self.multimodal_hashes = None + self.multimodal_uuids = None self.multimodal_positions = None self.multimodal_lengths = None self.multimodal_item_run_cu_offsets = None @@ -234,20 +235,30 @@ def test_core_result_is_returned_without_a_second_lookup(self, key: bytes | None mgr.impl.probe_first_new_block_key.assert_called_once() mgr.impl.probe_reuse.assert_not_called() - def test_multimodal_tokens_match_prepare_context(self) -> None: + @pytest.mark.parametrize("uuid", [None, "probe-item"]) + def test_multimodal_tokens_match_prepare_context(self, uuid: str | None) -> None: + """Probe and preparation preserve both digest identity and optional UUIDs.""" mgr = make_stub_manager() req = make_request(range(20)) req.multimodal_hashes = [[17] * 8] + if uuid is not None: + req.multimodal_uuids = [uuid] req.multimodal_positions = [3] req.multimodal_lengths = [2] probed, _ = probed_tokens_and_scope(mgr, req) prepared, _ = prepared_tokens_and_scope(mgr, req) assert probed == prepared - assert isinstance(probed[3], bytes) and len(probed[3]) == 32 + if uuid is None: + assert isinstance(probed[3], bytes) and len(probed[3]) == 32 + else: + assert probed[3].digest == prepared[3].digest == (17).to_bytes(4, "big") * 8 + assert probed[3].uuid == prepared[3].uuid == uuid req.multimodal_hashes[0][0] = 19 changed, _ = probed_tokens_and_scope(mgr, req) assert changed[3] != probed[3] assert changed[:3] + changed[4:] == probed[:3] + probed[4:] + if uuid is not None: + assert changed[3].uuid == uuid def test_augmentation_call_matches_prepare_context(self): """Multimodal requests key on content digests, so the probe has to use