From 98e190e57dbdbdd7551772f0e1c4e04064f08828 Mon Sep 17 00:00:00 2001 From: Guiju Zhang <7135567+cascade812@users.noreply.github.com> Date: Wed, 30 Sep 2026 11:12:55 -0700 Subject: [PATCH 1/5] Batched GPU-to-CPU demotion and source release Signed-off-by: Guiju Zhang <7135567+cascade812@users.noreply.github.com> --- .../kv_cache_manager_v2/kvCache.cpp | 16 + .../kv_cache_manager_v2/kvCache.h | 22 +- .../kv_cache_manager_v2/page.cpp | 73 +++ .../batch_manager/kv_cache_manager_v2/page.h | 40 +- .../kv_cache_manager_v2/storageManager.cpp | 165 ++++++ .../kv_cache_manager_v2/storageManager.h | 5 + .../kvCacheManagerV2ColdPageTest.cpp | 516 ++++++++++++++++++ 7 files changed, 826 insertions(+), 11 deletions(-) diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.cpp index 77f33f86176a..b8b994beb9e2 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.cpp @@ -213,6 +213,22 @@ CacheLevel KvCache::_lockLevel(Page const& page, BlockOrdinal ordinal) const return readOnly ? page.queryLockLevel() : kHotLevel; } +void KvCache::offloadSparsePages(std::vector> const& pages) +{ + KVCM2_API_GUARD(); + auto const apiLock = mManager->lockExclusive(); + if (!isActive()) + { + throw LogicError("Sparse history offload requires an active request"); + } + MigrationRecorder const migrationRecorder + = [this](std::vector> const& sources, std::vector const& slots, CacheLevel srcLevel, + CacheLevel dstLevel) { _recordMigratedSlots(sources, slots, srcLevel, dstLevel); }; + DropRecorder const dropRecorder = [this](std::vector> const& dropped, CacheLevel level) + { _recordDroppedPages(dropped, level); }; + storageManager()->offloadSparsePages(*this, pages, migrationRecorder, dropRecorder); +} + void KvCache::activate() { TLLM_CHECK_DEBUG(mStatus == Status::SUSPENDED); diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.h index fc9c13465508..5f336f57d1be 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.h +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.h @@ -27,6 +27,7 @@ #include "kv_cache_manager_v2/utils/funcGuard.h" #include "tensorrt_llm/common/assert.h" +#include #include #include #include @@ -213,6 +214,18 @@ class KvCache : public std::enable_shared_from_this void setCapacity(int capacity); void setHistoryLength(int historyLength); + //! Internal explicit demotion of complete, locked sparse history. Takes the manager's exclusive lock. + //! Duplicate pages and pages already in host history are ignored. Does not advance history length. + //! Caller must ensure every owner has finished the execution phase that requires these pages on GPU. + void offloadSparsePages(std::vector> const& pages); + + //! Changes when offload relocates an owned page, even if its numeric slot index stays the same. + //! Internal invalidation hook for page metadata; read under the manager's API lock. + uint64_t pageStorageVersion() const noexcept + { + return mPageStorageVersion; + } + // ---- Committing tokens ------------------------------------------------- // Commit tokens: finalises the oldest uncommitted block and makes it @@ -231,7 +244,7 @@ class KvCache : public std::enable_shared_from_this // Get base page indices (slot_id) for beamIdx × layerGroupId. // Returns a non-owning Span into the page-index buffer (owned by this KvCache, or by the // caller when set via setBasePageIndexBuf). The span is valid until the next resize(), - // setBasePageIndexBuf() or close(); its contents are also rewritten by suspend()/resume(). + // setBasePageIndexBuf() or close(); its contents are also rewritten by suspend()/resume() and offload. Span getBasePageIndices(LayerGroupId lgId, BeamIndex beamIdx = kDefaultBeamIndex) const; // Get aggregated (slot-level) page indices for one layer group + beam. @@ -457,9 +470,15 @@ class KvCache : public std::enable_shared_from_this private: friend class KvCacheIntrospection; + friend class UniqPageLock; friend std::vector batchedLockPages( KvCache& kvCache, std::vector const& targets); + void onPageStorageChanged() noexcept + { + ++mPageStorageVersion; + } + // Activate: lock active pages at their required levels. mCudaStream must already be set. // Internal — called by resume(). Not public (mirrors Python where activate() doesn't exist). void activate(); @@ -632,6 +651,7 @@ class KvCache : public std::enable_shared_from_this using LifeCyclePageIndexBuffers = TypedVec; using BeamPageIndexBuffers = TypedVec; BeamPageIndexBuffers mBasePageIndices; + uint64_t mPageStorageVersion = 0; TypedVec mBlocks; diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/page.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/page.cpp index 88a0138ff707..a2905cc56b8c 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/page.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/page.cpp @@ -23,6 +23,7 @@ #include "tensorrt_llm/common/assert.h" +#include #include namespace tensorrt_llm::batch_manager::kv_cache_manager_v2 @@ -301,6 +302,8 @@ UniqPageLock::UniqPageLock(SharedPtr h) { throw LogicError("Pages can only be locked on GPU or, for sparse attention, in level-1 host memory"); } + // Preserve readiness if the first shared lock fails before it can register an owner. + finishEvents.push_back(holder->page->readyEvent); } UniqPageLock::~UniqPageLock() @@ -310,6 +313,7 @@ UniqPageLock::~UniqPageLock() { Page& p = *page(); TLLM_CHECK_DEBUG(p.cacheLevel == p.queryLockLevel() && !p.scheduledForEviction()); + TLLM_CHECK_DEBUG(mOwners.empty()); // Set readyEvent to the merged finish events of all readers. For committed (read-only) // pages, this means the next reader will wait for prior reads to complete, which is // unnecessary but correct. See the CommittedPage comment in page.h for rationale. @@ -345,6 +349,69 @@ void UniqPageLock::notifyFinish(CachedCudaEvent event) } } +void UniqPageLock::prepareSparseOffload(KvCache const& requestingCache) +{ + Page const& p = *page(); + auto const* attn = std::get_if(&p.manager->getLifeCycle(p.lifeCycle)); + if (!attn || !attn->isSparse || !p.hasValidSlot() + || (p.cacheLevel != kHotLevel && p.cacheLevel != kSparseHistoryLevel) + || p.manager->numCacheLevels() <= kSparseHistoryLevel + || p.manager->cacheTier(kSparseHistoryLevel) != CacheTier::HOST_MEM) + { + throw LogicError("Offload requires a locked sparse attention page on GPU or in host history"); + } + if (p.isCommitted() && static_cast(p).numTokensInBlock != requestingCache.tokensPerBlock()) + { + throw LogicError("Cannot offload a partial committed page"); + } + bool requestingOwner = false; + for (auto const& owner : mOwners) + { + if (!owner.kvCache->isActive() || owner.lifeCycle != p.lifeCycle || owner.ordinal < BlockOrdinal{0} + || owner.ordinal >= BlockOrdinal{owner.kvCache->historyLength() / owner.kvCache->tokensPerBlock()}) + { + throw LogicError("Cannot offload a page outside an owner's complete history"); + } + requestingOwner |= owner.kvCache == &requestingCache; + } + if (!requestingOwner) + { + throw LogicError("The offloading request must own a lock on the page"); + } + finishEvents.reserve(1); +} + +void UniqPageLock::recordOffloadEvent(CachedCudaEvent const& event) +{ + // The copy stream already waited for every event being replaced here. + page()->readyEvent = event; + finishEvents.clear(); + finishEvents.push_back(event); +} + +Slot UniqPageLock::moveToSparseHistory(Slot&& hostSlot) +{ + Page& p = *page(); + TLLM_CHECK_DEBUG(p.cacheLevel == kHotLevel && !p.scheduledForEviction()); + Slot gpuSlot = p.exchangeSlot(std::move(hostSlot)); + p.cacheLevel = kSparseHistoryLevel; + for (auto const& owner : mOwners) + { + int const old = owner.kvCache->updateBasePageIndex( + owner.beamIndex, owner.ordinal, owner.lifeCycle, slotIdToPageIndexValue(p.slotId())); + TLLM_CHECK_DEBUG(old == slotIdToPageIndexValue(gpuSlot.slotId())); + owner.kvCache->onPageStorageChanged(); + } + return gpuSlot; +} + +void UniqPageLock::removeOwner(LockOwner const& owner) +{ + auto const it = std::find(mOwners.begin(), mOwners.end(), owner); + TLLM_CHECK_DEBUG(it != mOwners.end()); + mOwners.erase(it); +} + SharedPtr const& UniqPageLock::page() const { TLLM_CHECK_DEBUG(holder && holder->page); @@ -367,9 +434,14 @@ SharedPageLock::SharedPageLock(SharedPtr ul, KvCache& kvCache, Bea , mUser{&kvCache, beamIndex, ordinal, lc} { if (!skipWait) + { page()->readyEvent.waitInStream(reinterpret_cast(kvCache.cudaStream())); + } + mUniqLock->mOwners.push_back(mUser); + auto rollbackOwner = FuncGuard([this]() { mUniqLock->removeOwner(mUser); }); acquirePageIndex(); + rollbackOwner.cancel(); } SharedPageLock::~SharedPageLock() @@ -410,6 +482,7 @@ SharedPtr SharedPageLock::unlock() mUniqLock->notifyFinish(mUser.kvCache->finishEvent()); releasePageIndex(); + mUniqLock->removeOwner(mUser); auto p = page(); // copy shared_ptr before reset mUniqLock.reset(); return p; diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/page.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/page.h index 103f368dfa64..8b9cfcfd1507 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/page.h +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/page.h @@ -176,6 +176,17 @@ class PageHolder : public EnableSharedFromThis WeakPtr uniqLock; // non-null → LOCKED }; +//! Identifies one live SharedPageLock independently of the lock object's address. +struct LockOwner +{ + KvCache* kvCache; + BeamIndex beamIndex; + BlockOrdinal ordinal; + LifeCycleId lifeCycle; + + bool operator==(LockOwner const&) const = default; +}; + // --------------------------------------------------------------------------- // UniqPageLock — locks a page to prevent eviction (LOCKED status). // Owns finish events from all SharedPageLocks it issued. @@ -199,19 +210,28 @@ class UniqPageLock : public EnableSharedFromThis // Append a finish event, merging when count exceeds 32 to prevent unbounded growth. void notifyFinish(CachedCudaEvent event); + //! Validate complete sparse history for every owner and prepare non-allocating completion updates. + void prepareSparseOffload(KvCache const& requestingCache); + + //! Record a copy ordered after page readiness, finished readers, and all live owners' prior work. + void recordOffloadEvent(CachedCudaEvent const& event); + + //! Publish the host slot to every owner and return the fenced GPU slot. Caller holds the API lock. + [[nodiscard]] Slot moveToSparseHistory(Slot&& hostSlot); + + std::vector const& owners() const noexcept + { + return mOwners; + } + SharedPtr holder; std::vector finishEvents; -}; -// --------------------------------------------------------------------------- -// LockOwner — identifies who holds a SharedPageLock. -// --------------------------------------------------------------------------- -struct LockOwner -{ - KvCache* kvCache; - BeamIndex beamIndex; - BlockOrdinal ordinal; - LifeCycleId lifeCycle; +private: + friend class SharedPageLock; + void removeOwner(LockOwner const& owner); + + std::vector mOwners; }; // --------------------------------------------------------------------------- diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storageManager.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storageManager.cpp index 4dc97be91eae..a0d22311ba18 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storageManager.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storageManager.cpp @@ -20,6 +20,7 @@ #include "kv_cache_manager_v2/common.h" #include "kv_cache_manager_v2/copyEngine.h" #include "kv_cache_manager_v2/exceptions.h" +#include "kv_cache_manager_v2/kvCache.h" #include "kv_cache_manager_v2/page.h" #include "kv_cache_manager_v2/stagingBuffer.h" #include "kv_cache_manager_v2/utils/hostMem.h" @@ -1263,6 +1264,170 @@ void StorageManager::batchedMigrate( } } +void StorageManager::offloadSparsePages(KvCache& requestingCache, std::vector> const& pages, + MigrationRecorder const& migrationRecorder, DropRecorder const& dropRecorder) +{ + struct OffloadBatch + { + std::vector> pages; + std::vector> locks; + std::vector slots; + std::vector indices; + }; + + std::map batches; + std::set seen; + std::set ownerStreams; + for (auto const& page : pages) + { + if (!page || page->manager != this || requestingCache.storageManager() != this) + { + throw LogicError("Offload pages and requesting cache must belong to the same manager"); + } + if (!seen.insert(page.get()).second) + { + continue; + } + auto holder = page->holder.lock(); + auto lock = holder ? holder->uniqLock.lock() : nullptr; + if (!lock) + { + throw LogicError("Sparse history offload requires a locked page"); + } + lock->prepareSparseOffload(requestingCache); + if (page->cacheLevel == kSparseHistoryLevel) + { + continue; + } + auto& batch = batches[getMigrationBatchingLayerGroupId(kSparseHistoryLevel, kHotLevel, page->lifeCycle)]; + batch.pages.push_back(page); + for (auto const& owner : lock->owners()) + { + ownerStreams.insert(owner.kvCache->cudaStream()); + } + batch.locks.push_back(std::move(lock)); + } + if (batches.empty()) + { + return; + } + + TypedVec requirements(numPoolGroups(kSparseHistoryLevel), 0); + for (auto const& [layerGroup, batch] : batches) + { + requirements[getPoolGroupIndex(kSparseHistoryLevel, layerGroup)] += slotCountValueFromSize(batch.pages.size()); + } + prepareFreeSlots(kSparseHistoryLevel, requirements, migrationRecorder, dropRecorder); + auto releaseDestinations = FuncGuard( + [&]() + { + for (auto& [layerGroup, batch] : batches) + { + for (auto& slot : batch.slots) + { + if (slot.hasValidSlot()) + { + releaseSlot(layerGroup, kSparseHistoryLevel, std::move(slot)); + } + } + } + }); + for (auto& [layerGroup, batch] : batches) + { + auto& pool = poolGroup(kSparseHistoryLevel, getPoolGroupIndex(kSparseHistoryLevel, layerGroup)); + batch.slots = pool.allocateMultiple(slotCountValueFromSize(batch.pages.size())); + batch.indices.reserve(batch.pages.size()); + for (size_t i = 0; i < batch.pages.size(); ++i) + { + batch.indices.push_back({.dst = slotIdToPageIndexValue(batch.slots[i].slotId()), + .src = slotIdToPageIndexValue(batch.pages[i]->slotId())}); + } + } + + CUstream const stream = requestingCache.cudaStream(); + auto const cudaStream = reinterpret_cast(stream); + std::vector ownerEvents; + ownerEvents.reserve(ownerStreams.size()); + for (auto const ownerStream : ownerStreams) + { + ownerEvents.emplace_back(reinterpret_cast(ownerStream)); + ownerEvents.back().waitInStream(cudaStream); + } + for (auto const& [layerGroup, batch] : batches) + { + for (size_t i = 0; i < batch.pages.size(); ++i) + { + batch.pages[i]->readyEvent.waitInStream(cudaStream); + batch.slots[i].readyEvent.waitInStream(cudaStream); + for (auto const& event : batch.locks[i]->finishEvents) + { + event.waitInStream(cudaStream); + } + } + } + + // Install the fence even when a codec rejects after enqueueing only part of a batch. + CachedCudaEvent completion = CachedCudaEvent::makeNull(); + auto fenceCopies = FuncGuard( + [&]() + { + completion = CachedCudaEvent(cudaStream); + for (auto& [layerGroup, batch] : batches) + { + for (size_t i = 0; i < batch.pages.size(); ++i) + { + batch.slots[i].readyEvent = completion; + batch.locks[i]->recordOffloadEvent(completion); + } + } + }); + for (auto const& [layerGroup, batch] : batches) + { + submitMigrationBatch( + kSparseHistoryLevel, kHotLevel, layerGroup, batch.indices.data(), batch.indices.size(), stream); + } + fenceCopies.run(); + + // Subsequent host readers on every owner's stream must observe the completed copy. + for (auto const ownerStream : ownerStreams) + { + completion.waitInStream(reinterpret_cast(ownerStream)); + } + for (auto const& [layerGroup, batch] : batches) + { + if (migrationRecorder) + { + migrationRecorder(batch.pages, batch.slots, kHotLevel, kSparseHistoryLevel); + } + } + for (auto& [layerGroup, batch] : batches) + { + for (size_t i = 0; i < batch.pages.size(); ++i) + { + Slot source = batch.locks[i]->moveToSparseHistory(std::move(batch.slots[i])); + releaseSlot(batch.pages[i]->lifeCycle, kHotLevel, std::move(source)); + } + } + if (mEventSink) + { + for (auto const& [layerGroup, batch] : batches) + { + for (auto const& page : batch.pages) + { + if (page->isCommitted()) + { + auto const& committed = static_cast(*page); + auto const* block = committed.block; + if (block && !block->isOrphan() && block->holdsPage(committed)) + { + mEventSink->addCacheLevelUpdated(block->key, kHotLevel, kSparseHistoryLevel, page->lifeCycle); + } + } + } + } + } +} + int64_t StorageManager::prefetch( CacheLevel dstLevel, TypedVec>>> const& pages) { diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storageManager.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storageManager.h index 144f0f527803..684d34389b8c 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storageManager.h +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storageManager.h @@ -192,6 +192,11 @@ class StorageManager : public std::enable_shared_from_this void batchedMigrate( CacheLevel dstLevel, std::vector> const& pages, MigrationRecorder const& migrationRecorder); + //! Demote complete locked sparse pages on the requesting owner's stream, updating every live owner. + //! Caller holds the manager's exclusive lock. Allocation/copy failure preserves GPU ownership. + void offloadSparsePages(KvCache& requestingCache, std::vector> const& pages, + MigrationRecorder const& migrationRecorder = {}, DropRecorder const& dropRecorder = {}); + // Best-effort migration of grouped pages to a destination cache level. Returns how many pages // it moved off the disk tier, counted per migrated batch rather than per page. A throw reports // nothing, which in practice means slot preparation failed before anything moved. diff --git a/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp b/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp index f2026be52fa1..c80c897b11d3 100644 --- a/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp +++ b/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp @@ -18,6 +18,7 @@ #include "kvCacheManagerV2TestUtils.h" #include "tensorrt_llm/batch_manager/kv_cache_manager_v2/blockRadixTree.h" #include "tensorrt_llm/batch_manager/kv_cache_manager_v2/config.h" +#include "tensorrt_llm/batch_manager/kv_cache_manager_v2/eventManager.h" #include "tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.h" #include "tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCacheManager.h" #include "tensorrt_llm/batch_manager/kv_cache_manager_v2/storageManager.h" @@ -301,6 +302,91 @@ class AsyncRejectingColdPageCodec final : public IKvCacheColdPageCodec std::atomic mRelease{false}; }; +class ObservingColdPageCodec final : public IKvCacheColdPageCodec +{ +public: + bool configure(PoolGroupDesc const* descriptors, PoolGroupIndex count) noexcept override + { + return mCodec->configure(descriptors, count); + } + + size_t queryColdPageBytes(LayerGroupId layerGroup) const noexcept override + { + return mCodec->queryColdPageBytes(layerGroup); + } + + LayerGroupId getBatchingLayerGroupId(LayerGroupId layerGroup) const noexcept override + { + return mCodec->getBatchingLayerGroupId(layerGroup); + } + + PageIndexLocation queryPageIndexLocation(LayerGroupId layerGroup) const noexcept override + { + return mCodec->queryPageIndexLocation(layerGroup); + } + + bool encode(LayerGroupId layerGroup, void* destination, PageIndexPair const* indices, size_t count, + cudaStream_t stream) noexcept override + { + ++encodeCalls; + encodedPages += count; + encodeStream = stream; + bool const submitted = mCodec->encode(layerGroup, destination, indices, count, stream); + return submitted && encodeCalls != rejectEncodeCall; + } + + bool decode(LayerGroupId layerGroup, void const* source, PageIndexPair const* indices, size_t count, + cudaStream_t stream) noexcept override + { + return mCodec->decode(layerGroup, source, indices, count, stream); + } + + size_t encodeCalls = 0; + size_t encodedPages = 0; + size_t rejectEncodeCall = 0; + cudaStream_t encodeStream{}; + +private: + std::unique_ptr mCodec = createDefaultKvCacheColdPageCodec(); +}; + +class StreamGate +{ +public: + ~StreamGate() + { + release(); + if (mStream) + { + cudaStreamSynchronize(mStream); + } + } + + cudaError_t enqueue(cudaStream_t stream) + { + mStream = stream; + return cudaLaunchHostFunc(stream, wait, this); + } + + void release() noexcept + { + mRelease.store(true, std::memory_order_release); + } + +private: + static void CUDART_CB wait(void* data) + { + auto& gate = *static_cast(data); + while (!gate.mRelease.load(std::memory_order_acquire)) + { + std::this_thread::yield(); + } + } + + std::atomic mRelease{false}; + cudaStream_t mStream{}; +}; + SharedPtr makeCommittedPage(KvCacheManager& manager, StorageManager& storage, CacheLevel level, Slot& slot, LifeCycleId lifeCycle = LifeCycleId{0}, Priority priority = kPriorityDefault, int tokenBase = 0) { @@ -762,6 +848,436 @@ class KvCacheManagerV2PageLockTest : public ::testing::Test cudaStream_t mStream{}; }; +class KvCacheManagerV2SparseOffloadTest : public KvCacheManagerV2PageLockTest +{ +}; + +TEST_F(KvCacheManagerV2SparseOffloadTest, BatchesCompleteCoalescedPagesAndCountsPhysicalCopies) +{ + auto config = makeSplitColdGroupingConfig(); + config.enableStats = true; + for (auto& layerConfig : config.layers) + { + auto& layer = std::get(layerConfig); + layer.buffers.front().isSparse = true; + layer.buffers.push_back({.role = "value", .size = 2048, .tokensPerBlockOverride = 2, .isSparse = true}); + layer.buffers.push_back({.role = "scale", .size = 128, .isSparse = true}); + } + auto codec = std::make_unique(); + auto* observer = codec.get(); + auto manager = std::make_shared(std::move(config), nullptr, std::move(codec)); + auto const apiLock = manager->lockExclusive(); + auto& storage = manager->storage(); + ASSERT_EQ(storage.numLifeCycles(), LifeCycleId{2}); + ASSERT_EQ(storage.numPoolGroups(kHotLevel), PoolGroupIndex{1}); + ASSERT_GT(storage.numPools(PoolGroupIndex{0}), PoolIndex{1}); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(12, 8)); + + std::vector> pages; + std::vector> expected; + for (int ordinal = 0; ordinal < 2; ++ordinal) + { + for (LifeCycleId lc{0}; lc < storage.numLifeCycles(); ++lc) + { + auto page = pageAt(*cache, ordinal, lc); + auto const pg = storage.getPoolGroupIndex(kHotLevel, lc); + auto const& sizes = storage.slotSize(kHotLevel, pg); + std::vector bytes; + for (PoolIndex pool{0}; pool < sizes.size(); ++pool) + { + auto const pattern = static_cast(17 * pages.size() + pool.value() + 1); + auto const address = std::get(storage.slotAddress(kHotLevel, pg, page->slotId(), pool)); + ASSERT_EQ( + cudaMemsetAsync(reinterpret_cast(address), pattern, sizes[pool], mStream), cudaSuccess); + bytes.insert(bytes.end(), sizes[pool], pattern); + } + pages.push_back(std::move(page)); + expected.push_back(std::move(bytes)); + } + } + auto const gpuFree = storage.getStatistics(kHotLevel).free; + auto const hostFree = storage.getStatistics(kSparseHistoryLevel).free; + auto targets = pages; + targets.push_back(pages.front()); + cache->offloadSparsePages(targets); + EXPECT_EQ(observer->encodeCalls, 1); + EXPECT_EQ(observer->encodedPages, pages.size()); + EXPECT_EQ(observer->encodeStream, mStream); + EXPECT_EQ(storage.getStatistics(kHotLevel).free, gpuFree + pages.size()); + EXPECT_EQ(storage.getStatistics(kSparseHistoryLevel).free, hostFree - pages.size()); + EXPECT_EQ(cache->pageStorageVersion(), pages.size()); + + for (size_t i = 0; i < pages.size(); ++i) + { + auto const& page = pages[i]; + EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(page->status(), PageStatus::LOCKED); + EXPECT_FALSE(page->scheduledForEviction()); + page->readyEvent.synchronize(); + auto const pg = storage.getPoolGroupIndex(kSparseHistoryLevel, page->lifeCycle); + auto const address + = std::get(storage.slotAddress(kSparseHistoryLevel, pg, page->slotId(), PoolIndex{0})); + EXPECT_EQ(std::memcmp(reinterpret_cast(address), expected[i].data(), expected[i].size()), 0); + } + for (LifeCycleId lc{0}; lc < storage.numLifeCycles(); ++lc) + { + EXPECT_EQ(pageAt(*cache, 2, lc)->cacheLevel, kHotLevel); + auto const indices = cache->getBasePageIndices(lc); + for (int ordinal = 0; ordinal < 3; ++ordinal) + { + EXPECT_EQ(indices[ordinal], slotIdToPageIndexValue(pageAt(*cache, ordinal, lc)->slotId())); + } + } + auto const stats = manager->getAndResetIterationStats(); + for (LifeCycleId lc{0}; lc < storage.numLifeCycles(); ++lc) + { + EXPECT_EQ(stats.at(lc).iterOffloadBlocks, 2); + EXPECT_EQ(stats.at(lc).iterOffloadBytes, 2 * expected.front().size()); + } + cache->offloadSparsePages(targets); + EXPECT_EQ(observer->encodeCalls, 1); + EXPECT_EQ(cache->pageStorageVersion(), pages.size()); + EXPECT_TRUE(manager->getAndResetIterationStats().empty()); +} + +TEST_F(KvCacheManagerV2SparseOffloadTest, SharedOwnersPublishHostIndicesAndKeepHistoryPinned) +{ + auto manager = std::make_shared(sparseConfig()); + auto const apiLock = manager->lockExclusive(); + auto& storage = manager->storage(); + auto page = seedPrefix(*manager, kHotLevel); + auto first = manager->createKvCache({}, tokens()); + auto second = manager->createKvCache({}, tokens()); + std::vector externalIndices(1, kBadPageIndex.value()); + auto closeCaches = FuncGuard( + [&]() + { + first->close(); + second->close(); + }); + ASSERT_TRUE(first->resume(stream())); + ASSERT_TRUE(second->resume(stream())); + second->setBasePageIndexBuf(kDefaultBeamIndex, LifeCycleId{0}, externalIndices.data(), externalIndices.size()); + auto const gpuSlot = page->slotId(); + auto hostBlockers = storage.newSlots(kSparseHistoryLevel, TypedVec{1}); + auto releaseBlockers = FuncGuard([&]() + { storage.releaseSlot(LifeCycleId{0}, kSparseHistoryLevel, std::move(hostBlockers[LifeCycleId{0}].front())); }); + + EXPECT_THROW(storage.batchedMigrate(kSparseHistoryLevel, {page}, {}), LogicError); + first->offloadSparsePages({page, page}); + EXPECT_NE(page->slotId(), gpuSlot); + EXPECT_EQ(first->getBasePageIndices(LifeCycleId{0})[0], slotIdToPageIndexValue(page->slotId())); + EXPECT_EQ(externalIndices[0], slotIdToPageIndexValue(page->slotId())); + EXPECT_EQ(first->pageStorageVersion(), 1); + EXPECT_EQ(second->pageStorageVersion(), 1); + EXPECT_EQ(storage.getStatistics(kHotLevel).free, storage.getStatistics(kHotLevel).total); + EXPECT_FALSE(storage.isEvictable(*page)); + EXPECT_THROW(storage.batchedMigrate(kHotLevel, {page}, {}), LogicError); + first->suspend(); + EXPECT_EQ(page->status(), PageStatus::LOCKED); + ASSERT_TRUE(first->resume()); + EXPECT_EQ(pageAt(*first), page); + second->close(); + EXPECT_EQ(externalIndices[0], kBadPageIndex.value()); + EXPECT_EQ(page->status(), PageStatus::LOCKED); + EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); +} + +TEST_F(KvCacheManagerV2SparseOffloadTest, WaitsForLiveAndFinishedReadersBeforeRecyclingGpuSlot) +{ + for (bool const finishReader : {false, true}) + { + SCOPED_TRACE(finishReader); + auto manager = std::make_shared(sparseConfig()); + auto const apiLock = manager->lockExclusive(); + auto& storage = manager->storage(); + cudaStream_t readerStream{}; + ASSERT_EQ(cudaStreamCreateWithFlags(&readerStream, cudaStreamNonBlocking), cudaSuccess); + auto destroyReaderStream = FuncGuard([&]() { cudaStreamDestroy(readerStream); }); + auto page = seedPrefix(*manager, kHotLevel); + auto first = manager->createKvCache({}, tokens()); + auto second = manager->createKvCache({}, tokens()); + auto closeCaches = FuncGuard( + [&]() + { + first->close(); + second->close(); + }); + ASSERT_TRUE(first->resume(stream())); + ASSERT_TRUE(second->resume(reinterpret_cast(readerStream))); + auto const lc = page->lifeCycle; + auto const pg = storage.getPoolGroupIndex(kHotLevel, lc); + size_t const bytes = storage.slotSize(kHotLevel, pg)[PoolIndex{0}]; + auto const gpuSlot = page->slotId(); + auto const gpuAddress = std::get(storage.slotAddress(kHotLevel, pg, gpuSlot, PoolIndex{0})); + constexpr uint8_t kPattern = 0xA6; + ASSERT_EQ(cudaMemsetAsync(reinterpret_cast(gpuAddress), kPattern, bytes, mStream), cudaSuccess); + auto gpuBlocker = storage.newGpuSlots(TypedVec{1}); + auto hostScratch = storage.newSlots(kSparseHistoryLevel, TypedVec{1}); + auto releaseSlots = FuncGuard( + [&]() + { + storage.releaseSlot(lc, kHotLevel, std::move(gpuBlocker[lc].front())); + storage.releaseSlot(lc, kSparseHistoryLevel, std::move(hostScratch[lc].front())); + }); + // Warm the codec's descriptor/index staging before deliberately blocking a stream. + storage.copySlotData(lc, kSparseHistoryLevel, kHotLevel, hostScratch[lc].front().slotId(), gpuSlot, stream()); + ASSERT_EQ(cudaStreamSynchronize(mStream), cudaSuccess); + auto const readback = std::get(storage.slotAddress(kSparseHistoryLevel, + storage.getPoolGroupIndex(kSparseHistoryLevel, lc), hostScratch[lc].front().slotId(), PoolIndex{0})); + StreamGate gate; + ASSERT_EQ(gate.enqueue(readerStream), cudaSuccess); + ASSERT_EQ(cudaMemcpyAsync(reinterpret_cast(readback), reinterpret_cast(gpuAddress), bytes, + cudaMemcpyDeviceToHost, readerStream), + cudaSuccess); + if (finishReader) + { + second->suspend(); + } + + first->offloadSparsePages({page}); + EXPECT_FALSE(page->queryReady()); + EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); + auto recycled = storage.newGpuSlots(TypedVec{1}); + auto releaseRecycled + = FuncGuard([&]() { storage.releaseSlot(lc, kHotLevel, std::move(recycled[lc].front())); }); + EXPECT_EQ(recycled[lc].front().slotId(), gpuSlot); + EXPECT_FALSE(recycled[lc].front().queryReady()); + recycled[lc].front().readyEvent.waitInStream(reinterpret_cast(mStream)); + ASSERT_EQ(cudaMemsetAsync(reinterpret_cast(gpuAddress), 0, bytes, mStream), cudaSuccess); + gate.release(); + ASSERT_EQ(cudaStreamSynchronize(mStream), cudaSuccess); + auto const* readBytes = reinterpret_cast(readback); + EXPECT_TRUE(std::all_of(readBytes, readBytes + bytes, [](uint8_t value) { return value == kPattern; })); + page->readyEvent.synchronize(); + auto const hostAddress = std::get(storage.slotAddress( + kSparseHistoryLevel, storage.getPoolGroupIndex(kSparseHistoryLevel, lc), page->slotId(), PoolIndex{0})); + auto const* hostBytes = reinterpret_cast(hostAddress); + EXPECT_TRUE(std::all_of(hostBytes, hostBytes + bytes, [](uint8_t value) { return value == kPattern; })); + } +} + +TEST_F(KvCacheManagerV2SparseOffloadTest, HostOomLeavesEntireBatchOnGpu) +{ + auto config = sparseConfig(); + config.cacheTiers[1] = HostCacheTierConfig{2 << 20}; + auto codec = std::make_unique(); + auto* observer = codec.get(); + auto manager = std::make_shared(std::move(config), nullptr, std::move(codec)); + auto const apiLock = manager->lockExclusive(); + auto& storage = manager->storage(); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(8, 8)); + auto first = pageAt(*cache); + auto second = pageAt(*cache, 1); + auto const firstSlot = first->slotId(); + auto const secondSlot = second->slotId(); + EXPECT_THROW(cache->offloadSparsePages({first, second}), OutOfPagesError); + EXPECT_EQ(observer->encodeCalls, 0); + EXPECT_EQ(first->cacheLevel, kHotLevel); + EXPECT_EQ(second->cacheLevel, kHotLevel); + EXPECT_EQ(first->slotId(), firstSlot); + EXPECT_EQ(second->slotId(), secondSlot); + EXPECT_EQ(first->status(), PageStatus::LOCKED); + EXPECT_EQ(second->status(), PageStatus::LOCKED); + EXPECT_EQ(cache->pageStorageVersion(), 0); + EXPECT_EQ(storage.getStatistics(kSparseHistoryLevel).free, 1); + EXPECT_EQ(storage.getStatistics(kHotLevel).free, 0); + EXPECT_EQ(cache->getBasePageIndices(LifeCycleId{0})[0], slotIdToPageIndexValue(firstSlot)); + EXPECT_EQ(cache->getBasePageIndices(LifeCycleId{0})[1], slotIdToPageIndexValue(secondSlot)); +} + +TEST_F(KvCacheManagerV2SparseOffloadTest, AsynchronousRejectionFencesBothSlotsWithoutPublishingHostIndices) +{ + auto config = sparseConfig(); + config.enableStats = true; + auto codec = std::make_unique(AsyncRejectingColdPageCodec::Operation::kEncode); + auto* rejecting = codec.get(); + auto manager = std::make_shared(std::move(config), nullptr, std::move(codec)); + auto const apiLock = manager->lockExclusive(); + auto& storage = manager->storage(); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(4, 4)); + auto page = pageAt(*cache); + auto const gpuSlot = page->slotId(); + auto const lc = page->lifeCycle; + auto blocker = storage.newSlots(kSparseHistoryLevel, TypedVec{1}); + auto releaseBlocker + = FuncGuard([&]() { storage.releaseSlot(lc, kSparseHistoryLevel, std::move(blocker[lc].front())); }); + auto releaseCodec = FuncGuard([&]() { rejecting->release(); }); + EXPECT_THROW(cache->offloadSparsePages({page}), TllmException); + ASSERT_TRUE(rejecting->launched()); + EXPECT_EQ(page->cacheLevel, kHotLevel); + EXPECT_EQ(page->slotId(), gpuSlot); + EXPECT_EQ(page->status(), PageStatus::LOCKED); + EXPECT_FALSE(page->queryReady()); + EXPECT_EQ(cache->pageStorageVersion(), 0); + EXPECT_EQ(cache->getBasePageIndices(lc)[0], slotIdToPageIndexValue(gpuSlot)); + auto recycled = storage.newSlots(kSparseHistoryLevel, TypedVec{1}); + auto releaseRecycled + = FuncGuard([&]() { storage.releaseSlot(lc, kSparseHistoryLevel, std::move(recycled[lc].front())); }); + EXPECT_FALSE(recycled[lc].front().queryReady()); + EXPECT_TRUE(manager->getAndResetIterationStats().empty()); + rejecting->release(); + recycled[lc].front().readyEvent.synchronize(); +} + +TEST_F(KvCacheManagerV2SparseOffloadTest, LaterCodecBatchFailurePreservesAllSourcePages) +{ + auto config = makeSplitColdGroupingConfig(); + for (auto& layer : config.layers) + { + std::get(layer).buffers.front().isSparse = true; + } + std::get(config.layers[1]).buffers.front().size *= 2; + auto codec = std::make_unique(); + auto* observer = codec.get(); + observer->rejectEncodeCall = 2; + auto manager = std::make_shared(std::move(config), nullptr, std::move(codec)); + auto const apiLock = manager->lockExclusive(); + auto& storage = manager->storage(); + ASSERT_EQ(storage.numPoolGroups(kHotLevel), PoolGroupIndex{2}); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(4, 4)); + auto first = pageAt(*cache, 0, LifeCycleId{0}); + auto second = pageAt(*cache, 0, LifeCycleId{1}); + auto const firstSlot = first->slotId(); + auto const secondSlot = second->slotId(); + EXPECT_THROW(cache->offloadSparsePages({first, second}), TllmException); + EXPECT_EQ(observer->encodeCalls, 2); + EXPECT_EQ(first->cacheLevel, kHotLevel); + EXPECT_EQ(second->cacheLevel, kHotLevel); + EXPECT_EQ(first->slotId(), firstSlot); + EXPECT_EQ(second->slotId(), secondSlot); + EXPECT_EQ(cache->pageStorageVersion(), 0); + EXPECT_EQ(storage.getStatistics(kSparseHistoryLevel).free, storage.getStatistics(kSparseHistoryLevel).total); + observer->rejectEncodeCall = 0; + cache->offloadSparsePages({first, second}); + EXPECT_EQ(first->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(second->cacheLevel, kSparseHistoryLevel); +} + +TEST_F(KvCacheManagerV2SparseOffloadTest, RejectsWritablePartialAndDensePagesBeforeAnyCopy) +{ + for (bool const sparse : {false, true}) + { + for (int const history : {0, 2, 4}) + { + SCOPED_TRACE(sparse); + SCOPED_TRACE(history); + auto config = sparse ? sparseConfig() : makeTieredConfig(); + auto codec = std::make_unique(); + auto* observer = codec.get(); + auto manager = std::make_shared(std::move(config), nullptr, std::move(codec)); + auto const apiLock = manager->lockExclusive(); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(8, history)); + auto first = pageAt(*cache); + auto input = pageAt(*cache, 1); + EXPECT_THROW(cache->offloadSparsePages({first, input}), LogicError); + EXPECT_EQ(observer->encodeCalls, 0); + EXPECT_EQ(first->cacheLevel, kHotLevel); + EXPECT_EQ(input->cacheLevel, kHotLevel); + EXPECT_EQ(cache->pageStorageVersion(), 0); + } + } +} + +TEST_F(KvCacheManagerV2SparseOffloadTest, RejectsPartialCommittedPagesAndNonOwners) +{ + auto config = sparseConfig(); + config.commitMinSnapshot = true; + auto manager = std::make_shared(std::move(config)); + auto const apiLock = manager->lockExclusive(); + auto cache = manager->createKvCache(); + auto other = manager->createKvCache(); + auto closeCaches = FuncGuard( + [&]() + { + cache->close(); + other->close(); + }); + EXPECT_THROW(cache->offloadSparsePages({}), LogicError); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(other->resume(stream())); + ASSERT_TRUE(cache->resize(4, 2)); + cache->commit(tokens(2), /*isEnd=*/true); + auto partial = pageAt(*cache); + ASSERT_TRUE(partial->isCommitted()); + EXPECT_THROW(cache->offloadSparsePages({partial}), LogicError); + EXPECT_EQ(partial->cacheLevel, kHotLevel); + ASSERT_TRUE(other->resize(4, 4)); + auto complete = pageAt(*other); + EXPECT_THROW(cache->offloadSparsePages({complete}), LogicError); + EXPECT_EQ(complete->cacheLevel, kHotLevel); + EXPECT_THROW(other->offloadSparsePages({nullptr}), LogicError); +} + +TEST_F(KvCacheManagerV2SparseOffloadTest, EmitsCommittedTierChangeEvenWhenSlotIndexIsUnchanged) +{ + auto events = std::make_shared(128); + auto manager = std::make_shared(sparseConfig(), events); + auto const apiLock = manager->lockExclusive(); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(4, 4)); + cache->commit(tokens()); + auto page = pageAt(*cache); + ASSERT_TRUE(page->isCommitted()); + auto const originalIndex = cache->getBasePageIndices(LifeCycleId{0})[0]; + events->flushIterationEvents(); + events->getLatestEvents(/*timeoutMs=*/0); + + cache->offloadSparsePages({page, page}); + EXPECT_EQ(cache->getBasePageIndices(LifeCycleId{0})[0], originalIndex); + EXPECT_EQ(cache->pageStorageVersion(), 1); + events->flushIterationEvents(); + auto const updates = events->getLatestEvents(/*timeoutMs=*/0); + ASSERT_EQ(updates.size(), 1); + EXPECT_EQ(updates.front().layerGroupId, 0); + auto const* data = std::get_if(&updates.front().data); + ASSERT_NE(data, nullptr); + ASSERT_TRUE(data->cacheLevel.has_value()); + EXPECT_EQ(data->cacheLevel->oldValue, kHotLevel.value()); + EXPECT_EQ(data->cacheLevel->newValue, kSparseHistoryLevel.value()); +} + +TEST_F(KvCacheManagerV2SparseOffloadTest, DemotedUncommittedHistoryCanCommitAndBeReused) +{ + auto manager = std::make_shared(sparseConfig()); + auto const apiLock = manager->lockExclusive(); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(4, 4)); + cache->offloadSparsePages({pageAt(*cache)}); + auto const hostSlot = pageAt(*cache)->slotId(); + cache->commit(tokens()); + auto committed = pageAt(*cache); + ASSERT_TRUE(committed->isCommitted()); + EXPECT_EQ(committed->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(committed->slotId(), hostSlot); + auto reused = manager->createKvCache({}, tokens()); + auto closeReused = FuncGuard([&]() { reused->close(); }); + ASSERT_TRUE(reused->resume(stream())); + EXPECT_EQ(pageAt(*reused), committed); + cache->close(); + EXPECT_EQ(committed->status(), PageStatus::LOCKED); + EXPECT_EQ(reused->getBasePageIndices(LifeCycleId{0})[0], slotIdToPageIndexValue(hostSlot)); +} + TEST_F(KvCacheManagerV2PageLockTest, SparseHostPrefixStaysPinnedAcrossReuseAndResume) { auto manager = std::make_shared(sparseConfig()); From 25d01e521cf7e754e70760d4c830923dd3e8332f Mon Sep 17 00:00:00 2001 From: Guiju Zhang <7135567+cascade812@users.noreply.github.com> Date: Thu, 1 Oct 2026 12:13:14 -0700 Subject: [PATCH 2/5] trigger sparse offload during decode Signed-off-by: Guiju Zhang <7135567+cascade812@users.noreply.github.com> --- .../kv_cache_manager_v2/kvCache.cpp | 324 ++++++++--- .../kv_cache_manager_v2/kvCache.h | 29 +- .../kv_cache_manager_v2/lifeCycleRegistry.h | 3 +- .../kv_cache_manager_v2/storageManager.cpp | 55 +- .../batch_manager/kvCacheManagerV2.cpp | 8 +- .../kvCacheManagerV2ColdPageTest.cpp | 504 +++++++++++++++++- .../kv_cache/kv_cache_manager_v2.py | 4 +- .../runtime/kv_cache_manager_v2/__init__.pyi | 7 +- .../kv_cache/test_kv_cache_v2_scheduler.py | 33 +- 9 files changed, 832 insertions(+), 135 deletions(-) diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.cpp index b8b994beb9e2..4d5ebf26245e 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.cpp @@ -210,7 +210,72 @@ CacheLevel KvCache::_lockLevel(Page const& page, BlockOrdinal ordinal) const { bool const readOnly = page.isCommitted() || (ordinal != kBadBlockOrdinal && ordinal < BlockOrdinal{mHistoryLength / mTokensPerBlock}); - return readOnly ? page.queryLockLevel() : kHotLevel; + return mIsDecoding && readOnly ? page.queryLockLevel() : kHotLevel; +} + +void KvCache::_offloadSparseHistory(HalfOpenRange range, int historyLength) +{ + TLLM_CHECK_DEBUG(range.end <= BlockOrdinal{historyLength / mTokensPerBlock}); + std::vector> pages; + for (auto const& [lcId, lc] : mManager->lifeCycles()) + { + auto const* attn = std::get_if(&lc); + if (!attn || !attn->isSparse) + continue; + if (range.end > mBlocks.size()) + throw std::invalid_argument("Sparse history must already have allocated pages"); + for (BlockOrdinal ord = range.beg; ord < range.end; ++ord) + { + for (BeamIndex bi{0}; bi < mBeamWidth; ++bi) + { + auto page = _page(ord, bi, lcId); + TLLM_CHECK_DEBUG(page && page->status() == PageStatus::LOCKED); + if (page->cacheLevel == kSparseHistoryLevel) + continue; + auto const lock = page->holder.lock()->uniqLock.lock(); + for (auto const& owner : lock->owners()) + { + if (owner.kvCache != this && !owner.kvCache->isDecoding()) + throw LogicError("Cannot offload sparse history shared with a prefill request"); + } + pages.push_back(std::move(page)); + } + } + } + if (pages.empty()) + return; + int const oldHistoryLength = mHistoryLength; + auto restoreHistory = FuncGuard([&]() { mHistoryLength = oldHistoryLength; }); + mHistoryLength = historyLength; + offloadSparsePages(pages); +} + +void KvCache::_publishHistoryLength(int historyLength) +{ + if (mIsDecoding && historyLength / mTokensPerBlock != mHistoryLength / mTokensPerBlock) + onPageStorageChanged(); + mHistoryLength = historyLength; +} + +bool KvCache::enterDecode() +{ + KVCM2_API_GUARD(); + auto const apiLock = mManager->lockExclusive(); + if (!isActive()) + throw LogicError("Decode admission requires an active request"); + if (mIsDecoding) + return true; + try + { + _offloadSparseHistory({0, mHistoryLength / mTokensPerBlock}, mHistoryLength); + } + catch (OutOfPagesError const&) + { + return false; + } + mIsDecoding = true; + onPageStorageChanged(); + return true; } void KvCache::offloadSparsePages(std::vector> const& pages) @@ -236,7 +301,7 @@ void KvCache::activate() mFinishEvent.reset(); - // Cold sparse history stays on host; writable pages require GPU storage. + // Prefill restores GPU storage; decode can retain cold sparse history on host. auto activePages = _activePages(); std::vector targets; targets.reserve(activePages.size()); @@ -254,7 +319,14 @@ void KvCache::activate() } auto& holder = std::get>(*bp); TLLM_CHECK_DEBUG(holder); - targets.push_back({holder->page, ap.beamIdx, ap.ordinal, ap.lcId, _lockLevel(*holder->page, ap.ordinal)}); + // A reused partial block is only a copy source. resume() replaces it with a private GPU + // page before execution, so another owner's host lock need not be moved. + bool const partialCopySource = mNeverResumed && ap.ordinal != kBadBlockOrdinal + && numCommittedTokens() % mTokensPerBlock != 0 + && ap.ordinal == BlockOrdinal{numCommittedTokens() / mTokensPerBlock}; + CacheLevel const level + = partialCopySource ? holder->page->queryLockLevel() : _lockLevel(*holder->page, ap.ordinal); + targets.push_back({holder->page, ap.beamIdx, ap.ordinal, ap.lcId, level}); } { @@ -273,7 +345,7 @@ void KvCache::activate() } } -bool KvCache::resume(std::optional stream) +bool KvCache::resume(std::optional stream, std::optional isDecoding) { KVCM2_API_GUARD(); TLLM_CHECK(mStatus == Status::SUSPENDED); @@ -287,6 +359,11 @@ bool KvCache::resume(std::optional stream) TLLM_CHECK_DEBUG(!mFinishEvent.has_value()); auto const apiLock = mManager->lockExclusive(); + bool const oldIsDecoding = mIsDecoding; + if (mIsDecoding && isDecoding == false) + throw std::invalid_argument("Cannot return a decoding cache to prefill"); + auto restorePhase = FuncGuard([&]() { mIsDecoding = oldIsDecoding; }); + mIsDecoding = isDecoding.value_or(mIsDecoding); // Check utilization against threshold. auto const utilizations = mManager->storage().getUtilization(kHotLevel); @@ -303,6 +380,25 @@ bool KvCache::resume(std::optional stream) // Pre-allocate GPU slots for deferred copies (partial blocks + SSM) and scratch slots // before locking, so we never end up in a state where pages are locked but we can't allocate. TypedVec> deferredSlots(numLc); + bool deferredCopiesStarted = false; + auto releaseDeferredSlots = FuncGuard( + [&]() + { + std::optional completion; + for (LifeCycleId lc{0}; lc < numLc; ++lc) + { + if (deferredSlots[lc].has_value() && deferredSlots[lc]->hasValidSlot()) + { + if (deferredCopiesStarted) + { + if (!completion) + completion.emplace(reinterpret_cast(cudaStream())); + deferredSlots[lc]->readyEvent = *completion; + } + storageMgr.releaseSlot(lc, kHotLevel, std::move(*deferredSlots[lc])); + } + } + }); // Compute scratch slot deltas UNCONDITIONALLY (mirrors Python: _take_excess_scratch_slots // is called outside _never_resumed). @@ -388,16 +484,12 @@ bool KvCache::resume(std::optional stream) } catch (OutOfPagesError const&) { - // Release pre-allocated deferred slots on failure. - for (LifeCycleId lc{0}; lc < numLc; ++lc) - { - if (deferredSlots[lc].has_value()) - storageMgr.releaseSlot(lc, kHotLevel, std::move(*deferredSlots[lc])); - } // Scratch slots stay in mScratchSlots — they'll be freed by close() inside // a recordEventScope, matching Python behavior. return false; } + mStatus = Status::ACTIVE; + auto rollbackActivation = FuncGuard([&]() { _deactivate(); }); // Deferred copy: for partial blocks and SSM, copy from now-locked source pages // to pre-allocated GPU slots, then unlock sources and replace with new pages. @@ -446,6 +538,7 @@ bool KvCache::resume(std::optional stream) srcLocks.push_back(lock); CacheLevel const sourceLevel = lock->page()->cacheLevel; + deferredCopiesStarted = true; storageMgr.copySlotData(lcIdx, kHotLevel, sourceLevel, newSlot.slotId(), lock->page()->slotId(), cudaStr); if ((!ssmLcId.has_value() || lcIdx != *ssmLcId) && (recordManagerStats || recordRequestStats)) { @@ -511,17 +604,28 @@ bool KvCache::resume(std::optional stream) mBlocks[lastOrdinal].treeBlock = nullptr; } - // A freshly-created cache starts SUSPENDED and is activated by this same - // resume() call, so gate the counter on mNeverResumed: only a cache that was - // previously ACTIVE and got suspended counts as a preemption recovery. - // Without this, the counter would track request admissions, not preemption. - bool const firstActivation = mNeverResumed; + // Deferred copies survive a failed decode admission and must not be repeated on retry. mNeverResumed = false; - mStatus = Status::ACTIVE; - if (!firstActivation && _shouldRecordStats()) + if (mIsDecoding) + { + try + { + _offloadSparseHistory({0, mHistoryLength / mTokensPerBlock}, mHistoryLength); + } + catch (OutOfPagesError const&) + { + return false; + } + onPageStorageChanged(); + } + rollbackActivation.cancel(); + restorePhase.cancel(); + // Only a previously admitted request counts as a preemption recovery. + if (mHasResumed && _shouldRecordStats()) { mManager->recordRequestResumed(); } + mHasResumed = true; return true; } @@ -585,6 +689,15 @@ void KvCache::suspend() TLLM_CHECK_DEBUG(_checkSanity()); TLLM_CHECK_DEBUG(!mFinishEvent.has_value()); + _deactivate(); + if (_shouldRecordStats()) + { + mManager->recordRequestSuspended(); + } +} + +void KvCache::_deactivate() +{ // Copy data from external buffers back to internal vectors (mirrors Python's suspend). for (BeamIndex bi{0}; bi < mBeamWidth; ++bi) for (LifeCycleId lcId{0}; lcId < mBasePageIndices[bi].size(); ++lcId) @@ -613,10 +726,6 @@ void KvCache::suspend() _freeScratchSlots(); } mStatus = Status::SUSPENDED; - if (_shouldRecordStats()) - { - mManager->recordRequestSuspended(); - } } void KvCache::close() @@ -1123,6 +1232,12 @@ bool KvCache::resize(std::optional capacity, std::optional historyLeng throw std::invalid_argument("History length cannot be decreased"); if (newCap < newHist) throw std::invalid_argument("History length cannot exceed capacity"); + if (mIsDecoding && newHist > mCapacity) + { + for (auto const& [lcId, lc] : mManager->lifeCycles()) + if (auto const* attn = std::get_if(&lc); attn && attn->isSparse) + throw std::invalid_argument("Sparse decode history cannot include unallocated input tokens"); + } // Scratch reuse: enforce constraint. bool enableScratch = mEnableSwaScratchReuse; @@ -1137,10 +1252,21 @@ bool KvCache::resize(std::optional capacity, std::optional historyLeng } bool const recordGenerationAllocStats = mGenerationAllocReady && newCap > mCapacity; - if (!enableScratch && _shortcutSetCapacity(newCap) && _shortcutSetHistoryLength(newHist)) + if (!enableScratch && divUp(newCap, mTokensPerBlock) == divUp(mCapacity, mTokensPerBlock)) { - _refreshGenerationAllocReady(); - return true; + try + { + if (_shortcutSetHistoryLength(newHist)) + { + mCapacity = newCap; + _refreshGenerationAllocReady(); + return true; + } + } + catch (OutOfPagesError const&) + { + return false; + } } BlockOrdinal oldNumBlocks{divUp(mCapacity, mTokensPerBlock)}; @@ -1153,17 +1279,33 @@ bool KvCache::resize(std::optional capacity, std::optional historyLeng auto ssmLcId = mManager->lifeCycles().ssmLifeCycleId(); auto backupHolders = _unlockStaleBlocks(newHist); + // Compute scratch deltas. + auto [excessScratchSlots, deltaScratchSlots, scratchRanges] = _takeExcessScratchSlots(newCap, newHist); + auto restoreLocks = FuncGuard( + [&]() + { + _recoverExcessScratchSlots(excessScratchSlots); + _lockHeldBlocks(backupHolders); + }); + if (newNumBlocks < oldNumBlocks) { TLLM_CHECK_DEBUG_WITH_INFO(!hasScratchSlots(), "Cannot shrink while scratch slots exist"); + try + { + if (mIsDecoding) + _offloadSparseHistory({mHistoryLength / mTokensPerBlock, newHist / mTokensPerBlock}, newHist); + } + catch (OutOfPagesError const&) + { + return false; + } + restoreLocks.cancel(); _subtractPendingAllocationRange(newNumBlocks, oldNumBlocks); auto scope = recordEventScope(); _decreaseCapacity(newNumBlocks); } - // Compute scratch deltas. - auto [excessScratchSlots, deltaScratchSlots, scratchRanges] = _takeExcessScratchSlots(newCap, newHist); - if (newNumBlocks >= oldNumBlocks) { // Compute new normal slots needed per lifecycle. @@ -1225,8 +1367,6 @@ bool KvCache::resize(std::optional capacity, std::optional historyLeng } catch (OutOfPagesError const&) { - _recoverExcessScratchSlots(excessScratchSlots); - _lockHeldBlocks(backupHolders); return false; } } @@ -1235,6 +1375,27 @@ bool KvCache::resize(std::optional capacity, std::optional historyLeng newSlots.resize(numLc); } + // Reserve GPU growth before demoting history. A failed allocation or copy submission must + // leave the old capacity, watermark, and page ownership usable for a retry. + auto releaseNewSlots = FuncGuard( + [&]() + { + for (LifeCycleId lc{0}; lc < numLc; ++lc) + for (auto& slot : newSlots[lc]) + mManager->storage().releaseSlot(lc, kHotLevel, std::move(slot)); + }); + try + { + if (mIsDecoding) + _offloadSparseHistory({mHistoryLength / mTokensPerBlock, newHist / mTokensPerBlock}, newHist); + } + catch (OutOfPagesError const&) + { + return false; + } + releaseNewSlots.cancel(); + restoreLocks.cancel(); + // Wait on newly allocated slots. { std::vector readyEvents; @@ -1359,7 +1520,7 @@ bool KvCache::resize(std::optional capacity, std::optional historyLeng } mCapacity = newCap; - mHistoryLength = newHist; + _publishHistoryLength(newHist); _refreshGenerationAllocReady(); TLLM_CHECK_DEBUG(_checkSanity()); return true; @@ -1379,20 +1540,8 @@ void KvCache::setCapacity(int cap) void KvCache::setHistoryLength(int hist) { - bool success = resize(std::nullopt, hist); - TLLM_CHECK(success); - (void) success; -} - -bool KvCache::_shortcutSetCapacity(int newCap) -{ - if (newCap == mCapacity) - return true; - // No shortcut if block count changes. - if (divUp(newCap, mTokensPerBlock) != divUp(mCapacity, mTokensPerBlock)) - return false; - mCapacity = newCap; - return true; + if (!resize(std::nullopt, hist)) + throw OutOfPagesError("Not enough pages to advance history"); } bool KvCache::_shortcutSetHistoryLength(int newHist) @@ -1424,7 +1573,9 @@ bool KvCache::_shortcutSetHistoryLength(int newHist) if (changed) return false; } - mHistoryLength = newHist; + if (mIsDecoding) + _offloadSparseHistory({mHistoryLength / mTokensPerBlock, newHist / mTokensPerBlock}, newHist); + _publishHistoryLength(newHist); return true; } @@ -1646,6 +1797,16 @@ void KvCache::_commitBlock(int ord, bool isLast, bool commitSsm, bool moveSsm) // Existing block: rebase — reuse existing block's committed pages. // Mirrors Python's `elif tree_block.is_full and allow_seq_rebasing and is_full` path. std::vector reuseTasks; + std::vector missingPages; + std::vector> originalPages; + std::vector originalLocks; + auto restorePages = FuncGuard( + [&]() + { + for (auto& [lc, bp] : originalPages) + sb.pages[kDefaultBeamIndex][lc] = std::move(bp); + _lockHeldBlocks(originalLocks); + }); for (LifeCycleId lc{0}; lc < numLc; ++lc) { if (ssmLcId.has_value() && lc == *ssmLcId) @@ -1664,40 +1825,7 @@ void KvCache::_commitBlock(int ord, bool isLast, bool commitSsm, bool moveSsm) bool isLocked = std::holds_alternative(bp); if (existingPage == nullptr) { - // Existing page gone — put our uncommitted page into the tree block. - if (auto* lock = std::get_if(&bp)) - { - auto up = dynamicPointerCast(lock->page()); - if (up) - { - bp = std::monostate{}; - auto committed = up->convertToCommitted(newBlock, finishEvent(), numTokens); - if (newBlock->eventSink) - { - newBlock->eventSink->addStoredLifeCycle(*newBlock, lc); - } - bp = isLocked - ? BlockPage{committed->lock(*this, kDefaultBeamIndex, static_cast(ord), lc)} - : BlockPage{committed->hold()}; - } - } - else if (auto* holder = std::get_if>(&bp)) - { - if (*holder) - { - auto up = dynamicPointerCast((*holder)->page); - if (up) - { - bp = std::monostate{}; - auto committed = up->convertToCommitted(newBlock, finishEvent(), numTokens); - if (newBlock->eventSink) - { - newBlock->eventSink->addStoredLifeCycle(*newBlock, lc); - } - bp = committed->hold(); - } - } - } + missingPages.push_back(lc); } else { @@ -1705,8 +1833,12 @@ void KvCache::_commitBlock(int ord, bool isLast, bool commitSsm, bool moveSsm) if (isLocked) { auto holder = blockPageGetPage(bp)->hold(); + originalLocks.push_back( + {BlockOrdinal{ord}, kDefaultBeamIndex, lc, holder, holder->page->cacheLevel}); bp = std::move(holder); } + originalPages.emplace_back(lc, std::move(bp)); + bp = std::monostate{}; reuseTasks.push_back({existingPage->sharedFromThis(), kDefaultBeamIndex, static_cast(ord), lc, _lockLevel(*existingPage, static_cast(ord))}); } @@ -1719,6 +1851,24 @@ void KvCache::_commitBlock(int ord, bool isLast, bool commitSsm, bool moveSsm) LifeCycleId lc = reuseTasks[ri].lifeCycle; sb.pages[kDefaultBeamIndex][lc] = std::move(locks[ri]); } + if (mIsDecoding) + _offloadSparseHistory({ord, ord + 1}, mHistoryLength); + } + restorePages.cancel(); + // Publish missing lifecycle pages only after migrations and offload can no longer fail. + for (LifeCycleId lc : missingPages) + { + auto& bp = sb.pages[kDefaultBeamIndex][lc]; + auto up = dynamicPointerCast(blockPageGetPage(bp)); + if (!up) + continue; + bool const isLocked = std::holds_alternative(bp); + bp = std::monostate{}; + auto committed = up->convertToCommitted(newBlock, finishEvent(), numTokens); + if (newBlock->eventSink) + newBlock->eventSink->addStoredLifeCycle(*newBlock, lc); + bp = isLocked ? BlockPage{committed->lock(*this, kDefaultBeamIndex, BlockOrdinal{ord}, lc)} + : BlockPage{committed->hold()}; } // Don't clear SSM storage on rebase — the existing block may have a valid snapshot. sb.treeBlock = newBlock; @@ -1800,7 +1950,8 @@ void KvCache::commit(TokenSpan tokens, bool isEnd) bool const commitMinSnapshot = mManager->commitMinSnapshot(); auto ssmLcId = mManager->lifeCycles().ssmLifeCycleId(); - int const numCommitted = static_cast(mCommittedTokens.size()) + static_cast(tokens.size()); + int const oldNumCommittedTokens = numCommittedTokens(); + int const numCommitted = oldNumCommittedTokens + static_cast(tokens.size()); if (commitMinSnapshot) { if (mHistoryLength != static_cast(mCommittedTokens.size()) && mHistoryLength != numCommitted) @@ -1828,6 +1979,14 @@ void KvCache::commit(TokenSpan tokens, bool isEnd) bool const hasNewFullBlocks = newNumFullBlocks > numCommittedBlocksBefore; if (hasNewFullBlocks || hasPartialSnapshot) { + // Keep successfully published blocks, but do not retain newly appended tokens for a block + // whose rebase failed. Otherwise close() would retry that failed rebase during teardown. + auto restoreTokens = FuncGuard( + [&]() + { + int const publishedTokens = std::min(numCommitted, mNumCommittedBlocks * mTokensPerBlock); + mCommittedTokens.resize(std::max(oldNumCommittedTokens, publishedTokens)); + }); // Block whose end is the last committed token — where the SSM snapshot lives. int const ssmSnapshotOrdinal = (numCommitted - 1) / mTokensPerBlock; // Wrapped in recordEventScope() so SharedPageLock::unlock() shares one finish @@ -1854,6 +2013,7 @@ void KvCache::commit(TokenSpan tokens, bool isEnd) _snapshotPartialBlockToTree(partialOrdinal, /*commitSsm=*/ssmLcId.has_value()); } } + restoreTokens.cancel(); } if (isEnd && mCommitState != CommitState::USER_STOP) @@ -1998,8 +2158,8 @@ std::unique_ptr KvCache::planCommittedBlockDrop() BlockOrdinal windowStart; if (auto const* attn = std::get_if(&lc)) { - // Full-attention blocks may still be needed by later turns. - if (!attn->windowSize.has_value()) + // Full-attention and sparse history may still be needed by later turns. + if (attn->isSparse || !attn->windowSize.has_value()) continue; auto const staleRange = _getStaleRange(numCommittedTokens(), lc); windowStart = std::min(staleRange.end, end); diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.h index 5f336f57d1be..9e3ad64fa050 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.h +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.h @@ -188,8 +188,19 @@ class KvCache : public std::enable_shared_from_this // Resume: check utilization and lock active pages at their required storage levels. // Optionally sets a new CUDA stream; if nullopt, uses the existing one. + // isDecoding defaults to the current phase. A cache starts in prefill and cannot return to it + // after decode admission. Set true when admitting a suspended request directly to decode. // Returns false if utilization too high or out of memory. - bool resume(std::optional stream = std::nullopt); + bool resume(std::optional stream = std::nullopt, std::optional isDecoding = std::nullopt); + + // Enter decode only after prefill has submitted its final KV accesses. Reconciles all complete + // sparse history, including on retries with an unchanged watermark. Returns false on host OOM. + bool enterDecode(); + + bool isDecoding() const noexcept + { + return mIsDecoding; + } // Suspend: detach from CUDA stream, unlock pages → PageHolder. void suspend(); @@ -349,7 +360,7 @@ class KvCache : public std::enable_shared_from_this // Plan dropping SWA blocks needed only by the next conversation turn. // // The plan covers committed pages in each SWA life cycle's current attention - // window. Full-attention and attention-sink blocks are excluded because + // window. Sparse history, full-attention, and attention-sink blocks are excluded because // later turns may still need them. An SSM life cycle contributes its final block. Must be // called after stopCommitting(). Returns nullptr without creating a plan if // any required SWA page is unavailable. Mirrors Python's @@ -483,10 +494,17 @@ class KvCache : public std::enable_shared_from_this // Internal — called by resume(). Not public (mirrors Python where activate() doesn't exist). void activate(); - // Keep cold sparse history (and immutable reuse sources) in host memory. - // Writable pages require GPU storage; GPU history stays there until explicitly offloaded. + // Release active locks and scratch slots without recording a scheduler suspension. + void _deactivate(); + + // Prefill and writable pages require GPU storage. Decode keeps cold sparse history on host. CacheLevel _lockLevel(Page const& page, BlockOrdinal ordinal) const; + // Offload GPU pages in the supplied complete-history range, validating every live owner's phase. + // The candidate watermark is visible only under the exclusive API lock until offload succeeds. + void _offloadSparseHistory(HalfOpenRange range, int historyLength); + void _publishHistoryLength(int historyLength); + // Internal helpers. // Turn the per-block cache levels observed while holding the matched pages into logical token // counts. Called at the end of _setupForReuse, which collects them in the same walk. @@ -545,7 +563,6 @@ class KvCache : public std::enable_shared_from_this std::vector _activePages() const; SharedPtr _page(BlockOrdinal ordinal, BeamIndex beamIdx, LifeCycleId lcId) const; - bool _shortcutSetCapacity(int capacity); bool _shortcutSetHistoryLength(int historyLength); bool _shouldRecordManagerStats() const; bool _shouldRecordRequestStats() const; @@ -642,6 +659,7 @@ class KvCache : public std::enable_shared_from_this BeamIndex mBeamWidth; int mCapacity; int mHistoryLength; + bool mIsDecoding = false; std::optional mExpectedPromptLength; bool mGenerationAllocReady = false; @@ -673,6 +691,7 @@ class KvCache : public std::enable_shared_from_this // SSM pages: [beamIdx][lcId] — always initialized (empty entries = monostate). BeamBlockPages mSsmBlocks; bool mNeverResumed = true; + bool mHasResumed = false; // Successful admission, independent of completed deferred copies. PendingStats mPendingStats; diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/lifeCycleRegistry.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/lifeCycleRegistry.h index 7945ef221491..c62d9b40d2c2 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/lifeCycleRegistry.h +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/lifeCycleRegistry.h @@ -46,7 +46,8 @@ struct AttnLifeCycle { int numBlocks = divUp(historyLength, tokensPerBlock); BlockOrdinal start{std::min(numBlocks, numSinkBlocks)}; - if (!windowSize.has_value()) + // Sparse selection may revisit any history block, including outside the sliding window. + if (isSparse || !windowSize.has_value()) return {start, start}; // `+ 1` is intentional: attention always runs for >= 1 in-flight input // token at position `historyLength`, so the live window is diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storageManager.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storageManager.cpp index a0d22311ba18..b67d26cf4125 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storageManager.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storageManager.cpp @@ -1269,10 +1269,10 @@ void StorageManager::offloadSparsePages(KvCache& requestingCache, std::vector> pages; - std::vector> locks; - std::vector slots; - std::vector indices; + std::vector> srcPages; + std::vector> srcPageLocks; + std::vector dstSlots; + std::vector srcDstPageIndices; }; std::map batches; @@ -1300,12 +1300,12 @@ void StorageManager::offloadSparsePages(KvCache& requestingCache, std::vectorlifeCycle)]; - batch.pages.push_back(page); + batch.srcPages.push_back(page); for (auto const& owner : lock->owners()) { ownerStreams.insert(owner.kvCache->cudaStream()); } - batch.locks.push_back(std::move(lock)); + batch.srcPageLocks.push_back(std::move(lock)); } if (batches.empty()) { @@ -1315,7 +1315,8 @@ void StorageManager::offloadSparsePages(KvCache& requestingCache, std::vector requirements(numPoolGroups(kSparseHistoryLevel), 0); for (auto const& [layerGroup, batch] : batches) { - requirements[getPoolGroupIndex(kSparseHistoryLevel, layerGroup)] += slotCountValueFromSize(batch.pages.size()); + requirements[getPoolGroupIndex(kSparseHistoryLevel, layerGroup)] + += slotCountValueFromSize(batch.srcPages.size()); } prepareFreeSlots(kSparseHistoryLevel, requirements, migrationRecorder, dropRecorder); auto releaseDestinations = FuncGuard( @@ -1323,7 +1324,7 @@ void StorageManager::offloadSparsePages(KvCache& requestingCache, std::vectorslotId())}); + batch.srcDstPageIndices.push_back({.dst = slotIdToPageIndexValue(batch.dstSlots[i].slotId()), + .src = slotIdToPageIndexValue(batch.srcPages[i]->slotId())}); } } @@ -1355,11 +1356,11 @@ void StorageManager::offloadSparsePages(KvCache& requestingCache, std::vectorreadyEvent.waitInStream(cudaStream); - batch.slots[i].readyEvent.waitInStream(cudaStream); - for (auto const& event : batch.locks[i]->finishEvents) + batch.srcPages[i]->readyEvent.waitInStream(cudaStream); + batch.dstSlots[i].readyEvent.waitInStream(cudaStream); + for (auto const& event : batch.srcPageLocks[i]->finishEvents) { event.waitInStream(cudaStream); } @@ -1374,17 +1375,17 @@ void StorageManager::offloadSparsePages(KvCache& requestingCache, std::vectorrecordOffloadEvent(completion); + batch.dstSlots[i].readyEvent = completion; + batch.srcPageLocks[i]->recordOffloadEvent(completion); } } }); for (auto const& [layerGroup, batch] : batches) { - submitMigrationBatch( - kSparseHistoryLevel, kHotLevel, layerGroup, batch.indices.data(), batch.indices.size(), stream); + submitMigrationBatch(kSparseHistoryLevel, kHotLevel, layerGroup, batch.srcDstPageIndices.data(), + batch.srcDstPageIndices.size(), stream); } fenceCopies.run(); @@ -1397,22 +1398,22 @@ void StorageManager::offloadSparsePages(KvCache& requestingCache, std::vectormoveToSparseHistory(std::move(batch.slots[i])); - releaseSlot(batch.pages[i]->lifeCycle, kHotLevel, std::move(source)); + Slot source = batch.srcPageLocks[i]->moveToSparseHistory(std::move(batch.dstSlots[i])); + releaseSlot(batch.srcPages[i]->lifeCycle, kHotLevel, std::move(source)); } } if (mEventSink) { for (auto const& [layerGroup, batch] : batches) { - for (auto const& page : batch.pages) + for (auto const& page : batch.srcPages) { if (page->isCommitted()) { diff --git a/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp b/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp index fe4f4ac0f850..7fbd59ba4a3d 100644 --- a/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp +++ b/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp @@ -1705,15 +1705,17 @@ void KvCacheManagerV2Bindings::initBindings(nb::module_& m) nb::class_(m, "_KVCache") .def( "resume", - [](kv::KvCache& self, nb::object stream) + [](kv::KvCache& self, nb::object stream, std::optional isDecoding) { std::optional optStream; if (!stream.is_none()) optStream = reinterpret_cast(nb::cast(stream)); nb::gil_scoped_release rel; - return self.resume(optStream); + return self.resume(optStream, isDecoding); }, - nb::arg("cuda_stream") = nb::none()) + nb::arg("cuda_stream") = nb::none(), nb::arg("is_decoding") = nb::none()) + .def("enter_decode", &kv::KvCache::enterDecode, nb::call_guard()) + .def_prop_ro("is_decoding", &kv::KvCache::isDecoding) .def("suspend", &kv::KvCache::suspend, nb::call_guard()) .def( "prefetch", [](kv::KvCache& self, int target) { return self.prefetch(kv::CacheLevel{target}); }, diff --git a/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp b/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp index c80c897b11d3..9f3ad9601d28 100644 --- a/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp +++ b/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp @@ -976,6 +976,8 @@ TEST_F(KvCacheManagerV2SparseOffloadTest, SharedOwnersPublishHostIndicesAndKeepH EXPECT_EQ(storage.getStatistics(kHotLevel).free, storage.getStatistics(kHotLevel).total); EXPECT_FALSE(storage.isEvictable(*page)); EXPECT_THROW(storage.batchedMigrate(kHotLevel, {page}, {}), LogicError); + ASSERT_TRUE(first->enterDecode()); + ASSERT_TRUE(second->enterDecode()); first->suspend(); EXPECT_EQ(page->status(), PageStatus::LOCKED); ASSERT_TRUE(first->resume()); @@ -1263,6 +1265,7 @@ TEST_F(KvCacheManagerV2SparseOffloadTest, DemotedUncommittedHistoryCanCommitAndB ASSERT_TRUE(cache->resume(stream())); ASSERT_TRUE(cache->resize(4, 4)); cache->offloadSparsePages({pageAt(*cache)}); + ASSERT_TRUE(cache->enterDecode()); auto const hostSlot = pageAt(*cache)->slotId(); cache->commit(tokens()); auto committed = pageAt(*cache); @@ -1271,13 +1274,484 @@ TEST_F(KvCacheManagerV2SparseOffloadTest, DemotedUncommittedHistoryCanCommitAndB EXPECT_EQ(committed->slotId(), hostSlot); auto reused = manager->createKvCache({}, tokens()); auto closeReused = FuncGuard([&]() { reused->close(); }); - ASSERT_TRUE(reused->resume(stream())); + ASSERT_TRUE(reused->resume(stream(), true)); EXPECT_EQ(pageAt(*reused), committed); cache->close(); EXPECT_EQ(committed->status(), PageStatus::LOCKED); EXPECT_EQ(reused->getBasePageIndices(LifeCycleId{0})[0], slotIdToPageIndexValue(hostSlot)); } +class KvCacheManagerV2DecodeOffloadTest : public KvCacheManagerV2PageLockTest +{ +}; + +TEST_F(KvCacheManagerV2DecodeOffloadTest, PrefillRetainsSparseSwaHistoryAndRestoresGpuStorage) +{ + auto config = sparseConfig(); + config.swaScratchReuse = SwaScratchReuseConfig{}; + auto& layer = std::get(config.layers.front()); + layer.slidingWindowSize = 4; + layer.buffers.front().size = 4096; + auto manager = std::make_shared(std::move(config)); + auto const apiLock = manager->lockExclusive(); + auto& storage = manager->storage(); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + cache->setEnableSwaScratchReuse(true); + ASSERT_TRUE(cache->resize(12, 0)); + EXPECT_FALSE(cache->hasScratchSlots()); + ASSERT_TRUE(cache->resize(12, 8)); + EXPECT_FALSE(cache->isDecoding()); + EXPECT_EQ(cache->pageStorageVersion(), 0); + for (int ord = 0; ord < 3; ++ord) + { + ASSERT_NE(pageAt(*cache, ord), nullptr); + EXPECT_EQ(pageAt(*cache, ord)->cacheLevel, kHotLevel); + } + + for (bool prefetch : {false, true}) + { + cache->suspend(); + storage.forceEvict(kHotLevel, TypedVec{3}); + for (int ord = 0; ord < 3; ++ord) + EXPECT_EQ(pageAt(*cache, ord)->cacheLevel, kSparseHistoryLevel); + if (prefetch) + ASSERT_TRUE(cache->prefetch(kHotLevel)); + ASSERT_TRUE(cache->resume()); + for (int ord = 0; ord < 3; ++ord) + EXPECT_EQ(pageAt(*cache, ord)->cacheLevel, kHotLevel); + } + ASSERT_TRUE(cache->resize(12, 12)); + EXPECT_EQ(cache->pageStorageVersion(), 0); + ASSERT_TRUE(cache->enterDecode()); + for (int ord = 0; ord < 3; ++ord) + { + EXPECT_EQ(pageAt(*cache, ord)->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(pageAt(*cache, ord)->status(), PageStatus::LOCKED); + } +} + +TEST_F(KvCacheManagerV2DecodeOffloadTest, EntryScansUnchangedWatermarkAndResumeKeepsHostHistory) +{ + auto codec = std::make_unique(); + auto* observer = codec.get(); + auto manager = std::make_shared(sparseConfig(), nullptr, std::move(codec)); + auto const apiLock = manager->lockExclusive(); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(8, 8)); + EXPECT_EQ(observer->encodeCalls, 0); + ASSERT_TRUE(cache->enterDecode()); + EXPECT_TRUE(cache->isDecoding()); + EXPECT_EQ(cache->historyLength(), 8); + EXPECT_EQ(observer->encodeCalls, 1); + EXPECT_EQ(observer->encodedPages, 2); + EXPECT_EQ(cache->pageStorageVersion(), 3); + ASSERT_TRUE(cache->enterDecode()); + EXPECT_EQ(observer->encodeCalls, 1); + EXPECT_EQ(cache->pageStorageVersion(), 3); + EXPECT_THROW(cache->resize(8, 4), std::invalid_argument); + cache->suspend(); + EXPECT_THROW(cache->resume(std::nullopt, false), std::invalid_argument); + ASSERT_TRUE(cache->resume()); + EXPECT_EQ(observer->encodeCalls, 1); + EXPECT_EQ(cache->pageStorageVersion(), 4); + for (int ord = 0; ord < 2; ++ord) + { + auto const page = pageAt(*cache, ord); + EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(cache->getBasePageIndices(LifeCycleId{0})[ord], slotIdToPageIndexValue(page->slotId())); + } +} + +TEST_F(KvCacheManagerV2DecodeOffloadTest, SparseHistoryIsNotMarkedForTurnEndDrop) +{ + auto config = sparseConfig(); + std::get(config.layers.front()).slidingWindowSize = 4; + auto manager = std::make_shared(std::move(config)); + auto const apiLock = manager->lockExclusive(); + auto page = seedPrefix(*manager, kHotLevel); + auto cache = manager->createKvCache({}, tokens()); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream(), true)); + cache->stopCommitting(); + auto dropPlan = cache->planCommittedBlockDrop(); + EXPECT_EQ(page->plannedDropCount, 0); +} + +TEST_F(KvCacheManagerV2DecodeOffloadTest, CachedPrefillRestoresGpuOnResumePrefetchAndRebase) +{ + for (int path : {0, 1, 2}) + { + SCOPED_TRACE(path); + auto manager = std::make_shared(sparseConfig()); + auto const apiLock = manager->lockExclusive(); + auto page = seedPrefix(*manager, kSparseHistoryLevel); + auto cache = path == 2 ? manager->createKvCache() : manager->createKvCache({}, tokens()); + auto closeCache = FuncGuard([&]() { cache->close(); }); + if (path == 1) + { + ASSERT_TRUE(cache->prefetch(kHotLevel)); + EXPECT_EQ(page->cacheLevel, kHotLevel); + } + ASSERT_TRUE(cache->resume(stream())); + if (path == 2) + { + ASSERT_TRUE(cache->resize(4, 4)); + cache->commit(tokens()); + } + EXPECT_EQ(pageAt(*cache), page); + EXPECT_EQ(page->cacheLevel, kHotLevel); + EXPECT_EQ(cache->historyLength(), 4); + EXPECT_FALSE(cache->isDecoding()); + EXPECT_EQ(cache->pageStorageVersion(), 0); + EXPECT_EQ(cache->getBasePageIndices(LifeCycleId{0})[0], slotIdToPageIndexValue(page->slotId())); + } +} + +TEST_F(KvCacheManagerV2DecodeOffloadTest, MixedDenseSwaLocksAreRestoredAfterHostOom) +{ + auto config = sparseConfig(); + config.cacheTiers[0] = GpuCacheTierConfig{8 << 20}; + auto& sparse = std::get(config.layers.front()); + sparse.slidingWindowSize = 4; + auto dense = sparse; + dense.layerId = 1; + dense.buffers.front().isSparse = false; + config.layers.emplace_back(std::move(dense)); + auto manager = std::make_shared(std::move(config)); + auto const apiLock = manager->lockExclusive(); + auto& storage = manager->storage(); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(8, 4)); + auto const densePage = pageAt(*cache, 0, LifeCycleId{1}); + auto const denseSlot = densePage->slotId(); + ASSERT_TRUE(cache->enterDecode()); + EXPECT_EQ(densePage->cacheLevel, kHotLevel); + EXPECT_EQ(pageAt(*cache)->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(pageAt(*cache, 1)->cacheLevel, kHotLevel); + auto const version = cache->pageStorageVersion(); + auto const freeHost = manager->getStorageStatistics(kSparseHistoryLevel) + .at(storage.getPoolGroupIndex(kSparseHistoryLevel, LifeCycleId{0})) + .free; + auto blockers = storage.newSlots(kSparseHistoryLevel, TypedVec{freeHost, 0}); + auto releaseBlockers = FuncGuard( + [&]() + { + for (auto& slot : blockers[LifeCycleId{0}]) + storage.releaseSlot(LifeCycleId{0}, kSparseHistoryLevel, std::move(slot)); + }); + EXPECT_FALSE(cache->resize(8, 8)); + EXPECT_EQ(cache->historyLength(), 4); + EXPECT_EQ(cache->pageStorageVersion(), version); + EXPECT_EQ(pageAt(*cache, 0, LifeCycleId{1}), densePage); + EXPECT_EQ(densePage->status(), PageStatus::LOCKED); + EXPECT_EQ(densePage->cacheLevel, kHotLevel); + EXPECT_EQ(densePage->slotId(), denseSlot); + EXPECT_EQ(cache->getBasePageIndices(LifeCycleId{1})[0], slotIdToPageIndexValue(denseSlot)); + EXPECT_EQ(pageAt(*cache, 1)->cacheLevel, kHotLevel); +} + +TEST_F(KvCacheManagerV2DecodeOffloadTest, HistoryUpdatesOffloadOnlyNewFullPages) +{ + auto codec = std::make_unique(); + auto* observer = codec.get(); + auto manager = std::make_shared(sparseConfig(), nullptr, std::move(codec)); + auto const apiLock = manager->lockExclusive(); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(8, 4)); + ASSERT_TRUE(cache->enterDecode()); + auto const version = cache->pageStorageVersion(); + auto const firstHostSlot = pageAt(*cache)->slotId(); + ASSERT_TRUE(cache->resize(8, 7)); + EXPECT_EQ(observer->encodedPages, 1); + EXPECT_EQ(cache->pageStorageVersion(), version); + EXPECT_EQ(pageAt(*cache, 1)->cacheLevel, kHotLevel); + ASSERT_TRUE(cache->resize(8, 8)); + EXPECT_EQ(observer->encodeCalls, 2); + EXPECT_EQ(observer->encodedPages, 2); + EXPECT_EQ(cache->pageStorageVersion(), version + 2); + EXPECT_EQ(pageAt(*cache)->slotId(), firstHostSlot); + EXPECT_EQ(pageAt(*cache, 1)->cacheLevel, kSparseHistoryLevel); + ASSERT_TRUE(cache->resize(12)); + ASSERT_TRUE(cache->resize(12, 9)); + EXPECT_EQ(pageAt(*cache, 2)->cacheLevel, kHotLevel); + EXPECT_EQ(observer->encodedPages, 2); +} + +TEST_F(KvCacheManagerV2DecodeOffloadTest, SharedPrefillOwnerBlocksDemotionAndResumeReconcilesOlderGpuHistory) +{ + auto manager = std::make_shared(sparseConfig()); + auto const apiLock = manager->lockExclusive(); + auto page = seedPrefix(*manager, kHotLevel); + auto first = manager->createKvCache({}, tokens()); + auto second = manager->createKvCache({}, tokens()); + auto closeCaches = FuncGuard( + [&]() + { + first->close(); + second->close(); + }); + ASSERT_TRUE(first->resume(stream())); + ASSERT_TRUE(second->resume(stream())); + EXPECT_THROW(first->enterDecode(), LogicError); + EXPECT_FALSE(first->isDecoding()); + EXPECT_EQ(first->pageStorageVersion(), 0); + EXPECT_EQ(page->cacheLevel, kHotLevel); + second->suspend(); + ASSERT_TRUE(first->enterDecode()); + EXPECT_THROW(second->resume(), LogicError); + EXPECT_FALSE(second->isActive()); + EXPECT_FALSE(second->isDecoding()); + first->suspend(); + ASSERT_TRUE(second->resume()); + EXPECT_EQ(page->cacheLevel, kHotLevel); + EXPECT_THROW(first->resume(), LogicError); + EXPECT_FALSE(first->isActive()); + EXPECT_TRUE(first->isDecoding()); + second->close(); + ASSERT_TRUE(first->resume()); + EXPECT_EQ(first->historyLength(), 4); + EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); +} + +TEST_F(KvCacheManagerV2DecodeOffloadTest, HostOomDoesNotAdmitDecodeAndCanRetry) +{ + for (int admission : {0, 1, 2}) + { + SCOPED_TRACE(admission); + bool const suspended = admission != 0; + bool const firstResume = admission == 2; + int const history = firstResume ? 4 : 8; + auto manager = std::make_shared(sparseConfig()); + auto const apiLock = manager->lockExclusive(); + auto& storage = manager->storage(); + auto prefix = firstResume ? seedPrefix(*manager, kHotLevel) : nullptr; + auto cache = firstResume ? manager->createKvCache({}, tokens()) : manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + if (!firstResume) + { + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(history, history)); + if (suspended) + cache->suspend(); + } + manager->getAndResetIterationSuspendResumeStats(); + auto blockers = storage.newSlots(kSparseHistoryLevel, TypedVec{firstResume ? 2 : 1}); + auto releaseBlockers = FuncGuard( + [&]() + { + for (auto& slot : blockers[LifeCycleId{0}]) + storage.releaseSlot(LifeCycleId{0}, kSparseHistoryLevel, std::move(slot)); + }); + EXPECT_FALSE(suspended ? cache->resume(stream(), true) : cache->enterDecode()); + EXPECT_EQ(cache->isActive(), !suspended); + EXPECT_FALSE(cache->isDecoding()); + EXPECT_EQ(cache->historyLength(), history); + EXPECT_EQ(cache->capacity(), history); + EXPECT_EQ(cache->pageStorageVersion(), 0); + EXPECT_EQ(manager->getAndResetIterationSuspendResumeStats(), (std::pair{0, 0})); + for (int ord = 0; ord < history / 4; ++ord) + { + auto const page = pageAt(*cache, ord); + EXPECT_EQ(page->cacheLevel, kHotLevel); + EXPECT_EQ(cache->getBasePageIndices(LifeCycleId{0})[ord], + suspended ? kBadPageIndex.value() : slotIdToPageIndexValue(page->slotId())); + } + releaseBlockers.run(); + ASSERT_TRUE(suspended ? cache->resume(std::nullopt, true) : cache->enterDecode()); + EXPECT_TRUE(cache->isDecoding()); + for (int ord = 0; ord < history / 4; ++ord) + EXPECT_EQ(pageAt(*cache, ord)->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(manager->getAndResetIterationSuspendResumeStats().second, admission == 1 ? 1 : 0); + } +} + +TEST_F(KvCacheManagerV2DecodeOffloadTest, HistoryOomRollsBackShortcutGrowthAndShrink) +{ + for (auto const [oldCapacity, newCapacity] : {std::pair{8, 8}, std::pair{8, 12}, std::pair{12, 8}}) + { + SCOPED_TRACE(newCapacity); + SCOPED_TRACE(oldCapacity); + auto manager = std::make_shared(sparseConfig()); + auto const apiLock = manager->lockExclusive(); + auto& storage = manager->storage(); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(8, 4)); + ASSERT_TRUE(cache->enterDecode()); + ASSERT_TRUE(cache->resize(oldCapacity)); + auto const page = pageAt(*cache, 1); + auto const slotId = page->slotId(); + auto const version = cache->pageStorageVersion(); + auto const gpuFree = storage.getStatistics(kHotLevel).free; + auto blockers = storage.newSlots(kSparseHistoryLevel, TypedVec{1}); + auto releaseBlockers = FuncGuard( + [&]() + { + for (auto& slot : blockers[LifeCycleId{0}]) + storage.releaseSlot(LifeCycleId{0}, kSparseHistoryLevel, std::move(slot)); + }); + EXPECT_FALSE(cache->resize(newCapacity, 8)); + EXPECT_EQ(cache->capacity(), oldCapacity); + EXPECT_EQ(cache->historyLength(), 4); + EXPECT_TRUE(cache->isDecoding()); + EXPECT_EQ(cache->pageStorageVersion(), version); + EXPECT_EQ(pageAt(*cache, 1), page); + EXPECT_EQ(page->slotId(), slotId); + EXPECT_EQ(page->cacheLevel, kHotLevel); + EXPECT_EQ(cache->getBasePageIndices(LifeCycleId{0})[1], slotIdToPageIndexValue(slotId)); + EXPECT_EQ(storage.getStatistics(kHotLevel).free, gpuFree); + if (oldCapacity == 12) + EXPECT_EQ(pageAt(*cache, 2)->status(), PageStatus::LOCKED); + storage.releaseSlot(LifeCycleId{0}, kSparseHistoryLevel, std::move(blockers[LifeCycleId{0}].back())); + blockers[LifeCycleId{0}].clear(); + ASSERT_TRUE(cache->resize(newCapacity, 8)); + EXPECT_EQ(cache->capacity(), newCapacity); + EXPECT_EQ(cache->historyLength(), 8); + EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(cache->pageStorageVersion(), version + 2); + } +} + +TEST_F(KvCacheManagerV2DecodeOffloadTest, CodecRejectionDoesNotPublishEntryOrHistoryUpdate) +{ + for (bool entry : {false, true}) + { + SCOPED_TRACE(entry); + auto codec = std::make_unique(); + auto* observer = codec.get(); + auto manager = std::make_shared(sparseConfig(), nullptr, std::move(codec)); + auto const apiLock = manager->lockExclusive(); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + int const history = entry ? 8 : 4; + ASSERT_TRUE(cache->resize(8, history)); + if (!entry) + ASSERT_TRUE(cache->enterDecode()); + auto const version = cache->pageStorageVersion(); + auto const page = pageAt(*cache, 1); + auto const gpuSlot = page->slotId(); + observer->rejectEncodeCall = observer->encodeCalls + 1; + EXPECT_THROW(entry ? cache->enterDecode() : cache->resize(12, 8), TllmException); + EXPECT_EQ(cache->isDecoding(), !entry); + EXPECT_EQ(cache->capacity(), 8); + EXPECT_EQ(cache->historyLength(), history); + EXPECT_EQ(cache->pageStorageVersion(), version); + EXPECT_EQ(page->cacheLevel, kHotLevel); + EXPECT_EQ(page->slotId(), gpuSlot); + EXPECT_EQ(cache->getBasePageIndices(LifeCycleId{0})[1], slotIdToPageIndexValue(gpuSlot)); + observer->rejectEncodeCall = 0; + ASSERT_TRUE(entry ? cache->enterDecode() : cache->resize(12, 8)); + EXPECT_TRUE(cache->isDecoding()); + EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); + } +} + +TEST_F(KvCacheManagerV2DecodeOffloadTest, RebaseOffloadsOlderGpuHistoryAtUnchangedWatermark) +{ + auto manager = std::make_shared(sparseConfig()); + auto const apiLock = manager->lockExclusive(); + auto page = seedPrefix(*manager, kHotLevel); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(4, 4)); + ASSERT_TRUE(cache->enterDecode()); + auto const version = cache->pageStorageVersion(); + cache->commit(tokens()); + EXPECT_EQ(pageAt(*cache), page); + EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(cache->historyLength(), 4); + EXPECT_EQ(cache->pageStorageVersion(), version + 1); +} + +TEST_F(KvCacheManagerV2DecodeOffloadTest, PrefillRebaseCannotAdoptSharedHostIndices) +{ + auto manager = std::make_shared(sparseConfig()); + auto const apiLock = manager->lockExclusive(); + auto page = seedPrefix(*manager, kSparseHistoryLevel); + auto decoder = manager->createKvCache({}, tokens()); + auto prefill = manager->createKvCache(); + auto closeCaches = FuncGuard( + [&]() + { + decoder->close(); + prefill->close(); + }); + ASSERT_TRUE(decoder->resume(stream(), true)); + ASSERT_TRUE(prefill->resume(stream())); + ASSERT_TRUE(prefill->resize(4, 4)); + auto const privatePage = pageAt(*prefill); + EXPECT_THROW(prefill->commit(tokens()), LogicError); + EXPECT_FALSE(prefill->isDecoding()); + EXPECT_EQ(prefill->numCommittedTokens(), 0); + EXPECT_EQ(pageAt(*prefill), privatePage); + EXPECT_EQ(privatePage->cacheLevel, kHotLevel); + EXPECT_EQ(privatePage->status(), PageStatus::LOCKED); + EXPECT_EQ(prefill->getBasePageIndices(LifeCycleId{0})[0], slotIdToPageIndexValue(privatePage->slotId())); + EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(decoder->getBasePageIndices(LifeCycleId{0})[0], slotIdToPageIndexValue(page->slotId())); + EXPECT_NO_THROW(prefill->close()); +} + +TEST_F(KvCacheManagerV2DecodeOffloadTest, RebaseOomPreservesMissingLifecyclePagesAndCanRetryCommit) +{ + auto config = makeSplitColdGroupingConfig(); + for (auto& layer : config.layers) + std::get(layer).buffers.front().isSparse = true; + auto manager = std::make_shared(std::move(config)); + auto const apiLock = manager->lockExclusive(); + auto& storage = manager->storage(); + auto existing = seedPrefix(*manager, kHotLevel, LifeCycleId{1}); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(4, 4)); + ASSERT_TRUE(cache->enterDecode()); + auto first = pageAt(*cache, 0, LifeCycleId{0}); + auto second = pageAt(*cache, 0, LifeCycleId{1}); + auto const version = cache->pageStorageVersion(); + auto const pool = storage.getPoolGroupIndex(kSparseHistoryLevel, LifeCycleId{1}); + auto const freeHost = manager->getStorageStatistics(kSparseHistoryLevel).at(pool).free; + auto blockers = storage.newSlots(kSparseHistoryLevel, TypedVec{0, freeHost}); + auto releaseBlockers = FuncGuard( + [&]() + { + for (auto& slot : blockers[LifeCycleId{1}]) + storage.releaseSlot(LifeCycleId{1}, kSparseHistoryLevel, std::move(slot)); + }); + EXPECT_THROW(cache->commit(tokens()), OutOfPagesError); + EXPECT_EQ(cache->numCommittedTokens(), 0); + EXPECT_EQ(cache->numCommittedBlocks(), 0); + EXPECT_EQ(cache->historyLength(), 4); + EXPECT_EQ(cache->pageStorageVersion(), version); + EXPECT_EQ(pageAt(*cache, 0, LifeCycleId{0}), first); + EXPECT_EQ(pageAt(*cache, 0, LifeCycleId{1}), second); + for (auto const& page : {first, second}) + { + EXPECT_FALSE(page->isCommitted()); + EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(page->status(), PageStatus::LOCKED); + EXPECT_EQ(cache->getBasePageIndices(page->lifeCycle)[0], slotIdToPageIndexValue(page->slotId())); + } + EXPECT_EQ(existing->cacheLevel, kHotLevel); + releaseBlockers.run(); + cache->commit(tokens()); + EXPECT_EQ(cache->numCommittedTokens(), 4); + EXPECT_TRUE(pageAt(*cache, 0, LifeCycleId{0})->isCommitted()); + EXPECT_EQ(pageAt(*cache, 0, LifeCycleId{1}), existing); + EXPECT_EQ(existing->cacheLevel, kSparseHistoryLevel); +} + TEST_F(KvCacheManagerV2PageLockTest, SparseHostPrefixStaysPinnedAcrossReuseAndResume) { auto manager = std::make_shared(sparseConfig()); @@ -1287,8 +1761,8 @@ TEST_F(KvCacheManagerV2PageLockTest, SparseHostPrefixStaysPinnedAcrossReuseAndRe SlotId const hostSlot = page->slotId(); auto cache = manager->createKvCache({}, tokens()); auto closeCache = FuncGuard([&]() { cache->close(); }); - ASSERT_TRUE(cache->prefetch(kHotLevel)); - ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->prefetch(kSparseHistoryLevel)); + ASSERT_TRUE(cache->resume(stream(), true)); EXPECT_EQ(pageAt(*cache), page); EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); EXPECT_EQ(page->slotId(), hostSlot); @@ -1299,13 +1773,14 @@ TEST_F(KvCacheManagerV2PageLockTest, SparseHostPrefixStaysPinnedAcrossReuseAndRe auto second = manager->createKvCache({}, tokens()); auto closeSecond = FuncGuard([&]() { second->close(); }); - ASSERT_TRUE(second->prefetch(kHotLevel)); - ASSERT_TRUE(second->resume(stream())); + ASSERT_TRUE(second->prefetch(kSparseHistoryLevel)); + ASSERT_TRUE(second->resume(stream(), true)); EXPECT_EQ(pageAt(*second), page); cache->suspend(); EXPECT_EQ(page->status(), PageStatus::LOCKED); second->close(); EXPECT_EQ(page->status(), PageStatus::HELD); + ASSERT_TRUE(cache->prefetch(kHotLevel)); ASSERT_TRUE(cache->resume()); EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); EXPECT_EQ(cache->getBasePageIndices(LifeCycleId{0})[0], slotIdToPageIndexValue(hostSlot)); @@ -1380,7 +1855,7 @@ TEST_F(KvCacheManagerV2PageLockTest, PartialReuseCopiesSharedHostPrefixToPrivate auto full = manager->createKvCache({}, tokens()); auto closeFull = FuncGuard([&]() { full->close(); }); - ASSERT_TRUE(full->resume(stream())); + ASSERT_TRUE(full->resume(stream(), true)); auto partial = manager->createKvCache({}, tokens(2)); auto closePartial = FuncGuard([&]() { partial->close(); }); ASSERT_EQ(partial->historyLength(), 2); @@ -1416,7 +1891,7 @@ TEST_F(KvCacheManagerV2PageLockTest, DiskSparsePrefixRestoresToHostBeforeLocking EXPECT_THROW(makeShared(page->hold()), LogicError); auto cache = manager->createKvCache({}, tokens()); auto closeCache = FuncGuard([&]() { cache->close(); }); - ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resume(stream(), true)); EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); EXPECT_EQ(page->status(), PageStatus::LOCKED); auto const gpuStats = manager->getStorageStatistics(kHotLevel).at(PoolGroupIndex{0}); @@ -1442,7 +1917,7 @@ TEST_F(KvCacheManagerV2PageLockTest, WritableSparsePageRestoresToGpuButFullHisto TypedVec evictOne(storage.numPoolGroups(kHotLevel), 1); storage.forceEvict(kHotLevel, evictOne); ASSERT_EQ(page->cacheLevel, kSparseHistoryLevel); - ASSERT_TRUE(cache->resume()); + ASSERT_TRUE(cache->resume(std::nullopt, true)); EXPECT_EQ(page->cacheLevel, historyLength == 0 ? kHotLevel : kSparseHistoryLevel); EXPECT_EQ(page->status(), PageStatus::LOCKED); } @@ -1457,13 +1932,13 @@ TEST_F(KvCacheManagerV2PageLockTest, ResizeOomRestoresOriginalHostLock) auto page = seedPrefix(*manager, kSparseHistoryLevel); auto cache = manager->createKvCache({}, tokens()); auto closeCache = FuncGuard([&]() { cache->close(); }); - ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resume(stream(), true)); ASSERT_TRUE(cache->resize(8, 4)); SlotId const hostSlot = page->slotId(); auto const gpuFree = manager->getStorageStatistics(kHotLevel).at(PoolGroupIndex{0}).free; - // Advancing history unlocks block 0 before the request for three GPU pages - // exceeds the two-slot pool. Rollback must restore the original host lock. + // Sparse SWA history stays pinned. A failed request for three GPU pages + // must preserve the host lock and the previous eligible-history count. EXPECT_FALSE(cache->resize(20, 8)); EXPECT_EQ(cache->capacity(), 8); EXPECT_EQ(cache->historyLength(), 4); @@ -1488,7 +1963,7 @@ TEST_F(KvCacheManagerV2PageLockTest, MixedSparseAndDensePrefixUsesSeparateLockLe auto densePage = seedPrefix(*manager, kSparseHistoryLevel, LifeCycleId{1}); auto cache = manager->createKvCache({}, tokens()); auto closeCache = FuncGuard([&]() { cache->close(); }); - ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resume(stream(), true)); EXPECT_EQ(pageAt(*cache, 0, LifeCycleId{0}), sparsePage); EXPECT_EQ(pageAt(*cache, 0, LifeCycleId{1}), densePage); EXPECT_EQ(sparsePage->cacheLevel, kSparseHistoryLevel); @@ -1504,12 +1979,13 @@ TEST_F(KvCacheManagerV2PageLockTest, CommitRebasesOntoSharedHostPrefix) auto page = seedPrefix(*manager, kSparseHistoryLevel); auto first = manager->createKvCache({}, tokens()); auto closeFirst = FuncGuard([&]() { first->close(); }); - ASSERT_TRUE(first->resume(stream())); + ASSERT_TRUE(first->resume(stream(), true)); auto second = manager->createKvCache(); auto closeSecond = FuncGuard([&]() { second->close(); }); ASSERT_TRUE(second->resume(stream())); ASSERT_TRUE(second->resize(4, 4)); ASSERT_EQ(pageAt(*second)->cacheLevel, kHotLevel); + ASSERT_TRUE(second->enterDecode()); second->commit(tokens()); EXPECT_EQ(second->numCommittedBlocks(), 1); EXPECT_EQ(pageAt(*second), page); @@ -1526,7 +2002,7 @@ TEST_F(KvCacheManagerV2PageLockTest, ScratchSlotReturnsToGpuPool) auto page = seedPrefix(*manager, kSparseHistoryLevel); auto cache = manager->createKvCache({}, tokens()); auto closeCache = FuncGuard([&]() { cache->close(); }); - ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resume(stream(), true)); auto const gpuFree = manager->getStorageStatistics(kHotLevel).at(PoolGroupIndex{0}).free; auto const hostFree = manager->getStorageStatistics(kSparseHistoryLevel).at(PoolGroupIndex{0}).free; auto slots = manager->storage().newGpuSlots(TypedVec(LifeCycleId{1}, 1)); 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 e4e619c367af..b3ea320cfe11 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 @@ -3347,9 +3347,11 @@ def try_allocate_generation(self, req: LlmRequest) -> bool: return False if not kv_cache.is_active: - if not kv_cache.resume(self._stream.cuda_stream): + if not kv_cache.resume(self._stream.cuda_stream, is_decoding=True): return False self._restore_page_index_bufs(req.py_request_id, kv_cache) + elif not kv_cache.enter_decode(): + return False request_id = req.py_request_id draft_slots = self._generation_draft_slots(req) diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi index 7c9224ca6e4a..8b472c44afec 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi @@ -459,7 +459,12 @@ class _KVCache: def plan_committed_block_drop(self) -> PlannedDropHandle | None: ... def stop_committing(self) -> None: ... def suspend(self) -> None: ... - def resume(self, cuda_stream: CudaStream | None = None) -> bool: ... + def resume( + self, cuda_stream: CudaStream | None = None, is_decoding: bool | None = None + ) -> bool: ... + def enter_decode(self) -> bool: ... + @property + def is_decoding(self) -> bool: ... def prefetch(self, target: CacheLevel) -> bool: ... def get_scratch_desc(self, layer_group_id: LayerGroupId) -> ScratchDesc | None: ... @property diff --git a/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_scheduler.py b/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_scheduler.py index cc663b218a4d..f1fc920a35b7 100644 --- a/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_scheduler.py +++ b/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_scheduler.py @@ -22,13 +22,44 @@ import pytest -from tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2 import BlockReusePolicy +from tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2 import ( + BlockReusePolicy, + KVCacheManagerV2, +) from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequestState from tensorrt_llm.llmapi.llm_args import CapacitySchedulerPolicy, ContextChunkingPolicy pytestmark = pytest.mark.cpu_only +@pytest.mark.parametrize("active", [False, True]) +@pytest.mark.parametrize("admitted", [False, True]) +def test_generation_admits_decode_before_capacity_growth(active: bool, admitted: bool) -> None: + manager = object.__new__(KVCacheManagerV2) + cache = Mock(is_active=active, capacity=8) + cache.enter_decode.return_value = admitted + cache.resume.return_value = admitted + cache.resize.return_value = True + manager.kv_cache_map = {1: cache} + manager._stream = Mock(cuda_stream=123) + manager._restore_page_index_bufs = Mock() + manager._generation_draft_slots = Mock(return_value=0) + manager._allocated_draft_lens = {} + manager._has_cp_helix = False + manager._fill_fresh_kv_pages = Mock() + manager._log_window_crossing = Mock() + req = Mock(py_request_id=1) + + assert manager.try_allocate_generation(req) == admitted + admission = call.enter_decode() if active else call.resume(123, is_decoding=True) + assert cache.mock_calls == [admission] + ([call.resize(9)] if admitted else []) + if not active and admitted: + manager._restore_page_index_bufs.assert_called_once_with(1, cache) + else: + manager._restore_page_index_bufs.assert_not_called() + assert manager._allocated_draft_lens == ({1: 0} if admitted else {}) + + # --------------------------------------------------------------------------- # State value constants # --------------------------------------------------------------------------- From 2df6cb903e5c2c2b6ed5af72478a2a3a3d58c56c Mon Sep 17 00:00:00 2001 From: Guiju Zhang <7135567+cascade812@users.noreply.github.com> Date: Thu, 1 Oct 2026 21:24:41 -0700 Subject: [PATCH 3/5] CPU metadata and GPU publication Signed-off-by: Guiju Zhang <7135567+cascade812@users.noreply.github.com> --- .../kv_cache_manager_v2/AGENTS.md | 18 + .../kv_cache_manager_v2/CMakeLists.txt | 1 + .../kv_cache_manager_v2/batch.cpp | 424 ++++++++++ .../batch_manager/kv_cache_manager_v2/batch.h | 135 +++ .../kv_cache_manager_v2/kvCache.cpp | 141 +++- .../kv_cache_manager_v2/kvCache.h | 87 +- .../kv_cache_manager_v2/kvCacheManager.cpp | 7 + .../kv_cache_manager_v2/kvCacheManager.h | 3 + .../kv_cache_manager_v2/page.cpp | 3 + .../batch_manager/kvCacheManagerV2.cpp | 131 +++ .../kvCacheManagerV2ColdPageTest.cpp | 775 +++++++++++++++++- .../kv_cache/kv_cache_manager_v2.py | 45 + .../runtime/kv_cache_manager_v2/__init__.py | 8 + .../runtime/kv_cache_manager_v2/__init__.pyi | 98 +++ .../kv_cache/test_kv_cache_v2_scheduler.py | 122 +++ .../test_kv_cache_manager_v2.py | 144 ++++ 16 files changed, 2107 insertions(+), 35 deletions(-) create mode 100644 cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/batch.cpp create mode 100644 cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/batch.h diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/AGENTS.md b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/AGENTS.md index 34f1d81f90d3..5e43f816534a 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/AGENTS.md +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/AGENTS.md @@ -25,6 +25,7 @@ layouts or ownership models must not override the implementation. - `blockRadixTree.*`: the shared prefix-reuse tree and SHA-256 block keys. - `page.*`, `kvCache.*`, and `kvCacheManager.*`: page lifecycle, per-request cache state, and the top-level manager. +- `batch.*`: stable request rows, dirty tracking, and raw GPU metadata publication. - `storage/`, `storageManager.*`, `evictionController.*`, and `copyEngine.*`: pools, eviction ownership, migration, and data movement. - `lifeCycleRegistry.*`: layer-group/lifecycle mapping, including attention and @@ -79,6 +80,14 @@ slots, schedules pages for eviction, migrates pages between levels, and resizes pools. `CopyEngine` performs the actual batched transfers; C++ code calls it directly and must not round-trip through Python bindings. +`Batch` groups requests from one manager across all layer groups. Request changes +invalidate their stable rows; `publish()` uploads final raw indices and eligible +history counts, including post-rollback state. Device addresses remain fixed. +Publish and wait for readiness outside graph capture; after submitting readers +or replaying a graph, call `recordRead()` before mutating requests or membership. +Publication waits for offload and prior readers, and retains each staging buffer +until its upload completes. DLPack views keep the allocation alive, not KV pages. + The dependency direction is broadly: ```text @@ -146,6 +155,11 @@ KvCache |- KvCacheManager (shared; cache keeps manager alive) `- per-beam/per-block page holders and locks +Batch +|- KvCacheManager (shared) +|- request rows (non-owning; exclusive membership, driven by one owning thread) +`- fixed device metadata and event-protected staging buffers + BlockRadixTree `- roots -> child Blocks (strong ownership through next maps) `- lifecycle page entries (raw observer links) @@ -159,6 +173,10 @@ Eviction controller strong ones merely to simplify access. - `KvCache` keeps its `KvCacheManager` alive. The manager's registry of living caches must not create the reverse strong-reference cycle. +- Closing a request removes its `Batch` row; closing a batch detaches its live + requests without closing them. Removed rows stay dirty until publication clears + them. Batch destruction waits for uploads and recorded readers before freeing + memory; exported arrays can extend the allocation lifetime beyond `close()`. - A committed page is referenced by the radix tree without making the tree its permanent owner. Eviction queues may be the only strong owner of a droppable page, so never store a raw pointer past the operation that obtained it. diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/CMakeLists.txt b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/CMakeLists.txt index a2a737db9821..5033b0036ce8 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/CMakeLists.txt +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/CMakeLists.txt @@ -39,6 +39,7 @@ if(CMAKE_SYSTEM_PROCESSOR MATCHES "aarch64|arm64|ARM64") endif() set(KV_CACHE_MANAGER_V2_SRCS + kv_cache_manager_v2/batch.cpp kv_cache_manager_v2/common.cpp kv_cache_manager_v2/config.cpp kv_cache_manager_v2/lifeCycleRegistry.cpp diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/batch.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/batch.cpp new file mode 100644 index 000000000000..54bfada69bf2 --- /dev/null +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/batch.cpp @@ -0,0 +1,424 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "kv_cache_manager_v2/batch.h" +#include "kv_cache_manager_v2/exceptions.h" +#include "kv_cache_manager_v2/kvCache.h" +#include "kv_cache_manager_v2/kvCacheManager.h" +#include "kv_cache_manager_v2/utils/funcGuard.h" +#include "kv_cache_manager_v2/utils/optionalGilRelease.h" + +#include +#include +#include +#include + +namespace tensorrt_llm::batch_manager::kv_cache_manager_v2 +{ +namespace +{ + +void checkOutsideCapture(CudaStream stream) +{ + CUstreamCaptureStatus status; + cuCheck(cuStreamIsCapturing(reinterpret_cast(stream), &status)); + if (status != CU_STREAM_CAPTURE_STATUS_NONE) + { + throw LogicError("Batch publication and reader fences must run outside CUDA graph capture"); + } +} + +} // namespace + +Batch::Batch(std::shared_ptr manager, int maxRows, int maxBlocks, int maxBeamWidth) + : mManager(std::move(manager)) + , mMaxRows(maxRows) + , mMaxBlocks(maxBlocks) + , mMaxBeamWidth(maxBeamWidth) +{ + KVCM2_API_GUARD(); + if (!mManager || maxRows <= 0 || maxBlocks <= 0 || maxBeamWidth != 1) + { + throw std::invalid_argument("Batch requires a manager, positive dimensions and beam width 1"); + } + auto const apiLock = mManager->lockExclusive(); + mNumLayerGroups = mManager->lifeCycles().size().value(); + size_t const rows = static_cast(mNumLayerGroups) * mMaxRows * mMaxBeamWidth; + size_t const columns = static_cast(mMaxBlocks) + 1; + if (rows == 0 || rows > std::numeric_limits::max() / sizeof(int32_t) / columns) + { + throw std::invalid_argument("Batch metadata dimensions overflow"); + } + mTableElements = rows * mMaxBlocks; + mTotalBytes = rows * columns * sizeof(int32_t); + mRows.resize(mMaxRows, nullptr); + mDirty.resize(mMaxRows, 1); + cuCheck(cuCtxGetDevice(&mDeviceId)); + CUdeviceptr ptr = 0; + cuCheck(cuMemAlloc(&ptr, mTotalBytes)); + mDeviceMemory.reset(reinterpret_cast(ptr)); + // Two generations cover the usual publication/consumer overlap without + // allocating or registering pinned memory in the publication path. + mUploads.push_back({std::make_unique(mTotalBytes)}); + mUploads.push_back({std::make_unique(mTotalBytes)}); +} + +Batch::~Batch() +{ + KVCM2_POISON_ON_EXCEPT( + [this]() + { + OptionalGilRelease const gilRelease; + close(); + auto const apiLock = mManager->lockExclusive(); + mReady.synchronize(); + for (auto const& reader : mReaders) + { + reader.synchronize(); + } + for (auto const& upload : mUploads) + { + upload.completion.synchronize(); + } + }); +} + +void Batch::checkOpen() const +{ + if (mClosed) + { + throw LogicError("Batch is closed"); + } +} + +void Batch::checkPublished() const +{ + checkOpen(); + if (!mPublished) + { + throw LogicError("Batch metadata changed; publish before consumption"); + } +} + +void Batch::markDirty(int row) noexcept +{ + mDirty[row] = 1; + mPublished = false; +} + +int Batch::add(KvCache& cache, std::optional row) +{ + KVCM2_API_GUARD(); + auto const apiLock = mManager->lockExclusive(); + checkOpen(); + if (&cache.manager() != mManager.get() || cache.isClosed()) + { + throw LogicError("Batch members must be live requests from the same manager"); + } + if (cache.mPageStorageBatch != nullptr) + { + if (cache.mPageStorageBatch == this && (!row || row == cache.mPageStorageRow)) + { + return *cache.mPageStorageRow; + } + throw LogicError("A request can belong to only one batch and row"); + } + int const index = row.value_or(static_cast(std::find(mRows.begin(), mRows.end(), nullptr) - mRows.begin())); + if (index < 0 || index >= mMaxRows || mRows[index] != nullptr) + { + throw std::invalid_argument("Batch row is unavailable"); + } + if (cache.numBlocks().value() > mMaxBlocks || cache.beamWidth().value() > mMaxBeamWidth) + { + throw std::invalid_argument("Request exceeds batch dimensions"); + } + mRows[index] = &cache; + cache.mPageStorageBatch = this; + cache.mPageStorageRow = index; + cache.onPageStorageChanged(); + return index; +} + +void Batch::remove(KvCache& cache) +{ + auto const apiLock = mManager->lockExclusive(); + if (cache.mPageStorageBatch == nullptr) + { + return; + } + if (cache.mPageStorageBatch != this) + { + throw LogicError("Request belongs to another batch"); + } + int const row = *cache.mPageStorageRow; + mRows[row] = nullptr; + markDirty(row); + cache.mPageStorageBatch = nullptr; + cache.mPageStorageRow.reset(); + cache.onPageStorageChanged(); +} + +void Batch::close() +{ + auto const apiLock = mManager->lockExclusive(); + if (mClosed || Poison::poisoned()) + { + return; + } + for (auto* cache : mRows) + { + if (cache != nullptr) + { + remove(*cache); + } + } + mClosed = true; + mPublished = false; +} + +size_t Batch::rowOffset(int group, int row) const noexcept +{ + return (static_cast(group) * mMaxRows + row) * mMaxBeamWidth * mMaxBlocks; +} + +size_t Batch::countOffset(int group, int row) const noexcept +{ + return mTableElements + (static_cast(group) * mMaxRows + row) * mMaxBeamWidth; +} + +std::vector Batch::dirtyRows() const +{ + auto const apiLock = mManager->lockShared(); + checkOpen(); + std::vector rows; + for (int row = 0; row < mMaxRows; ++row) + { + if (mDirty[row]) + { + rows.push_back(row); + } + } + return rows; +} + +std::vector Batch::publish(CudaStream stream) +{ + KVCM2_API_GUARD(); + auto const apiLock = mManager->lockExclusive(); + checkOpen(); + auto const cudaStream = reinterpret_cast(stream); + checkOutsideCapture(stream); + auto rows = dirtyRows(); + if (rows.empty()) + { + mReady.waitInStream(stream); + return rows; + } + + struct RequestSnapshot + { + KvCache* cache; + uint64_t version; + std::vector groups; + }; + + std::vector snapshots; + snapshots.reserve(rows.size()); + for (int row : rows) + { + auto* cache = mRows[row]; + RequestSnapshot request{cache, cache ? cache->pageStorageVersion() : 0, {}}; + if (cache != nullptr) + { + if (cache->numBlocks().value() > mMaxBlocks || cache->beamWidth().value() != mMaxBeamWidth) + { + throw std::invalid_argument("Request exceeds batch dimensions"); + } + request.groups.reserve(mNumLayerGroups); + for (int group = 0; group < mNumLayerGroups; ++group) + { + request.groups.push_back(cache->getPageStorageSnapshot(LayerGroupId{group})); + } + } + snapshots.push_back(std::move(request)); + } + + // A busy staging buffer cannot be overwritten by the CPU just by queueing a + // stream wait. Reuse only completed buffers; retain each in-flight generation. + std::unique_ptr staging; + for (auto it = mUploads.begin(); it != mUploads.end();) + { + if (it->completion.queryComplete()) + { + staging = std::move(it->staging); + mUploads.erase(it); + break; + } + else + { + ++it; + } + } + if (!staging) + { + staging = std::make_unique(mTotalBytes); + } + auto* host = reinterpret_cast(staging->address()); + for (size_t i = 0; i < rows.size(); ++i) + { + for (int group = 0; group < mNumLayerGroups; ++group) + { + int32_t* indices = host + rowOffset(group, rows[i]); + std::fill_n(indices, mMaxBlocks, kBadPageIndex.value()); + int32_t& count = host[countOffset(group, rows[i])]; + count = 0; + if (snapshots[i].cache != nullptr) + { + auto const& snapshot = snapshots[i].groups[group]; + std::copy(snapshot.basePageIndices().begin(), snapshot.basePageIndices().end(), indices); + count = snapshot.eligibleHistoryBlocks(); + } + } + } + mUploads.push_back({std::move(staging)}); + auto& upload = mUploads.back(); + auto fence = FuncGuard( + [&]() + { + mReady = CachedCudaEvent(stream); + upload.completion = mReady; + }); + mReady.waitInStream(stream); + for (auto const& reader : mReaders) + { + reader.waitInStream(stream); + } + mReaders.clear(); + for (auto const& request : snapshots) + { + for (auto const& snapshot : request.groups) + { + snapshot.waitReady(stream); + } + } + auto const device = reinterpret_cast(mDeviceMemory.get()); + for (int row : rows) + { + for (int group = 0; group < mNumLayerGroups; ++group) + { + size_t const offset = rowOffset(group, row); + cuCheck(cuMemcpyHtoDAsync( + device + offset * sizeof(int32_t), host + offset, mMaxBlocks * sizeof(int32_t), cudaStream)); + size_t const count = countOffset(group, row); + cuCheck(cuMemcpyHtoDAsync(device + count * sizeof(int32_t), host + count, sizeof(int32_t), cudaStream)); + } + } + fence.run(); + for (size_t i = 0; i < rows.size(); ++i) + { + auto const& request = snapshots[i]; + if (request.cache != nullptr) + { + TLLM_CHECK(request.cache->acknowledgePageStorage(request.version)); + } + mDirty[rows[i]] = 0; + } + mPublished = true; + return rows; +} + +void Batch::waitReady(CudaStream stream) const +{ + KVCM2_API_GUARD(); + auto const apiLock = mManager->lockShared(); + checkPublished(); + checkOutsideCapture(stream); + mReady.waitInStream(stream); +} + +void Batch::recordRead(CudaStream stream) +{ + KVCM2_API_GUARD(); + auto const apiLock = mManager->lockExclusive(); + checkOpen(); + checkOutsideCapture(stream); + std::erase_if(mReaders, [](auto const& event) { return event.queryComplete(); }); + CachedCudaEvent completion(stream); + mReaders.push_back(completion); + for (auto* cache : mRows) + { + if (cache != nullptr && cache->isActive()) + { + completion.waitInStream(reinterpret_cast(cache->cudaStream())); + } + } +} + +std::vector> Batch::resize(std::vector> const& capacities, + std::vector> const& historyLengths, CudaStream stream) +{ + KVCM2_API_GUARD(); + auto const apiLock = mManager->lockExclusive(); + checkOpen(); + checkOutsideCapture(stream); + if (capacities.size() != mRows.size() || historyLengths.size() != mRows.size()) + { + throw std::invalid_argument("Batch resize arguments must be indexed by stable row"); + } + for (int row = 0; row < mMaxRows; ++row) + { + if ((capacities[row] + && (*capacities[row] < 0 + || static_cast(*capacities[row]) + > static_cast(mMaxBlocks) * mManager->tokensPerBlock())) + || (historyLengths[row] && *historyLengths[row] < 0) + || (mRows[row] == nullptr && (capacities[row] || historyLengths[row]))) + { + throw std::invalid_argument("Invalid batch resize capacity, history, or empty row"); + } + } + std::vector> results(mMaxRows); + for (int row = 0; row < mMaxRows; ++row) + { + if (auto* cache = mRows[row]) + { + results[row] = cache->resize(capacities[row], historyLengths[row]); + } + } + publish(stream); + return results; +} + +MemAddress Batch::pageTableAddress(LayerGroupId group) const +{ + if (group.value() < 0 || group.value() >= mNumLayerGroups) + { + throw std::out_of_range("Invalid batch layer group"); + } + return reinterpret_cast(mDeviceMemory.get()) + rowOffset(group.value(), 0) * sizeof(int32_t); +} + +MemAddress Batch::numBlocksAddress(LayerGroupId group) const +{ + if (group.value() < 0 || group.value() >= mNumLayerGroups) + { + throw std::out_of_range("Invalid batch layer group"); + } + return reinterpret_cast(mDeviceMemory.get()) + countOffset(group.value(), 0) * sizeof(int32_t); +} + +} // namespace tensorrt_llm::batch_manager::kv_cache_manager_v2 diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/batch.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/batch.h new file mode 100644 index 000000000000..e5c2970f63eb --- /dev/null +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/batch.h @@ -0,0 +1,135 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#pragma once + +#include "kv_cache_manager_v2/common.h" +#include "kv_cache_manager_v2/lifeCycleRegistry.h" +#include "kv_cache_manager_v2/stagingBuffer.h" +#include "kv_cache_manager_v2/utils/cudaEvent.h" + +#include +#include +#include +#include +#include + +namespace tensorrt_llm::batch_manager::kv_cache_manager_v2 +{ + +class KvCache; +class KvCacheManager; + +//! Stable request rows and raw device metadata, shared across all layer groups. +//! +//! Driven by one owning thread, like its member requests. Membership is non-owning: +//! closing/destroying a request detaches it; closing/destroying a batch detaches its +//! live members without closing them. Mutable state uses the manager's API lock. +//! Publish outside graph capture, then waitReady on a consumer stream. Submit all +//! reads and call recordRead before mutating requests, rows, or tables. +class Batch : public std::enable_shared_from_this +{ +public: + Batch(std::shared_ptr manager, int maxRows, int maxBlocks, int maxBeamWidth = 1); + ~Batch(); + + Batch(Batch const&) = delete; + Batch& operator=(Batch const&) = delete; + + //! Attach a request at a chosen row or the first free row. Membership is exclusive. + int add(KvCache& cache, std::optional row = std::nullopt); + //! Detach a member, leaving its row dirty so the next publication clears it. + void remove(KvCache& cache); + //! Detach all members. Exported arrays retain their allocation until the last owner dies. + void close(); + + //! Upload final dirty rows and counts; return the row slots uploaded. + //! Staging buffers are retained until their asynchronous copies complete. + std::vector publish(CudaStream stream); + //! Wait for publication and KV readiness. Reject unpublished changes. + void waitReady(CudaStream stream) const; + //! Fence submitted reads before table overwrite and request storage release. + void recordRead(CudaStream stream); + //! Resize by stable row slot, then publish once. Empty rows return nullopt. + std::vector> resize(std::vector> const& capacities, + std::vector> const& historyLengths, CudaStream stream); + + std::vector dirtyRows() const; + + int maxRows() const noexcept + { + return mMaxRows; + } + + int maxBlocks() const noexcept + { + return mMaxBlocks; + } + + int maxBeamWidth() const noexcept + { + return mMaxBeamWidth; + } + + int numLayerGroups() const noexcept + { + return mNumLayerGroups; + } + + int deviceId() const noexcept + { + return mDeviceId; + } + + //! Internal device views: [row, beam, block] and [row, beam], respectively. + //! Addresses stay fixed for the batch lifetime; callers must obey the read contract. + MemAddress pageTableAddress(LayerGroupId group) const; + MemAddress numBlocksAddress(LayerGroupId group) const; + +private: + friend class KvCache; + void markDirty(int row) noexcept; + void checkOpen() const; + void checkPublished() const; + size_t rowOffset(int group, int row) const noexcept; + size_t countOffset(int group, int row) const noexcept; + + struct Upload + { + std::unique_ptr staging; + CachedCudaEvent completion = CachedCudaEvent::makeNull(); + }; + + std::shared_ptr mManager; + int mMaxRows; + int mMaxBlocks; + int mMaxBeamWidth; + int mNumLayerGroups = 0; + int mDeviceId = 0; + size_t mTableElements = 0; + size_t mTotalBytes = 0; + CudaUniqPtr mDeviceMemory; + std::vector mRows; + std::vector mDirty; + std::list mUploads; + CachedCudaEvent mReady = CachedCudaEvent::makeNull(); + std::vector mReaders; + bool mPublished = false; + bool mClosed = false; +}; + +} // namespace tensorrt_llm::batch_manager::kv_cache_manager_v2 diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.cpp index 4d5ebf26245e..ccd522526f45 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.cpp @@ -16,6 +16,7 @@ */ #include "kv_cache_manager_v2/kvCache.h" +#include "kv_cache_manager_v2/batch.h" #include "kv_cache_manager_v2/blockRadixTree.h" #include "kv_cache_manager_v2/common.h" #include "kv_cache_manager_v2/exceptions.h" @@ -726,6 +727,7 @@ void KvCache::_deactivate() _freeScratchSlots(); } mStatus = Status::SUSPENDED; + onPageStorageChanged(); } void KvCache::close() @@ -764,7 +766,13 @@ void KvCache::close() auto scope = recordEventScope(); _clearBlocks(); } + if (mPageStorageBatch != nullptr) + { + mPageStorageBatch->remove(*this); + } mStatus = Status::CLOSED; + mPageStorageRow.reset(); + onPageStorageChanged(); mManager->unregisterKvCache(this); } @@ -2628,7 +2636,14 @@ bool KvCache::_checkSanity() const } else { - TLLM_CHECK_DEBUG(std::holds_alternative(bp)); + if (mStatus == Status::ACTIVE) + { + TLLM_CHECK_DEBUG(std::holds_alternative(bp)); + } + else + { + TLLM_CHECK_DEBUG(std::holds_alternative>(bp)); + } auto page = blockPageGetPage(bp); TLLM_CHECK_DEBUG(dynamicPointerCast(page) != nullptr); } @@ -2643,6 +2658,120 @@ bool KvCache::_checkSanity() const // Page index tables // --------------------------------------------------------------------------- +void PageStorageSnapshot::waitReady(CudaStream stream) const +{ + for (auto const& event : mReadyEvents) + event.waitInStream(stream); +} + +void KvCache::onPageStorageChanged() noexcept +{ + ++mPageStorageVersion; + mPageStorageDirty = true; + if (mPageStorageBatch != nullptr) + { + mPageStorageBatch->markDirty(*mPageStorageRow); + } +} + +uint64_t KvCache::pageStorageVersion() const +{ + auto const apiLock = mManager->lockShared(); + return mPageStorageVersion; +} + +bool KvCache::pageStorageDirty() const +{ + auto const apiLock = mManager->lockShared(); + return mPageStorageDirty; +} + +bool KvCache::acknowledgePageStorage(uint64_t version) +{ + KVCM2_API_GUARD(); + auto const apiLock = mManager->lockExclusive(); + if (version != mPageStorageVersion) + return false; + mPageStorageDirty = false; + return true; +} + +void KvCache::bindPageStorageRow(std::optional row) +{ + KVCM2_API_GUARD(); + auto const apiLock = mManager->lockExclusive(); + if (mPageStorageBatch != nullptr) + { + throw LogicError("Use Batch membership APIs to change a batched request's row"); + } + if (row && (*row < 0 || mStatus == Status::CLOSED)) + throw LogicError("Page storage rows must be nonnegative and bound to a live request"); + mPageStorageRow = row; + onPageStorageChanged(); +} + +std::optional KvCache::pageStorageRow() const +{ + auto const apiLock = mManager->lockShared(); + return mPageStorageRow; +} + +PageStorageSnapshot KvCache::getPageStorageSnapshot(LayerGroupId lgId, BeamIndex beamIdx) const +{ + auto const apiLock = mManager->lockShared(); + auto const& buf = mBasePageIndices.at(beamIdx).at(lgId); + PageStorageSnapshot snapshot; + snapshot.mVersion = mPageStorageVersion; + snapshot.mRow = mPageStorageRow; + auto const numBlocks = mBlocks.stdSize(); + if (numBlocks != 0) + { + std::visit([&](auto const& indices) + { snapshot.mBasePageIndices.assign(indices.data(), indices.data() + numBlocks); }, + buf); + } + snapshot.mCacheLevels.resize(numBlocks); + if (!isActive()) + { + // Held pages may be evicted or relocated; only active locks expose usable slot indices. + std::fill(snapshot.mBasePageIndices.begin(), snapshot.mBasePageIndices.end(), kBadPageIndex.value()); + return snapshot; + } + + auto const* attn = std::get_if(&mManager->lifeCycles()[lgId]); + if (mIsDecoding && attn && attn->isSparse) + snapshot.mEligibleHistoryBlocks = mHistoryLength / mTokensPerBlock; + snapshot.mReadyEvents.reserve(numBlocks); + for (BlockOrdinal ord{0}; ord < mBlocks.size(); ++ord) + { + int const index = snapshot.mBasePageIndices[toSizeT(ord)]; + auto const& page = blockPageGetPage(mBlocks[ord].pages.at(beamIdx).at(lgId)); + if (ord.value() < snapshot.mEligibleHistoryBlocks) + { + TLLM_CHECK_WITH_INFO(page && index != kBadPageIndex.value() && page->cacheLevel == kSparseHistoryLevel + && page->hasValidSlot() && index == slotIdToPageIndexValue(page->slotId()), + "Eligible sparse history must have a locked host mapping"); + } + if (index == kBadPageIndex.value()) + continue; + // Dense SWA scratch indices refer to GPU slots without a Page object. + snapshot.mCacheLevels[toSizeT(ord)] = page ? page->cacheLevel : kHotLevel; + if (page) + snapshot.mReadyEvents.push_back(page->readyEvent); + } + return snapshot; +} + +void KvCache::recordPageStorageRead(CudaStream stream) +{ + KVCM2_API_GUARD(); + auto const apiLock = mManager->lockExclusive(); + if (!isActive()) + throw LogicError("Page storage reads must finish submission before the request is deactivated"); + CachedCudaEvent completion(stream); + completion.waitInStream(reinterpret_cast(cudaStream())); +} + void KvCache::_checkPageIndexBufferCapacity(BlockOrdinal newNumBlocks) const { for (auto const& beamIndices : mBasePageIndices) @@ -2662,6 +2791,8 @@ void KvCache::_checkPageIndexBufferCapacity(BlockOrdinal newNumBlocks) const void KvCache::_resizePageIndexBuffers(BlockOrdinal newNumBlocks) { + // External buffers keep their full allocation; the number of published entries still changes. + onPageStorageChanged(); for (BeamIndex bi{0}; bi < mBeamWidth; ++bi) { for (LifeCycleId lcId{0}; lcId < mBasePageIndices[bi].size(); ++lcId) @@ -2700,7 +2831,7 @@ int KvCache::updateBasePageIndex(BeamIndex bi, BlockOrdinal ord, LifeCycleId lc, if (ord == kBadBlockOrdinal) return kBadPageIndex.value(); // SSM pages use BAD_BLOCK_ORDINAL auto& buf = mBasePageIndices[bi][lc]; - return std::visit( + int const old = std::visit( [&](auto& b) -> int { using T = std::decay_t; @@ -2720,6 +2851,9 @@ int KvCache::updateBasePageIndex(BeamIndex bi, BlockOrdinal ord, LifeCycleId lc, } }, buf); + if (old != value) + onPageStorageChanged(); + return old; } Span KvCache::getBasePageIndices(LayerGroupId lgId, BeamIndex beamIdx) const @@ -2764,6 +2898,7 @@ std::vector KvCache::getAggregatedPageIndices(LayerGroupId lgId, BeamIndex void KvCache::setBasePageIndexBuf(BeamIndex beamIdx, LayerGroupId lgId, int32_t* buf, int len) { + auto const apiLock = mManager->lockExclusive(); auto& slot = mBasePageIndices[beamIdx][lgId]; BlockOrdinal const numBlocks = mBlocks.size(); @@ -2776,6 +2911,7 @@ void KvCache::setBasePageIndexBuf(BeamIndex beamIdx, LayerGroupId lgId, int32_t* auto const n = std::min(toSizeT(numBlocks), static_cast(ext->len)); std::vector vec(ext->ptr, ext->ptr + n); slot = std::move(vec); + onPageStorageChanged(); } // If already a vector, nothing to do. return; @@ -2804,6 +2940,7 @@ void KvCache::setBasePageIndexBuf(BeamIndex beamIdx, LayerGroupId lgId, int32_t* std::copy(oldData, oldData + copyLen, buf); std::fill(buf + copyLen, buf + len, kBadPageIndex.value()); slot = Span{buf, len}; + onPageStorageChanged(); } int KvCache::getSsmBlockBaseIndex(LayerGroupId lgId, BeamIndex beamIdx) const diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.h index 9e3ad64fa050..998a6f4c9bc1 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.h +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.h @@ -38,6 +38,7 @@ namespace tensorrt_llm::batch_manager::kv_cache_manager_v2 { // Forward declarations. +class Batch; class KvCacheIntrospection; class KvCacheManager; class StorageManager; @@ -152,6 +153,55 @@ class PlannedDropHandle std::optional>> mPageRefs; }; +// Independent host metadata for one request's layer group and beam. Events retain copy readiness, +// not storage ownership: use the indices only while the request is active and its version matches. +class PageStorageSnapshot +{ +public: + uint64_t version() const noexcept + { + return mVersion; + } + + std::optional row() const noexcept + { + return mRow; + } + + std::vector const& basePageIndices() const noexcept + { + return mBasePageIndices; + } + + // A missing level accompanies BAD_PAGE_INDEX. Valid indices address slots in the indicated level. + std::vector> const& cacheLevels() const noexcept + { + return mCacheLevels; + } + + int eligibleHistoryBlocks() const noexcept + { + return mEligibleHistoryBlocks; + } + + std::vector const& readyEvents() const noexcept + { + return mReadyEvents; + } + + // Queue copy-completion dependencies without blocking the CPU or publishing device metadata. + void waitReady(CudaStream stream) const; + +private: + friend class KvCache; + uint64_t mVersion = 0; + std::optional mRow; + std::vector mBasePageIndices; + std::vector> mCacheLevels; + int mEligibleHistoryBlocks = 0; + std::vector mReadyEvents; +}; + // --------------------------------------------------------------------------- // KvCache — manages the per-sequence KV cache state. // Mirrors Python's _KVCache. @@ -230,12 +280,27 @@ class KvCache : public std::enable_shared_from_this //! Caller must ensure every owner has finished the execution phase that requires these pages on GPU. void offloadSparsePages(std::vector> const& pages); - //! Changes when offload relocates an owned page, even if its numeric slot index stays the same. - //! Internal invalidation hook for page metadata; read under the manager's API lock. - uint64_t pageStorageVersion() const noexcept - { - return mPageStorageVersion; - } + // CPU-side invalidation for page indices, levels, readiness, eligibility and row/buffer bindings. + // Several changes leave one pending refresh. Acknowledging an older version never clears it. + uint64_t pageStorageVersion() const; + bool pageStorageDirty() const; + // Acknowledge only after refreshing every group/beam from snapshots of this same version. + bool acknowledgePageStorage(uint64_t version); + + // Associate a consumer's stable row with this request; nullopt detaches it. Every bind + // requires a refresh, including reuse of the same row number by a new consumer. Close detaches it. + // Batch members must use Batch::add/remove instead of rebinding directly. + void bindPageStorageRow(std::optional row); + std::optional pageStorageRow() const; + + // Copies raw indices (no expansion or BAD-to-zero conversion) and readiness under the API lock. + // Eligibility is zero for prefill, inactive requests and non-sparse layer groups. + PageStorageSnapshot getPageStorageSnapshot(LayerGroupId lgId, BeamIndex beamIdx = kDefaultBeamIndex) const; + + // Call after submitting reads on another stream, before mutating/suspending/closing this cache. + // Joins those reads into the request stream so its normal unlock/commit fences protect storage. + // Snapshot acquisition, read submission and this call belong to the request's owning thread. + void recordPageStorageRead(CudaStream stream); // ---- Committing tokens ------------------------------------------------- @@ -475,6 +540,7 @@ class KvCache : public std::enable_shared_from_this // ---- Internal callbacks (called by SharedPageLock) ---------------------- + // Caller holds the manager's exclusive API lock, including when updating another owner's table. int updateBasePageIndex(BeamIndex bi, BlockOrdinal ord, LifeCycleId lc, int value); std::optional id; // opaque identifier (mirrors Python's id field) @@ -482,13 +548,11 @@ class KvCache : public std::enable_shared_from_this private: friend class KvCacheIntrospection; friend class UniqPageLock; + friend class Batch; friend std::vector batchedLockPages( KvCache& kvCache, std::vector const& targets); - void onPageStorageChanged() noexcept - { - ++mPageStorageVersion; - } + void onPageStorageChanged() noexcept; // Activate: lock active pages at their required levels. mCudaStream must already be set. // Internal — called by resume(). Not public (mirrors Python where activate() doesn't exist). @@ -669,7 +733,10 @@ class KvCache : public std::enable_shared_from_this using LifeCyclePageIndexBuffers = TypedVec; using BeamPageIndexBuffers = TypedVec; BeamPageIndexBuffers mBasePageIndices; + Batch* mPageStorageBatch = nullptr; // Non-owning; both destructors detach membership. uint64_t mPageStorageVersion = 0; + bool mPageStorageDirty = true; + std::optional mPageStorageRow; TypedVec mBlocks; diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCacheManager.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCacheManager.cpp index 15c852ba8c17..180a72afb001 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCacheManager.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCacheManager.cpp @@ -300,6 +300,13 @@ int KvCacheManager::getPageIndexScale(LayerId layerId, DataRole role) const return mStorage->mSlotToPageIndices.at(attr.lifeCycleId).at(attr.poolIndex); } +bool KvCacheManager::isSparse(LayerId layerId, DataRole role) const +{ + auto const& attr = mStorage->getBufferAttr(layerId, role); + auto const* attn = std::get_if(&mLifeCycles[attr.lifeCycleId]); + return attn && attn->isSparse; +} + PageIndexConverter KvCacheManager::getPageIndexConverter(LayerId layerId, DataRole role) const { auto const& attr = mStorage->getBufferAttr(layerId, role); diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCacheManager.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCacheManager.h index 3d15d175244b..20a94559582e 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCacheManager.h +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCacheManager.h @@ -166,6 +166,9 @@ class KvCacheManager : public std::enable_shared_from_this int getPageStride(LayerId layerId, DataRole role) const; size_t getPageIndexUpperBound(LayerId layerId, DataRole role) const; + // Whether this buffer belongs to a sparse-attention lifecycle. Unknown buffers throw. + bool isSparse(LayerId layerId, DataRole role) const; + // Scale factor: base_page_index * scale → kernel page index. int getPageIndexScale(LayerId layerId, DataRole role) const; diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/page.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/page.cpp index a2905cc56b8c..f071bed08ec3 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/page.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/page.cpp @@ -387,6 +387,9 @@ void UniqPageLock::recordOffloadEvent(CachedCudaEvent const& event) page()->readyEvent = event; finishEvents.clear(); finishEvents.push_back(event); + // A rejected copy can change readiness without changing the source slot. + for (auto const& owner : mOwners) + owner.kvCache->onPageStorageChanged(); } Slot UniqPageLock::moveToSparseHistory(Slot&& hostSlot) diff --git a/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp b/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp index 7fbd59ba4a3d..1140a5b66586 100644 --- a/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp +++ b/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp @@ -17,6 +17,7 @@ #include "tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.h" +#include "kv_cache_manager_v2/batch.h" #include "kv_cache_manager_v2/blockRadixTree.h" #include "kv_cache_manager_v2/coldPageCodec.h" #include "kv_cache_manager_v2/common.h" @@ -67,6 +68,14 @@ namespace tensorrt_llm::nanobind::batch_manager namespace { +//! A DLPack view keeps the batch's device allocation alive without importing torch. +struct BatchDeviceArray +{ + std::shared_ptr batch; + kv::LayerGroupId group; + bool counts; +}; + // Exposed via introspection sub-module for tests. class TestPaddingColdPageCodec final : public kv::IKvCacheColdPageCodec { @@ -1701,6 +1710,111 @@ void KvCacheManagerV2Bindings::initBindings(nb::module_& m) nb::class_(m, "PlannedDropHandle") .def("drop", &kv::PlannedDropHandle::drop, nb::call_guard()); + nb::class_(m, "PageStorageSnapshot") + .def_prop_ro("version", &kv::PageStorageSnapshot::version) + .def_prop_ro("row", &kv::PageStorageSnapshot::row) + .def_prop_ro("base_page_indices", &kv::PageStorageSnapshot::basePageIndices) + .def_prop_ro("cache_levels", + [](kv::PageStorageSnapshot const& self) + { + std::vector> levels; + levels.reserve(self.cacheLevels().size()); + for (auto const& level : self.cacheLevels()) + levels.push_back(level ? std::optional{level->value()} : std::nullopt); + return levels; + }) + .def_prop_ro("eligible_history_blocks", &kv::PageStorageSnapshot::eligibleHistoryBlocks) + .def("wait_ready", &kv::PageStorageSnapshot::waitReady, nb::arg("cuda_stream"), + nb::call_guard()); + + nb::class_(m, "BatchDeviceArray") + .def("__dlpack_device__", + [](BatchDeviceArray const& self) + { return std::make_pair(nb::device::cuda::value, self.batch->deviceId()); }) + .def( + "__dlpack__", + [](BatchDeviceArray const& self, std::optional stream, nb::kwargs kwargs) + { + if (kwargs.contains("copy") && !kwargs["copy"].is_none() && nb::cast(kwargs["copy"])) + { + throw std::invalid_argument("Batch arrays support only zero-copy DLPack export"); + } + if (kwargs.contains("dl_device") && !kwargs["dl_device"].is_none() + && nb::cast>(kwargs["dl_device"]) + != std::make_pair(nb::device::cuda::value, self.batch->deviceId())) + { + throw std::invalid_argument("Batch arrays cannot be exported to a different device"); + } + // DLPack 1/2 denote the legacy/per-thread default CUDA streams. + auto cudaStream = stream.value_or(1); + if (cudaStream == 0 || cudaStream < -1) + { + throw std::invalid_argument("Invalid DLPack CUDA stream"); + } + if (cudaStream != -1) + { + auto const consumerStream = cudaStream == 1 ? CU_STREAM_LEGACY + : cudaStream == 2 ? CU_STREAM_PER_THREAD + : reinterpret_cast(cudaStream); + nb::gil_scoped_release release; + self.batch->waitReady(reinterpret_cast(consumerStream)); + } + else + { + nb::gil_scoped_release release; + if (!self.batch->dirtyRows().empty()) + { + throw kv::LogicError("Publish Batch metadata before exporting it"); + } + } + std::vector shape{ + static_cast(self.batch->maxRows()), static_cast(self.batch->maxBeamWidth())}; + auto address + = self.counts ? self.batch->numBlocksAddress(self.group) : self.batch->pageTableAddress(self.group); + if (!self.counts) + { + shape.push_back(self.batch->maxBlocks()); + } + return nb::ndarray(reinterpret_cast(address), shape.size(), + shape.data(), nb::cast(self.batch), nullptr, nb::dtype(), nb::device::cuda::value, + self.batch->deviceId()); + }, + nb::arg("stream").none() = nb::none(), nb::arg("kwargs")); + + nb::class_(m, "Batch") + .def(nb::init, int, int, int>(), nb::arg("manager"), nb::arg("max_rows"), + nb::arg("max_blocks"), nb::arg("max_beam_width") = 1, nb::call_guard()) + .def("add", &kv::Batch::add, nb::arg("kv_cache"), nb::arg("row").none() = std::nullopt, + nb::call_guard()) + .def("remove", &kv::Batch::remove, nb::arg("kv_cache"), nb::call_guard()) + .def("close", &kv::Batch::close, nb::call_guard()) + .def("publish", &kv::Batch::publish, nb::arg("cuda_stream"), nb::call_guard()) + .def("wait_ready", &kv::Batch::waitReady, nb::arg("cuda_stream"), nb::call_guard()) + .def("record_read", &kv::Batch::recordRead, nb::arg("cuda_stream"), nb::call_guard()) + .def("resize", &kv::Batch::resize, nb::arg("capacities"), nb::arg("history_lengths"), nb::arg("cuda_stream"), + nb::call_guard()) + .def_prop_ro("dirty_rows", &kv::Batch::dirtyRows, nb::call_guard()) + .def_prop_ro("max_rows", &kv::Batch::maxRows) + .def_prop_ro("max_blocks", &kv::Batch::maxBlocks) + .def_prop_ro("max_beam_width", &kv::Batch::maxBeamWidth) + .def_prop_ro("num_layer_groups", &kv::Batch::numLayerGroups) + .def( + "page_table", + [](std::shared_ptr self, int group) + { + self->pageTableAddress(kv::LayerGroupId{group}); + return BatchDeviceArray{std::move(self), kv::LayerGroupId{group}, false}; + }, + nb::arg("layer_group_id")) + .def( + "num_blocks", + [](std::shared_ptr self, int group) + { + self->numBlocksAddress(kv::LayerGroupId{group}); + return BatchDeviceArray{std::move(self), kv::LayerGroupId{group}, true}; + }, + nb::arg("layer_group_id")); + // ---- KvCache ----------------------------------------------------------- nb::class_(m, "_KVCache") .def( @@ -1716,6 +1830,20 @@ void KvCacheManagerV2Bindings::initBindings(nb::module_& m) nb::arg("cuda_stream") = nb::none(), nb::arg("is_decoding") = nb::none()) .def("enter_decode", &kv::KvCache::enterDecode, nb::call_guard()) .def_prop_ro("is_decoding", &kv::KvCache::isDecoding) + .def_prop_ro("page_storage_version", &kv::KvCache::pageStorageVersion, nb::call_guard()) + .def_prop_ro("page_storage_dirty", &kv::KvCache::pageStorageDirty, nb::call_guard()) + .def_prop_ro("page_storage_row", &kv::KvCache::pageStorageRow, nb::call_guard()) + .def("bind_page_storage_row", &kv::KvCache::bindPageStorageRow, nb::arg("row").none(), + nb::call_guard()) + .def("acknowledge_page_storage", &kv::KvCache::acknowledgePageStorage, nb::arg("version"), + nb::call_guard()) + .def( + "get_page_storage_snapshot", + [](kv::KvCache const& self, int layerGroupId, int beamIdx) + { return self.getPageStorageSnapshot(kv::LayerGroupId{layerGroupId}, kv::BeamIndex{beamIdx}); }, + nb::arg("layer_group_id"), nb::arg("beam_id") = 0, nb::call_guard()) + .def("record_page_storage_read", &kv::KvCache::recordPageStorageRead, nb::arg("cuda_stream"), + nb::call_guard()) .def("suspend", &kv::KvCache::suspend, nb::call_guard()) .def( "prefetch", [](kv::KvCache& self, int target) { return self.prefetch(kv::CacheLevel{target}); }, @@ -1866,6 +1994,7 @@ void KvCacheManagerV2Bindings::initBindings(nb::module_& m) kv::LayerGroupId const typedLayerGroupId{layerGroupId}; if (bufObj.is_none()) { + nb::gil_scoped_release release; self.setBasePageIndexBuf(typedBeamIdx, typedLayerGroupId, nullptr, 0); return; } @@ -1885,6 +2014,7 @@ void KvCacheManagerV2Bindings::initBindings(nb::module_& m) } cleanup{&view}; if (std::string(view.format) != "i" || view.ndim != 1) throw std::invalid_argument("set_base_page_index_buf: buffer must be 1-D int32 ('i')"); + nb::gil_scoped_release release; self.setBasePageIndexBuf(typedBeamIdx, typedLayerGroupId, static_cast(view.buf), static_cast(view.len / sizeof(int32_t))); }, @@ -2302,6 +2432,7 @@ void KvCacheManagerV2Bindings::initBindings(nb::module_& m) nb::arg("config"), nb::arg("event_manager").none() = nb::none(), nb::arg("cold_page_codec").none() = nb::none()) .def("shutdown", &kv::KvCacheManager::shutdown, nb::call_guard()) + .def("is_sparse", &kv::KvCacheManager::isSparse, nb::arg("layer_id"), nb::arg("data_role")) .def( "clear_reusable_blocks", &kv::KvCacheManager::clearReusableBlocks, nb::call_guard()) .def( diff --git a/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp b/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp index 9f3ad9601d28..a3573d2a9bb1 100644 --- a/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp +++ b/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp @@ -16,6 +16,7 @@ */ #include "kvCacheManagerV2TestUtils.h" +#include "tensorrt_llm/batch_manager/kv_cache_manager_v2/batch.h" #include "tensorrt_llm/batch_manager/kv_cache_manager_v2/blockRadixTree.h" #include "tensorrt_llm/batch_manager/kv_cache_manager_v2/config.h" #include "tensorrt_llm/batch_manager/kv_cache_manager_v2/eventManager.h" @@ -852,6 +853,693 @@ class KvCacheManagerV2SparseOffloadTest : public KvCacheManagerV2PageLockTest { }; +class KvCacheManagerV2PageStorageTest : public KvCacheManagerV2PageLockTest +{ +}; + +class KvCacheManagerV2BatchTest : public KvCacheManagerV2PageLockTest +{ +protected: + CudaStream batchStream() const + { + return reinterpret_cast(stream()); + } + + std::vector read(MemAddress address, size_t size) + { + cuCheck(cuStreamSynchronize(stream())); + std::vector result(size); + cuCheck(cuMemcpyDtoH(result.data(), address, size * sizeof(int32_t))); + return result; + } +}; + +TEST_F(KvCacheManagerV2BatchTest, PublishesRawMixedTierRowsAndSkipsUnchangedRows) +{ + auto config = makeSplitColdGroupingConfig(); + auto& sparse = std::get(config.layers.front()); + sparse.buffers.front().isSparse = true; + sparse.buffers.push_back({.role = "value", .size = 2048, .tokensPerBlockOverride = 2, .isSparse = true}); + auto coalesced = sparse; + coalesced.layerId = 2; + config.layers.push_back(std::move(coalesced)); + auto manager = std::make_shared(std::move(config)); + Batch batch(manager, 3, 4); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(12, 6)); + EXPECT_EQ(batch.add(*cache, 2), 2); + auto const sparseGroup = manager->getLayerGroupId(0); + auto const denseGroup = manager->getLayerGroupId(1); + auto const address = batch.pageTableAddress(sparseGroup); + EXPECT_EQ(batch.publish(batchStream()), (std::vector{0, 1, 2})); + EXPECT_EQ(read(batch.numBlocksAddress(sparseGroup), 3), (std::vector{0, 0, 0})); + EXPECT_EQ(batch.publish(batchStream()), (std::vector{})); + EXPECT_FALSE(cache->pageStorageDirty()); + + ASSERT_TRUE(cache->enterDecode()); + EXPECT_EQ(batch.dirtyRows(), (std::vector{2})); + EXPECT_THROW(batch.waitReady(batchStream()), LogicError); + EXPECT_EQ(batch.publish(batchStream()), (std::vector{2})); + for (auto const group : {sparseGroup, denseGroup}) + { + auto expected = std::vector(12, kBadPageIndex.value()); + auto const snapshot = cache->getPageStorageSnapshot(group); + std::copy(snapshot.basePageIndices().begin(), snapshot.basePageIndices().end(), expected.begin() + 8); + EXPECT_EQ(read(batch.pageTableAddress(group), 12), expected); + } + EXPECT_EQ(read(batch.numBlocksAddress(sparseGroup), 3), (std::vector{0, 0, 1})); + EXPECT_EQ(read(batch.numBlocksAddress(denseGroup), 3), (std::vector{0, 0, 0})); + EXPECT_EQ(batch.pageTableAddress(sparseGroup), address); + + cache->suspend(); + EXPECT_EQ(batch.publish(batchStream()), (std::vector{2})); + EXPECT_EQ(read(address, 12), (std::vector(12, kBadPageIndex.value()))); + EXPECT_EQ(read(batch.numBlocksAddress(sparseGroup), 3), (std::vector{0, 0, 0})); + ASSERT_TRUE(cache->resume()); + EXPECT_EQ(batch.publish(batchStream()), (std::vector{2})); + EXPECT_EQ(read(batch.numBlocksAddress(sparseGroup), 3).back(), 1); +} + +TEST_F(KvCacheManagerV2BatchTest, PublishesIncrementalOffloadAndClearsReusedRows) +{ + auto config = makeSplitColdGroupingConfig(); + auto& sparse = std::get(config.layers.front()); + sparse.buffers.front().isSparse = true; + sparse.buffers.push_back({.role = "value", .size = 2048, .tokensPerBlockOverride = 2, .isSparse = true}); + auto coalesced = sparse; + coalesced.layerId = 2; + config.layers.push_back(std::move(coalesced)); + auto codec = std::make_unique(); + auto* observer = codec.get(); + auto manager = std::make_shared(std::move(config), nullptr, std::move(codec)); + auto const apiLock = manager->lockExclusive(); + auto& storage = manager->storage(); + auto const sparseGroup = manager->getLayerGroupId(0); + auto const denseGroup = manager->getLayerGroupId(1); + auto const gpuPool = storage.getPoolGroupIndex(kHotLevel, sparseGroup); + auto const hostPool = storage.getPoolGroupIndex(kSparseHistoryLevel, sparseGroup); + auto const& sizes = storage.slotSize(kHotLevel, gpuPool); + ASSERT_EQ(sizes.size(), PoolIndex{1}); + size_t const bytes = sizes[PoolIndex{0}]; + Batch batch(manager, 2, 4); + auto const tableAddress = batch.pageTableAddress(sparseGroup); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(12, 6)); + for (int ordinal = 0; ordinal < 3; ++ordinal) + { + auto const page = pageAt(*cache, ordinal, sparseGroup); + auto const address + = std::get(storage.slotAddress(kHotLevel, gpuPool, page->slotId(), PoolIndex{0})); + ASSERT_EQ(cudaMemsetAsync(reinterpret_cast(address), 0x40 + ordinal, bytes, mStream), cudaSuccess); + } + batch.add(*cache, 1); + auto checkPublished = [&](int eligible) + { + for (auto const group : {sparseGroup, denseGroup}) + { + auto expected = std::vector(8, kBadPageIndex.value()); + for (int ordinal = 0; ordinal < 3; ++ordinal) + { + auto const page = pageAt(*cache, ordinal, group); + expected[4 + ordinal] = slotIdToPageIndexValue(page->slotId()); + EXPECT_EQ( + page->cacheLevel, group == sparseGroup && ordinal < eligible ? kSparseHistoryLevel : kHotLevel); + } + EXPECT_EQ(read(batch.pageTableAddress(group), 8), expected); + EXPECT_EQ( + read(batch.numBlocksAddress(group), 2), (std::vector{0, group == sparseGroup ? eligible : 0})); + } + for (int ordinal = 0; ordinal < eligible; ++ordinal) + { + auto const page = pageAt(*cache, ordinal, sparseGroup); + auto const address = std::get( + storage.slotAddress(kSparseHistoryLevel, hostPool, page->slotId(), PoolIndex{0})); + auto const* data = reinterpret_cast(address); + EXPECT_TRUE(std::all_of(data, data + bytes, [ordinal](uint8_t value) { return value == 0x40 + ordinal; })); + } + }; + batch.publish(batchStream()); + checkPublished(0); + auto const freeGpuBefore = manager->getStorageStatistics(kHotLevel)[gpuPool].free; + ASSERT_TRUE(cache->enterDecode()); + batch.publish(batchStream()); + checkPublished(1); + auto const firstHostSlot = pageAt(*cache, 0, sparseGroup)->slotId(); + auto const version = cache->pageStorageVersion(); + EXPECT_EQ(batch.resize({std::nullopt, 12}, {std::nullopt, 7}, batchStream()), + (std::vector>{std::nullopt, true})); + EXPECT_EQ(cache->pageStorageVersion(), version); + EXPECT_EQ(observer->encodedPages, 1); + checkPublished(1); + EXPECT_EQ(batch.resize({std::nullopt, 12}, {std::nullopt, 8}, batchStream()), + (std::vector>{std::nullopt, true})); + checkPublished(2); + EXPECT_EQ(observer->encodedPages, 2); + EXPECT_EQ(pageAt(*cache, 0, sparseGroup)->slotId(), firstHostSlot); + EXPECT_EQ(manager->getStorageStatistics(kHotLevel)[gpuPool].free, freeGpuBefore + 2); + EXPECT_EQ(batch.publish(batchStream()), (std::vector{})); + + cache->close(); + auto replacement = manager->createKvCache(); + auto closeReplacement = FuncGuard([&]() { replacement->close(); }); + ASSERT_TRUE(replacement->resume(stream())); + ASSERT_TRUE(replacement->resize(4, 0)); + batch.add(*replacement, 1); + EXPECT_EQ(batch.publish(batchStream()), (std::vector{1})); + auto expected = std::vector(8, kBadPageIndex.value()); + expected[4] = slotIdToPageIndexValue(pageAt(*replacement, 0, sparseGroup)->slotId()); + EXPECT_EQ(read(tableAddress, 8), expected); + EXPECT_EQ(read(batch.numBlocksAddress(sparseGroup), 2), (std::vector{0, 0})); + EXPECT_EQ(batch.pageTableAddress(sparseGroup), tableAddress); +} + +TEST_F(KvCacheManagerV2BatchTest, MembershipCloseAndPublicationFailuresPreserveRows) +{ + auto manager = std::make_shared(sparseConfig()); + Batch batch(manager, 3, 1); + Batch other(manager, 3, 1); + auto first = manager->createKvCache(); + auto second = manager->createKvCache(); + auto closeCaches = FuncGuard( + [&]() + { + first->close(); + if (second) + { + second->close(); + } + }); + EXPECT_EQ(batch.add(*first, 2), 2); + EXPECT_EQ(batch.add(*first), 2); + EXPECT_EQ(batch.add(*second), 0); + EXPECT_THROW(other.add(*first), LogicError); + EXPECT_THROW(first->bindPageStorageRow(1), LogicError); + EXPECT_THROW(batch.add(*second, 2), LogicError); + batch.publish(batchStream()); + first->close(); + EXPECT_EQ(batch.dirtyRows(), (std::vector{2})); + EXPECT_EQ(second->pageStorageRow(), 0); + batch.remove(*second); + EXPECT_EQ(other.add(*second, 1), 1); + other.close(); + EXPECT_FALSE(second->pageStorageRow().has_value()); + EXPECT_FALSE(second->isClosed()); + EXPECT_EQ(batch.add(*second, 1), 1); + ASSERT_TRUE(second->resume(stream())); + ASSERT_TRUE(second->resize(8, 0)); + EXPECT_THROW(batch.publish(batchStream()), std::invalid_argument); + EXPECT_THROW(batch.waitReady(batchStream()), LogicError); + EXPECT_TRUE(second->pageStorageDirty()); + EXPECT_FALSE(batch.dirtyRows().empty()); + ASSERT_TRUE(second->resize(4, 0)); + batch.publish(batchStream()); + EXPECT_FALSE(second->pageStorageDirty()); + second.reset(); + EXPECT_EQ(batch.dirtyRows(), (std::vector{1})); + batch.publish(batchStream()); + EXPECT_EQ(read(batch.pageTableAddress(LifeCycleId{0}), 3), (std::vector(3, -1))); + closeCaches.cancel(); +} + +TEST_F(KvCacheManagerV2BatchTest, BatchedResizePublishesFinalStatesAfterPartialOom) +{ + auto config = sparseConfig(); + config.cacheTiers[0] = GpuCacheTierConfig{8 << 20}; + auto manager = std::make_shared(std::move(config)); + Batch batch(manager, 3, 4); + auto first = manager->createKvCache(); + auto second = manager->createKvCache(); + auto closeCaches = FuncGuard( + [&]() + { + first->close(); + second->close(); + }); + ASSERT_TRUE(first->resume(stream())); + ASSERT_TRUE(second->resume(stream())); + ASSERT_TRUE(first->resize(4, 0)); + ASSERT_TRUE(second->resize(4, 0)); + batch.add(*first, 0); + batch.add(*second, 2); + auto const results = batch.resize({8, std::nullopt, 12}, {0, std::nullopt, 0}, batchStream()); + EXPECT_EQ(results, (std::vector>{true, std::nullopt, false})); + EXPECT_EQ(first->capacity(), 8); + EXPECT_EQ(second->capacity(), 4); + EXPECT_TRUE(batch.dirtyRows().empty()); + auto expected = std::vector(12, -1); + auto const firstIndices = first->getBasePageIndices(LifeCycleId{0}); + auto const secondIndices = second->getBasePageIndices(LifeCycleId{0}); + std::copy(firstIndices.data(), firstIndices.data() + 2, expected.begin()); + expected[8] = secondIndices[0]; + EXPECT_EQ(read(batch.pageTableAddress(LifeCycleId{0}), 12), expected); +} + +TEST_F(KvCacheManagerV2BatchTest, SharedOffloadInvalidatesEveryOwner) +{ + auto manager = std::make_shared(sparseConfig()); + auto const apiLock = manager->lockExclusive(); + auto page = seedPrefix(*manager, kHotLevel); + auto first = manager->createKvCache({}, tokens()); + auto second = manager->createKvCache({}, tokens()); + auto closeCaches = FuncGuard( + [&]() + { + first->close(); + second->close(); + }); + ASSERT_TRUE(first->resume(stream())); + ASSERT_TRUE(second->resume(stream())); + Batch batch(manager, 2, 1); + batch.add(*first); + batch.add(*second); + batch.publish(batchStream()); + first->offloadSparsePages({page}); + EXPECT_EQ(batch.dirtyRows(), (std::vector{0, 1})); + ASSERT_TRUE(first->enterDecode()); + ASSERT_TRUE(second->enterDecode()); + EXPECT_EQ(batch.publish(batchStream()), (std::vector{0, 1})); + EXPECT_EQ(read(batch.pageTableAddress(LifeCycleId{0}), 2), + (std::vector(2, slotIdToPageIndexValue(page->slotId())))); + EXPECT_EQ(read(batch.numBlocksAddress(LifeCycleId{0}), 2), (std::vector{1, 1})); +} + +TEST_F(KvCacheManagerV2BatchTest, RetainsStagingUntilUploadAndOrdersTableReuseAfterReaders) +{ + auto config = sparseConfig(); + std::get(config.layers.front()).buffers.front().size = 4096; + auto manager = std::make_shared(std::move(config)); + Batch batch(manager, 1, 2); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(4, 0)); + batch.add(*cache); + batch.publish(batchStream()); + auto const expected = read(batch.pageTableAddress(LifeCycleId{0}), 2); + HostMem readback(2 * sizeof(int32_t)); + cudaStream_t readerStream{}; + ASSERT_EQ(cudaStreamCreateWithFlags(&readerStream, cudaStreamNonBlocking), cudaSuccess); + auto destroyReader = FuncGuard([&]() { cudaStreamDestroy(readerStream); }); + StreamGate uploadGate; + StreamGate readerGate; + auto releaseGates = FuncGuard( + [&]() + { + uploadGate.release(); + readerGate.release(); + }); + ASSERT_EQ(uploadGate.enqueue(mStream), cudaSuccess); + EXPECT_THROW(cache->bindPageStorageRow(std::nullopt), LogicError); + cache->setBasePageIndexBuf(kDefaultBeamIndex, LifeCycleId{0}, nullptr, 0); + // Commit changes lock ownership/readiness even when the numeric slot stays unchanged. + cache->commit(tokens()); + batch.publish(batchStream()); + batch.waitReady(reinterpret_cast(readerStream)); + ASSERT_EQ(readerGate.enqueue(readerStream), cudaSuccess); + ASSERT_EQ(cudaMemcpyAsync(reinterpret_cast(readback.address()), + reinterpret_cast(batch.pageTableAddress(LifeCycleId{0})), 2 * sizeof(int32_t), + cudaMemcpyDeviceToHost, readerStream), + cudaSuccess); + batch.recordRead(reinterpret_cast(readerStream)); + cache->close(); + batch.publish(batchStream()); + CachedCudaEvent cleared(batchStream()); + EXPECT_FALSE(cleared.queryComplete()); + uploadGate.release(); + EXPECT_FALSE(cleared.queryComplete()); + readerGate.release(); + cleared.synchronize(); + auto const* oldIndices = reinterpret_cast(readback.address()); + EXPECT_EQ(std::vector(oldIndices, oldIndices + 2), expected); + EXPECT_EQ(read(batch.pageTableAddress(LifeCycleId{0}), 2), (std::vector{-1, -1})); +} + +TEST_F(KvCacheManagerV2BatchTest, DeviceAddressesSurviveGraphReplayAcrossPublication) +{ + auto manager = std::make_shared(sparseConfig()); + Batch batch(manager, 1, 2); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(4, 0)); + batch.add(*cache); + batch.publish(batchStream()); + batch.waitReady(batchStream()); + auto const table = batch.pageTableAddress(LifeCycleId{0}); + void* output = nullptr; + ASSERT_EQ(cudaMalloc(&output, 2 * sizeof(int32_t)), cudaSuccess); + auto freeOutput = FuncGuard([&]() { cudaFree(output); }); + cudaGraph_t graph{}; + cudaGraphExec_t executable{}; + auto destroyGraph = FuncGuard( + [&]() + { + if (executable) + { + cudaGraphExecDestroy(executable); + } + if (graph) + { + cudaGraphDestroy(graph); + } + }); + ASSERT_EQ(cudaStreamBeginCapture(mStream, cudaStreamCaptureModeThreadLocal), cudaSuccess); + EXPECT_THROW(batch.publish(batchStream()), LogicError); + ASSERT_EQ(cudaMemcpyAsync( + output, reinterpret_cast(table), 2 * sizeof(int32_t), cudaMemcpyDeviceToDevice, mStream), + cudaSuccess); + ASSERT_EQ(cudaStreamEndCapture(mStream, &graph), cudaSuccess); + ASSERT_EQ(cudaGraphInstantiateWithFlags(&executable, graph, 0), cudaSuccess); + for (bool active : {true, false, true}) + { + if (active && !cache->isActive()) + { + ASSERT_TRUE(cache->resume()); + } + if (!active) + { + cache->suspend(); + } + batch.publish(batchStream()); + batch.waitReady(batchStream()); + ASSERT_EQ(cudaGraphLaunch(executable, mStream), cudaSuccess); + batch.recordRead(batchStream()); + auto expected = std::vector{-1, -1}; + if (active) + { + expected[0] = cache->getBasePageIndices(LifeCycleId{0})[0]; + } + EXPECT_EQ(read(reinterpret_cast(output), 2), expected); + EXPECT_EQ(batch.pageTableAddress(LifeCycleId{0}), table); + } +} + +TEST_F(KvCacheManagerV2BatchTest, PublicationWaitsForOffloadAndReaderFencesProtectHostSlots) +{ + for (bool commit : {false, true}) + { + SCOPED_TRACE(commit); + auto manager = std::make_shared(sparseConfig()); + Batch batch(manager, 1, 1); + auto const apiLock = manager->lockExclusive(); + auto& storage = manager->storage(); + auto const lc = LifeCycleId{0}; + auto const gpuPool = storage.getPoolGroupIndex(kHotLevel, lc); + auto const hostPool = storage.getPoolGroupIndex(kSparseHistoryLevel, lc); + size_t const bytes = storage.slotSize(kHotLevel, gpuPool)[PoolIndex{0}]; + cudaStream_t readerStream{}; + ASSERT_EQ(cudaStreamCreateWithFlags(&readerStream, cudaStreamNonBlocking), cudaSuccess); + auto destroyReader = FuncGuard([&]() { cudaStreamDestroy(readerStream); }); + void* gpuReadback = nullptr; + void* hostReadback = nullptr; + ASSERT_EQ(cudaMalloc(&gpuReadback, bytes), cudaSuccess); + auto freeGpuReadback = FuncGuard([&]() { cudaFree(gpuReadback); }); + ASSERT_EQ(cudaMallocHost(&hostReadback, bytes), cudaSuccess); + auto freeHostReadback = FuncGuard([&]() { cudaFreeHost(hostReadback); }); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(4, 4)); + auto const gpuSlot = pageAt(*cache)->slotId(); + auto const gpuAddress = std::get(storage.slotAddress(kHotLevel, gpuPool, gpuSlot, PoolIndex{0})); + constexpr uint8_t kPattern = 0xD3; + ASSERT_EQ(cudaMemsetAsync(reinterpret_cast(gpuAddress), kPattern, bytes, mStream), cudaSuccess); + auto blocker = storage.newSlots(kSparseHistoryLevel, TypedVec{1}); + auto releaseBlocker + = FuncGuard([&]() { storage.releaseSlot(lc, kSparseHistoryLevel, std::move(blocker[lc].front())); }); + // Warm transfer staging before deliberately delaying the copy. + storage.copySlotData(lc, kSparseHistoryLevel, kHotLevel, blocker[lc].front().slotId(), gpuSlot, stream()); + batch.add(*cache); + batch.publish(batchStream()); + ASSERT_EQ(cudaStreamSynchronize(mStream), cudaSuccess); + StreamGate copyGate; + StreamGate readerGate; + auto releaseGates = FuncGuard( + [&]() + { + copyGate.release(); + readerGate.release(); + }); + ASSERT_EQ(copyGate.enqueue(mStream), cudaSuccess); + ASSERT_TRUE(cache->enterDecode()); + auto const snapshot = cache->getPageStorageSnapshot(lc); + ASSERT_EQ(snapshot.eligibleHistoryBlocks(), 1); + ASSERT_EQ(snapshot.readyEvents().size(), 1); + EXPECT_FALSE(snapshot.readyEvents().front().queryComplete()); + batch.publish(reinterpret_cast(readerStream)); + batch.waitReady(reinterpret_cast(readerStream)); + CachedCudaEvent published(reinterpret_cast(readerStream)); + EXPECT_FALSE(published.queryComplete()); + ASSERT_EQ(readerGate.enqueue(readerStream), cudaSuccess); + auto const hostSlot = pageAt(*cache)->slotId(); + auto const hostAddress + = std::get(storage.slotAddress(kSparseHistoryLevel, hostPool, hostSlot, PoolIndex{0})); + ASSERT_EQ(cudaMemcpyAsync(gpuReadback, reinterpret_cast(hostAddress), bytes, + cudaMemcpyHostToDevice, readerStream), + cudaSuccess); + ASSERT_EQ(cudaMemcpyAsync(hostReadback, gpuReadback, bytes, cudaMemcpyDeviceToHost, readerStream), cudaSuccess); + batch.recordRead(reinterpret_cast(readerStream)); + if (commit) + cache->commit(tokens()); + cache->close(); + batch.publish(batchStream()); + manager->clearReusableBlocks(); + auto recycled = storage.newSlots(kSparseHistoryLevel, TypedVec{1}); + auto releaseRecycled + = FuncGuard([&]() { storage.releaseSlot(lc, kSparseHistoryLevel, std::move(recycled[lc].front())); }); + EXPECT_EQ(recycled[lc].front().slotId(), hostSlot); + EXPECT_FALSE(recycled[lc].front().queryReady()); + copyGate.release(); + published.synchronize(); + EXPECT_FALSE(recycled[lc].front().queryReady()); + readerGate.release(); + recycled[lc].front().readyEvent.synchronize(); + EXPECT_EQ(read(batch.pageTableAddress(lc), 1), (std::vector{-1})); + EXPECT_EQ(read(batch.numBlocksAddress(lc), 1), (std::vector{0})); + auto const* readBytes = static_cast(hostReadback); + EXPECT_TRUE(std::all_of(readBytes, readBytes + bytes, [](uint8_t v) { return v == kPattern; })); + } +} + +TEST_F(KvCacheManagerV2PageStorageTest, QueriesSparseBuffersAndRejectsUnknownBuffers) +{ + auto config = makeSplitColdGroupingConfig(); + auto& sparse = std::get(config.layers[0]); + sparse.buffers.front().isSparse = true; + sparse.buffers.push_back({.role = "value", .size = 2048, .tokensPerBlockOverride = 2, .isSparse = true}); + config.layers.emplace_back(SsmLayerConfig{.layerId = 2, .buffers = {{"state", 4096}}}); + config.commitMinSnapshot = true; + auto manager = std::make_shared(std::move(config)); + EXPECT_TRUE(manager->isSparse(0, "key")); + EXPECT_TRUE(manager->isSparse(0, "value")); + EXPECT_FALSE(manager->isSparse(1, "key")); + EXPECT_FALSE(manager->isSparse(2, "state")); + EXPECT_THROW(manager->isSparse(0, "missing"), std::out_of_range); + EXPECT_THROW(manager->isSparse(3, "key"), std::out_of_range); +} + +TEST_F(KvCacheManagerV2PageStorageTest, SnapshotsRawMixedTierIndicesAndDecodeEligibility) +{ + auto config = makeSplitColdGroupingConfig(); + auto& sparse = std::get(config.layers[0]); + sparse.buffers.front().isSparse = true; + sparse.buffers.push_back({.role = "value", .size = 2048, .tokensPerBlockOverride = 2, .isSparse = true}); + auto coalesced = sparse; + coalesced.layerId = 2; + config.layers.emplace_back(std::move(coalesced)); + auto manager = std::make_shared(std::move(config)); + auto const apiLock = manager->lockExclusive(); + auto const sparseGroup = manager->getLayerGroupId(0); + auto const denseGroup = manager->getLayerGroupId(1); + EXPECT_EQ(manager->getLayerGroupId(2), sparseGroup); + EXPECT_GT(manager->getPageIndexScale(0, "value"), 1); + auto cache = manager->createKvCache(); + std::vector externalIndices(8, kBadPageIndex.value()); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(12, 6)); + cache->setBasePageIndexBuf(kDefaultBeamIndex, sparseGroup, externalIndices.data(), externalIndices.size()); + auto const prefill = cache->getPageStorageSnapshot(sparseGroup); + EXPECT_EQ(prefill.eligibleHistoryBlocks(), 0); + EXPECT_EQ(prefill.cacheLevels(), (std::vector>(3, kHotLevel))); + + ASSERT_TRUE(cache->enterDecode()); + auto const decode = cache->getPageStorageSnapshot(sparseGroup); + EXPECT_EQ(decode.eligibleHistoryBlocks(), 1); + ASSERT_EQ(decode.basePageIndices().size(), 3); + EXPECT_EQ(decode.basePageIndices(), (std::vector(externalIndices.begin(), externalIndices.begin() + 3))); + EXPECT_EQ( + decode.cacheLevels(), (std::vector>{kSparseHistoryLevel, kHotLevel, kHotLevel})); + EXPECT_EQ(cache->getPageStorageSnapshot(denseGroup).eligibleHistoryBlocks(), 0); + EXPECT_EQ(prefill.cacheLevels()[0], kHotLevel); + EXPECT_GT(decode.version(), prefill.version()); + for (int ord = 0; ord < 3; ++ord) + EXPECT_EQ(decode.basePageIndices()[ord], slotIdToPageIndexValue(pageAt(*cache, ord, sparseGroup)->slotId())); + + cache->suspend(); + auto const suspended = cache->getPageStorageSnapshot(sparseGroup); + EXPECT_EQ(suspended.eligibleHistoryBlocks(), 0); + EXPECT_EQ(suspended.basePageIndices(), (std::vector(3, kBadPageIndex.value()))); + EXPECT_EQ(suspended.cacheLevels(), (std::vector>(3, std::nullopt))); + EXPECT_TRUE(suspended.readyEvents().empty()); + ASSERT_TRUE(cache->resume()); + EXPECT_EQ(cache->getPageStorageSnapshot(sparseGroup).eligibleHistoryBlocks(), 1); + ASSERT_TRUE(cache->resize(12, 8)); + EXPECT_EQ(cache->getPageStorageSnapshot(sparseGroup).eligibleHistoryBlocks(), 2); + EXPECT_EQ(pageAt(*cache, 2, sparseGroup)->cacheLevel, kHotLevel); +} + +TEST_F(KvCacheManagerV2PageStorageTest, DirtyAcknowledgmentTracksBindingsCommitAndRequestLifetime) +{ + auto manager = std::make_shared(sparseConfig()); + auto cache = manager->createKvCache(); + std::vector externalIndices(2, kBadPageIndex.value()); + auto closeCache = FuncGuard([&]() { cache->close(); }); + EXPECT_TRUE(cache->pageStorageDirty()); + EXPECT_FALSE(cache->pageStorageRow().has_value()); + EXPECT_THROW(cache->bindPageStorageRow(-1), LogicError); + cache->bindPageStorageRow(7); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(8, 4)); + auto const prefill = cache->getPageStorageSnapshot(LifeCycleId{0}); + EXPECT_EQ(prefill.row(), 7); + ASSERT_TRUE(cache->acknowledgePageStorage(prefill.version())); + EXPECT_FALSE(cache->pageStorageDirty()); + ASSERT_TRUE(cache->resize(8, 5)); + EXPECT_FALSE(cache->pageStorageDirty()); + ASSERT_TRUE(cache->enterDecode()); + EXPECT_TRUE(cache->pageStorageDirty()); + EXPECT_FALSE(cache->acknowledgePageStorage(prefill.version())); + auto const decode = cache->getPageStorageSnapshot(LifeCycleId{0}); + ASSERT_TRUE(cache->acknowledgePageStorage(decode.version())); + ASSERT_TRUE(cache->enterDecode()); + EXPECT_FALSE(cache->pageStorageDirty()); + + cache->bindPageStorageRow(7); + EXPECT_TRUE(cache->pageStorageDirty()); + EXPECT_FALSE(cache->acknowledgePageStorage(decode.version())); + cache->bindPageStorageRow(9); + cache->setBasePageIndexBuf(kDefaultBeamIndex, LifeCycleId{0}, externalIndices.data(), externalIndices.size()); + auto const rebound = cache->getPageStorageSnapshot(LifeCycleId{0}); + EXPECT_EQ(rebound.row(), 9); + EXPECT_EQ(rebound.basePageIndices(), decode.basePageIndices()); + ASSERT_TRUE(cache->acknowledgePageStorage(rebound.version())); + cache->setBasePageIndexBuf(kDefaultBeamIndex, LifeCycleId{0}, nullptr, 0); + EXPECT_TRUE(cache->pageStorageDirty()); + + auto const beforeCommit = cache->getPageStorageSnapshot(LifeCycleId{0}); + ASSERT_TRUE(cache->acknowledgePageStorage(beforeCommit.version())); + cache->commit(tokens()); + EXPECT_TRUE(cache->pageStorageDirty()); + auto const committed = cache->getPageStorageSnapshot(LifeCycleId{0}); + EXPECT_EQ(committed.basePageIndices(), beforeCommit.basePageIndices()); + EXPECT_EQ(committed.cacheLevels(), beforeCommit.cacheLevels()); + ASSERT_TRUE(cache->acknowledgePageStorage(committed.version())); + cache->suspend(); + EXPECT_TRUE(cache->pageStorageDirty()); + EXPECT_EQ(cache->pageStorageRow(), 9); + ASSERT_TRUE(cache->resume()); + EXPECT_EQ(cache->getPageStorageSnapshot(LifeCycleId{0}).eligibleHistoryBlocks(), 1); + cache->close(); + EXPECT_TRUE(cache->pageStorageDirty()); + EXPECT_FALSE(cache->pageStorageRow().has_value()); + auto const closed = cache->getPageStorageSnapshot(LifeCycleId{0}); + EXPECT_TRUE(closed.basePageIndices().empty()); + EXPECT_EQ(closed.eligibleHistoryBlocks(), 0); + EXPECT_THROW(cache->bindPageStorageRow(9), LogicError); + EXPECT_THROW(cache->recordPageStorageRead(reinterpret_cast(mStream)), LogicError); + + auto reused = manager->createKvCache({}, tokens()); + auto closeReused = FuncGuard([&]() { reused->close(); }); + EXPECT_TRUE(reused->pageStorageDirty()); + EXPECT_FALSE(reused->pageStorageRow().has_value()); + EXPECT_EQ(reused->getPageStorageSnapshot(LifeCycleId{0}).basePageIndices(), (std::vector{-1})); + ASSERT_TRUE(reused->resume(stream(), true)); + EXPECT_EQ(reused->getPageStorageSnapshot(LifeCycleId{0}).eligibleHistoryBlocks(), 1); +} + +TEST_F(KvCacheManagerV2PageStorageTest, ReadinessAndReaderFencesSurviveCommitAndClose) +{ + for (bool commit : {false, true}) + { + SCOPED_TRACE(commit); + auto manager = std::make_shared(sparseConfig()); + auto const apiLock = manager->lockExclusive(); + auto& storage = manager->storage(); + auto const lc = LifeCycleId{0}; + auto const gpuPool = storage.getPoolGroupIndex(kHotLevel, lc); + auto const hostPool = storage.getPoolGroupIndex(kSparseHistoryLevel, lc); + size_t const bytes = storage.slotSize(kHotLevel, gpuPool)[PoolIndex{0}]; + cudaStream_t readerStream{}; + ASSERT_EQ(cudaStreamCreateWithFlags(&readerStream, cudaStreamNonBlocking), cudaSuccess); + auto destroyReader = FuncGuard([&]() { cudaStreamDestroy(readerStream); }); + void* gpuReadback = nullptr; + void* hostReadback = nullptr; + ASSERT_EQ(cudaMalloc(&gpuReadback, bytes), cudaSuccess); + auto freeGpuReadback = FuncGuard([&]() { cudaFree(gpuReadback); }); + ASSERT_EQ(cudaMallocHost(&hostReadback, bytes), cudaSuccess); + auto freeHostReadback = FuncGuard([&]() { cudaFreeHost(hostReadback); }); + auto cache = manager->createKvCache(); + auto closeCache = FuncGuard([&]() { cache->close(); }); + ASSERT_TRUE(cache->resume(stream())); + ASSERT_TRUE(cache->resize(4, 4)); + auto const gpuSlot = pageAt(*cache)->slotId(); + auto const gpuAddress = std::get(storage.slotAddress(kHotLevel, gpuPool, gpuSlot, PoolIndex{0})); + constexpr uint8_t kPattern = 0xD3; + ASSERT_EQ(cudaMemsetAsync(reinterpret_cast(gpuAddress), kPattern, bytes, mStream), cudaSuccess); + auto blocker = storage.newSlots(kSparseHistoryLevel, TypedVec{1}); + auto releaseBlocker + = FuncGuard([&]() { storage.releaseSlot(lc, kSparseHistoryLevel, std::move(blocker[lc].front())); }); + // Warm transfer staging before deliberately delaying the copy. + storage.copySlotData(lc, kSparseHistoryLevel, kHotLevel, blocker[lc].front().slotId(), gpuSlot, stream()); + ASSERT_EQ(cudaStreamSynchronize(mStream), cudaSuccess); + StreamGate copyGate; + StreamGate readerGate; + auto releaseGates = FuncGuard( + [&]() + { + copyGate.release(); + readerGate.release(); + }); + ASSERT_EQ(copyGate.enqueue(mStream), cudaSuccess); + ASSERT_TRUE(cache->enterDecode()); + auto const snapshot = cache->getPageStorageSnapshot(lc); + ASSERT_EQ(snapshot.eligibleHistoryBlocks(), 1); + ASSERT_EQ(snapshot.readyEvents().size(), 1); + EXPECT_FALSE(snapshot.readyEvents().front().queryComplete()); + snapshot.waitReady(reinterpret_cast(readerStream)); + ASSERT_EQ(readerGate.enqueue(readerStream), cudaSuccess); + auto const hostSlot = pageAt(*cache)->slotId(); + auto const hostAddress + = std::get(storage.slotAddress(kSparseHistoryLevel, hostPool, hostSlot, PoolIndex{0})); + ASSERT_EQ(cudaMemcpyAsync(gpuReadback, reinterpret_cast(hostAddress), bytes, + cudaMemcpyHostToDevice, readerStream), + cudaSuccess); + ASSERT_EQ(cudaMemcpyAsync(hostReadback, gpuReadback, bytes, cudaMemcpyDeviceToHost, readerStream), cudaSuccess); + cache->recordPageStorageRead(reinterpret_cast(readerStream)); + if (commit) + cache->commit(tokens()); + cache->close(); + manager->clearReusableBlocks(); + auto recycled = storage.newSlots(kSparseHistoryLevel, TypedVec{1}); + auto releaseRecycled + = FuncGuard([&]() { storage.releaseSlot(lc, kSparseHistoryLevel, std::move(recycled[lc].front())); }); + EXPECT_EQ(recycled[lc].front().slotId(), hostSlot); + EXPECT_FALSE(recycled[lc].front().queryReady()); + copyGate.release(); + snapshot.readyEvents().front().synchronize(); + EXPECT_FALSE(recycled[lc].front().queryReady()); + readerGate.release(); + recycled[lc].front().readyEvent.synchronize(); + auto const* readBytes = static_cast(hostReadback); + EXPECT_TRUE(std::all_of(readBytes, readBytes + bytes, [](uint8_t v) { return v == kPattern; })); + } +} + TEST_F(KvCacheManagerV2SparseOffloadTest, BatchesCompleteCoalescedPagesAndCountsPhysicalCopies) { auto config = makeSplitColdGroupingConfig(); @@ -902,13 +1590,14 @@ TEST_F(KvCacheManagerV2SparseOffloadTest, BatchesCompleteCoalescedPagesAndCounts auto const hostFree = storage.getStatistics(kSparseHistoryLevel).free; auto targets = pages; targets.push_back(pages.front()); + auto const versionBeforeOffload = cache->pageStorageVersion(); cache->offloadSparsePages(targets); EXPECT_EQ(observer->encodeCalls, 1); EXPECT_EQ(observer->encodedPages, pages.size()); EXPECT_EQ(observer->encodeStream, mStream); EXPECT_EQ(storage.getStatistics(kHotLevel).free, gpuFree + pages.size()); EXPECT_EQ(storage.getStatistics(kSparseHistoryLevel).free, hostFree - pages.size()); - EXPECT_EQ(cache->pageStorageVersion(), pages.size()); + EXPECT_GT(cache->pageStorageVersion(), versionBeforeOffload); for (size_t i = 0; i < pages.size(); ++i) { @@ -937,9 +1626,10 @@ TEST_F(KvCacheManagerV2SparseOffloadTest, BatchesCompleteCoalescedPagesAndCounts EXPECT_EQ(stats.at(lc).iterOffloadBlocks, 2); EXPECT_EQ(stats.at(lc).iterOffloadBytes, 2 * expected.front().size()); } + auto const versionAfterOffload = cache->pageStorageVersion(); cache->offloadSparsePages(targets); EXPECT_EQ(observer->encodeCalls, 1); - EXPECT_EQ(cache->pageStorageVersion(), pages.size()); + EXPECT_EQ(cache->pageStorageVersion(), versionAfterOffload); EXPECT_TRUE(manager->getAndResetIterationStats().empty()); } @@ -967,12 +1657,22 @@ TEST_F(KvCacheManagerV2SparseOffloadTest, SharedOwnersPublishHostIndicesAndKeepH { storage.releaseSlot(LifeCycleId{0}, kSparseHistoryLevel, std::move(hostBlockers[LifeCycleId{0}].front())); }); EXPECT_THROW(storage.batchedMigrate(kSparseHistoryLevel, {page}, {}), LogicError); + auto const firstVersion = first->pageStorageVersion(); + auto const secondVersion = second->pageStorageVersion(); + ASSERT_TRUE(first->acknowledgePageStorage(firstVersion)); + ASSERT_TRUE(second->acknowledgePageStorage(secondVersion)); first->offloadSparsePages({page, page}); EXPECT_NE(page->slotId(), gpuSlot); EXPECT_EQ(first->getBasePageIndices(LifeCycleId{0})[0], slotIdToPageIndexValue(page->slotId())); EXPECT_EQ(externalIndices[0], slotIdToPageIndexValue(page->slotId())); - EXPECT_EQ(first->pageStorageVersion(), 1); - EXPECT_EQ(second->pageStorageVersion(), 1); + EXPECT_GT(first->pageStorageVersion(), firstVersion); + EXPECT_GT(second->pageStorageVersion(), secondVersion); + EXPECT_TRUE(first->pageStorageDirty()); + EXPECT_TRUE(second->pageStorageDirty()); + auto const snapshot = second->getPageStorageSnapshot(LifeCycleId{0}); + EXPECT_EQ(snapshot.basePageIndices(), externalIndices); + EXPECT_EQ(snapshot.cacheLevels()[0], kSparseHistoryLevel); + EXPECT_EQ(snapshot.eligibleHistoryBlocks(), 0); EXPECT_EQ(storage.getStatistics(kHotLevel).free, storage.getStatistics(kHotLevel).total); EXPECT_FALSE(storage.isEvictable(*page)); EXPECT_THROW(storage.batchedMigrate(kHotLevel, {page}, {}), LogicError); @@ -1079,6 +1779,7 @@ TEST_F(KvCacheManagerV2SparseOffloadTest, HostOomLeavesEntireBatchOnGpu) auto second = pageAt(*cache, 1); auto const firstSlot = first->slotId(); auto const secondSlot = second->slotId(); + auto const version = cache->pageStorageVersion(); EXPECT_THROW(cache->offloadSparsePages({first, second}), OutOfPagesError); EXPECT_EQ(observer->encodeCalls, 0); EXPECT_EQ(first->cacheLevel, kHotLevel); @@ -1087,7 +1788,7 @@ TEST_F(KvCacheManagerV2SparseOffloadTest, HostOomLeavesEntireBatchOnGpu) EXPECT_EQ(second->slotId(), secondSlot); EXPECT_EQ(first->status(), PageStatus::LOCKED); EXPECT_EQ(second->status(), PageStatus::LOCKED); - EXPECT_EQ(cache->pageStorageVersion(), 0); + EXPECT_EQ(cache->pageStorageVersion(), version); EXPECT_EQ(storage.getStatistics(kSparseHistoryLevel).free, 1); EXPECT_EQ(storage.getStatistics(kHotLevel).free, 0); EXPECT_EQ(cache->getBasePageIndices(LifeCycleId{0})[0], slotIdToPageIndexValue(firstSlot)); @@ -1114,13 +1815,14 @@ TEST_F(KvCacheManagerV2SparseOffloadTest, AsynchronousRejectionFencesBothSlotsWi auto releaseBlocker = FuncGuard([&]() { storage.releaseSlot(lc, kSparseHistoryLevel, std::move(blocker[lc].front())); }); auto releaseCodec = FuncGuard([&]() { rejecting->release(); }); + auto const version = cache->pageStorageVersion(); EXPECT_THROW(cache->offloadSparsePages({page}), TllmException); ASSERT_TRUE(rejecting->launched()); EXPECT_EQ(page->cacheLevel, kHotLevel); EXPECT_EQ(page->slotId(), gpuSlot); EXPECT_EQ(page->status(), PageStatus::LOCKED); EXPECT_FALSE(page->queryReady()); - EXPECT_EQ(cache->pageStorageVersion(), 0); + EXPECT_GT(cache->pageStorageVersion(), version); EXPECT_EQ(cache->getBasePageIndices(lc)[0], slotIdToPageIndexValue(gpuSlot)); auto recycled = storage.newSlots(kSparseHistoryLevel, TypedVec{1}); auto releaseRecycled @@ -1154,13 +1856,14 @@ TEST_F(KvCacheManagerV2SparseOffloadTest, LaterCodecBatchFailurePreservesAllSour auto second = pageAt(*cache, 0, LifeCycleId{1}); auto const firstSlot = first->slotId(); auto const secondSlot = second->slotId(); + auto const version = cache->pageStorageVersion(); EXPECT_THROW(cache->offloadSparsePages({first, second}), TllmException); EXPECT_EQ(observer->encodeCalls, 2); EXPECT_EQ(first->cacheLevel, kHotLevel); EXPECT_EQ(second->cacheLevel, kHotLevel); EXPECT_EQ(first->slotId(), firstSlot); EXPECT_EQ(second->slotId(), secondSlot); - EXPECT_EQ(cache->pageStorageVersion(), 0); + EXPECT_GT(cache->pageStorageVersion(), version); EXPECT_EQ(storage.getStatistics(kSparseHistoryLevel).free, storage.getStatistics(kSparseHistoryLevel).total); observer->rejectEncodeCall = 0; cache->offloadSparsePages({first, second}); @@ -1187,11 +1890,12 @@ TEST_F(KvCacheManagerV2SparseOffloadTest, RejectsWritablePartialAndDensePagesBef ASSERT_TRUE(cache->resize(8, history)); auto first = pageAt(*cache); auto input = pageAt(*cache, 1); + auto const version = cache->pageStorageVersion(); EXPECT_THROW(cache->offloadSparsePages({first, input}), LogicError); EXPECT_EQ(observer->encodeCalls, 0); EXPECT_EQ(first->cacheLevel, kHotLevel); EXPECT_EQ(input->cacheLevel, kHotLevel); - EXPECT_EQ(cache->pageStorageVersion(), 0); + EXPECT_EQ(cache->pageStorageVersion(), version); } } } @@ -1242,9 +1946,10 @@ TEST_F(KvCacheManagerV2SparseOffloadTest, EmitsCommittedTierChangeEvenWhenSlotIn events->flushIterationEvents(); events->getLatestEvents(/*timeoutMs=*/0); + auto const version = cache->pageStorageVersion(); cache->offloadSparsePages({page, page}); EXPECT_EQ(cache->getBasePageIndices(LifeCycleId{0})[0], originalIndex); - EXPECT_EQ(cache->pageStorageVersion(), 1); + EXPECT_GT(cache->pageStorageVersion(), version); events->flushIterationEvents(); auto const updates = events->getLatestEvents(/*timeoutMs=*/0); ASSERT_EQ(updates.size(), 1); @@ -1303,7 +2008,7 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, PrefillRetainsSparseSwaHistoryAndResto EXPECT_FALSE(cache->hasScratchSlots()); ASSERT_TRUE(cache->resize(12, 8)); EXPECT_FALSE(cache->isDecoding()); - EXPECT_EQ(cache->pageStorageVersion(), 0); + EXPECT_EQ(cache->getPageStorageSnapshot(LifeCycleId{0}).eligibleHistoryBlocks(), 0); for (int ord = 0; ord < 3; ++ord) { ASSERT_NE(pageAt(*cache, ord), nullptr); @@ -1323,7 +2028,7 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, PrefillRetainsSparseSwaHistoryAndResto EXPECT_EQ(pageAt(*cache, ord)->cacheLevel, kHotLevel); } ASSERT_TRUE(cache->resize(12, 12)); - EXPECT_EQ(cache->pageStorageVersion(), 0); + EXPECT_EQ(cache->getPageStorageSnapshot(LifeCycleId{0}).eligibleHistoryBlocks(), 0); ASSERT_TRUE(cache->enterDecode()); for (int ord = 0; ord < 3; ++ord) { @@ -1343,21 +2048,23 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, EntryScansUnchangedWatermarkAndResumeK ASSERT_TRUE(cache->resume(stream())); ASSERT_TRUE(cache->resize(8, 8)); EXPECT_EQ(observer->encodeCalls, 0); + auto const prefillVersion = cache->pageStorageVersion(); ASSERT_TRUE(cache->enterDecode()); EXPECT_TRUE(cache->isDecoding()); EXPECT_EQ(cache->historyLength(), 8); EXPECT_EQ(observer->encodeCalls, 1); EXPECT_EQ(observer->encodedPages, 2); - EXPECT_EQ(cache->pageStorageVersion(), 3); + EXPECT_GT(cache->pageStorageVersion(), prefillVersion); + auto const decodeVersion = cache->pageStorageVersion(); ASSERT_TRUE(cache->enterDecode()); EXPECT_EQ(observer->encodeCalls, 1); - EXPECT_EQ(cache->pageStorageVersion(), 3); + EXPECT_EQ(cache->pageStorageVersion(), decodeVersion); EXPECT_THROW(cache->resize(8, 4), std::invalid_argument); cache->suspend(); EXPECT_THROW(cache->resume(std::nullopt, false), std::invalid_argument); ASSERT_TRUE(cache->resume()); EXPECT_EQ(observer->encodeCalls, 1); - EXPECT_EQ(cache->pageStorageVersion(), 4); + EXPECT_GT(cache->pageStorageVersion(), decodeVersion); for (int ord = 0; ord < 2; ++ord) { auto const page = pageAt(*cache, ord); @@ -1406,7 +2113,7 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, CachedPrefillRestoresGpuOnResumePrefet EXPECT_EQ(page->cacheLevel, kHotLevel); EXPECT_EQ(cache->historyLength(), 4); EXPECT_FALSE(cache->isDecoding()); - EXPECT_EQ(cache->pageStorageVersion(), 0); + EXPECT_EQ(cache->getPageStorageSnapshot(LifeCycleId{0}).eligibleHistoryBlocks(), 0); EXPECT_EQ(cache->getBasePageIndices(LifeCycleId{0})[0], slotIdToPageIndexValue(page->slotId())); } } @@ -1435,6 +2142,8 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, MixedDenseSwaLocksAreRestoredAfterHost EXPECT_EQ(pageAt(*cache)->cacheLevel, kSparseHistoryLevel); EXPECT_EQ(pageAt(*cache, 1)->cacheLevel, kHotLevel); auto const version = cache->pageStorageVersion(); + auto const before = cache->getPageStorageSnapshot(LifeCycleId{0}); + ASSERT_TRUE(cache->acknowledgePageStorage(version)); auto const freeHost = manager->getStorageStatistics(kSparseHistoryLevel) .at(storage.getPoolGroupIndex(kSparseHistoryLevel, LifeCycleId{0})) .free; @@ -1447,7 +2156,13 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, MixedDenseSwaLocksAreRestoredAfterHost }); EXPECT_FALSE(cache->resize(8, 8)); EXPECT_EQ(cache->historyLength(), 4); - EXPECT_EQ(cache->pageStorageVersion(), version); + EXPECT_GT(cache->pageStorageVersion(), version); + EXPECT_TRUE(cache->pageStorageDirty()); + EXPECT_FALSE(cache->acknowledgePageStorage(version)); + auto const after = cache->getPageStorageSnapshot(LifeCycleId{0}); + EXPECT_EQ(after.basePageIndices(), before.basePageIndices()); + EXPECT_EQ(after.cacheLevels(), before.cacheLevels()); + EXPECT_EQ(after.eligibleHistoryBlocks(), before.eligibleHistoryBlocks()); EXPECT_EQ(pageAt(*cache, 0, LifeCycleId{1}), densePage); EXPECT_EQ(densePage->status(), PageStatus::LOCKED); EXPECT_EQ(densePage->cacheLevel, kHotLevel); @@ -1476,7 +2191,7 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, HistoryUpdatesOffloadOnlyNewFullPages) ASSERT_TRUE(cache->resize(8, 8)); EXPECT_EQ(observer->encodeCalls, 2); EXPECT_EQ(observer->encodedPages, 2); - EXPECT_EQ(cache->pageStorageVersion(), version + 2); + EXPECT_GT(cache->pageStorageVersion(), version); EXPECT_EQ(pageAt(*cache)->slotId(), firstHostSlot); EXPECT_EQ(pageAt(*cache, 1)->cacheLevel, kSparseHistoryLevel); ASSERT_TRUE(cache->resize(12)); @@ -1500,9 +2215,10 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, SharedPrefillOwnerBlocksDemotionAndRes }); ASSERT_TRUE(first->resume(stream())); ASSERT_TRUE(second->resume(stream())); + auto const version = first->pageStorageVersion(); EXPECT_THROW(first->enterDecode(), LogicError); EXPECT_FALSE(first->isDecoding()); - EXPECT_EQ(first->pageStorageVersion(), 0); + EXPECT_EQ(first->pageStorageVersion(), version); EXPECT_EQ(page->cacheLevel, kHotLevel); second->suspend(); ASSERT_TRUE(first->enterDecode()); @@ -1555,7 +2271,7 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, HostOomDoesNotAdmitDecodeAndCanRetry) EXPECT_FALSE(cache->isDecoding()); EXPECT_EQ(cache->historyLength(), history); EXPECT_EQ(cache->capacity(), history); - EXPECT_EQ(cache->pageStorageVersion(), 0); + EXPECT_EQ(cache->getPageStorageSnapshot(LifeCycleId{0}).eligibleHistoryBlocks(), 0); EXPECT_EQ(manager->getAndResetIterationSuspendResumeStats(), (std::pair{0, 0})); for (int ord = 0; ord < history / 4; ++ord) { @@ -1617,7 +2333,7 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, HistoryOomRollsBackShortcutGrowthAndSh EXPECT_EQ(cache->capacity(), newCapacity); EXPECT_EQ(cache->historyLength(), 8); EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); - EXPECT_EQ(cache->pageStorageVersion(), version + 2); + EXPECT_GT(cache->pageStorageVersion(), version); } } @@ -1638,6 +2354,7 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, CodecRejectionDoesNotPublishEntryOrHis if (!entry) ASSERT_TRUE(cache->enterDecode()); auto const version = cache->pageStorageVersion(); + auto const before = cache->getPageStorageSnapshot(LifeCycleId{0}); auto const page = pageAt(*cache, 1); auto const gpuSlot = page->slotId(); observer->rejectEncodeCall = observer->encodeCalls + 1; @@ -1645,7 +2362,11 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, CodecRejectionDoesNotPublishEntryOrHis EXPECT_EQ(cache->isDecoding(), !entry); EXPECT_EQ(cache->capacity(), 8); EXPECT_EQ(cache->historyLength(), history); - EXPECT_EQ(cache->pageStorageVersion(), version); + EXPECT_GT(cache->pageStorageVersion(), version); + auto const after = cache->getPageStorageSnapshot(LifeCycleId{0}); + EXPECT_EQ(after.basePageIndices(), before.basePageIndices()); + EXPECT_EQ(after.cacheLevels(), before.cacheLevels()); + EXPECT_EQ(after.eligibleHistoryBlocks(), before.eligibleHistoryBlocks()); EXPECT_EQ(page->cacheLevel, kHotLevel); EXPECT_EQ(page->slotId(), gpuSlot); EXPECT_EQ(cache->getBasePageIndices(LifeCycleId{0})[1], slotIdToPageIndexValue(gpuSlot)); @@ -1671,7 +2392,7 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, RebaseOffloadsOlderGpuHistoryAtUnchang EXPECT_EQ(pageAt(*cache), page); EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); EXPECT_EQ(cache->historyLength(), 4); - EXPECT_EQ(cache->pageStorageVersion(), version + 1); + EXPECT_GT(cache->pageStorageVersion(), version); } TEST_F(KvCacheManagerV2DecodeOffloadTest, PrefillRebaseCannotAdoptSharedHostIndices) @@ -1720,6 +2441,8 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, RebaseOomPreservesMissingLifecyclePage auto first = pageAt(*cache, 0, LifeCycleId{0}); auto second = pageAt(*cache, 0, LifeCycleId{1}); auto const version = cache->pageStorageVersion(); + auto const before = cache->getPageStorageSnapshot(LifeCycleId{0}); + ASSERT_TRUE(cache->acknowledgePageStorage(version)); auto const pool = storage.getPoolGroupIndex(kSparseHistoryLevel, LifeCycleId{1}); auto const freeHost = manager->getStorageStatistics(kSparseHistoryLevel).at(pool).free; auto blockers = storage.newSlots(kSparseHistoryLevel, TypedVec{0, freeHost}); @@ -1733,7 +2456,13 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, RebaseOomPreservesMissingLifecyclePage EXPECT_EQ(cache->numCommittedTokens(), 0); EXPECT_EQ(cache->numCommittedBlocks(), 0); EXPECT_EQ(cache->historyLength(), 4); - EXPECT_EQ(cache->pageStorageVersion(), version); + EXPECT_GT(cache->pageStorageVersion(), version); + EXPECT_TRUE(cache->pageStorageDirty()); + EXPECT_FALSE(cache->acknowledgePageStorage(version)); + auto const after = cache->getPageStorageSnapshot(LifeCycleId{0}); + EXPECT_EQ(after.basePageIndices(), before.basePageIndices()); + EXPECT_EQ(after.cacheLevels(), before.cacheLevels()); + EXPECT_EQ(after.eligibleHistoryBlocks(), before.eligibleHistoryBlocks()); EXPECT_EQ(pageAt(*cache, 0, LifeCycleId{0}), first); EXPECT_EQ(pageAt(*cache, 0, LifeCycleId{1}), second); for (auto const& page : {first, second}) 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 b3ea320cfe11..f8fbda1a5fe3 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 @@ -56,6 +56,7 @@ GPU_LEVEL, AttentionLayerConfig, AttnLifeCycle, + Batch, BatchDesc, BufferConfig, CacheLevel, @@ -1105,6 +1106,8 @@ def _settle_context_cursor(req: LlmRequest, reuse: int, tokens_per_block: int) - class KVCacheManagerV2(BaseResourceManager): + # Sparse managers attach a metadata batch during initialization. + sparse_metadata_batch: Batch | None = None # Filled lazily by _cold_pool_group_membership(); the grouping is fixed after construction. # Declared on the class so it is present even when an instance is built without running __init__. _cold_pool_group_membership_cache: Optional[tuple[tuple[int, frozenset[int]], ...]] = None @@ -1256,6 +1259,7 @@ def __init__( self._stream = ( execution_stream if execution_stream is not None else torch.cuda.current_stream() ) + self.sparse_metadata_batch: Batch | None = None logger.info(f"[KVCacheManager] execution_stream: {self._stream}") # Materialize an exact per-local-layer vector for cache and attention consumers. @@ -1748,6 +1752,10 @@ def create_cold_page_codec(cache_config: object) -> Optional[object]: self.index_mapper = IndexMapper(index_mapper_capacity, max_beam_width) self._early_freed_index_requests: set[int] = set() self._prepare_page_table_tensor(index_mapper_capacity) + if any(self.impl.is_sparse(buf.layer_id, buf.role) for buf in self.impl.all_buffer_ids): + self.sparse_metadata_batch = Batch( + self.impl, index_mapper_capacity, self.max_blocks_per_seq, self.max_beam_width + ) self._log_kv_cache_pool_lifecycle_mapping() self._reserve_guard_page() @@ -3456,6 +3464,9 @@ def _restore_page_index_bufs(self, request_id: int, kv_cache) -> None: ] kv_cache.set_base_page_index_buf(i, pool_idx, memoryview(buffer.numpy())) + if self.sparse_metadata_batch is not None: + self.sparse_metadata_batch.add(kv_cache, index) + def _resume_and_restore(self, req_id: int, kv_cache) -> bool: """Resume a suspended KV cache and restore its page index buffers. @@ -3899,6 +3910,7 @@ def prepare_resources(self, scheduled_batch: ScheduledRequests): # Mirror the main manager. Under one-model spec decoding the # scheduler may already have created the context cache. self._prepare_draft_resources(scheduled_batch) + self._publish_sparse_metadata() return # KV pages are allocated in `KVCacheV2Scheduler`, so by this point every @@ -3906,6 +3918,18 @@ def prepare_resources(self, scheduled_batch: ScheduledRequests): # That is what makes this the place to drive the connector. if self.kv_connector_manager is not None: self._run_kv_connector_hooks(scheduled_batch) + self._publish_sparse_metadata() + + def _publish_sparse_metadata(self) -> None: + """Refresh stable GPU rows on the execution stream before model work. + + Batch rows match IndexMapper slots, including holes. Sparse consumers use + ``sparse_metadata_batch`` for raw tables and eligible-history counts. + Readers on another stream must use Batch.wait_ready/record_read. + """ + if self.sparse_metadata_batch is not None: + self.sparse_metadata_batch.record_read(self._stream.cuda_stream) + self.sparse_metadata_batch.publish(self._stream.cuda_stream) def _run_kv_connector_hooks(self, scheduled_batch: ScheduledRequests) -> None: """Serve final-batch queries for connectors without source reservations.""" @@ -3961,6 +3985,8 @@ def report_batch_to_connector( if finalize_prefix_reservations and self._connector_reservations_enabled(): self._accept_connector_prefix_reservations(scheduled_batch) self.kv_connector_manager.build_scheduler_output(scheduled_batch, self) + # Connector acceptance can resize requests after prepare_resources. + self._publish_sparse_metadata() # ---- KV connector prefix ---- @@ -5214,6 +5240,9 @@ def release_index_slot(self, request_id: int) -> None: # mirrored, and the target may release the same request twice. return if kv_cache is not None: + if self.sparse_metadata_batch is not None: + self.sparse_metadata_batch.record_read(self._stream.cuda_stream) + self.sparse_metadata_batch.remove(kv_cache) for i in range(self.max_beam_width): for pool_idx in range(self.num_pools): kv_cache.set_base_page_index_buf(i, pool_idx, None) @@ -5503,6 +5532,10 @@ def check_invalid_values_in_kv_cache(self, fill_with_zero: bool = False) -> bool return bool(has_invalid_values) def shutdown(self): + if self.sparse_metadata_batch is not None: + self.sparse_metadata_batch.record_read(self._stream.cuda_stream) + self.sparse_metadata_batch.close() + self.sparse_metadata_batch = None for kv_cache in self.kv_cache_map.values(): kv_cache.close() self.kv_cache_map.clear() @@ -5737,6 +5770,16 @@ def copy_batch_block_offsets( num_seqs: int, max_blocks: Optional[int] = None, ): + self._publish_sparse_metadata() + if self.sparse_metadata_batch is not None and any( + self.kv_cache_map[req_id].is_decoding + and self.kv_cache_map[req_id].history_length >= self.tokens_per_block + for req_id in request_ids + ): + raise RuntimeError( + "Offloaded sparse history requires Batch metadata and sparse fetch; " + "dense attention offsets cannot address host slots" + ) # max_blocks is accepted for signature parity with KVCacheManager; the # device-side copy op here already scales with allocated blocks only. assert beam_width == 1, "beam_width must be 1 for KVCacheManagerV2" @@ -5813,6 +5856,8 @@ def _create_kv_cache( self.impl.mark_stats_excluded(request_id) kv_cache.discard_pending_stats() index = self.index_mapper.add_new_sequence(request_id) + if self.sparse_metadata_batch is not None: + self.sparse_metadata_batch.add(kv_cache, index) for i in range(self.max_beam_width): for pool_idx in range(self.num_pools): buffer: torch.Tensor = self.host_kv_cache_block_offsets[ diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py index a8a71706f7ca..90878bdcff90 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.py @@ -133,11 +133,15 @@ class _KVCacheManagerConfigFieldSpec: KVCacheStoredData = _cpp.KVCacheStoredData KVCacheUpdatedData = _cpp.KVCacheUpdatedData KvCacheStatus = _cpp.KvCacheStatus +LogicError = _cpp.LogicError OutOfMemoryError = _cpp.OutOfMemoryError OutOfPagesError = _cpp.OutOfPagesError PageIndexConverter = _cpp.PageIndexConverter PageIndexMode = _cpp.PageIndexMode PageStatus = _cpp.PageStatus +PageStorageSnapshot = _cpp.PageStorageSnapshot +Batch = _cpp.Batch +BatchDeviceArray = _cpp.BatchDeviceArray PlannedDropHandle = _cpp.PlannedDropHandle PoolDesc = _cpp.PoolDesc PoolGroupDesc = _cpp.PoolGroupDesc @@ -232,6 +236,7 @@ def typed_range(*args: int) -> range: "LayerGroupId", "LayerId", "LifeCycleId", + "LogicError", "MemAddress", "NDEBUG", "OutOfPagesError", @@ -240,6 +245,9 @@ def typed_range(*args: int) -> range: "PoolGroupPeakBlockStats", "PageIndexMode", "PageStatus", + "PageStorageSnapshot", + "Batch", + "BatchDeviceArray", "PoolDesc", "PoolGroupDesc", "PoolGroupIndex", diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi index 8b472c44afec..4b584d61c381 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi @@ -47,6 +47,9 @@ class CuError(Exception): error_code: Any +class LogicError(Exception): + """An operation violates the cache or batch's state and usage requirements.""" + class OutOfMemoryError(Exception): ... class OutOfPagesError(OutOfMemoryError): ... @@ -391,6 +394,83 @@ KvCacheStatus: TypeAlias = _Status IndexSeq = array.array[int] | memoryview[int] +class BatchDeviceArray: + """DLPack view of int32 CUDA metadata; the view keeps its allocation alive. + + Treat these arrays as read-only. Convert with ``torch.from_dlpack(view)`` or + another DLPack consumer, after publication. Holding a view does not pin KV pages. + """ + + def __dlpack_device__(self) -> tuple[int, int]: ... + def __dlpack__(self, stream: int | None = None, **kwargs: object) -> object: ... + +class Batch: + """Stable request slots and raw GPU metadata across all layer groups. + + Membership is non-owning and exclusive. Closing/destroying a request removes + it; closing/destroying the batch detaches live requests without closing them. + Use the requests' owning thread. ``publish`` runs outside graph capture after + mutations; ``wait_ready`` orders readers after upload and KV-copy completion. + Call ``record_read`` after submitting reads (or graph replay), before mutating, + suspending, removing, or closing requests. Device addresses remain stable. + """ + + def __init__( + self, manager: KVCacheManager, max_rows: int, max_blocks: int, max_beam_width: int = 1 + ) -> None: + """Allocate fixed device tables. Currently supports beam width 1.""" + @property + def max_rows(self) -> int: ... + @property + def max_blocks(self) -> int: ... + @property + def max_beam_width(self) -> int: ... + @property + def num_layer_groups(self) -> int: ... + @property + def dirty_rows(self) -> list[int]: ... + def add(self, kv_cache: _KVCache, row: int | None = None) -> int: ... + def remove(self, kv_cache: _KVCache) -> None: ... + def close(self) -> None: ... + def publish(self, cuda_stream: CudaStream) -> list[int]: + """Queue dirty rows and return their slots. Failures stay dirty; retry before reading.""" + def wait_ready(self, cuda_stream: CudaStream) -> None: ... + def record_read(self, cuda_stream: CudaStream) -> None: ... + def resize( + self, + capacities: list[int | None], + history_lengths: list[int | None], + cuda_stream: CudaStream, + ) -> list[bool | None]: + """Lists use stable row slots. Return per-request success (None for holes), then publish.""" + def page_table(self, layer_group_id: LayerGroupId) -> BatchDeviceArray: + """Raw slot IDs, shape [max_rows, max_beam_width, max_blocks]; preserves BAD_PAGE_INDEX.""" + def num_blocks(self, layer_group_id: LayerGroupId) -> BatchDeviceArray: + """Eligible history counts, shape [max_rows, max_beam_width]; zero for inactive/dense rows.""" + +class PageStorageSnapshot: + """Copied host metadata for one layer group and beam; indices are raw mixed-tier slot IDs. + + ``BAD_PAGE_INDEX`` is preserved and has no cache level. Eligibility is zero for + prefill, inactive requests and dense groups. Readiness events are retained internally; + they do not pin storage. Use indices only while the request is active and the version + matches. Submit reads on the request's stream, or call ``record_page_storage_read`` + after submission on another stream, before mutating or closing the request. + """ + + @property + def version(self) -> int: ... + @property + def row(self) -> int | None: ... + @property + def base_page_indices(self) -> list[int]: ... + @property + def cache_levels(self) -> list[CacheLevel | None]: ... + @property + def eligible_history_blocks(self) -> int: ... + def wait_ready(self, cuda_stream: CudaStream) -> None: + """Queue copy-completion waits without blocking the CPU or uploading metadata.""" + class _KVCache: Status: ClassVar[Type[_Status]] id: Any @@ -465,6 +545,22 @@ class _KVCache: def enter_decode(self) -> bool: ... @property def is_decoding(self) -> bool: ... + @property + def page_storage_version(self) -> int: ... + @property + def page_storage_dirty(self) -> bool: ... + @property + def page_storage_row(self) -> int | None: ... + def bind_page_storage_row(self, row: int | None) -> None: + """Bind a standalone consumer's row; Batch members must use Batch.add/remove.""" + def acknowledge_page_storage(self, version: int) -> bool: + """Clear dirty state after all groups/beams use this same version, if it is still current.""" + def get_page_storage_snapshot( + self, layer_group_id: LayerGroupId, beam_id: BeamIndex = DEFAULT_BEAM_INDEX + ) -> PageStorageSnapshot: + """Read the final state under the manager lock, including after a failed operation's rollback.""" + def record_page_storage_read(self, cuda_stream: CudaStream) -> None: + """Join submitted reader work into the active request's stream before any cache mutation.""" def prefetch(self, target: CacheLevel) -> bool: ... def get_scratch_desc(self, layer_group_id: LayerGroupId) -> ScratchDesc | None: ... @property @@ -594,6 +690,8 @@ class KVCacheManager: def __del__(self) -> None: ... def shutdown(self) -> None: ... def clear_reusable_blocks(self) -> None: ... + def is_sparse(self, layer_id: LayerId, data_role: DataRole) -> bool: + """Whether the named buffer uses sparse attention. Rejects unknown buffers.""" def get_mem_pool_base_address( self, layer_id: LayerId, data_role: DataRole, index_mode: PageIndexMode | None = None ) -> MemAddress: ... diff --git a/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_scheduler.py b/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_scheduler.py index f1fc920a35b7..1072dff66d98 100644 --- a/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_scheduler.py +++ b/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_scheduler.py @@ -60,6 +60,128 @@ def test_generation_admits_decode_before_capacity_growth(active: bool, admitted: assert manager._allocated_draft_lens == ({1: 0} if admitted else {}) +@pytest.mark.parametrize("is_draft", [False, True]) +def test_sparse_metadata_publishes_after_preparation(is_draft: bool) -> None: + manager = object.__new__(KVCacheManagerV2) + order = Mock() + manager._disagg_receive_ready = {} + manager.is_draft = is_draft + manager._stream = Mock(cuda_stream=123) + manager.sparse_metadata_batch = Mock() + manager.kv_connector_manager = Mock() + manager._prepare_draft_resources = order.prepare_draft + manager._run_kv_connector_hooks = order.connector + order.attach_mock(manager.sparse_metadata_batch.record_read, "record_read") + order.attach_mock(manager.sparse_metadata_batch.publish, "publish") + scheduled = Mock(context_requests=[]) + manager.prepare_resources(scheduled) + prepare = call.prepare_draft(scheduled) if is_draft else call.connector(scheduled) + assert order.mock_calls == [prepare, call.record_read(123), call.publish(123)] + + +def test_sparse_metadata_republishes_after_connector_acceptance() -> None: + manager = object.__new__(KVCacheManagerV2) + order = Mock() + manager.is_draft = False + manager._stream = Mock(cuda_stream=123) + manager.sparse_metadata_batch = Mock() + manager.kv_connector_manager = Mock() + manager._connector_reservations_enabled = Mock(return_value=True) + manager._accept_connector_prefix_reservations = order.accept + order.attach_mock(manager.kv_connector_manager.build_scheduler_output, "report") + order.attach_mock(manager.sparse_metadata_batch.record_read, "record_read") + order.attach_mock(manager.sparse_metadata_batch.publish, "publish") + scheduled = Mock() + manager.report_batch_to_connector(scheduled) + assert order.mock_calls == [ + call.accept(scheduled), + call.report(scheduled, manager), + call.record_read(123), + call.publish(123), + ] + + +def test_sparse_host_indices_cannot_reach_dense_attention_offsets() -> None: + manager = object.__new__(KVCacheManagerV2) + manager._stream = Mock(cuda_stream=123) + manager.sparse_metadata_batch = Mock() + manager.tokens_per_block = 4 + manager.kv_cache_map = {7: Mock(is_decoding=True, history_length=4)} + with patch( + "tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2." + "copy_batch_block_offsets_to_device" + ) as dense_copy: + with pytest.raises(RuntimeError, match="Offloaded sparse history"): + manager.copy_batch_block_offsets(Mock(), [7], 1, 0, 1) + dense_copy.assert_not_called() + manager.sparse_metadata_batch.publish.assert_called_once_with(123) + + +def test_sparse_publication_failure_stops_metadata_preparation() -> None: + manager = object.__new__(KVCacheManagerV2) + manager._stream = Mock(cuda_stream=123) + manager.sparse_metadata_batch = Mock() + manager.sparse_metadata_batch.publish.side_effect = RuntimeError("upload failed") + with patch( + "tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2." + "copy_batch_block_offsets_to_device" + ) as dense_copy: + with pytest.raises(RuntimeError, match="upload failed"): + manager.copy_batch_block_offsets(Mock(), [7], 1, 0, 1) + dense_copy.assert_not_called() + + +def test_dense_metadata_preparation_uses_existing_offsets() -> None: + manager = object.__new__(KVCacheManagerV2) + manager._stream = Mock(cuda_stream=123) + manager._use_per_layer_page_tables = False + manager.index_mapper = Mock() + manager.index_mapper.get_copy_index.return_value = Mock(shape=(1,)) + manager.host_kv_cache_block_offsets = Mock() + manager.index_scales = Mock() + manager.kv_offset = Mock() + destination = Mock() + with patch( + "tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2." + "copy_batch_block_offsets_to_device" + ) as dense_copy: + manager.copy_batch_block_offsets(destination, [7], 1, 0, 1) + dense_copy.assert_called_once_with( + manager.host_kv_cache_block_offsets, + destination, + manager.index_mapper.get_copy_index.return_value, + manager.index_scales, + manager.kv_offset, + 123, + ) + + +def test_sparse_index_slot_release_detaches_batch_before_reuse() -> None: + manager = object.__new__(KVCacheManagerV2) + manager.is_draft = False + manager._stream = Mock(cuda_stream=123) + manager.max_beam_width = 1 + manager.num_pools = 1 + manager._early_freed_index_requests = set() + cache = Mock() + manager.kv_cache_map = {7: cache} + manager.sparse_metadata_batch = Mock() + manager.index_mapper = Mock() + order = Mock() + order.attach_mock(manager.sparse_metadata_batch.record_read, "record_read") + order.attach_mock(manager.sparse_metadata_batch.remove, "remove") + order.attach_mock(cache.set_base_page_index_buf, "detach_buffer") + order.attach_mock(manager.index_mapper.remove_sequence, "release_slot") + manager.release_index_slot(7) + assert order.mock_calls == [ + call.record_read(123), + call.remove(cache), + call.detach_buffer(0, 0, None), + call.release_slot(7), + ] + assert manager._early_freed_index_requests == {7} + + # --------------------------------------------------------------------------- # State value constants # --------------------------------------------------------------------------- diff --git a/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_manager_v2.py b/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_manager_v2.py index 4a4bb5c4e6aa..324ddd9f7096 100755 --- a/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_manager_v2.py +++ b/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_manager_v2.py @@ -38,6 +38,7 @@ DEFAULT_BEAM_INDEX, GPU_LEVEL, AttentionLayerConfig, + Batch, BatchDesc, BufferConfig, BufferId, @@ -53,6 +54,7 @@ KVCacheManagerConfig, LayerGroupId, LayerId, + LogicError, MemAddress, OutOfPagesError, PageIndexMode, @@ -89,6 +91,7 @@ DEFAULT_BEAM_INDEX, GPU_LEVEL, AttentionLayerConfig, + Batch, BatchDesc, BufferConfig, BufferId, @@ -104,6 +107,7 @@ KVCacheManagerConfig, LayerGroupId, LayerId, + LogicError, MemAddress, OutOfPagesError, PageIndexMode, @@ -456,6 +460,146 @@ def values(stats): class TestNoBatching(TestKVCacheManagerV2): + def test_batch_publishes_sparse_gpu_metadata(self) -> None: + import torch + + self.manager = KVCacheManager( + KVCacheManagerConfig( + tokens_per_block=4, + cache_tiers=[GpuCacheTierConfig(4 << 20), HostCacheTierConfig(4 << 20)], + layers=[ + AttentionLayerConfig(0, [BufferConfig("key", 4096, is_sparse=True)]), + AttentionLayerConfig(1, [BufferConfig("key", 4096)]), + ], + ) + ) + batch = Batch(self.manager, max_rows=3, max_blocks=4) + cache = self.manager.create_kv_cache() + sparse_group = self.manager.get_layer_group_id(0) + dense_group = self.manager.get_layer_group_id(1) + stream = torch.cuda.Stream() + with torch.cuda.stream(stream): + try: + self.assertTrue(cache.resume(stream.cuda_stream)) + self.assertTrue(cache.resize(12, 6)) + self.assertEqual(batch.add(cache, row=2), 2) + with self.assertRaises(LogicError): + torch.from_dlpack(batch.page_table(sparse_group)) + self.assertEqual(batch.publish(stream.cuda_stream), [0, 1, 2]) + table = torch.from_dlpack(batch.page_table(sparse_group)) + counts = torch.from_dlpack(batch.num_blocks(sparse_group)) + dense_counts = torch.from_dlpack(batch.num_blocks(dense_group)) + self.assertEqual(tuple(table.shape), (3, 1, 4)) + self.assertEqual(table.dtype, torch.int32) + self.assertEqual(counts.cpu().tolist(), [[0], [0], [0]]) + address = table.data_ptr() + batch.record_read(stream.cuda_stream) + + old_version = cache.page_storage_version + self.assertTrue(cache.enter_decode()) + self.assertFalse(cache.acknowledge_page_storage(old_version)) + self.assertEqual(batch.dirty_rows, [2]) + self.assertEqual(batch.publish(stream.cuda_stream), [2]) + batch.wait_ready(stream.cuda_stream) + expected = [ + [[BAD_PAGE_INDEX] * 4], + [[BAD_PAGE_INDEX] * 4], + [list(cache.get_base_page_indices(sparse_group)) + [BAD_PAGE_INDEX]], + ] + self.assertEqual(table.cpu().tolist(), expected) + self.assertEqual(counts.cpu().tolist(), [[0], [0], [1]]) + self.assertEqual(dense_counts.cpu().tolist(), [[0], [0], [0]]) + self.assertFalse(cache.page_storage_dirty) + self.assertEqual(batch.publish(stream.cuda_stream), []) + batch.record_read(stream.cuda_stream) + + cache.suspend() + self.assertEqual(batch.publish(stream.cuda_stream), [2]) + self.assertEqual(counts.cpu().tolist(), [[0], [0], [0]]) + self.assertEqual(table.cpu().tolist(), [[[BAD_PAGE_INDEX] * 4]] * 3) + self.assertTrue(cache.resume()) + self.assertEqual(batch.publish(stream.cuda_stream), [2]) + self.assertEqual(counts.cpu().tolist(), [[0], [0], [1]]) + self.assertEqual(table.data_ptr(), address) + batch.record_read(stream.cuda_stream) + cache.close() + self.assertIsNone(cache.page_storage_row) + self.assertEqual(batch.publish(stream.cuda_stream), [2]) + self.assertEqual(table.cpu().tolist(), [[[BAD_PAGE_INDEX] * 4]] * 3) + self.assertEqual(counts.cpu().tolist(), [[0], [0], [0]]) + batch.close() + self.assertEqual(table.data_ptr(), address) + self.assertEqual(table.cpu().tolist(), [[[BAD_PAGE_INDEX] * 4]] * 3) + with self.assertRaises(LogicError): + torch.from_dlpack(batch.page_table(sparse_group)) + finally: + stream.synchronize() + batch.close() + cache.close() + + def test_sparse_page_storage_metadata(self) -> None: + self.manager = KVCacheManager( + KVCacheManagerConfig( + tokens_per_block=4, + cache_tiers=[GpuCacheTierConfig(4 << 20), HostCacheTierConfig(4 << 20)], + layers=[ + AttentionLayerConfig(0, [BufferConfig("key", 4096, is_sparse=True)]), + AttentionLayerConfig(1, [BufferConfig("key", 4096)]), + ], + ) + ) + self.assertTrue(self.manager.is_sparse(0, "key")) + self.assertFalse(self.manager.is_sparse(1, "key")) + sparse_group = self.manager.get_layer_group_id(0) + dense_group = self.manager.get_layer_group_id(1) + cache = self.manager.create_kv_cache() + with TemporaryCudaStream([]) as stream_holder: + stream = cast(CudaStream, stream_holder.handle) + try: + self.assertTrue(cache.resume(stream)) + self.assertTrue(cache.resize(12, 6)) + cache.bind_page_storage_row(3) + prefill = cache.get_page_storage_snapshot(sparse_group) + self.assertEqual(prefill.row, 3) + self.assertEqual(prefill.eligible_history_blocks, 0) + self.assertEqual(prefill.cache_levels, [0, 0, 0]) + self.assertTrue(cache.acknowledge_page_storage(prefill.version)) + self.assertFalse(cache.page_storage_dirty) + self.assertTrue(cache.enter_decode()) + self.assertTrue(cache.page_storage_dirty) + self.assertFalse(cache.acknowledge_page_storage(prefill.version)) + decode = cache.get_page_storage_snapshot(sparse_group) + self.assertEqual(decode.version, cache.page_storage_version) + self.assertEqual(decode.eligible_history_blocks, 1) + self.assertEqual(decode.cache_levels, [1, 0, 0]) + self.assertEqual( + decode.base_page_indices, list(cache.get_base_page_indices(sparse_group)) + ) + self.assertEqual( + cache.get_page_storage_snapshot(dense_group).eligible_history_blocks, 0 + ) + copied_indices = decode.base_page_indices + copied_indices[0] = BAD_PAGE_INDEX + self.assertNotEqual(decode.base_page_indices[0], BAD_PAGE_INDEX) + decode.wait_ready(stream) + cache.record_page_storage_read(stream) + self.assertTrue(cache.acknowledge_page_storage(decode.version)) + cache.commit([0, 1, 2, 3]) + self.assertTrue(cache.page_storage_dirty) + cache.suspend() + suspended = cache.get_page_storage_snapshot(sparse_group) + self.assertEqual(suspended.base_page_indices, [BAD_PAGE_INDEX] * 3) + self.assertEqual(suspended.cache_levels, [None] * 3) + self.assertEqual(suspended.eligible_history_blocks, 0) + self.assertEqual(cache.page_storage_row, 3) + self.assertTrue(cache.resume()) + cache.bind_page_storage_row(None) + self.assertIsNone(cache.page_storage_row) + finally: + cache.close() + stream_holder.take_finish_event().synchronize() + self.assertEqual(cache.get_page_storage_snapshot(sparse_group).base_page_indices, []) + class Request(NamedTuple): id: int kv_cache: _KVCache From 54d3af5f288315f275c91ca76372965926c18f9b Mon Sep 17 00:00:00 2001 From: Guiju Zhang <7135567+cascade812@users.noreply.github.com> Date: Fri, 2 Oct 2026 15:59:04 -0700 Subject: [PATCH 4/5] Defer sparse history offload for shared prefill pages Allow decode admission while shared prefill owners still require GPU pages. Retry deferred offloads at decode, history-update, and Batch publication boundaries, and publish only the contiguous host-eligible history count. Cover owner transitions, unchanged-history retries, transfer failures, and CUDA ordering. Signed-off-by: Guiju Zhang <7135567+cascade812@users.noreply.github.com> --- .../kv_cache_manager_v2/batch.cpp | 9 + .../batch_manager/kv_cache_manager_v2/batch.h | 3 +- .../kv_cache_manager_v2/kvCache.cpp | 61 +++-- .../kv_cache_manager_v2/kvCache.h | 11 +- .../kvCacheManagerV2ColdPageTest.cpp | 249 +++++++++++++++++- .../test_kv_cache_manager_v2.py | 66 +++++ 6 files changed, 369 insertions(+), 30 deletions(-) diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/batch.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/batch.cpp index 54bfada69bf2..9bd792e7859b 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/batch.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/batch.cpp @@ -222,6 +222,15 @@ std::vector Batch::publish(CudaStream stream) checkOpen(); auto const cudaStream = reinterpret_cast(stream); checkOutsideCapture(stream); + // Owner release can unblock history without changing a request's watermark or metadata version. + // Retry before collecting dirty rows: one shared-page move can invalidate several rows. + for (auto* cache : mRows) + { + if (cache != nullptr && cache->isActive() && cache->mIsDecoding && cache->mHasDeferredSparseOffload) + { + cache->_offloadSparseHistory({0, 0}, cache->mHistoryLength); + } + } auto rows = dirtyRows(); if (rows.empty()) { diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/batch.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/batch.h index e5c2970f63eb..0f078474b750 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/batch.h +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/batch.h @@ -57,7 +57,8 @@ class Batch : public std::enable_shared_from_this //! Detach all members. Exported arrays retain their allocation until the last owner dies. void close(); - //! Upload final dirty rows and counts; return the row slots uploaded. + //! Retry deferred sparse offloads, then upload final dirty rows and counts; return uploaded rows. + //! Offload failures propagate and retain pending work for retry. //! Staging buffers are retained until their asynchronous copies complete. std::vector publish(CudaStream stream); //! Wait for publication and KV readiness. Reject unpublished changes. diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.cpp index ccd522526f45..bf88ecfdc508 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.cpp @@ -217,6 +217,11 @@ CacheLevel KvCache::_lockLevel(Page const& page, BlockOrdinal ordinal) const void KvCache::_offloadSparseHistory(HalfOpenRange range, int historyLength) { TLLM_CHECK_DEBUG(range.end <= BlockOrdinal{historyLength / mTokensPerBlock}); + if (mHasDeferredSparseOffload) + { + range = {0, historyLength / mTokensPerBlock}; + } + bool deferred = false; std::vector> pages; for (auto const& [lcId, lc] : mManager->lifeCycles()) { @@ -234,21 +239,36 @@ void KvCache::_offloadSparseHistory(HalfOpenRange range, int histo if (page->cacheLevel == kSparseHistoryLevel) continue; auto const lock = page->holder.lock()->uniqLock.lock(); + bool needsGpu = false; for (auto const& owner : lock->owners()) { - if (owner.kvCache != this && !owner.kvCache->isDecoding()) - throw LogicError("Cannot offload sparse history shared with a prefill request"); + int const ownerHistory = owner.kvCache == this ? historyLength : owner.kvCache->historyLength(); + if ((owner.kvCache != this && !owner.kvCache->isDecoding()) + || owner.ordinal >= BlockOrdinal{ownerHistory / owner.kvCache->tokensPerBlock()}) + { + needsGpu = true; + break; + } + } + if (needsGpu) + { + deferred = true; + continue; } pages.push_back(std::move(page)); } } } - if (pages.empty()) - return; - int const oldHistoryLength = mHistoryLength; - auto restoreHistory = FuncGuard([&]() { mHistoryLength = oldHistoryLength; }); - mHistoryLength = historyLength; - offloadSparsePages(pages); + // Preserve retry work if allocation or copy submission fails. + mHasDeferredSparseOffload = true; + if (!pages.empty()) + { + int const oldHistoryLength = mHistoryLength; + auto restoreHistory = FuncGuard([&]() { mHistoryLength = oldHistoryLength; }); + mHistoryLength = historyLength; + offloadSparsePages(pages); + } + mHasDeferredSparseOffload = deferred; } void KvCache::_publishHistoryLength(int historyLength) @@ -264,7 +284,7 @@ bool KvCache::enterDecode() auto const apiLock = mManager->lockExclusive(); if (!isActive()) throw LogicError("Decode admission requires an active request"); - if (mIsDecoding) + if (mIsDecoding && !mHasDeferredSparseOffload) return true; try { @@ -274,8 +294,11 @@ bool KvCache::enterDecode() { return false; } - mIsDecoding = true; - onPageStorageChanged(); + if (!mIsDecoding) + { + mIsDecoding = true; + onPageStorageChanged(); + } return true; } @@ -727,6 +750,7 @@ void KvCache::_deactivate() _freeScratchSlots(); } mStatus = Status::SUSPENDED; + mHasDeferredSparseOffload = false; onPageStorageChanged(); } @@ -771,6 +795,7 @@ void KvCache::close() mPageStorageBatch->remove(*this); } mStatus = Status::CLOSED; + mHasDeferredSparseOffload = false; mPageStorageRow.reset(); onPageStorageChanged(); mManager->unregisterKvCache(this); @@ -1554,7 +1579,7 @@ void KvCache::setHistoryLength(int hist) bool KvCache::_shortcutSetHistoryLength(int newHist) { - if (newHist == mHistoryLength) + if (newHist == mHistoryLength && !mHasDeferredSparseOffload) return true; // Check if stale range changes for any lifecycle. for (auto [lcId, lc] : mManager->lifeCycles()) @@ -2739,18 +2764,20 @@ PageStorageSnapshot KvCache::getPageStorageSnapshot(LayerGroupId lgId, BeamIndex } auto const* attn = std::get_if(&mManager->lifeCycles()[lgId]); - if (mIsDecoding && attn && attn->isSparse) - snapshot.mEligibleHistoryBlocks = mHistoryLength / mTokensPerBlock; + int const completeSparseHistory = mIsDecoding && attn && attn->isSparse ? mHistoryLength / mTokensPerBlock : 0; snapshot.mReadyEvents.reserve(numBlocks); for (BlockOrdinal ord{0}; ord < mBlocks.size(); ++ord) { int const index = snapshot.mBasePageIndices[toSizeT(ord)]; auto const& page = blockPageGetPage(mBlocks[ord].pages.at(beamIdx).at(lgId)); - if (ord.value() < snapshot.mEligibleHistoryBlocks) + // The scalar count exposes only the contiguous host prefix, stopping at any deferred GPU page. + if (ord.value() == snapshot.mEligibleHistoryBlocks && ord.value() < completeSparseHistory && page + && index != kBadPageIndex.value() && page->cacheLevel == kSparseHistoryLevel) { - TLLM_CHECK_WITH_INFO(page && index != kBadPageIndex.value() && page->cacheLevel == kSparseHistoryLevel - && page->hasValidSlot() && index == slotIdToPageIndexValue(page->slotId()), + TLLM_CHECK_WITH_INFO(page->status() == PageStatus::LOCKED && page->hasValidSlot() + && index == slotIdToPageIndexValue(page->slotId()), "Eligible sparse history must have a locked host mapping"); + ++snapshot.mEligibleHistoryBlocks; } if (index == kBadPageIndex.value()) continue; diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.h index 998a6f4c9bc1..c22e69702889 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.h +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.h @@ -179,6 +179,7 @@ class PageStorageSnapshot return mCacheLevels; } + // Contiguous complete sparse history on host; stops at the first missing or GPU mapping. int eligibleHistoryBlocks() const noexcept { return mEligibleHistoryBlocks; @@ -243,8 +244,9 @@ class KvCache : public std::enable_shared_from_this // Returns false if utilization too high or out of memory. bool resume(std::optional stream = std::nullopt, std::optional isDecoding = std::nullopt); - // Enter decode only after prefill has submitted its final KV accesses. Reconciles all complete - // sparse history, including on retries with an unchanged watermark. Returns false on host OOM. + // Enter decode only after prefill has submitted its final KV accesses. Offloads complete sparse + // history, deferring pages still needed on GPU by another owner. Retries deferred pages even + // with an unchanged watermark. Returns false on host OOM. bool enterDecode(); bool isDecoding() const noexcept @@ -564,7 +566,8 @@ class KvCache : public std::enable_shared_from_this // Prefill and writable pages require GPU storage. Decode keeps cold sparse history on host. CacheLevel _lockLevel(Page const& page, BlockOrdinal ordinal) const; - // Offload GPU pages in the supplied complete-history range, validating every live owner's phase. + // Offload GPU pages in the supplied complete-history range and retry deferred history. + // Pages stay on GPU until they belong to every live owner's complete decode history. // The candidate watermark is visible only under the exclusive API lock until offload succeeds. void _offloadSparseHistory(HalfOpenRange range, int historyLength); void _publishHistoryLength(int historyLength); @@ -724,6 +727,8 @@ class KvCache : public std::enable_shared_from_this int mCapacity; int mHistoryLength; bool mIsDecoding = false; + // Retry by scanning current blocks; deferred work does not retain pages or other requests. + bool mHasDeferredSparseOffload = false; std::optional mExpectedPromptLength; bool mGenerationAllocReady = false; diff --git a/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp b/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp index a3573d2a9bb1..32f49ebb6d46 100644 --- a/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp +++ b/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp @@ -1127,6 +1127,131 @@ TEST_F(KvCacheManagerV2BatchTest, SharedOffloadInvalidatesEveryOwner) EXPECT_EQ(read(batch.numBlocksAddress(LifeCycleId{0}), 2), (std::vector{1, 1})); } +TEST_F(KvCacheManagerV2BatchTest, DeferredHistoryGapClosesWithoutAdvancingWatermark) +{ + for (bool const closeOwner : {false, true}) + { + SCOPED_TRACE(closeOwner); + auto config = sparseConfig(); + config.cacheTiers[0] = GpuCacheTierConfig{8 << 20}; + auto codec = std::make_unique(); + auto* observer = codec.get(); + auto manager = std::make_shared(std::move(config), nullptr, std::move(codec)); + auto const apiLock = manager->lockExclusive(); + auto page = seedPrefix(*manager, kHotLevel); + auto decoder = manager->createKvCache({}, tokens()); + auto prefill = manager->createKvCache({}, tokens()); + auto closeCaches = FuncGuard( + [&]() + { + decoder->close(); + prefill->close(); + }); + ASSERT_TRUE(decoder->resume(stream())); + ASSERT_TRUE(prefill->resume(stream())); + ASSERT_TRUE(decoder->resize(12, 8)); + Batch batch(manager, 1, 3); + batch.add(*decoder); + auto const group = manager->getLayerGroupId(0); + ASSERT_TRUE(decoder->enterDecode()); + EXPECT_EQ(page->cacheLevel, kHotLevel); + EXPECT_EQ(pageAt(*decoder, 1)->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(pageAt(*decoder, 2)->cacheLevel, kHotLevel); + auto const snapshot = decoder->getPageStorageSnapshot(group); + EXPECT_EQ(snapshot.eligibleHistoryBlocks(), 0); + EXPECT_EQ(snapshot.cacheLevels(), + (std::vector>{kHotLevel, kSparseHistoryLevel, kHotLevel})); + EXPECT_EQ(batch.publish(batchStream()), (std::vector{0})); + EXPECT_EQ(read(batch.pageTableAddress(group), 3), snapshot.basePageIndices()); + EXPECT_EQ(read(batch.numBlocksAddress(group), 1), (std::vector{0})); + EXPECT_EQ(observer->encodedPages, 1); + auto const hostSlot = pageAt(*decoder, 1)->slotId(); + auto const version = decoder->pageStorageVersion(); + + if (closeOwner) + { + prefill->close(); + } + else + { + prefill->suspend(); + } + EXPECT_EQ(page->cacheLevel, kHotLevel); + EXPECT_EQ(decoder->pageStorageVersion(), version); + EXPECT_TRUE(batch.dirtyRows().empty()); + EXPECT_EQ(batch.publish(batchStream()), (std::vector{0})); + EXPECT_EQ(decoder->historyLength(), 8); + EXPECT_GT(decoder->pageStorageVersion(), version); + EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(pageAt(*decoder, 1)->slotId(), hostSlot); + EXPECT_EQ(observer->encodedPages, 2); + auto const finalSnapshot = decoder->getPageStorageSnapshot(group); + EXPECT_EQ(finalSnapshot.eligibleHistoryBlocks(), 2); + EXPECT_EQ(read(batch.numBlocksAddress(group), 1), (std::vector{2})); + EXPECT_EQ(read(batch.pageTableAddress(group), 3), finalSnapshot.basePageIndices()); + EXPECT_EQ(batch.publish(batchStream()), (std::vector{})); + EXPECT_EQ(observer->encodeCalls, 2); + } +} + +TEST_F(KvCacheManagerV2BatchTest, DeferredOffloadRetainsWorkAfterHostOomAndCodecRejection) +{ + auto codec = std::make_unique(); + auto* observer = codec.get(); + auto manager = std::make_shared(sparseConfig(), nullptr, std::move(codec)); + auto const apiLock = manager->lockExclusive(); + auto& storage = manager->storage(); + auto page = seedPrefix(*manager, kHotLevel); + auto decoder = manager->createKvCache({}, tokens()); + auto prefill = manager->createKvCache({}, tokens()); + auto closeCaches = FuncGuard( + [&]() + { + decoder->close(); + prefill->close(); + }); + ASSERT_TRUE(decoder->resume(stream())); + ASSERT_TRUE(prefill->resume(stream())); + auto const freeHost = storage.getStatistics(kSparseHistoryLevel).free; + auto blockers = storage.newSlots(kSparseHistoryLevel, TypedVec{freeHost}); + auto releaseBlockers = FuncGuard( + [&]() + { + for (auto& slot : blockers[LifeCycleId{0}]) + { + storage.releaseSlot(LifeCycleId{0}, kSparseHistoryLevel, std::move(slot)); + } + }); + ASSERT_TRUE(decoder->enterDecode()); + Batch batch(manager, 1, 1); + batch.add(*decoder); + auto const group = manager->getLayerGroupId(0); + batch.publish(batchStream()); + auto const gpuSlot = page->slotId(); + prefill->close(); + EXPECT_THROW(batch.publish(batchStream()), OutOfPagesError); + EXPECT_EQ(observer->encodeCalls, 0); + EXPECT_EQ(page->slotId(), gpuSlot); + EXPECT_EQ(page->cacheLevel, kHotLevel); + EXPECT_EQ(read(batch.numBlocksAddress(group), 1), (std::vector{0})); + releaseBlockers.run(); + + observer->rejectEncodeCall = 1; + EXPECT_THROW(batch.publish(batchStream()), TllmException); + EXPECT_EQ(page->cacheLevel, kHotLevel); + EXPECT_EQ(page->slotId(), gpuSlot); + EXPECT_TRUE(decoder->isDecoding()); + EXPECT_EQ(decoder->historyLength(), 4); + EXPECT_EQ(decoder->getPageStorageSnapshot(group).eligibleHistoryBlocks(), 0); + EXPECT_EQ(read(batch.pageTableAddress(group), 1), (std::vector{slotIdToPageIndexValue(gpuSlot)})); + EXPECT_EQ(read(batch.numBlocksAddress(group), 1), (std::vector{0})); + observer->rejectEncodeCall = 0; + EXPECT_EQ(batch.publish(batchStream()), (std::vector{0})); + EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(read(batch.numBlocksAddress(group), 1), (std::vector{1})); + EXPECT_EQ(storage.getStatistics(kSparseHistoryLevel).free, freeHost - 1); +} + TEST_F(KvCacheManagerV2BatchTest, RetainsStagingUntilUploadAndOrdersTableReuseAfterReaders) { auto config = sparseConfig(); @@ -1690,8 +1815,10 @@ TEST_F(KvCacheManagerV2SparseOffloadTest, SharedOwnersPublishHostIndicesAndKeepH TEST_F(KvCacheManagerV2SparseOffloadTest, WaitsForLiveAndFinishedReadersBeforeRecyclingGpuSlot) { - for (bool const finishReader : {false, true}) + for (auto const [deferred, finishReader] : + {std::pair{false, false}, std::pair{false, true}, std::pair{true, false}, std::pair{true, true}}) { + SCOPED_TRACE(deferred); SCOPED_TRACE(finishReader); auto manager = std::make_shared(sparseConfig()); auto const apiLock = manager->lockExclusive(); @@ -1730,6 +1857,11 @@ TEST_F(KvCacheManagerV2SparseOffloadTest, WaitsForLiveAndFinishedReadersBeforeRe ASSERT_EQ(cudaStreamSynchronize(mStream), cudaSuccess); auto const readback = std::get(storage.slotAddress(kSparseHistoryLevel, storage.getPoolGroupIndex(kSparseHistoryLevel, lc), hostScratch[lc].front().slotId(), PoolIndex{0})); + if (deferred) + { + ASSERT_TRUE(first->enterDecode()); + EXPECT_EQ(page->cacheLevel, kHotLevel); + } StreamGate gate; ASSERT_EQ(gate.enqueue(readerStream), cudaSuccess); ASSERT_EQ(cudaMemcpyAsync(reinterpret_cast(readback), reinterpret_cast(gpuAddress), bytes, @@ -1740,7 +1872,15 @@ TEST_F(KvCacheManagerV2SparseOffloadTest, WaitsForLiveAndFinishedReadersBeforeRe second->suspend(); } - first->offloadSparsePages({page}); + if (deferred) + { + ASSERT_TRUE(finishReader ? first->resize(4, 4) : second->enterDecode()); + EXPECT_EQ(first->historyLength(), 4); + } + else + { + first->offloadSparsePages({page}); + } EXPECT_FALSE(page->queryReady()); EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); auto recycled = storage.newGpuSlots(TypedVec{1}); @@ -2200,7 +2340,7 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, HistoryUpdatesOffloadOnlyNewFullPages) EXPECT_EQ(observer->encodedPages, 2); } -TEST_F(KvCacheManagerV2DecodeOffloadTest, SharedPrefillOwnerBlocksDemotionAndResumeReconcilesOlderGpuHistory) +TEST_F(KvCacheManagerV2DecodeOffloadTest, SharedPrefillOwnerDefersDemotionAcrossDecodeResume) { auto manager = std::make_shared(sparseConfig()); auto const apiLock = manager->lockExclusive(); @@ -2216,10 +2356,11 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, SharedPrefillOwnerBlocksDemotionAndRes ASSERT_TRUE(first->resume(stream())); ASSERT_TRUE(second->resume(stream())); auto const version = first->pageStorageVersion(); - EXPECT_THROW(first->enterDecode(), LogicError); - EXPECT_FALSE(first->isDecoding()); - EXPECT_EQ(first->pageStorageVersion(), version); + ASSERT_TRUE(first->enterDecode()); + EXPECT_TRUE(first->isDecoding()); + EXPECT_GT(first->pageStorageVersion(), version); EXPECT_EQ(page->cacheLevel, kHotLevel); + EXPECT_EQ(first->getPageStorageSnapshot(LifeCycleId{0}).eligibleHistoryBlocks(), 0); second->suspend(); ASSERT_TRUE(first->enterDecode()); EXPECT_THROW(second->resume(), LogicError); @@ -2228,15 +2369,105 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, SharedPrefillOwnerBlocksDemotionAndRes first->suspend(); ASSERT_TRUE(second->resume()); EXPECT_EQ(page->cacheLevel, kHotLevel); - EXPECT_THROW(first->resume(), LogicError); - EXPECT_FALSE(first->isActive()); + ASSERT_TRUE(first->resume()); + EXPECT_TRUE(first->isActive()); EXPECT_TRUE(first->isDecoding()); + EXPECT_EQ(page->cacheLevel, kHotLevel); second->close(); - ASSERT_TRUE(first->resume()); + ASSERT_TRUE(first->enterDecode()); EXPECT_EQ(first->historyLength(), 4); EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); } +TEST_F(KvCacheManagerV2DecodeOffloadTest, SharedOwnersEnterDecodeWithoutMutuallyBlocking) +{ + auto config = sparseConfig(); + config.enableStats = true; + auto codec = std::make_unique(); + auto* observer = codec.get(); + auto manager = std::make_shared(std::move(config), nullptr, std::move(codec)); + auto const apiLock = manager->lockExclusive(); + auto page = seedPrefix(*manager, kHotLevel); + auto first = manager->createKvCache({}, tokens()); + auto second = manager->createKvCache({}, tokens()); + auto third = manager->createKvCache({}, tokens()); + auto closeCaches = FuncGuard( + [&]() + { + first->close(); + second->close(); + third->close(); + }); + for (auto const& cache : {first, second, third}) + { + ASSERT_TRUE(cache->resume(stream())); + EXPECT_EQ(pageAt(*cache), page); + } + auto const gpuSlot = page->slotId(); + auto const hostFree = manager->storage().getStatistics(kSparseHistoryLevel).free; + ASSERT_TRUE(first->enterDecode()); + ASSERT_TRUE(second->enterDecode()); + auto const version = first->pageStorageVersion(); + ASSERT_TRUE(third->resize(4, 4)); + ASSERT_TRUE(first->enterDecode()); + ASSERT_TRUE(second->resize(4, 4)); + EXPECT_EQ(first->pageStorageVersion(), version); + EXPECT_EQ(page->cacheLevel, kHotLevel); + EXPECT_EQ(page->slotId(), gpuSlot); + EXPECT_EQ(observer->encodeCalls, 0); + EXPECT_EQ(manager->storage().getStatistics(kSparseHistoryLevel).free, hostFree); + EXPECT_EQ(first->getPageStorageSnapshot(LifeCycleId{0}).eligibleHistoryBlocks(), 0); + + ASSERT_TRUE(third->enterDecode()); + EXPECT_EQ(observer->encodedPages, 1); + EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); + for (auto const& cache : {first, second, third}) + { + EXPECT_EQ(cache->historyLength(), 4); + auto const snapshot = cache->getPageStorageSnapshot(LifeCycleId{0}); + EXPECT_EQ(snapshot.eligibleHistoryBlocks(), 1); + EXPECT_EQ(snapshot.basePageIndices()[0], slotIdToPageIndexValue(page->slotId())); + ASSERT_TRUE(cache->enterDecode()); + } + EXPECT_EQ(observer->encodeCalls, 1); + EXPECT_EQ(manager->getAndResetIterationStats().at(LifeCycleId{0}).iterOffloadBlocks, 1); +} + +TEST_F(KvCacheManagerV2DecodeOffloadTest, DeferredOwnerCanCloseBeforePrefillOwners) +{ + for (bool const closeAll : {false, true}) + { + SCOPED_TRACE(closeAll); + auto codec = std::make_unique(); + auto* observer = codec.get(); + auto manager = std::make_shared(sparseConfig(), nullptr, std::move(codec)); + auto const apiLock = manager->lockExclusive(); + auto page = seedPrefix(*manager, kHotLevel); + auto decoder = manager->createKvCache({}, tokens()); + auto prefill = manager->createKvCache({}, tokens()); + auto closeCaches = FuncGuard( + [&]() + { + decoder->close(); + prefill->close(); + }); + ASSERT_TRUE(decoder->resume(stream())); + ASSERT_TRUE(prefill->resume(stream())); + ASSERT_TRUE(decoder->enterDecode()); + decoder->close(); + if (closeAll) + { + prefill->close(); + EXPECT_EQ(observer->encodeCalls, 0); + prefill = manager->createKvCache({}, tokens()); + ASSERT_TRUE(prefill->resume(stream())); + } + ASSERT_TRUE(prefill->enterDecode()); + EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(observer->encodedPages, 1); + } +} + TEST_F(KvCacheManagerV2DecodeOffloadTest, HostOomDoesNotAdmitDecodeAndCanRetry) { for (int admission : {0, 1, 2}) diff --git a/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_manager_v2.py b/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_manager_v2.py index 324ddd9f7096..61e98ed86df7 100755 --- a/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_manager_v2.py +++ b/tests/unittest/kv_cache_manager_v2_tests/test_kv_cache_manager_v2.py @@ -537,6 +537,72 @@ def test_batch_publishes_sparse_gpu_metadata(self) -> None: batch.close() cache.close() + @parameterized.expand(["decode", "suspend", "close"]) + def test_shared_prefill_defers_sparse_offload(self, release: str) -> None: + import torch + + self.manager = KVCacheManager( + KVCacheManagerConfig( + tokens_per_block=4, + cache_tiers=[GpuCacheTierConfig(4 << 20), HostCacheTierConfig(4 << 20)], + layers=[AttentionLayerConfig(0, [BufferConfig("key", 4096, is_sparse=True)])], + ) + ) + decoder = self.manager.create_kv_cache() + prefill = None + batch = Batch(self.manager, max_rows=1, max_blocks=3) + group = self.manager.get_layer_group_id(0) + stream = torch.cuda.Stream() + with torch.cuda.stream(stream): + try: + self.assertTrue(decoder.resume(stream.cuda_stream)) + self.assertTrue(decoder.resize(12, 8)) + decoder.commit([0, 1, 2, 3]) + prefill = self.manager.create_kv_cache(ReuseScope(), [0, 1, 2, 3]) + self.assertTrue(prefill.resume(stream.cuda_stream)) + self.assertEqual( + decoder.get_base_page_indices(group)[0], prefill.get_base_page_indices(group)[0] + ) + self.assertTrue(decoder.enter_decode()) + self.assertTrue(decoder.is_decoding) + deferred = decoder.get_page_storage_snapshot(group) + self.assertEqual(deferred.cache_levels, [0, 1, 0]) + self.assertEqual(deferred.eligible_history_blocks, 0) + batch.add(decoder) + self.assertEqual(batch.publish(stream.cuda_stream), [0]) + table = torch.from_dlpack(batch.page_table(group)) + counts = torch.from_dlpack(batch.num_blocks(group)) + address = table.data_ptr() + self.assertEqual(table.cpu().tolist(), [[deferred.base_page_indices]]) + self.assertEqual(counts.cpu().tolist(), [[0]]) + batch.record_read(stream.cuda_stream) + self.assertEqual(batch.publish(stream.cuda_stream), []) + + if release == "decode": + self.assertTrue(prefill.enter_decode()) + elif release == "suspend": + prefill.suspend() + else: + prefill.close() + self.assertEqual(batch.publish(stream.cuda_stream), [0]) + batch.wait_ready(stream.cuda_stream) + self.assertEqual(decoder.history_length, 8) + completed = decoder.get_page_storage_snapshot(group) + self.assertEqual(completed.cache_levels, [1, 1, 0]) + self.assertEqual(completed.eligible_history_blocks, 2) + self.assertGreater(completed.version, deferred.version) + self.assertEqual(table.data_ptr(), address) + self.assertEqual(table.cpu().tolist(), [[completed.base_page_indices]]) + self.assertEqual(counts.cpu().tolist(), [[2]]) + batch.record_read(stream.cuda_stream) + self.assertEqual(batch.publish(stream.cuda_stream), []) + finally: + stream.synchronize() + batch.close() + decoder.close() + if prefill is not None: + prefill.close() + def test_sparse_page_storage_metadata(self) -> None: self.manager = KVCacheManager( KVCacheManagerConfig( From c7bcc26c5911f5baab013c9025252ec9e7c618e7 Mon Sep 17 00:00:00 2001 From: Guiju Zhang <7135567+cascade812@users.noreply.github.com> Date: Sun, 4 Oct 2026 21:11:52 -0700 Subject: [PATCH 5/5] [None][fix] initialize CUDA stream pool before test gates Signed-off-by: Guiju Zhang <7135567+cascade812@users.noreply.github.com> --- .../unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp | 3 +++ 1 file changed, 3 insertions(+) diff --git a/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp b/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp index 32f49ebb6d46..7cbfc5a68ed8 100644 --- a/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp +++ b/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp @@ -23,6 +23,7 @@ #include "tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCache.h" #include "tensorrt_llm/batch_manager/kv_cache_manager_v2/kvCacheManager.h" #include "tensorrt_llm/batch_manager/kv_cache_manager_v2/storageManager.h" +#include "tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/cudaEvent.h" #include "tensorrt_llm/batch_manager/kv_cache_manager_v2/utils/funcGuard.h" #include "tensorrt_llm/common/tllmException.h" @@ -365,6 +366,8 @@ class StreamGate cudaError_t enqueue(cudaStream_t stream) { + // Event merging must not create streams while a gate is held: stream creation can wait for host callbacks. + CudaStreamPool::instance(); mStream = stream; return cudaLaunchHostFunc(stream, wait, this); }