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 fe4f4ac0f850..a3df4b5f416a 100644 --- a/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp +++ b/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp @@ -232,7 +232,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) @@ -373,8 +378,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; @@ -441,12 +453,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]); } @@ -455,8 +467,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; } @@ -467,13 +485,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; @@ -885,6 +905,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) @@ -2596,7 +2663,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) @@ -2605,9 +2673,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 43e993e62aae..61e49e966de4 100644 --- a/docs/source/features/kvcache.md +++ b/docs/source/features/kvcache.md @@ -203,7 +203,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:** @@ -222,16 +222,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 ed0cc2a89afe..908e3e42bb30 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 @@ -861,6 +861,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, @@ -895,6 +896,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(), @@ -905,13 +907,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, ) ) @@ -922,6 +933,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, @@ -949,6 +961,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 @@ -4415,13 +4432,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/_utils.py b/tensorrt_llm/_utils.py index c875715c4501..da55f06a2449 100644 --- a/tensorrt_llm/_utils.py +++ b/tensorrt_llm/_utils.py @@ -1197,9 +1197,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) @@ -1209,14 +1213,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 06bae3975faf..ef084b6639b8 100644 --- a/tensorrt_llm/inputs/multimodal.py +++ b/tensorrt_llm/inputs/multimodal.py @@ -215,9 +215,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 a8a71706f7ca..8419a269a27e 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py @@ -171,12 +171,13 @@ class _KVCacheManagerConfigFieldSpec: LayerId = int LifeCycleId = int MemAddress = int +MmItemContext = _cpp.MmItemContext PoolGroupIndex = int PoolIndex = int Priority = int SlidingWindowSize = Optional[int] TokenId = int -TokenIdExt = Union[int, bytes] +TokenIdExt = Union[int, bytes, MmItemContext] BAD_PAGE_INDEX = -1 DEFAULT_BEAM_INDEX = 0 @@ -233,6 +234,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 7c9224ca6e4a..5ee6ef43d2b2 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi @@ -96,7 +96,15 @@ class AttnLifeCycle: 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: ... @@ -279,7 +287,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) @@ -374,6 +384,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/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 e83fb5641178..f3a7db6489c6 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 @@ -107,6 +107,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 @@ -235,20 +236,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 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 7e1e82f71047..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 @@ -11,7 +11,11 @@ from tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2 import KVCacheManagerV2 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 import gen_multimodal_cache_key_tokens +from tensorrt_llm.runtime.kv_cache_manager_v2 import ( + ReuseScope, + gen_multimodal_cache_key_tokens, + sequence_to_blockchain_keys, +) pytestmark = pytest.mark.cpu_only @@ -48,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) @@ -74,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], @@ -100,16 +128,42 @@ 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): 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( @@ -163,7 +217,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) @@ -188,8 +242,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(): @@ -214,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, @@ -247,6 +328,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, @@ -268,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}, 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 0f7e3a9bb01e..8f9e05ab0a53 100644 --- a/tests/unittest/kv_cache_manager_v2_tests/kernels.py +++ b/tests/unittest/kv_cache_manager_v2_tests/kernels.py @@ -30,9 +30,15 @@ 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 import CudaStream, LayerId, MemAddress, TokenIdExt + from kv_cache_manager_v2 import CudaStream, LayerId, MemAddress, MmItemContext, TokenIdExt else: - from tensorrt_llm.runtime.kv_cache_manager_v2 import CudaStream, LayerId, MemAddress, TokenIdExt + from tensorrt_llm.runtime.kv_cache_manager_v2 import ( + CudaStream, + LayerId, + MemAddress, + MmItemContext, + TokenIdExt, + ) _TEST_DIR = os.path.dirname(os.path.abspath(__file__)) # cuda_test_utils supplies temporary_sys_path, so its own path entry is added and @@ -231,7 +237,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 f427394b56f2..1bf8cf2d5b1d 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 @@ -47,6 +47,7 @@ CacheLevel, CudaStream, KVCacheManager, + MmItemContext, ReuseScope, TokenId, _introspection, @@ -57,6 +58,7 @@ CacheLevel, CudaStream, KVCacheManager, + MmItemContext, ReuseScope, TokenId, _introspection, @@ -274,7 +276,11 @@ def test_event_data_pickle_round_trip_preserves_value_semantics(): [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", ) diff = KVCacheEventDiff(0, 1) @@ -955,7 +961,8 @@ def test_v2_kv_cache_event_manager_derives_mm_keys_across_blocks( 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, MmItemContext(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) @@ -972,17 +979,37 @@ 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 @@ -994,7 +1021,9 @@ def test_v2_kv_cache_event_manager_preserves_mm_keys_after_life_cycle_removal(re ) 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([MmItemContext(mm_hash, uuid), 1001], [2, 2]) block_key = _block_key(block) _add_stored_block(event_manager, block) @@ -1004,6 +1033,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, } ] 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."""