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/7] 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/7] 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/7] 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/7] 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/7] [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); } From 2d0da3fba0b90a9a266b81fb393ee5c35a8ab760 Mon Sep 17 00:00:00 2001 From: Guiju Zhang <7135567+cascade812@users.noreply.github.com> Date: Mon, 5 Oct 2026 13:44:07 -0700 Subject: [PATCH 6/7] [None][fix] promote shared host KV pages for prefill Restore host-locked sparse prefixes into one shared GPU allocation for prefill admission, resume, prefetch, and rebasing. Order transfers after all readers, update every owner, and fence source-slot reuse. Retain deferred offload for decoding owners so publication retries after prefill releases its GPU requirement, even without history growth. Cover GPU exhaustion, copy failures, shared metadata, and reader ordering. Validation: 145 isolated native tests and 40 Python runtime checks passed on A30; 12 performance cases skipped. Full production package and CI wrapper validation remain unverified. Signed-off-by: Guiju Zhang <7135567+cascade812@users.noreply.github.com> --- .../kv_cache_manager_v2/AGENTS.md | 8 + .../kv_cache_manager_v2/page.cpp | 53 ++- .../batch_manager/kv_cache_manager_v2/page.h | 9 +- .../kv_cache_manager_v2/storageManager.cpp | 89 +++-- .../kv_cache_manager_v2/storageManager.h | 10 +- .../kvCacheManagerV2ColdPageTest.cpp | 364 +++++++++++++++++- 6 files changed, 483 insertions(+), 50 deletions(-) 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 5e43f816534a..750c67b07361 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 @@ -140,6 +140,14 @@ moving a shared host source. Rollback records the original lock level, and `ScratchSlotLock` remains GPU-only. This does not itself demote GPU history; offload requires a separate transfer and GPU-slot ownership handoff. +Prefill admission, prefetch, and rebasing promote a host-locked sparse page into +one shared GPU slot. Promotion waits for all owners' prior work and finished +readers, updates every owner's page indices and metadata version, and releases +the source host slot with a copy-completion fence. Allocation or copy failure +preserves host ownership. Decoding owners retain deferred-offload work so that +`Batch.publish()` retries demotion after the prefill owner enters decode, +suspends, or closes, even without history growth. + ## Ownership and lifetime The high-level ownership shape is: 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 f071bed08ec3..f71015cc1c90 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 @@ -381,7 +381,25 @@ void UniqPageLock::prepareSparseOffload(KvCache const& requestingCache) finishEvents.reserve(1); } -void UniqPageLock::recordOffloadEvent(CachedCudaEvent const& event) +void UniqPageLock::prepareSparsePromotion() +{ + Page const& p = *page(); + if (!p.hasValidSlot() || p.cacheLevel != kSparseHistoryLevel || p.queryLockLevel() != kSparseHistoryLevel + || mOwners.empty()) + { + throw LogicError("Promotion requires a locked sparse host page"); + } + for (auto const& owner : mOwners) + { + if (!owner.kvCache->isActive() || owner.lifeCycle != p.lifeCycle) + { + throw LogicError("Promotion requires active owners of the same sparse page"); + } + } + finishEvents.reserve(1); +} + +void UniqPageLock::recordMigrationEvent(CachedCudaEvent const& event) { // The copy stream already waited for every event being replaced here. page()->readyEvent = event; @@ -392,20 +410,27 @@ void UniqPageLock::recordOffloadEvent(CachedCudaEvent const& event) owner.kvCache->onPageStorageChanged(); } -Slot UniqPageLock::moveToSparseHistory(Slot&& hostSlot) +Slot UniqPageLock::moveToCacheLevel(CacheLevel destination, Slot&& slot) { Page& p = *page(); - TLLM_CHECK_DEBUG(p.cacheLevel == kHotLevel && !p.scheduledForEviction()); - Slot gpuSlot = p.exchangeSlot(std::move(hostSlot)); - p.cacheLevel = kSparseHistoryLevel; + TLLM_CHECK_DEBUG(!p.scheduledForEviction()); + TLLM_CHECK_DEBUG((p.cacheLevel == kHotLevel && destination == kSparseHistoryLevel) + || (p.cacheLevel == kSparseHistoryLevel && destination == kHotLevel)); + Slot source = p.exchangeSlot(std::move(slot)); + p.cacheLevel = destination; 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())); + TLLM_CHECK_DEBUG(old == slotIdToPageIndexValue(source.slotId())); + if (destination == kHotLevel && owner.kvCache->mIsDecoding) + { + // Retry even without history growth once the prefill owner releases its GPU requirement. + owner.kvCache->mHasDeferredSparseOffload = true; + } owner.kvCache->onPageStorageChanged(); } - return gpuSlot; + return source; } void UniqPageLock::removeOwner(LockOwner const& owner) @@ -531,6 +556,7 @@ std::vector batchedLockPages(KvCache& kvCache, std::vector destinations; TypedVec>> pagesByLevel(storeMgr->numCacheLevels()); + std::vector> lockedHostPages; for (auto const& target : targets) { auto const& page = target.page; @@ -539,13 +565,18 @@ std::vector batchedLockPages(KvCache& kvCache, std::vectorstatus() == PageStatus::LOCKED && page->cacheLevel != level) - throw LogicError("Cannot migrate a page locked by another owner"); auto const [it, inserted] = destinations.emplace(page.get(), level); if (!inserted && it->second != level) throw LogicError("Conflicting lock levels for a shared page"); if (inserted) + { pagesByLevel[level].push_back(page); + if (page->status() == PageStatus::LOCKED && page->cacheLevel != level) + { + TLLM_CHECK_DEBUG(level == kHotLevel && page->cacheLevel == kSparseHistoryLevel); + lockedHostPages.push_back(page); + } + } } // Protect every destination group while any group is allocating. On failure, @@ -578,6 +609,10 @@ std::vector batchedLockPages(KvCache& kvCache, std::vectorcacheLevel != level) ++requirements[storeMgr->getPoolGroupIndex(level, page->lifeCycle)]; storeMgr->prepareFreeSlots(level, requirements, migrationRecorder, dropRecorder); + if (level == kHotLevel && !lockedHostPages.empty()) + { + storeMgr->promoteSparsePages(kvCache.cudaStream(), lockedHostPages, migrationRecorder, dropRecorder); + } storeMgr->batchedMigrate(level, pages, migrationRecorder); } 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 8b9cfcfd1507..c60e8eafc411 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 @@ -213,11 +213,14 @@ class UniqPageLock : public EnableSharedFromThis //! Validate complete sparse history for every owner and prepare non-allocating completion updates. void prepareSparseOffload(KvCache const& requestingCache); + //! Validate a locked host page and prepare non-allocating completion updates for shared promotion. + void prepareSparsePromotion(); + //! Record a copy ordered after page readiness, finished readers, and all live owners' prior work. - void recordOffloadEvent(CachedCudaEvent const& event); + void recordMigrationEvent(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); + //! Publish a GPU/host handoff to every owner and return the fenced source slot. Caller holds the API lock. + [[nodiscard]] Slot moveToCacheLevel(CacheLevel destination, Slot&& slot); std::vector const& owners() const noexcept { 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 b67d26cf4125..9e2f440be329 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 @@ -1267,7 +1267,27 @@ void StorageManager::batchedMigrate( void StorageManager::offloadSparsePages(KvCache& requestingCache, std::vector> const& pages, MigrationRecorder const& migrationRecorder, DropRecorder const& dropRecorder) { - struct OffloadBatch + if (requestingCache.storageManager() != this) + { + throw LogicError("Offload pages and requesting cache must belong to the same manager"); + } + _migrateLockedSparsePages( + requestingCache.cudaStream(), kSparseHistoryLevel, pages, &requestingCache, migrationRecorder, dropRecorder); +} + +void StorageManager::promoteSparsePages(CUstream stream, std::vector> const& pages, + MigrationRecorder const& migrationRecorder, DropRecorder const& dropRecorder) +{ + _migrateLockedSparsePages(stream, kHotLevel, pages, nullptr, migrationRecorder, dropRecorder); +} + +void StorageManager::_migrateLockedSparsePages(CUstream stream, CacheLevel dstLevel, + std::vector> const& pages, KvCache const* offloadingCache, + MigrationRecorder const& migrationRecorder, DropRecorder const& dropRecorder) +{ + CacheLevel const srcLevel = dstLevel == kHotLevel ? kSparseHistoryLevel : kHotLevel; + + struct MigrationBatch { std::vector> srcPages; std::vector> srcPageLocks; @@ -1275,14 +1295,14 @@ void StorageManager::offloadSparsePages(KvCache& requestingCache, std::vector srcDstPageIndices; }; - std::map batches; + std::map batches; std::set seen; std::set ownerStreams; for (auto const& page : pages) { - if (!page || page->manager != this || requestingCache.storageManager() != this) + if (!page || page->manager != this) { - throw LogicError("Offload pages and requesting cache must belong to the same manager"); + throw LogicError("Sparse migration pages must belong to the same manager"); } if (!seen.insert(page.get()).second) { @@ -1292,14 +1312,21 @@ void StorageManager::offloadSparsePages(KvCache& requestingCache, std::vectoruniqLock.lock() : nullptr; if (!lock) { - throw LogicError("Sparse history offload requires a locked page"); + throw LogicError("Sparse migration requires a locked page"); + } + if (offloadingCache) + { + lock->prepareSparseOffload(*offloadingCache); + } + else + { + lock->prepareSparsePromotion(); } - lock->prepareSparseOffload(requestingCache); - if (page->cacheLevel == kSparseHistoryLevel) + if (page->cacheLevel == dstLevel) { continue; } - auto& batch = batches[getMigrationBatchingLayerGroupId(kSparseHistoryLevel, kHotLevel, page->lifeCycle)]; + auto& batch = batches[getMigrationBatchingLayerGroupId(dstLevel, srcLevel, page->lifeCycle)]; batch.srcPages.push_back(page); for (auto const& owner : lock->owners()) { @@ -1312,13 +1339,12 @@ void StorageManager::offloadSparsePages(KvCache& requestingCache, std::vector requirements(numPoolGroups(kSparseHistoryLevel), 0); + TypedVec requirements(numPoolGroups(dstLevel), 0); for (auto const& [layerGroup, batch] : batches) { - requirements[getPoolGroupIndex(kSparseHistoryLevel, layerGroup)] - += slotCountValueFromSize(batch.srcPages.size()); + requirements[getPoolGroupIndex(dstLevel, layerGroup)] += slotCountValueFromSize(batch.srcPages.size()); } - prepareFreeSlots(kSparseHistoryLevel, requirements, migrationRecorder, dropRecorder); + prepareFreeSlots(dstLevel, requirements, migrationRecorder, dropRecorder); auto releaseDestinations = FuncGuard( [&]() { @@ -1328,14 +1354,14 @@ void StorageManager::offloadSparsePages(KvCache& requestingCache, std::vector(stream); std::vector ownerEvents; ownerEvents.reserve(ownerStreams.size()); @@ -1378,18 +1403,18 @@ void StorageManager::offloadSparsePages(KvCache& requestingCache, std::vectorrecordOffloadEvent(completion); + batch.srcPageLocks[i]->recordMigrationEvent(completion); } } }); for (auto const& [layerGroup, batch] : batches) { - submitMigrationBatch(kSparseHistoryLevel, kHotLevel, layerGroup, batch.srcDstPageIndices.data(), - batch.srcDstPageIndices.size(), stream); + submitMigrationBatch( + dstLevel, srcLevel, layerGroup, batch.srcDstPageIndices.data(), batch.srcDstPageIndices.size(), stream); } fenceCopies.run(); - // Subsequent host readers on every owner's stream must observe the completed copy. + // Subsequent readers on every owner's stream must observe the completed copy. for (auto const ownerStream : ownerStreams) { completion.waitInStream(reinterpret_cast(ownerStream)); @@ -1398,15 +1423,15 @@ void StorageManager::offloadSparsePages(KvCache& requestingCache, std::vectormoveToSparseHistory(std::move(batch.dstSlots[i])); - releaseSlot(batch.srcPages[i]->lifeCycle, kHotLevel, std::move(source)); + Slot source = batch.srcPageLocks[i]->moveToCacheLevel(dstLevel, std::move(batch.dstSlots[i])); + releaseSlot(batch.srcPages[i]->lifeCycle, srcLevel, std::move(source)); } } if (mEventSink) @@ -1421,7 +1446,7 @@ void StorageManager::offloadSparsePages(KvCache& requestingCache, std::vectorisOrphan() && block->holdsPage(committed)) { - mEventSink->addCacheLevelUpdated(block->key, kHotLevel, kSparseHistoryLevel, page->lifeCycle); + mEventSink->addCacheLevelUpdated(block->key, srcLevel, dstLevel, page->lifeCycle); } } } @@ -1494,6 +1519,24 @@ int64_t StorageManager::prefetch( for (auto& [migrationPath, migrationPages] : migrationGroups) { CacheLevel const srcLevel = migrationPath.first; + if (dstLevel == kHotLevel && srcLevel == kSparseHistoryLevel) + { + std::vector> lockedPages; + for (auto const& page : migrationPages) + { + if (page->status() == PageStatus::LOCKED) + { + lockedPages.push_back(page); + } + } + if (!lockedPages.empty()) + { + TemporaryCudaStream stream({}); + auto scope = stream.enter(); + promoteSparsePages(stream.get(), lockedPages); + std::erase_if(migrationPages, [](auto const& page) { return page->cacheLevel == kHotLevel; }); + } + } _batchedMigrate(dstLevel, srcLevel, migrationPages, /*updateSrc=*/true); // Per batch, after it landed: the pages are already grouped by source level, so this costs // nothing per page and never credits a batch that did not run. 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 684d34389b8c..acb5e418e9a8 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 @@ -197,9 +197,13 @@ class StorageManager : public std::enable_shared_from_this void offloadSparsePages(KvCache& requestingCache, std::vector> const& pages, MigrationRecorder const& migrationRecorder = {}, DropRecorder const& dropRecorder = {}); + //! Promote locked sparse host pages into one shared GPU allocation per page, updating all owners. + //! Caller holds the manager's exclusive lock. Allocation/copy failure preserves host ownership. + void promoteSparsePages(CUstream stream, 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. + // it moved off the disk tier. If a later batch fails, earlier batches may already have migrated. int64_t prefetch( CacheLevel dstLevel, TypedVec>>> const& pages); @@ -382,6 +386,8 @@ class StorageManager : public std::enable_shared_from_this std::optional> _batchedMigrate(CacheLevel dstLevel, CacheLevel srcLevel, std::vector> const& srcPages, bool updateSrc, MigrationRecorder const& migrationRecorder = {}, bool defrag = false); + void _migrateLockedSparsePages(CUstream stream, CacheLevel dstLevel, std::vector> const& pages, + KvCache const* offloadingCache, MigrationRecorder const& migrationRecorder, DropRecorder const& dropRecorder); [[nodiscard]] LayerGroupId getMigrationBatchingLayerGroupId( CacheLevel dstLevel, CacheLevel srcLevel, LifeCycleId lifeCycle) const; void submitMigrationBatch(CacheLevel dstLevel, CacheLevel srcLevel, LayerGroupId batchingLayerGroupId, diff --git a/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp b/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp index 7cbfc5a68ed8..d099a75513fa 100644 --- a/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp +++ b/cpp/tests/unit_tests/batch_manager/kvCacheManagerV2ColdPageTest.cpp @@ -340,12 +340,18 @@ class ObservingColdPageCodec final : public IKvCacheColdPageCodec 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); + ++decodeCalls; + decodedPages += count; + bool const submitted = mCodec->decode(layerGroup, source, indices, count, stream); + return submitted && decodeCalls != rejectDecodeCall; } size_t encodeCalls = 0; size_t encodedPages = 0; size_t rejectEncodeCall = 0; + size_t decodeCalls = 0; + size_t decodedPages = 0; + size_t rejectDecodeCall = 0; cudaStream_t encodeStream{}; private: @@ -2343,6 +2349,98 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, HistoryUpdatesOffloadOnlyNewFullPages) EXPECT_EQ(observer->encodedPages, 2); } +TEST_F(KvCacheManagerV2BatchTest, PrefillPromotesOneSharedPageAndReleaseRetriesOffload) +{ + for (int path : {0, 1, 2, 3}) + { + SCOPED_TRACE(path); + 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 first = manager->createKvCache({}, tokens()); + auto second = manager->createKvCache({}, tokens()); + std::shared_ptr prefill; + Batch batch(manager, 3, 1); + auto closeCaches = FuncGuard( + [&]() + { + if (prefill) + { + prefill->close(); + } + first->close(); + second->close(); + }); + ASSERT_TRUE(first->resume(stream())); + ASSERT_TRUE(second->resume(stream())); + if (path == 1) + { + prefill = manager->createKvCache({}, tokens()); + ASSERT_TRUE(prefill->resume(stream())); + prefill->suspend(); + } + auto const lc = page->lifeCycle; + auto const pg = storage.getPoolGroupIndex(kHotLevel, lc); + size_t const bytes = storage.slotSize(kHotLevel, pg)[PoolIndex{0}]; + auto const originalAddress + = std::get(storage.slotAddress(kHotLevel, pg, page->slotId(), PoolIndex{0})); + ASSERT_EQ(cudaMemsetAsync(reinterpret_cast(originalAddress), 0x3C, bytes, mStream), cudaSuccess); + ASSERT_TRUE(first->enterDecode()); + ASSERT_TRUE(second->enterDecode()); + ASSERT_EQ(page->cacheLevel, kSparseHistoryLevel); + batch.add(*first, 0); + batch.add(*second, 1); + batch.publish(batchStream()); + auto const firstVersion = first->pageStorageVersion(); + auto const secondVersion = second->pageStorageVersion(); + + if (!prefill) + { + prefill = path == 2 ? manager->createKvCache() : manager->createKvCache({}, tokens()); + } + if (path == 3) + { + ASSERT_TRUE(prefill->prefetch(kHotLevel)); + } + ASSERT_TRUE(prefill->resume(stream())); + if (path == 2) + { + ASSERT_TRUE(prefill->resize(4, 4)); + prefill->commit(tokens()); + } + EXPECT_EQ(page->cacheLevel, kHotLevel); + EXPECT_EQ(observer->decodeCalls, 1); + EXPECT_EQ(observer->decodedPages, 1); + EXPECT_GT(first->pageStorageVersion(), firstVersion); + EXPECT_GT(second->pageStorageVersion(), secondVersion); + for (auto const& cache : {first, second, prefill}) + { + EXPECT_EQ(pageAt(*cache), page); + EXPECT_EQ(cache->getBasePageIndices(lc)[0], slotIdToPageIndexValue(page->slotId())); + } + EXPECT_EQ(storage.getStatistics(kHotLevel).free, storage.getStatistics(kHotLevel).total - 1); + EXPECT_EQ(storage.getStatistics(kSparseHistoryLevel).free, storage.getStatistics(kSparseHistoryLevel).total); + batch.add(*prefill, 2); + EXPECT_EQ(batch.publish(batchStream()), (std::vector{0, 1, 2})); + EXPECT_EQ( + read(batch.pageTableAddress(lc), 3), (std::vector(3, slotIdToPageIndexValue(page->slotId())))); + EXPECT_EQ(read(batch.numBlocksAddress(lc), 3), (std::vector{0, 0, 0})); + auto const address = std::get(storage.slotAddress(kHotLevel, pg, page->slotId(), PoolIndex{0})); + auto const data = read(address, bytes / sizeof(int32_t)); + EXPECT_TRUE(std::all_of(data.begin(), data.end(), [](int32_t value) { return value == 0x3C3C3C3C; })); + + prefill->close(); + batch.publish(batchStream()); + EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(observer->encodeCalls, 2); + EXPECT_EQ(storage.getStatistics(kHotLevel).free, storage.getStatistics(kHotLevel).total); + EXPECT_EQ(read(batch.numBlocksAddress(lc), 3), (std::vector{1, 1, 0})); + } +} + TEST_F(KvCacheManagerV2DecodeOffloadTest, SharedPrefillOwnerDefersDemotionAcrossDecodeResume) { auto manager = std::make_shared(sparseConfig()); @@ -2366,11 +2464,12 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, SharedPrefillOwnerDefersDemotionAcross EXPECT_EQ(first->getPageStorageSnapshot(LifeCycleId{0}).eligibleHistoryBlocks(), 0); second->suspend(); ASSERT_TRUE(first->enterDecode()); - EXPECT_THROW(second->resume(), LogicError); - EXPECT_FALSE(second->isActive()); + ASSERT_TRUE(second->resume()); + EXPECT_TRUE(second->isActive()); EXPECT_FALSE(second->isDecoding()); + EXPECT_EQ(page->cacheLevel, kHotLevel); + EXPECT_EQ(pageAt(*second), page); first->suspend(); - ASSERT_TRUE(second->resume()); EXPECT_EQ(page->cacheLevel, kHotLevel); ASSERT_TRUE(first->resume()); EXPECT_TRUE(first->isActive()); @@ -2382,6 +2481,242 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, SharedPrefillOwnerDefersDemotionAcross EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); } +TEST_F(KvCacheManagerV2DecodeOffloadTest, PromotionGpuOomPreservesHostPageAndAllowsRetry) +{ + auto config = sparseConfig(); + config.maxUtilForResume = 1.0f; + 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 page = seedPrefix(*manager, kHotLevel); + auto decoder = manager->createKvCache({}, tokens()); + auto prefill = manager->createKvCache({}, tokens()); + auto closeCaches = FuncGuard( + [&]() + { + prefill->close(); + decoder->close(); + }); + ASSERT_TRUE(decoder->resume(stream())); + ASSERT_TRUE(decoder->enterDecode()); + auto const hostSlot = page->slotId(); + auto const version = decoder->pageStorageVersion(); + auto const lc = page->lifeCycle; + auto blockers = storage.newGpuSlots(TypedVec{storage.getStatistics(kHotLevel).free}); + auto releaseBlockers = FuncGuard( + [&]() + { + for (auto& slot : blockers[lc]) + { + storage.releaseSlot(lc, kHotLevel, std::move(slot)); + } + }); + EXPECT_FALSE(prefill->resume(stream())); + EXPECT_FALSE(prefill->isActive()); + EXPECT_TRUE(decoder->isActive()); + EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(page->slotId(), hostSlot); + EXPECT_EQ(decoder->pageStorageVersion(), version); + EXPECT_EQ(decoder->getBasePageIndices(lc)[0], slotIdToPageIndexValue(hostSlot)); + EXPECT_EQ(observer->decodeCalls, 0); + releaseBlockers.run(); + ASSERT_TRUE(prefill->resume()); + EXPECT_EQ(pageAt(*prefill), page); + EXPECT_EQ(page->cacheLevel, kHotLevel); + EXPECT_EQ(observer->decodedPages, 1); +} + +TEST_F(KvCacheManagerV2DecodeOffloadTest, PromotionCodecBatchFailurePreservesAllHostPages) +{ + auto config = makeSplitColdGroupingConfig(); + config.enableStats = true; + 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(); + auto manager = std::make_shared(std::move(config), nullptr, std::move(codec)); + auto const apiLock = manager->lockExclusive(); + auto& storage = manager->storage(); + auto decoder = manager->createKvCache(); + std::shared_ptr prefill; + auto closeCaches = FuncGuard( + [&]() + { + if (prefill) + { + prefill->close(); + } + decoder->close(); + }); + ASSERT_TRUE(decoder->resume(stream())); + ASSERT_TRUE(decoder->resize(4, 4)); + decoder->commit(tokens()); + ASSERT_TRUE(decoder->enterDecode()); + prefill = manager->createKvCache({}, tokens()); + auto const first = pageAt(*decoder, 0, LifeCycleId{0}); + auto const second = pageAt(*decoder, 0, LifeCycleId{1}); + auto const firstSlot = first->slotId(); + auto const secondSlot = second->slotId(); + auto const version = decoder->pageStorageVersion(); + manager->getAndResetIterationStats(); + observer->rejectDecodeCall = 2; + EXPECT_THROW(prefill->resume(stream()), TllmException); + EXPECT_FALSE(prefill->isActive()); + EXPECT_EQ(observer->decodeCalls, 2); + EXPECT_EQ(first->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(second->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(first->slotId(), firstSlot); + EXPECT_EQ(second->slotId(), secondSlot); + EXPECT_EQ(decoder->getBasePageIndices(LifeCycleId{0})[0], slotIdToPageIndexValue(firstSlot)); + EXPECT_EQ(decoder->getBasePageIndices(LifeCycleId{1})[0], slotIdToPageIndexValue(secondSlot)); + EXPECT_GT(decoder->pageStorageVersion(), version); + EXPECT_EQ(storage.getStatistics(kHotLevel).free, storage.getStatistics(kHotLevel).total); + EXPECT_TRUE(manager->getAndResetIterationStats().empty()); + observer->rejectDecodeCall = 0; + ASSERT_TRUE(prefill->resume()); + EXPECT_EQ(pageAt(*prefill, 0, LifeCycleId{0}), first); + EXPECT_EQ(pageAt(*prefill, 0, LifeCycleId{1}), second); + EXPECT_EQ(first->cacheLevel, kHotLevel); + EXPECT_EQ(second->cacheLevel, kHotLevel); + EXPECT_EQ(storage.getStatistics(kSparseHistoryLevel).free, storage.getStatistics(kSparseHistoryLevel).total); + auto const stats = manager->getAndResetIterationStats(); + EXPECT_EQ(stats.at(LifeCycleId{0}).iterOnboardBlocks, 1); + EXPECT_EQ(stats.at(LifeCycleId{1}).iterOnboardBlocks, 1); +} + +TEST_F(KvCacheManagerV2DecodeOffloadTest, PromotionRejectionFencesSourceAndRecycledDestination) +{ + auto codec = std::make_unique(AsyncRejectingColdPageCodec::Operation::kDecode); + auto* rejecting = 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( + [&]() + { + prefill->close(); + decoder->close(); + }); + ASSERT_TRUE(decoder->resume(stream())); + ASSERT_TRUE(decoder->enterDecode()); + auto const lc = page->lifeCycle; + auto const hostSlot = page->slotId(); + auto blocker = storage.newGpuSlots(TypedVec{1}); + auto releaseBlocker = FuncGuard([&]() { storage.releaseSlot(lc, kHotLevel, std::move(blocker[lc].front())); }); + auto releaseCodec = FuncGuard([&]() { rejecting->release(); }); + EXPECT_THROW(prefill->resume(stream()), TllmException); + ASSERT_TRUE(rejecting->launched()); + EXPECT_FALSE(prefill->isActive()); + EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); + EXPECT_EQ(page->slotId(), hostSlot); + EXPECT_FALSE(page->queryReady()); + EXPECT_EQ(decoder->getBasePageIndices(lc)[0], slotIdToPageIndexValue(hostSlot)); + auto recycled = storage.newGpuSlots(TypedVec{1}); + auto releaseRecycled = FuncGuard([&]() { storage.releaseSlot(lc, kHotLevel, std::move(recycled[lc].front())); }); + EXPECT_FALSE(recycled[lc].front().queryReady()); + rejecting->release(); + recycled[lc].front().readyEvent.synchronize(); +} + +TEST_F(KvCacheManagerV2DecodeOffloadTest, PromotionWaitsForLiveAndFinishedReadersBeforeRecyclingHostSlot) +{ + 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 prefill = manager->createKvCache({}, tokens()); + auto closeCaches = FuncGuard( + [&]() + { + prefill->close(); + first->close(); + second->close(); + }); + ASSERT_TRUE(first->resume(stream())); + ASSERT_TRUE(second->resume(reinterpret_cast(readerStream))); + auto const lc = page->lifeCycle; + auto const gpuPool = storage.getPoolGroupIndex(kHotLevel, lc); + auto const hostPool = storage.getPoolGroupIndex(kSparseHistoryLevel, lc); + size_t const bytes = storage.slotSize(kHotLevel, gpuPool)[PoolIndex{0}]; + auto const originalAddress + = std::get(storage.slotAddress(kHotLevel, gpuPool, page->slotId(), PoolIndex{0})); + ASSERT_EQ(cudaMemsetAsync(reinterpret_cast(originalAddress), 0x3C, bytes, mStream), cudaSuccess); + ASSERT_TRUE(first->enterDecode()); + ASSERT_TRUE(second->enterDecode()); + auto const hostSlot = page->slotId(); + auto const hostAddress + = std::get(storage.slotAddress(kSparseHistoryLevel, hostPool, hostSlot, PoolIndex{0})); + auto readback = storage.newGpuSlots(TypedVec{1}); + auto hostBlocker = storage.newSlots(kSparseHistoryLevel, TypedVec{1}); + auto releaseSlots = FuncGuard( + [&]() + { + storage.releaseSlot(lc, kHotLevel, std::move(readback[lc].front())); + storage.releaseSlot(lc, kSparseHistoryLevel, std::move(hostBlocker[lc].front())); + }); + auto const readbackAddress = std::get( + storage.slotAddress(kHotLevel, gpuPool, readback[lc].front().slotId(), PoolIndex{0})); + // Warm the H2D codec before holding a CUDA callback. + storage.copySlotData(lc, kHotLevel, kSparseHistoryLevel, readback[lc].front().slotId(), hostSlot, stream()); + ASSERT_EQ(cudaStreamSynchronize(mStream), cudaSuccess); + std::pair overwrite{hostAddress, bytes}; + StreamGate gate; + ASSERT_EQ(gate.enqueue(readerStream), cudaSuccess); + ASSERT_EQ(cudaMemcpyAsync(reinterpret_cast(readbackAddress), reinterpret_cast(hostAddress), + bytes, cudaMemcpyHostToDevice, readerStream), + cudaSuccess); + if (finishReader) + { + second->suspend(); + } + ASSERT_TRUE(prefill->resume(stream())); + EXPECT_EQ(page->cacheLevel, kHotLevel); + EXPECT_FALSE(page->queryReady()); + EXPECT_EQ(pageAt(*prefill), page); + 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()); + recycled[lc].front().readyEvent.waitInStream(reinterpret_cast(mStream)); + ASSERT_EQ(cudaLaunchHostFunc( + mStream, + [](void* data) + { + auto const& [address, size] = *static_cast const*>(data); + std::memset(reinterpret_cast(address), 0, size); + }, + &overwrite), + cudaSuccess); + gate.release(); + ASSERT_EQ(cudaStreamSynchronize(mStream), cudaSuccess); + auto const promotedAddress + = std::get(storage.slotAddress(kHotLevel, gpuPool, page->slotId(), PoolIndex{0})); + for (auto const address : {readbackAddress, promotedAddress}) + { + std::vector data(bytes); + cuCheck(cuMemcpyDtoH(data.data(), address, bytes)); + EXPECT_TRUE(std::all_of(data.begin(), data.end(), [](uint8_t value) { return value == 0x3C; })); + } + } +} + TEST_F(KvCacheManagerV2DecodeOffloadTest, SharedOwnersEnterDecodeWithoutMutuallyBlocking) { auto config = sparseConfig(); @@ -2629,7 +2964,7 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, RebaseOffloadsOlderGpuHistoryAtUnchang EXPECT_GT(cache->pageStorageVersion(), version); } -TEST_F(KvCacheManagerV2DecodeOffloadTest, PrefillRebaseCannotAdoptSharedHostIndices) +TEST_F(KvCacheManagerV2DecodeOffloadTest, PrefillRebasePromotesSharedHostPage) { auto manager = std::make_shared(sparseConfig()); auto const apiLock = manager->lockExclusive(); @@ -2645,17 +2980,20 @@ TEST_F(KvCacheManagerV2DecodeOffloadTest, PrefillRebaseCannotAdoptSharedHostIndi 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_NO_THROW(prefill->commit(tokens())); 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(prefill->numCommittedTokens(), 4); + EXPECT_EQ(pageAt(*prefill), page); + EXPECT_EQ(page->cacheLevel, kHotLevel); + EXPECT_EQ(page->status(), PageStatus::LOCKED); + EXPECT_EQ(prefill->getBasePageIndices(LifeCycleId{0})[0], slotIdToPageIndexValue(page->slotId())); EXPECT_EQ(decoder->getBasePageIndices(LifeCycleId{0})[0], slotIdToPageIndexValue(page->slotId())); + EXPECT_EQ(manager->storage().getStatistics(kHotLevel).free, manager->storage().getStatistics(kHotLevel).total - 1); + EXPECT_EQ(manager->storage().getStatistics(kSparseHistoryLevel).free, + manager->storage().getStatistics(kSparseHistoryLevel).total); EXPECT_NO_THROW(prefill->close()); + ASSERT_TRUE(decoder->enterDecode()); + EXPECT_EQ(page->cacheLevel, kSparseHistoryLevel); } TEST_F(KvCacheManagerV2DecodeOffloadTest, RebaseOomPreservesMissingLifecyclePagesAndCanRetryCommit) From 999811c2c58cad268f5ba68cc1a94ee0a0dd81f1 Mon Sep 17 00:00:00 2001 From: Guiju Zhang <7135567+cascade812@users.noreply.github.com> Date: Mon, 5 Oct 2026 15:39:36 -0700 Subject: [PATCH 7/7] [None][fix] check sparse KV residency before dense offsets Check page-storage snapshots after sparse metadata publication instead of inferring offload from decode history length. Allow GPU-resident sparse history and reject mapped host pages across scheduled requests and sparse layer groups. Add regressions for GPU-only and mixed mappings, invalid slots, per-layer offsets, and publication-triggered offload. Validation: reproduced four false-rejection cases before the fix; 18 focused mock cases and pre-commit checks pass. Full pytest collection remains blocked by missing tensorrt_llm.bindings. Signed-off-by: Guiju Zhang <7135567+cascade812@users.noreply.github.com> --- .../kv_cache/kv_cache_manager_v2.py | 20 +++- .../kv_cache/test_kv_cache_v2_scheduler.py | 108 +++++++++++++++++- 2 files changed, 122 insertions(+), 6 deletions(-) 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 f8fbda1a5fe3..a430603d5ebe 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 @@ -1752,7 +1752,16 @@ 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_layer_group_ids = tuple( + sorted( + { + self.layer_to_pool_mapping_dict[buf.layer_id] + for buf in self.impl.all_buffer_ids + if self.impl.is_sparse(buf.layer_id, buf.role) + } + ) + ) + if self._sparse_layer_group_ids: self.sparse_metadata_batch = Batch( self.impl, index_mapper_capacity, self.max_blocks_per_seq, self.max_beam_width ) @@ -5771,10 +5780,15 @@ def copy_batch_block_offsets( max_blocks: Optional[int] = None, ): self._publish_sparse_metadata() + # Sparse history can remain on GPU while offload is deferred. Check all + # mapped pages after publication, which can retry the deferred offload. 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 + level is not None and level != GPU_LEVEL for req_id in request_ids + for layer_group_id in self._sparse_layer_group_ids + for level in self.kv_cache_map[req_id] + .get_page_storage_snapshot(layer_group_id) + .cache_levels ): raise RuntimeError( "Offloaded sparse history requires Batch metadata and sparse fetch; " 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 1072dff66d98..390c586cdb9c 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 @@ -101,19 +101,121 @@ def test_sparse_metadata_republishes_after_connector_acceptance() -> None: ] -def test_sparse_host_indices_cannot_reach_dense_attention_offsets() -> None: +@pytest.fixture +def sparse_offset_manager() -> KVCacheManagerV2: manager = object.__new__(KVCacheManagerV2) manager._stream = Mock(cuda_stream=123) manager.sparse_metadata_batch = Mock() + manager._sparse_layer_group_ids = (1, 3) manager.tokens_per_block = 4 - manager.kv_cache_map = {7: Mock(is_decoding=True, history_length=4)} + manager.kv_cache_map = {req_id: Mock(is_decoding=True, history_length=8) for req_id in (7, 8)} + for cache in manager.kv_cache_map.values(): + cache.get_page_storage_snapshot.return_value = Mock( + cache_levels=[0, 0], eligible_history_blocks=0 + ) + manager._use_per_layer_page_tables = False + manager._copy_batch_block_offsets_per_layer = Mock() + manager.index_mapper = Mock() + manager.index_mapper.get_copy_index.return_value = Mock(shape=(2,)) + manager.host_kv_cache_block_offsets = Mock() + manager.index_scales = Mock() + manager.kv_offset = Mock() + return manager + + +@pytest.mark.parametrize("per_layer", [False, True]) +@pytest.mark.parametrize( + "cache_levels", + [pytest.param([0, 0], id="gpu"), pytest.param([None, 0, None], id="invalid-slots")], +) +def test_sparse_gpu_resident_history_uses_dense_attention_offsets( + sparse_offset_manager: KVCacheManagerV2, cache_levels: list[int | None], per_layer: bool +) -> None: + manager = sparse_offset_manager + manager._use_per_layer_page_tables = per_layer + manager.kv_cache_map[8].get_page_storage_snapshot.return_value.cache_levels = cache_levels + unscheduled_cache = Mock(is_decoding=True, history_length=8) + unscheduled_cache.get_page_storage_snapshot.return_value = Mock(cache_levels=[1, 1]) + manager.kv_cache_map[9] = unscheduled_cache + 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, 8], 1, 0, 2) + if per_layer: + manager._copy_batch_block_offsets_per_layer.assert_called_once_with( + destination, [7, 8], manager.index_mapper.get_copy_index.return_value, 0, 2 + ) + dense_copy.assert_not_called() + else: + 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, + ) + manager._copy_batch_block_offsets_per_layer.assert_not_called() + for req_id in (7, 8): + assert manager.kv_cache_map[req_id].get_page_storage_snapshot.call_args_list == [ + call(1), + call(3), + ] + unscheduled_cache.get_page_storage_snapshot.assert_not_called() + manager.sparse_metadata_batch.publish.assert_called_once_with(123) + + +@pytest.mark.parametrize( + ("cache_levels", "eligible_history_blocks"), + [ + pytest.param([1, 1], 2, id="host"), + pytest.param([0, 1], 0, id="host-after-gpu"), + pytest.param([None, 1], 0, id="host-after-invalid-slot"), + ], +) +def test_sparse_host_indices_cannot_reach_dense_attention_offsets( + sparse_offset_manager: KVCacheManagerV2, + cache_levels: list[int | None], + eligible_history_blocks: int, +) -> None: + manager = sparse_offset_manager + snapshots = { + 1: Mock(cache_levels=[0, 0], eligible_history_blocks=0), + 3: Mock(cache_levels=cache_levels, eligible_history_blocks=eligible_history_blocks), + } + manager.kv_cache_map[8].get_page_storage_snapshot.side_effect = snapshots.__getitem__ 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) + manager.copy_batch_block_offsets(Mock(), [7, 8], 1, 0, 2) + dense_copy.assert_not_called() + manager._copy_batch_block_offsets_per_layer.assert_not_called() + manager.sparse_metadata_batch.publish.assert_called_once_with(123) + + +def test_sparse_dense_offsets_check_residency_after_publication( + sparse_offset_manager: KVCacheManagerV2, +) -> None: + manager = sparse_offset_manager + snapshot = manager.kv_cache_map[8].get_page_storage_snapshot.return_value + + def offload_history(stream: int) -> None: + snapshot.cache_levels = [1, 1] + snapshot.eligible_history_blocks = 2 + + manager.sparse_metadata_batch.publish.side_effect = offload_history + 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, 8], 1, 0, 2) dense_copy.assert_not_called() + manager._copy_batch_block_offsets_per_layer.assert_not_called() manager.sparse_metadata_batch.publish.assert_called_once_with(123)