Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -193,8 +193,8 @@ BlockKey Hasher::digest() const
// genMultimodalCacheKeyTokens
// ---------------------------------------------------------------------------

std::vector<TokenIdExt> genMultimodalCacheKeyTokens(
int idOffset, std::vector<uint8_t> const& multiModalDataDigest, int numTokens, int tokenOffset)
std::vector<TokenIdExt> genMultimodalCacheKeyTokens(int idOffset, std::vector<uint8_t> const& multiModalDataDigest,
int numTokens, int tokenOffset, std::optional<std::string> uuid)
{
TLLM_CHECK(numTokens > 0);
TLLM_CHECK(tokenOffset >= 0);
Expand All @@ -207,7 +207,7 @@ std::vector<TokenIdExt> genMultimodalCacheKeyTokens(
{
Digest digest;
std::memcpy(digest.data(), multiModalDataDigest.data(), kDIGEST_LEN);
result.emplace_back(digest);
result.emplace_back(digest, std::move(uuid));
}
else
{
Expand Down Expand Up @@ -330,13 +330,13 @@ Block::Block(BlockKey k, std::vector<TokenIdExt> toks, NodeBase* prevNode)
{
if (prevNode->type() == Type::kBLOCK)
{
mLastTokenDigest = static_cast<Block const*>(prevNode)->getLastTokenDigest();
mLastMmItemContext = static_cast<Block const*>(prevNode)->getLastMmItemContext();
}
for (auto iter = tokens.rbegin(); iter != tokens.rend(); ++iter)
{
if (iter->isDigest())
{
mLastTokenDigest = std::make_shared<Digest const>(iter->digest());
mLastMmItemContext = iter->sharedMmItemContext();
break;
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -160,8 +160,8 @@ inline auto sequenceToBlockchainKeys(
}

// Generate multi-modal token IDs (mirrors gen_multimodal_cache_key_tokens in Python).
std::vector<TokenIdExt> genMultimodalCacheKeyTokens(
int idOffset, std::vector<uint8_t> const& multiModalDataDigest, int numTokens, int tokenOffset = 0);
std::vector<TokenIdExt> genMultimodalCacheKeyTokens(int idOffset, std::vector<uint8_t> const& multiModalDataDigest,
int numTokens, int tokenOffset = 0, std::optional<std::string> uuid = std::nullopt);

// ---------------------------------------------------------------------------
// NodeBase — common base for RootBlock and Block (nodes in the radix tree).
Expand Down Expand Up @@ -273,10 +273,10 @@ struct Block : NodeBase, EnableSharedFromThis<Block>
return storage.size();
}

//! Latest digest token in this block's prefix, when requested by the event sink.
std::shared_ptr<Digest const> const& getLastTokenDigest() const noexcept
//! Latest multimodal item context in this block's prefix, when requested by the event sink.
std::shared_ptr<MmItemContext const> const& getLastMmItemContext() const noexcept
{
return mLastTokenDigest;
return mLastMmItemContext;
}

bool isFull() const noexcept
Expand Down Expand Up @@ -344,7 +344,7 @@ struct Block : NodeBase, EnableSharedFromThis<Block>
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<Digest const> mLastTokenDigest;
std::shared_ptr<MmItemContext const> mLastMmItemContext;
};

// ---------------------------------------------------------------------------
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -567,10 +567,10 @@ std::optional<KVCacheStoredBlockData> EventManager::storedBlockFromBlock(
}

std::vector<MmKey> 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 const*>(block.prev)->getLastTokenDigest().get();
itemContext = static_cast<Block const*>(block.prev)->getLastMmItemContext().get();
}
bool inMmRun = false;
std::vector<UniqueToken> tokens;
Expand All @@ -582,13 +582,14 @@ std::optional<KVCacheStoredBlockData> 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<char const*>(itemDigest->data()), itemDigest->size()),
token.tokenId() - *mMmTokenIdOffset, std::nullopt, false});
mmKeys.push_back({std::string(reinterpret_cast<char const*>(itemContext->digest.data()),
itemContext->digest.size()),
token.tokenId() - *mMmTokenIdOffset, itemContext->uuid,
itemContext->uuid.has_value() ? MmKeyUuidMode::kAdditive : MmKeyUuidMode::kNone});
}
inMmRun = true;
}
Expand All @@ -605,9 +606,11 @@ std::optional<KVCacheStoredBlockData> EventManager::storedBlockFromBlock(
tokens.push_back(std::move(uniqueToken));
if (mMmTokenIdOffset.has_value())
{
itemDigest = &token.digest();
mmKeys.push_back({std::string(reinterpret_cast<char const*>(itemDigest->data()), itemDigest->size()), 0,
std::nullopt, false});
itemContext = &token.mmItemContext();
mmKeys.push_back(
{std::string(reinterpret_cast<char const*>(itemContext->digest.data()), itemContext->digest.size()),
0, itemContext->uuid,
itemContext->uuid.has_value() ? MmKeyUuidMode::kAdditive : MmKeyUuidMode::kNone});
inMmRun = true;
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -64,17 +64,24 @@ struct KVCacheCreatedData
}
};

enum class MmKeyUuidMode : uint8_t
{
kNone,
kReplacesHash,
kAdditive,
};

struct MmKey
{
std::string hash;
int startOffset = 0;
std::optional<std::string> 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;
}
};

Expand Down
73 changes: 53 additions & 20 deletions cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/tokenIdExt.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -24,14 +24,16 @@
#include <cstdint>
#include <cstdio>
#include <deque>
#include <memory>
#include <mutex>
#include <utility>

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.
//
Expand Down Expand Up @@ -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<std::mutex> const lock(mMutex);
return allocLocked(digest);
return allocLocked(std::make_shared<MmItemContext const>(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<std::mutex> 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<std::mutex> const lock(mMutex);
return mStore.at(idx);
return *liveSlotLocked(idx);
}

// Shared ownership of the context at slot `idx`.
[[nodiscard]] std::shared_ptr<MmItemContext const> getShared(uint32_t idx) const
{
std::lock_guard<std::mutex> const lock(mMutex);
return liveSlotLocked(idx);
}

// Clear slot `idx` and reclaim trailing free slots.
Expand All @@ -106,6 +115,7 @@ class DigestPool
return; // the default / moved-from sentinel index — nothing to free
}
std::lock_guard<std::mutex> const lock(mMutex);
mStore[idx].reset();
mInUse.clear(idx);
if (idx < mMinFreeHint)
{
Expand All @@ -123,6 +133,14 @@ class DigestPool
private:
DigestPool() = default;

// Precondition: caller holds mMutex. Validate that `idx` names a live slot.
[[nodiscard]] std::shared_ptr<MmItemContext const> 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).
Expand All @@ -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<MmItemContext const> context)
{
size_t const cap = capacity();
size_t idx = mMinFreeHint;
Expand All @@ -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
Expand Down Expand Up @@ -191,9 +209,9 @@ class DigestPool
static constexpr size_t kSlackLow = 64;

mutable std::mutex mMutex;
std::deque<Digest> 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<std::shared_ptr<MmItemContext const>> 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
Expand All @@ -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<std::string> uuid)
: TokenIdExt(MmItemContext{digestValue, std::move(uuid)})
{
}

TokenIdExt::TokenIdExt(MmItemContext context)
: mBits(DigestPool::instance().alloc(std::move(context)) | kTagMask)
{
}

Expand All @@ -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<MmItemContext const> TokenIdExt::sharedMmItemContext() const
{
return DigestPool::instance().getShared(digestIndex());
}

namespace detail
{

Expand Down
29 changes: 26 additions & 3 deletions cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/tokenIdExt.h
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,9 @@
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <memory>
#include <optional>
#include <string>
#include <type_traits>
#include <vector>

Expand Down Expand Up @@ -51,14 +54,27 @@ struct alignas(kDIGEST_LEN) Digest : std::array<std::byte, kDIGEST_LEN>
}
};

//! Multimodal item metadata carried by a digest token. UUID does not
//! participate in equality or radix-tree hashing.
struct MmItemContext
{
Digest digest;
std::optional<std::string> uuid;

bool operator==(MmItemContext const& other) const noexcept
{
return digest == other.digest;
}
};

// ---------------------------------------------------------------------------
// TokenIdExt — 4-byte self-describing token handle (RAII value type).
//
// One uint32_t; the high bit tags the low 31 bits:
// - 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.
Expand Down Expand Up @@ -89,8 +105,9 @@ class TokenIdExt
TLLM_CHECK_DEBUG(id >= 0 && id <= static_cast<TokenId>(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<std::string> uuid = std::nullopt);
explicit TokenIdExt(MmItemContext context);

~TokenIdExt()
{
Expand Down Expand Up @@ -155,6 +172,12 @@ class TokenIdExt
return static_cast<TokenId>(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<MmItemContext const> sharedMmItemContext() const;

// The pooled 32-byte digest. Precondition: isDigest().
[[nodiscard]] Digest const& digest() const;

Expand Down
10 changes: 10 additions & 0 deletions cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -259,6 +259,16 @@ void initBindings(nb::module_& m)
}
return hashes;
})
.def_prop_ro("multimodal_uuids",
[](GenLlmReq& self)
{
std::optional<std::vector<std::optional<std::string>>> uuids = std::nullopt;
if (self.getMultimodalUuids())
{
uuids = *self.getMultimodalUuids().value();
}
return uuids;
})
.def_prop_ro("multimodal_positions",
[](GenLlmReq& self)
{
Expand Down
Loading
Loading