diff --git a/cpp/include/tensorrt_llm/batch_manager/cacheTransceiver.h b/cpp/include/tensorrt_llm/batch_manager/cacheTransceiver.h index 80a06ea5ddf7..a3b8cd348ce4 100644 --- a/cpp/include/tensorrt_llm/batch_manager/cacheTransceiver.h +++ b/cpp/include/tensorrt_llm/batch_manager/cacheTransceiver.h @@ -237,7 +237,6 @@ class CacheTransceiver : public BaseCacheTransceiver executor::kv_cache::CacheState::AttentionType attentionType = executor::kv_cache::CacheState::AttentionType::kDEFAULT, std::optional cacheTransceiverConfig = std::nullopt, - rnn_state_manager::RnnStateManager* rnnStateManager = nullptr, std::vector const& rnnLayerNumPerPP = {}); CacheTransceiver(kv_cache_manager::BaseKVCacheManager* cacheManager, std::vector numKvHeadsPerLayer, @@ -246,11 +245,10 @@ class CacheTransceiver : public BaseCacheTransceiver executor::kv_cache::CacheState::AttentionType attentionType = executor::kv_cache::CacheState::AttentionType::kDEFAULT, std::optional cacheTransceiverConfig = std::nullopt, - rnn_state_manager::RnnStateManager* rnnStateManager = nullptr, std::vector const& rnnLayerNumPerPP = {}) : CacheTransceiver(cacheManager, executor::kv_cache::CacheState::ModelConfig{numKvHeadsPerLayer, sizePerHead, tokensPerBlock}, worldConfig, - attentionLayerNumPerPP, dataType, attentionType, cacheTransceiverConfig, rnnStateManager, rnnLayerNumPerPP) + attentionLayerNumPerPP, dataType, attentionType, cacheTransceiverConfig, rnnLayerNumPerPP) { } @@ -306,7 +304,6 @@ class CacheTransceiver : public BaseCacheTransceiver std::vector> mCacheTransBufferManagers; std::vector mCacheTransBufferManagerPtrs; - rnn_state_manager::RnnStateManager* mRnnStateManager{nullptr}; // TODO(shreyasm): update this to use same container as kv by using base trans buffers instead std::unique_ptr mRnnCacheTransBufferManager{nullptr}; diff --git a/cpp/include/tensorrt_llm/batch_manager/rnnCacheFormatter.h b/cpp/include/tensorrt_llm/batch_manager/rnnCacheFormatter.h index 1bc3fadf62f4..b57d08f2c12b 100644 --- a/cpp/include/tensorrt_llm/batch_manager/rnnCacheFormatter.h +++ b/cpp/include/tensorrt_llm/batch_manager/rnnCacheFormatter.h @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); @@ -34,14 +34,11 @@ class TransferSession; namespace rnn_state_manager { -class RnnStateManager; class RnnCacheTransBufferManager; } // namespace rnn_state_manager /// @brief RNN Cache Formatter for formatting/unformatting RNN states during transfer. -/// Supports two operating modes: -/// - Slot mode: uses RnnStateManager (for CppMambaCacheManager, separate tensor storage) -/// - Unified pool mode: uses BaseKVCacheManager (for CppMambaHybridCacheManager, block-indexed pool) +/// Uses unified pool mode via BaseKVCacheManager (for CppMambaHybridCacheManager, block-indexed pool). class RnnCacheFormatter : public kv_cache_manager::BaseCacheFormatter { public: @@ -49,12 +46,6 @@ class RnnCacheFormatter : public kv_cache_manager::BaseCacheFormatter using CacheState = executor::kv_cache::CacheState; using RequestIdType = tensorrt_llm::batch_manager::RequestIdType; - /// @brief Constructor for slot-based mode (CppMambaCacheManager with RnnStateManager). - /// @param rnnStateManager Pointer to the RNN state manager. - /// @param rnnCacheTransBufferManager Pointer to the RNN cache transfer buffer manager. - RnnCacheFormatter(rnn_state_manager::RnnStateManager* rnnStateManager, - rnn_state_manager::RnnCacheTransBufferManager* rnnCacheTransBufferManager); - /// @brief Constructor for unified pool mode (CppMambaHybridCacheManager). /// @param kvCacheManager Pointer to the KV cache manager with unified pool. /// @param rnnCacheTransBufferManager Pointer to the RNN cache transfer buffer manager. @@ -81,39 +72,13 @@ class RnnCacheFormatter : public kv_cache_manager::BaseCacheFormatter CacheState const& selfConfig, SizeType32 selfIdx, CacheState const& destConfig, std::vector const& counterPartRanks) const override; - /// @brief Returns the KV cache manager (non-null in unified pool mode). + /// @brief Returns the KV cache manager. [[nodiscard]] kv_cache_manager::BaseKVCacheManager* getCacheManager() const noexcept override { return mKvCacheManager; } - /// @brief Get the RNN state manager (non-null in slot mode). - /// @return Pointer to the RNN state manager. - [[nodiscard]] rnn_state_manager::RnnStateManager* getRnnStateManager() const noexcept - { - return mRnnStateManager; - } - - /// @brief Check if operating in unified pool mode. - [[nodiscard]] bool isUnifiedPoolMode() const noexcept - { - return mKvCacheManager != nullptr; - } - private: - /// @brief Format logic for slot-based path (RnnStateManager). - void formatSlotMode(TransferSession& session); - - /// @brief Unformat logic for slot-based path (RnnStateManager). - void unformatSlotMode(TransferSession& session); - - /// @brief Format logic for unified pool path (BaseKVCacheManager). - void formatUnifiedPoolMode(TransferSession& session); - - /// @brief Unformat logic for unified pool path (BaseKVCacheManager). - void unformatUnifiedPoolMode(TransferSession& session); - - rnn_state_manager::RnnStateManager* mRnnStateManager{nullptr}; rnn_state_manager::RnnCacheTransBufferManager* mRnnCacheTransBufferManager; kv_cache_manager::BaseKVCacheManager* mKvCacheManager{nullptr}; }; diff --git a/cpp/tensorrt_llm/batch_manager/cacheTransceiver.cpp b/cpp/tensorrt_llm/batch_manager/cacheTransceiver.cpp index 2c69af0059be..7ef2f4cd4def 100644 --- a/cpp/tensorrt_llm/batch_manager/cacheTransceiver.cpp +++ b/cpp/tensorrt_llm/batch_manager/cacheTransceiver.cpp @@ -293,9 +293,8 @@ CacheTransceiver::CacheTransceiver(kv_cache_manager::BaseKVCacheManager* cacheMa std::vector const& attentionLayerNumPerPP, nvinfer1::DataType dataType, executor::kv_cache::CacheState::AttentionType attentionType, std::optional cacheTransceiverConfig, - rnn_state_manager::RnnStateManager* rnnStateManager, std::vector const& rnnLayerNumPerPP) + std::vector const& rnnLayerNumPerPP) : mCacheTransceiverConfig{cacheTransceiverConfig} - , mRnnStateManager{rnnStateManager} { using tensorrt_llm::batch_manager::kv_cache_manager::CacheFormatter; if (useMPI()) @@ -364,28 +363,9 @@ CacheTransceiver::CacheTransceiver(kv_cache_manager::BaseKVCacheManager* cacheMa std::make_unique(cacheManager, maxNumTokens, true)); } - // RNN specific setup - if (mRnnStateManager != nullptr) - { - TLLM_LOG_DEBUG("Setting up RNN cache transfer components."); - TLLM_CHECK(!rnnLayerNumPerPP.empty()); - - mRnnCacheTransBufferManager - = std::make_unique(mRnnStateManager, maxNumTokens); - - auto rnnModelCfg = mRnnStateManager->getRnnCacheStateModelConfig(); - - auto const convStateDataType = mRnnStateManager->getConvStateDataType(); - auto const ssmStateDataType = mRnnStateManager->getSsmStateDataType(); - - mCacheState->setRnnConfig(rnnModelCfg, rnnLayerNumPerPP, convStateDataType, ssmStateDataType); - - TLLM_LOG_INFO("RNN cache transfer components initialized."); - } - // Unified pool path (CppMambaHybridCacheManager): build RnnModelConfig from - // LinearAttentionMetadata. Detected by rnnLayerNumPerPP set but no RnnStateManager. - if (mRnnStateManager == nullptr && !rnnLayerNumPerPP.empty()) + // LinearAttentionMetadata. Detected by rnnLayerNumPerPP being non-empty. + if (!rnnLayerNumPerPP.empty()) { auto const& blockManager = cacheManager->getBlockManager(); auto const& linearMeta = blockManager.getLinearAttentionMetadata(); @@ -502,11 +482,6 @@ CacheTransceiver::CacheTransceiver(kv_cache_manager::BaseKVCacheManager* cacheMa auto makeRnnFormatter = [this, cacheManager]() -> std::unique_ptr { - if (mRnnStateManager != nullptr && mRnnCacheTransBufferManager != nullptr) - { - // Slot-based path (CppMambaCacheManager) - return std::make_unique(mRnnStateManager, mRnnCacheTransBufferManager.get()); - } // Unified pool path (CppMambaHybridCacheManager) if (mCacheState->hasRnnConfig() && mRnnCacheTransBufferManager != nullptr) { diff --git a/cpp/tensorrt_llm/batch_manager/cacheTransferLayer.cpp b/cpp/tensorrt_llm/batch_manager/cacheTransferLayer.cpp index c013fd75c6e4..d0a54dbb7d3c 100644 --- a/cpp/tensorrt_llm/batch_manager/cacheTransferLayer.cpp +++ b/cpp/tensorrt_llm/batch_manager/cacheTransferLayer.cpp @@ -54,8 +54,7 @@ void CacheTransferLayer::validateSupport(executor::DataTransceiverState const& p if (mRnnFormatter && selfHasRnn) { - // Both slot-based (CppMambaCacheManager) and unified pool (CppMambaHybridCacheManager) - // paths now use RnnCacheFormatter. + // Unified pool path (CppMambaHybridCacheManager) uses RnnCacheFormatter. if (peerHasRnn) { TLLM_CHECK_WITH_INFO(mRnnFormatter->inquireSupport(mCacheState, peerState.getCacheState().value()), diff --git a/cpp/tensorrt_llm/batch_manager/rnnCacheFormatter.cpp b/cpp/tensorrt_llm/batch_manager/rnnCacheFormatter.cpp index 05742af17a28..c5e4262e34ab 100644 --- a/cpp/tensorrt_llm/batch_manager/rnnCacheFormatter.cpp +++ b/cpp/tensorrt_llm/batch_manager/rnnCacheFormatter.cpp @@ -20,7 +20,6 @@ #include "tensorrt_llm/batch_manager/dataTransceiver.h" #include "tensorrt_llm/batch_manager/kvCacheManager.h" #include "tensorrt_llm/batch_manager/kvCacheUtils.h" -#include "tensorrt_llm/batch_manager/rnnStateManager.h" #include "tensorrt_llm/common/assert.h" #include "tensorrt_llm/common/dataType.h" #include "tensorrt_llm/common/logger.h" @@ -33,15 +32,6 @@ namespace tensorrt_llm::batch_manager { using CacheState = executor::kv_cache::CacheState; -RnnCacheFormatter::RnnCacheFormatter(rnn_state_manager::RnnStateManager* rnnStateManager, - rnn_state_manager::RnnCacheTransBufferManager* rnnCacheTransBufferManager) - : mRnnStateManager{rnnStateManager} - , mRnnCacheTransBufferManager{rnnCacheTransBufferManager} -{ - TLLM_CHECK(mRnnStateManager != nullptr); - TLLM_CHECK(mRnnCacheTransBufferManager != nullptr); -} - RnnCacheFormatter::RnnCacheFormatter(kv_cache_manager::BaseKVCacheManager* kvCacheManager, rnn_state_manager::RnnCacheTransBufferManager* rnnCacheTransBufferManager) : mRnnCacheTransBufferManager{rnnCacheTransBufferManager} @@ -52,470 +42,6 @@ RnnCacheFormatter::RnnCacheFormatter(kv_cache_manager::BaseKVCacheManager* kvCac } void RnnCacheFormatter::format(TransferSession& session) -{ - if (isUnifiedPoolMode()) - { - formatUnifiedPoolMode(session); - } - else - { - formatSlotMode(session); - } -} - -void RnnCacheFormatter::unformat(TransferSession& session) -{ - if (isUnifiedPoolMode()) - { - unformatUnifiedPoolMode(session); - } - else - { - unformatSlotMode(session); - } -} - -void RnnCacheFormatter::formatSlotMode(TransferSession& session) -{ - NVTX3_SCOPED_RANGE(RnnCacheFormatter_format); - session.setTime(TransferSession::kTimeFormatter); - - auto const& llmRequest = session.getLlmRequest(); - TLLM_LOG_DEBUG( - mpi::MpiComm::world().getRank(), "Start sending RNN state for request ID: %ld.", llmRequest.mRequestId); - TLLM_CHECK_WITH_INFO(llmRequest.mSamplingConfig.beamWidth == 1, "Currently, only beam width 1 is supported."); - - auto const& connections = session.getConnections(); - auto const& selfConfig = session.getSelfState().getCacheState().value(); - auto const& destConfig = session.getOtherState().getCacheState().value(); - auto const selfIdx = session.getSelfState().getCommState().value().getSelfIdx(); - auto& bufferManager = session.getBufferManager(); - - auto targetInfo = executor::kv_cache::targetIRanksForRnn(destConfig, selfConfig, selfIdx); - if (!cache_formatter_utils::needSendCache(selfConfig, destConfig, selfIdx, targetInfo)) - { - return; - } - - auto pickUpConnections = cache_formatter_utils::pickSendConnections( - connections.size(), selfConfig, selfIdx, destConfig, session.getCounterPartRanks(), targetInfo); - auto const targetNum = pickUpConnections.size(); - if (targetNum == 0) - { - TLLM_LOG_DEBUG("No targets to send RNN state to for request ID: %ld", llmRequest.mRequestId); - return; - } - - auto const slotIdx = mRnnStateManager->getCacheIndex(llmRequest.mRequestId); - int deviceId; - TLLM_CUDA_CHECK(cudaGetDevice(&deviceId)); - - auto const& selfParallel = selfConfig.getParallelConfig(); - auto const selfTPNum = selfParallel.mTensorParallelism; - auto const selfPPRank = selfIdx / selfTPNum; - auto const& selfLayersPerPP = selfConfig.getRnnCacheState().mLayerNumPerPP; - SizeType32 const numLocalLayers = selfLayersPerPP[selfPPRank]; - - if (common::getEnvTryZCopyForKVCacheTransfer() && destConfig == selfConfig) - { - TLLM_LOG_DEBUG("Try using zero-copy for the RNN cache."); - NVTX3_SCOPED_RANGE(RnnZeroCopySend); - - TLLM_CHECK(pickUpConnections.size() == 1); - - TLLM_CUDA_CHECK(cudaSetDevice(deviceId)); - for (size_t i = 0; i < pickUpConnections.size(); i++) - { - for (SizeType32 layer = 0; layer < numLocalLayers; layer++) - { - - // Get conv state for this layer: shape is [maxBatchSize, convDim, dConv-1] - auto convState = mRnnStateManager->getConvStates(mRnnStateManager->getGlobalLayerNum(layer)); - // Slice out the specific slot: shape becomes [convDim, dConv-1] - auto slotConv = runtime::ITensor::slice(convState, slotIdx, 1); - slotConv->squeeze(0); - - // Receive conv state - // llmRequest.updateKvCacheSize(slotConv->getSizeInBytes()); - session.send(pickUpConnections[i], slotConv->data(), slotConv->getSizeInBytes()); - - // Get SSM state for this layer: shape is [maxBatchSize, numHeads, headDim, dState] - auto ssmState = mRnnStateManager->getSsmStates(mRnnStateManager->getGlobalLayerNum(layer)); - // Slice out the specific slot: shape becomes [numHeads, headDim, dState] - auto slotSsm = runtime::ITensor::slice(ssmState, slotIdx, 1); - slotSsm->squeeze(0); - - // Receive SSM state - // llmRequest.updateKvCacheSize(slotSsm->getSizeInBytes()); - session.send(pickUpConnections[i], slotSsm->data(), slotSsm->getSizeInBytes()); - } - } - TLLM_LOG_DEBUG(mpi::MpiComm::world().getRank(), "End the sending of RNN cache for the request ID: %ld.", - llmRequest.mRequestId); - - return; - } - - // Calculate buffer sizes for each target - // Each target gets: conv states + ssm states for overlapping layers - auto const& modelConfig = selfConfig.getRnnModelConfig(); - auto const maxBatchSize = mRnnStateManager->getMaxBatchSize(); - int const selfTPSizePerDPGroup = selfConfig.getParallelConfig().mEnableAttentionDP - ? selfTPNum / selfConfig.getParallelConfig().mDPsize - : selfTPNum; - SizeType32 convDimLocal = modelConfig.mConvDimSize / selfTPSizePerDPGroup; - SizeType32 numHeadsLocal = modelConfig.mNumHeads / selfTPSizePerDPGroup; - - size_t convBytesPerLayer - = convDimLocal * (modelConfig.mDConv - 1) * common::getDTypeSize(selfConfig.getConvStateDataType()); - convBytesPerLayer = (convBytesPerLayer + 15) & ~static_cast(15); - size_t ssmBytesPerLayer = numHeadsLocal * modelConfig.mHeadDim * modelConfig.mDState - * common::getDTypeSize(selfConfig.getSsmStateDataType()); - - int peerDuplicateHeadFactor = targetInfo.mPeerDupHeadFactor; - auto bufferTargetNum = targetNum / peerDuplicateHeadFactor; - - std::vector bufferSizesPerTarget(targetNum, 0); - - for (size_t i = 0; i < targetNum; i++) - { - SizeType32 layersForTarget = targetInfo.getPeerPPDomainLayerNum(static_cast(i)); - bufferSizesPerTarget[i] = layersForTarget * (convBytesPerLayer + ssmBytesPerLayer) * peerDuplicateHeadFactor - / targetInfo.mDomainTPSize; - } - - auto cacheBufferId = mRnnCacheTransBufferManager->assignBufferIndexForSend(); - auto allocationResult = mRnnCacheTransBufferManager->getOrAllocateSendBuffers( - cacheBufferId, static_cast(bufferTargetNum), bufferSizesPerTarget, bufferManager); - auto& outputBuffers = std::get<0>(allocationResult); - auto& bufferCoverTargetNum = std::get<1>(allocationResult); - auto& onlyUseDynamicBuffer = std::get<2>(allocationResult); - - TLLM_CHECK(cacheBufferId.has_value() || onlyUseDynamicBuffer); - - auto const* agentConnection - = dynamic_cast(connections[pickUpConnections[0]]); - if (agentConnection != nullptr) - { - TLLM_CHECK_WITH_INFO(bufferCoverTargetNum == bufferTargetNum, "Agent needs all RNN send buffers pre-allocated"); - TLLM_CHECK(onlyUseDynamicBuffer == false); - } - - std::vector inputConvBlocks; - std::vector inputSsmBlocks; - - auto convStates = mRnnStateManager->getConvStates(); // [numLocalLayers, maxBatchSize, convDim, dConv-1] - auto ssmStates = mRnnStateManager->getSsmStates(); // [numLocalLayers, maxBatchSize, numHeads, headDim, dState] - - inputConvBlocks.push_back(convStates); - inputSsmBlocks.push_back(ssmStates); - - tensorrt_llm::executor::rnn_cache::splitRnnConvStateDispatch( - inputConvBlocks, outputBuffers, slotIdx, maxBatchSize, destConfig, selfConfig, selfIdx, bufferManager); - - // Conv and SSM use same output buffer. So need to track convBytesPerLayer to compute offset. - tensorrt_llm::executor::rnn_cache::splitRnnSsmStateDispatch(inputSsmBlocks, outputBuffers, slotIdx, maxBatchSize, - convBytesPerLayer, destConfig, selfConfig, selfIdx, bufferManager); - - bufferManager.getStream().synchronize(); - session.setTime(TransferSession::kTimePreprocess); - - auto preAllocSendBuffer = mRnnCacheTransBufferManager->getSendBuffer(cacheBufferId); - - sendAllBuffers(session, deviceId, outputBuffers, bufferCoverTargetNum, preAllocSendBuffer, bufferManager, - targetInfo, pickUpConnections); - - session.setTime(TransferSession::kTimeTransmissions); - - mRnnCacheTransBufferManager->freeBufferIndexForSend(cacheBufferId); - session.setTime(TransferSession::kTimePostprocess); - - TLLM_LOG_DEBUG( - mpi::MpiComm::world().getRank(), "End sending RNN state for request ID: %ld.", llmRequest.mRequestId); -} - -void RnnCacheFormatter::unformatSlotMode(TransferSession& session) -{ - NVTX3_SCOPED_RANGE(RnnCacheFormatter_unformat); - session.setTime(TransferSession::kTimeFormatter); - - auto& llmRequest = session.getLlmRequest(); - TLLM_LOG_DEBUG( - mpi::MpiComm::world().getRank(), "Start receiving RNN state for request ID: %ld.", llmRequest.mRequestId); - TLLM_CHECK_WITH_INFO(llmRequest.mSamplingConfig.beamWidth == 1, "Currently, only beam width 1 is supported."); - - auto const& connections = session.getConnections(); - auto const& selfConfig = session.getSelfState().getCacheState().value(); - auto const& destConfig = session.getOtherState().getCacheState().value(); - auto const selfIdx = session.getSelfState().getCommState().value().getSelfIdx(); - auto& bufferManager = session.getBufferManager(); - - auto sourceInfo = executor::kv_cache::targetIRanksForRnn(destConfig, selfConfig, selfIdx); - int deviceId; - TLLM_CUDA_CHECK(cudaGetDevice(&deviceId)); - - auto pickRecvConnResult = cache_formatter_utils::pickRecvConnections( - connections.size(), selfConfig, selfIdx, destConfig, session.getCounterPartRanks(), sourceInfo); - auto pickUpConnections = std::get<0>(pickRecvConnResult); - auto localRankIndices = std::get<1>(pickRecvConnResult); - auto const sourceNum = pickUpConnections.size(); - - if (sourceNum == 0) - { - TLLM_LOG_DEBUG("No sources to receive RNN state from for request ID: %ld", llmRequest.mRequestId); - return; - } - - if (common::getEnvDisaggLayerwise()) - { - TLLM_LOG_ERROR("Layer-wise RNN cache transfer is not supported yet"); - return; - } - - // Since allocation happens earlier - auto const slotIdx = mRnnStateManager->getCacheIndex(llmRequest.mRequestId); - - auto const& selfParallel = selfConfig.getParallelConfig(); - auto const selfTPNum = selfParallel.mTensorParallelism; - auto const selfPPRank = selfIdx / selfTPNum; - auto const& selfLayersPerPP = selfConfig.getRnnCacheState().mLayerNumPerPP; - SizeType32 const numLocalLayers = selfLayersPerPP[selfPPRank]; - - if (common::getEnvTryZCopyForKVCacheTransfer() && destConfig == selfConfig) - { - TLLM_LOG_DEBUG("try zcopy for RNN cache"); - NVTX3_SCOPED_RANGE(RnnZeroCopyRecv); - - TLLM_CHECK(sourceNum == 1); - - TLLM_CUDA_CHECK(cudaSetDevice(deviceId)); - for (size_t i = 0; i < sourceNum; i++) - { - for (SizeType32 layer = 0; layer < numLocalLayers; layer++) - { - - // Get conv state for this layer: shape is [maxBatchSize, convDim, dConv-1] - auto convState = mRnnStateManager->getConvStates(mRnnStateManager->getGlobalLayerNum(layer)); - // Slice out the specific slot: shape becomes [convDim, dConv-1] - auto slotConv = runtime::ITensor::slice(convState, slotIdx, 1); - slotConv->squeeze(0); - - // Send conv state - // llmRequest.updateKvCacheSize(slotConv->getSizeInBytes()); - session.recv(pickUpConnections[i], slotConv->data(), slotConv->getSizeInBytes()); - - // Get SSM state for this layer: shape is [maxBatchSize, numHeads, headDim, dState] - auto ssmState = mRnnStateManager->getSsmStates(mRnnStateManager->getGlobalLayerNum(layer)); - // Slice out the specific slot: shape becomes [numHeads, headDim, dState] - auto slotSsm = runtime::ITensor::slice(ssmState, slotIdx, 1); - slotSsm->squeeze(0); - - // Send SSM state - // llmRequest.updateKvCacheSize(slotSsm->getSizeInBytes()); - session.recv(pickUpConnections[i], slotSsm->data(), slotSsm->getSizeInBytes()); - } - } - TLLM_LOG_DEBUG( - mpi::MpiComm::world().getRank(), "End receiving RNN cache for request ID: %ld.", llmRequest.mRequestId); - return; - } - - // Calculate buffer sizes - auto const& modelConfig = selfConfig.getRnnModelConfig(); - int const selfTPSizePerDPGroup = selfParallel.mEnableAttentionDP ? selfTPNum / selfParallel.mDPsize : selfTPNum; - SizeType32 selfConvDimLocal = modelConfig.mConvDimSize / selfTPSizePerDPGroup; - int const selfNumHeadsLocal = modelConfig.mNumHeads / selfTPSizePerDPGroup; - - size_t convBytesPerLayer - = selfConvDimLocal * (modelConfig.mDConv - 1) * common::getDTypeSize(selfConfig.getConvStateDataType()); - convBytesPerLayer = (convBytesPerLayer + 15) & ~static_cast(15); - size_t ssmBytesPerLayer = selfNumHeadsLocal * modelConfig.mHeadDim * modelConfig.mDState - * common::getDTypeSize(selfConfig.getSsmStateDataType()); - - std::vector bufferSizesPerSource(sourceNum, 0); - size_t validTpSources = sourceNum / sourceInfo.mDomainPPSize; - - // Compute source conv bytes for SSM offset - size_t sourceConvBytesPerLayer = convBytesPerLayer / validTpSources; - - for (size_t i = 0; i < sourceNum; i++) - { - SizeType32 layersFromSource = sourceInfo.getPeerPPDomainLayerNum(static_cast(localRankIndices[i])); - bufferSizesPerSource[i] = layersFromSource * (convBytesPerLayer + ssmBytesPerLayer) / validTpSources; - } - - // Allocate receive buffers - size_t remainNoCoverSourceNum = 0; - size_t bufferCoverSourceNum = 0; - std::optional cacheBufferId = std::nullopt; - - auto preAssignedRnnId - = connections[pickUpConnections[0]]->getPreAssignedBufferId(static_cast(BufferKind::kRNN)); - if (preAssignedRnnId.has_value()) - { - cacheBufferId = static_cast(*preAssignedRnnId); - } - else - { - cacheBufferId = mRnnCacheTransBufferManager->assignBufferIndexForRecv(); - } - - auto allocationResult = mRnnCacheTransBufferManager->getOrAllocateRecvBuffers( - cacheBufferId, static_cast(sourceNum), bufferSizesPerSource, bufferManager); - auto& recvBuffers = std::get<0>(allocationResult); - auto& bufferCoverSourceNumTmp = std::get<1>(allocationResult); - auto& onlyUseDynamicBuffer = std::get<2>(allocationResult); - - TLLM_CHECK(cacheBufferId.has_value() || onlyUseDynamicBuffer); - - if (preAssignedRnnId.has_value()) - { - TLLM_CHECK_WITH_INFO(bufferCoverSourceNumTmp == sourceNum, "Agent needs all RNN recv buffers pre-allocated"); - TLLM_CHECK(onlyUseDynamicBuffer == false); - } - - bufferCoverSourceNum = bufferCoverSourceNumTmp; - remainNoCoverSourceNum = sourceNum > bufferCoverSourceNum ? sourceNum - bufferCoverSourceNum : 0; - - bufferManager.getStream().synchronize(); - session.setTime(TransferSession::kTimePreprocess); - - // Get pre-allocated buffer for chunked receive - runtime::ITensor::SharedPtr preAllocRecvBuffer = nullptr; - if (cacheBufferId.has_value()) - { - preAllocRecvBuffer = mRnnCacheTransBufferManager->getRecvBuffer(cacheBufferId); - TLLM_CHECK(preAllocRecvBuffer != nullptr); - } - - auto recvBufferFun = [&](int devId, size_t srcIdx) - { - NVTX3_SCOPED_RANGE(recvBufferFun); - TLLM_CUDA_CHECK(cudaSetDevice(devId)); - TLLM_CHECK(recvBuffers.size() > srcIdx); - auto startTime = LlmRequest::getSteadyClockNow(); - size_t size = 0; - - if (srcIdx >= remainNoCoverSourceNum) - { - // Fast path: buffer is pre-allocated, receive directly - auto& buffer = recvBuffers[srcIdx]; - size = buffer->getSizeInBytes(); - TLLM_LOG_DEBUG( - mpi::MpiComm::world().getRank(), " start recv srcIdx: %lu size:%lu", srcIdx, buffer->getSizeInBytes()); - session.recv(pickUpConnections[srcIdx], buffer->data(), buffer->getSizeInBytes()); - TLLM_LOG_DEBUG( - mpi::MpiComm::world().getRank(), " recv srcIdx: %lu size:%lu", srcIdx, buffer->getSizeInBytes()); - } - else - { - // Slow path: chunked receive for buffers that couldn't be pre-allocated - auto recvBufferIdx = bufferCoverSourceNum == 0 ? 0 : srcIdx % bufferCoverSourceNum + remainNoCoverSourceNum; - auto recvBufferUsed = bufferCoverSourceNum == 0 ? preAllocRecvBuffer : recvBuffers[recvBufferIdx]; - - size_t remainRecvSize = recvBuffers[srcIdx]->getSize(); - size_t needRecvSize = recvBuffers[srcIdx]->getSize(); - - while (remainRecvSize > 0) - { - TLLM_CHECK(recvBufferUsed != nullptr); - auto recvBufferEleSize = recvBufferUsed->getSize(); - auto recvSize = std::min(remainRecvSize, recvBufferEleSize); - auto recvSlice = runtime::ITensor::slice(recvBufferUsed, 0, recvSize); - auto copySlice = runtime::ITensor::slice(recvBuffers[srcIdx], needRecvSize - remainRecvSize, recvSize); - size += recvSlice->getSizeInBytes(); - session.recv(pickUpConnections[srcIdx], recvSlice->data(), recvSlice->getSizeInBytes()); - // Use cudaMemcpyAsync since we're copying bytes - TLLM_CUDA_CHECK(cudaMemcpyAsync(copySlice->data(), recvSlice->data(), recvSlice->getSizeInBytes(), - cudaMemcpyDeviceToDevice, bufferManager.getStream().get())); - bufferManager.getStream().synchronize(); - remainRecvSize -= recvSize; - } - } - - auto endTime = LlmRequest::getSteadyClockNow(); - session.appendMeasure(startTime, endTime, size); - }; - - // Dispatch receives (sequential or parallel based on env var) - if (sourceNum > 1) - { - if (!common::getEnvEnableReceiveKVCacheParallel()) - { - TLLM_LOG_DEBUG("Sequential receive for RNN cache."); - for (size_t i = 0; i < sourceNum; i++) - { - recvBufferFun(deviceId, i); - } - } - else - { - // Parallel receive with controlled concurrency - auto concurrencyNum = std::min(std::max(static_cast(1), bufferCoverSourceNum), sourceNum); - auto remainRecvNum = sourceNum; - - while (remainRecvNum > 0) - { - auto recvConcurrencyNum = std::min(remainRecvNum, concurrencyNum); - - // Avoid leaving a tiny remainder - if (remainRecvNum > concurrencyNum && remainRecvNum < (2 * concurrencyNum)) - { - recvConcurrencyNum = remainRecvNum - concurrencyNum; - } - - std::vector> futures; - futures.reserve(recvConcurrencyNum); - for (size_t i = 0; i < recvConcurrencyNum; i++) - { - size_t idx = i + (sourceNum - remainRecvNum); - TLLM_CHECK(idx < sourceNum); - futures.push_back(std::async(std::launch::async, recvBufferFun, deviceId, idx)); - } - for (auto& future : futures) - { - future.get(); - } - remainRecvNum -= recvConcurrencyNum; - } - } - } - else - { - recvBufferFun(deviceId, 0); - } - session.setTime(TransferSession::kTimeTransmissions); - - // Unpack received buffers into RNN states - std::vector outputConvBlocks; - std::vector outputSsmBlocks; - - auto const maxBatchSize = mRnnStateManager->getMaxBatchSize(); - auto convStates = mRnnStateManager->getConvStates(); // [numLocalLayers, maxBatchSize, convDim, dConv-1] - auto ssmStates = mRnnStateManager->getSsmStates(); // [numLocalLayers, maxBatchSize, numHeads, headDim, dState] - - outputConvBlocks.push_back(convStates); - outputSsmBlocks.push_back(ssmStates); - - tensorrt_llm::executor::rnn_cache::concatRnnConvStateDispatch( - recvBuffers, outputConvBlocks, slotIdx, maxBatchSize, destConfig, selfConfig, selfIdx, bufferManager); - - tensorrt_llm::executor::rnn_cache::concatRnnSsmStateDispatch(recvBuffers, outputSsmBlocks, slotIdx, maxBatchSize, - sourceConvBytesPerLayer, destConfig, selfConfig, selfIdx, bufferManager); - - bufferManager.getStream().synchronize(); - - if (cacheBufferId.has_value()) - { - mRnnCacheTransBufferManager->freeBufferIndexForRecv(cacheBufferId); - } - session.setTime(TransferSession::kTimePostprocess); - - TLLM_LOG_DEBUG( - mpi::MpiComm::world().getRank(), "End receiving RNN state for request ID: %ld.", llmRequest.mRequestId); -} - -void RnnCacheFormatter::formatUnifiedPoolMode(TransferSession& session) { NVTX3_SCOPED_RANGE(RnnCacheFormatter_formatUnifiedPool); session.setTime(TransferSession::kTimeFormatter); @@ -673,7 +199,7 @@ void RnnCacheFormatter::formatUnifiedPoolMode(TransferSession& session) bufferManager.getStream().synchronize(); session.setTime(TransferSession::kTimePreprocess); - // Send (same protocol as slot mode). + // Send buffers to targets. int deviceId; TLLM_CUDA_CHECK(cudaGetDevice(&deviceId)); auto preAllocSendBuffer = mRnnCacheTransBufferManager->getSendBuffer(cacheBufferId); @@ -689,7 +215,7 @@ void RnnCacheFormatter::formatUnifiedPoolMode(TransferSession& session) llmRequest.mRequestId); } -void RnnCacheFormatter::unformatUnifiedPoolMode(TransferSession& session) +void RnnCacheFormatter::unformat(TransferSession& session) { NVTX3_SCOPED_RANGE(RnnCacheFormatter_unformatUnifiedPool); session.setTime(TransferSession::kTimeFormatter); @@ -819,7 +345,7 @@ void RnnCacheFormatter::unformatUnifiedPoolMode(TransferSession& session) bufferSizesPerSource[t] = ssmBufBytes + convBufBytes; } - // Use pre-assigned buffer ID from NIXL connection if available (same as slot mode). + // Use pre-assigned buffer ID from NIXL connection if available. std::optional cacheBufferId = std::nullopt; auto preAssignedRnnId = connections[rnnRecvConns[0]]->getPreAssignedBufferId(static_cast(BufferKind::kRNN)); diff --git a/cpp/tensorrt_llm/batch_manager/rnnCacheTransBuffer.cpp b/cpp/tensorrt_llm/batch_manager/rnnCacheTransBuffer.cpp index 8fa9508cbefe..37af8e31baf9 100644 --- a/cpp/tensorrt_llm/batch_manager/rnnCacheTransBuffer.cpp +++ b/cpp/tensorrt_llm/batch_manager/rnnCacheTransBuffer.cpp @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); @@ -27,52 +27,6 @@ namespace tensorrt_llm::batch_manager::rnn_state_manager { -size_t RnnCacheTransBufferManager::computeTransferBufferSize( - RnnStateManager* rnnStateManager, std::optional maxNumTokens) -{ - SizeType32 numLocalLayers = rnnStateManager->getNumLocalLayers(); - - // Get the tensor for one layer to determine per-slot dimensions - // Conv state shape per layer: [maxBatchSize, convDim_local, dConv-1] - // SSM state shape per layer: [maxBatchSize, numHeads_local, headDim, dState] - // The tensors are shaped [maxBatchSize, ...], so one slot = total_size / maxBatchSize - auto convState = rnnStateManager->getConvStates(rnnStateManager->getGlobalLayerNum(0)); // Get first layer's tensor - auto ssmState = rnnStateManager->getSsmStates(rnnStateManager->getGlobalLayerNum(0)); - - auto convShape = convState->getShape(); - auto ssmShape = ssmState->getShape(); - - // Compute elements per slot per layer (divide total volume by batch size) - size_t convElemsPerSlotPerLayer = runtime::ITensor::volume(convShape) / convShape.d[0]; - size_t ssmElemsPerSlotPerLayer = runtime::ITensor::volume(ssmShape) / ssmShape.d[0]; - - size_t convDtypeSize = common::getDTypeSize(rnnStateManager->getConvStateDataType()); - size_t ssmDtypeSize = common::getDTypeSize(rnnStateManager->getSsmStateDataType()); - - size_t convBytesPerSlotPerLayer = convElemsPerSlotPerLayer * convDtypeSize; - size_t ssmBytesPerSlotPerLayer = ssmElemsPerSlotPerLayer * ssmDtypeSize; - - size_t bufferSizePerSlot = numLocalLayers * (convBytesPerSlotPerLayer + ssmBytesPerSlotPerLayer); - - TLLM_LOG_DEBUG( - "RNN computeTransferBufferSize: numLocalLayers=%d, convBytesPerLayer=%lu, ssmBytesPerLayer=%lu, " - "totalPerSlot=%lu", - numLocalLayers, convBytesPerSlotPerLayer, ssmBytesPerSlotPerLayer, bufferSizePerSlot); - - return bufferSizePerSlot > 0 ? bufferSizePerSlot : common::getEnvMemSizeForKVCacheTransferBuffer(); -} - -RnnCacheTransBufferManager::RnnCacheTransBufferManager( - RnnStateManager* rnnStateManager, std::optional maxNumTokens) - : BaseTransBufferManager(computeTransferBufferSize(rnnStateManager, maxNumTokens), - nvinfer1::DataType::kUINT8, // Use byte buffer for mixed dtypes - maxNumTokens) - , mRnnStateManager{rnnStateManager} -{ - TLLM_CHECK(mRnnStateManager != nullptr); - TLLM_LOG_INFO("RnnCacheTransBufferManager created for RNN cache"); -} - size_t RnnCacheTransBufferManager::computeTransferBufferSizeFromPool( kv_cache_manager::BaseKVCacheManager* kvCacheManager, executor::kv_cache::CacheState const& cacheState, std::optional maxNumTokens) @@ -142,7 +96,6 @@ RnnCacheTransBufferManager::RnnCacheTransBufferManager(kv_cache_manager::BaseKVC executor::kv_cache::CacheState const& cacheState, std::optional maxNumTokens) : BaseTransBufferManager(computeTransferBufferSizeFromPool(kvCacheManager, cacheState, maxNumTokens), nvinfer1::DataType::kUINT8, maxNumTokens) - , mRnnStateManager{nullptr} { TLLM_CHECK(kvCacheManager != nullptr); TLLM_LOG_INFO("RnnCacheTransBufferManager created for unified pool RNN cache"); diff --git a/cpp/tensorrt_llm/batch_manager/rnnCacheTransBuffer.h b/cpp/tensorrt_llm/batch_manager/rnnCacheTransBuffer.h index e6ffa06db994..124525184815 100644 --- a/cpp/tensorrt_llm/batch_manager/rnnCacheTransBuffer.h +++ b/cpp/tensorrt_llm/batch_manager/rnnCacheTransBuffer.h @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 * * Licensed under the Apache License, Version 2.0 (the "License"); @@ -18,7 +18,6 @@ #pragma once #include "tensorrt_llm/batch_manager/baseTransBuffer.h" -#include "tensorrt_llm/batch_manager/rnnStateManager.h" #include "tensorrt_llm/executor/dataTransceiverState.h" #include "tensorrt_llm/executor/executor.h" #include "tensorrt_llm/runtime/bufferManager.h" @@ -43,11 +42,6 @@ class RnnCacheTransBufferManager : public BaseTransBufferManager using SizeType32 = tensorrt_llm::runtime::SizeType32; using CacheState = executor::kv_cache::CacheState; - /// @brief Constructor for slot-based path (CppMambaCacheManager with RnnStateManager). - /// @param rnnStateManager Pointer to the RNN state manager. - /// @param maxNumTokens Optional maximum number of tokens for buffer sizing. - RnnCacheTransBufferManager(RnnStateManager* rnnStateManager, std::optional maxNumTokens = std::nullopt); - /// @brief Constructor for unified pool path (CppMambaHybridCacheManager). /// Computes buffer sizes from the KV cache manager's recurrent state pool metadata. /// @param kvCacheManager Pointer to the KV cache manager with unified pool. @@ -63,26 +57,15 @@ class RnnCacheTransBufferManager : public BaseTransBufferManager static size_t preAllocBufferSize( size_t rnnStateSizeBytes, std::optional const& cacheTransceiverConfig); - /// @brief Get the RNN state manager. - [[nodiscard]] RnnStateManager* getRnnStateManager() const noexcept - { - return mRnnStateManager; - } - [[nodiscard]] BufferKind getBufferKind() const override { return BufferKind::kRNN; } private: - /// @brief Compute transfer buffer size from RNN state configuration. - static size_t computeTransferBufferSize(RnnStateManager* rnnStateManager, std::optional maxNumTokens); - /// @brief Compute transfer buffer size from unified pool metadata. static size_t computeTransferBufferSizeFromPool(kv_cache_manager::BaseKVCacheManager* kvCacheManager, executor::kv_cache::CacheState const& cacheState, std::optional maxNumTokens); - - RnnStateManager* mRnnStateManager{nullptr}; }; } // namespace tensorrt_llm::batch_manager::rnn_state_manager diff --git a/cpp/tensorrt_llm/executor/cache_transmission/rnnCacheSplitConcat.cu b/cpp/tensorrt_llm/executor/cache_transmission/rnnCacheSplitConcat.cu index b697f2eeb3d6..fb6837e8d188 100644 --- a/cpp/tensorrt_llm/executor/cache_transmission/rnnCacheSplitConcat.cu +++ b/cpp/tensorrt_llm/executor/cache_transmission/rnnCacheSplitConcat.cu @@ -964,8 +964,7 @@ void concatRnnSsmStateDispatch(std::vector const& i // SSM portion: [numHeads, headDim, dState] — split by heads // Conv portion: [section0_dim, section1_dim, ...] x [dConv-1] — section-aware split // -// The input/output pointer arrays follow the same pattern as the existing -// RnnStateManager kernels: input pointers → output pointers → prefixLayerNum. +// The input/output pointer arrays follow the pattern: input pointers → output pointers → prefixLayerNum. /** * @brief Kernel to split SSM state from unified pool blocks to per-target buffers. diff --git a/cpp/tensorrt_llm/nanobind/batch_manager/cacheTransceiver.cpp b/cpp/tensorrt_llm/nanobind/batch_manager/cacheTransceiver.cpp index 5257c7adb1c9..6e9b0037baf3 100644 --- a/cpp/tensorrt_llm/nanobind/batch_manager/cacheTransceiver.cpp +++ b/cpp/tensorrt_llm/nanobind/batch_manager/cacheTransceiver.cpp @@ -120,11 +120,11 @@ void tb::CacheTransceiverBindings::initBindings(nb::module_& m) .def(nb::init, SizeType32, SizeType32, runtime::WorldConfig, std::vector, nvinfer1::DataType, executor::kv_cache::CacheState::AttentionType, std::optional, - tb::rnn_state_manager::RnnStateManager*, std::vector>(), + std::vector>(), nb::arg("cache_manager"), nb::arg("num_kv_heads_per_layer"), nb::arg("size_per_head"), nb::arg("tokens_per_block"), nb::arg("world_config"), nb::arg("attention_layer_num_per_pp"), nb::arg("dtype"), nb::arg("attention_type"), nb::arg("cache_transceiver_config") = std::nullopt, - nb::arg("rnn_state_manager") = nullptr, nb::arg("rnn_layer_num_per_pp") = std::vector{}); + nb::arg("rnn_layer_num_per_pp") = std::vector{}); nb::class_(m, "CacheTransceiverComm") .def( diff --git a/cpp/tests/unit_tests/batch_manager/rnnCacheFormatterTest.cpp b/cpp/tests/unit_tests/batch_manager/rnnCacheFormatterTest.cpp index 282475bc007c..8840130df1ea 100644 --- a/cpp/tests/unit_tests/batch_manager/rnnCacheFormatterTest.cpp +++ b/cpp/tests/unit_tests/batch_manager/rnnCacheFormatterTest.cpp @@ -194,7 +194,7 @@ TEST_F(RnnTargetIRanksTest, inquireSupport) // Use reinterpret_cast to pass a non-null dummy pointer (formatter only stores it, doesn't use it in // inquireSupport) - tbm::RnnCacheFormatter formatter(reinterpret_cast(0x1), + tbm::RnnCacheFormatter formatter(reinterpret_cast(0x1), reinterpret_cast(0x2)); // Same TP, different PP -> should be supported @@ -272,7 +272,7 @@ TEST_F(HybridModelCounterpartsTest, DifferentPPDistributionKvRnn) auto genState = makeHybridState(/*kvNumLayers=*/10, /*rnnNumLayers=*/6, tp, /*pp=*/2, {5, 5}, {3, 3}); // Use dummy formatter pointers (we only need them to call getCounterparts) - tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), + tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), reinterpret_cast(0x2)); // ============= Test from Context rank 0 (PP=0, TP=0) ============= @@ -337,7 +337,7 @@ TEST_F(HybridModelCounterpartsTest, AsymmetricKvRnnDistribution) auto contextState = makeHybridState(/*kvNumLayers=*/8, /*rnnNumLayers=*/4, tp, /*pp=*/1, {8}, {4}); auto genState = makeHybridState(/*kvNumLayers=*/8, /*rnnNumLayers=*/4, tp, /*pp=*/4, {2, 2, 2, 2}, {2, 2, 0, 0}); - tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), + tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), reinterpret_cast(0x2)); // ============= Test from Context rank 0 (PP=0, TP=0) ============= @@ -420,7 +420,7 @@ TEST_F(HybridModelCounterpartsTest, DisjointKvRnnCounterparts) auto genState = makeHybridState( /*kvNumLayers=*/4, /*rnnNumLayers=*/4, tp, /*pp=*/4, {2, 2, 0, 0}, {0, 0, 2, 2}); - tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), + tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), reinterpret_cast(0x2)); // ============= Test from Context rank 0 (PP=0, TP=0) - has KV only ============= @@ -586,7 +586,7 @@ TEST_F(HybridModelCounterpartsTest, InterleavedLayers) auto contextState = makeHybridState(/*kvNumLayers=*/8, /*rnnNumLayers=*/8, tp, /*pp=*/1, {8}, {8}); auto genState = makeHybridState(/*kvNumLayers=*/8, /*rnnNumLayers=*/8, tp, /*pp=*/2, {4, 4}, {4, 4}); - tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), + tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), reinterpret_cast(0x2)); SizeType32 contextRank0 = 0; @@ -612,7 +612,7 @@ TEST_F(HybridModelCounterpartsTest, RnnMorePPThanKv) auto contextState = makeHybridState(/*kvNumLayers=*/4, /*rnnNumLayers=*/8, tp, /*pp=*/1, {4}, {8}); auto genState = makeHybridState(/*kvNumLayers=*/4, /*rnnNumLayers=*/8, tp, /*pp=*/4, {2, 2, 0, 0}, {2, 2, 2, 2}); - tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), + tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), reinterpret_cast(0x2)); SizeType32 contextRank0 = 0; @@ -639,7 +639,7 @@ TEST_F(HybridModelCounterpartsTest, KvMorePPThanRnn) auto contextState = makeHybridState(/*kvNumLayers=*/8, /*rnnNumLayers=*/4, tp, /*pp=*/1, {8}, {4}); auto genState = makeHybridState(/*kvNumLayers=*/8, /*rnnNumLayers=*/4, tp, /*pp=*/4, {2, 2, 2, 2}, {2, 2, 0, 0}); - tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), + tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), reinterpret_cast(0x2)); SizeType32 contextRank0 = 0; @@ -672,7 +672,7 @@ TEST_F(HybridModelCounterpartsTest, RnnOnlyModel) auto contextState = makeHybridState(/*kvNumLayers=*/0, /*rnnNumLayers=*/8, tp, /*pp=*/1, {0}, {8}); auto genState = makeHybridState(/*kvNumLayers=*/0, /*rnnNumLayers=*/8, tp, /*pp=*/2, {0, 0}, {4, 4}); - tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), + tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), reinterpret_cast(0x2)); SizeType32 contextRank0 = 0; @@ -702,7 +702,7 @@ TEST_F(HybridModelCounterpartsTest, LargeScaleMixedLayers) auto contextState = makeHybridState(/*kvNumLayers=*/32, /*rnnNumLayers=*/16, tp, /*pp=*/1, {32}, {16}); auto genState = makeHybridState(/*kvNumLayers=*/32, /*rnnNumLayers=*/16, tp, /*pp=*/4, {8, 8, 8, 8}, {8, 8, 0, 0}); - tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), + tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), reinterpret_cast(0x2)); // Context rank 0 (PP=0, TP=0) @@ -754,7 +754,7 @@ TEST_F(HybridModelCounterpartsTest, ContextPPGreaterThanGenPP) auto genState = makeHybridState( /*kvNumLayers=*/16, /*rnnNumLayers=*/8, tp, /*pp=*/1, {16}, {8}); - tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), + tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), reinterpret_cast(0x2)); // ============= Test from Context rank 0 (PP=0, TP=0) ============= @@ -812,7 +812,7 @@ TEST_F(HybridModelCounterpartsTest, AsymmetricContextPPGreaterThanGenPP) auto genState = makeHybridState( /*kvNumLayers=*/8, /*rnnNumLayers=*/6, tp, /*pp=*/2, {4, 4}, {3, 3}); - tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), + tbm::RnnCacheFormatter rnnFormatter(reinterpret_cast(0x1), reinterpret_cast(0x2)); // ============= Test from Context rank 0 (PP=0, TP=0) - KV only ============= diff --git a/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py b/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py index 0e1c0f7274de..c124d678de5d 100644 --- a/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py +++ b/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py @@ -21,10 +21,7 @@ PoolView, ) from tensorrt_llm._torch.disaggregation.resource.utils import get_physical_pool -from tensorrt_llm._torch.pyexecutor.mamba_cache_manager import ( - MambaHybridCacheManager, - PythonMambaCacheManager, -) +from tensorrt_llm._torch.pyexecutor.mamba_cache_manager import MambaHybridCacheManager from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm._utils import get_size_in_bytes, nvtx_range from tensorrt_llm.bindings import DataType @@ -92,10 +89,6 @@ def extract( def _build_layer_group_for_mamba( manager: MambaHybridCacheManager, pool_group_idx: int ) -> MambaLayerGroup: - assert isinstance(manager._impl, PythonMambaCacheManager), ( - "CppMambaCacheManager is not supported with Python transceiver, please set TRTLLM_USE_CPP_MAMBA=0" - ) - mamba_layer_offsets = { int(global_layer_id): int(local_layer_id) for global_layer_id, local_layer_id in manager._impl.mamba_layer_offsets.items() diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index b3854da74de4..72548e58cfa3 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -30,10 +30,7 @@ from tensorrt_llm._torch.distributed.communicator import Distributed from tensorrt_llm._torch.pyexecutor.kv_cache_transceiver import KvCacheTransceiver from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest -from tensorrt_llm._torch.pyexecutor.mamba_cache_manager import ( - MambaHybridCacheManager, - PythonMambaCacheManager, -) +from tensorrt_llm._torch.pyexecutor.mamba_cache_manager import MambaHybridCacheManager from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm._utils import nvtx_range from tensorrt_llm.bindings import LlmRequestState @@ -141,9 +138,6 @@ def _exchange_rank_info(self): endpoints = cast(list, self._dist.allgather(self._transfer_worker.sender_endpoint)) layer_num = len(self._kv_cache_manager.pp_layers) if isinstance(self._kv_cache_manager, MambaHybridCacheManager): - assert isinstance(self._kv_cache_manager._impl, PythonMambaCacheManager), ( - "CppMambaCacheManager is not supported with Python transceiver, please set TRTLLM_USE_CPP_MAMBA=0" - ) layer_num += len(self._kv_cache_manager._impl.mamba_layer_offsets) layer_num_per_pp = cast(list, getattr(self._dist, "pp_allgather")(layer_num)) self._transfer_worker.populate_instance_and_rank_info( diff --git a/tensorrt_llm/_torch/model_config.py b/tensorrt_llm/_torch/model_config.py index 0752742c629a..eefd56e9a6bd 100644 --- a/tensorrt_llm/_torch/model_config.py +++ b/tensorrt_llm/_torch/model_config.py @@ -51,32 +51,6 @@ } -def _unified_kv_pool_includes_mamba( - is_disagg: bool, spec_config: Optional['SpeculativeConfig']) -> bool: - """Whether the KV cache pool will include mamba layers for a hybrid model. - - True for the default Python ``MambaHybridCacheManager`` route, where - mamba state is allocated alongside attention KV inside one pool (with - zero KV heads on mamba layers). False for the V1-route managers used - when: - - * disaggregated serving forces the C++ mamba manager - (``TRTLLM_USE_CPP_MAMBA=1`` enables the same path locally), or - * ``TRTLLM_USE_PY_MAMBA=1`` forces the Python mamba manager locally - (agg-mode override), or - * one-model speculative decoding splits mamba and attention into - separate caches. - - Single source of truth for the binding-side layer-counting decision; do - not duplicate the predicate at call sites. - """ - use_split_pool = is_disagg \ - or os.environ.get('TRTLLM_USE_CPP_MAMBA', '0') == '1' \ - or os.environ.get('TRTLLM_USE_PY_MAMBA', '0') == '1' - use_spec = spec_config is not None - return not (use_split_pool or use_spec) - - @contextlib.contextmanager def config_file_lock(timeout: int = 10): """ diff --git a/tensorrt_llm/_torch/modules/mamba/gdn_mixer.py b/tensorrt_llm/_torch/modules/mamba/gdn_mixer.py index f81f92cbf945..b59d5a909708 100644 --- a/tensorrt_llm/_torch/modules/mamba/gdn_mixer.py +++ b/tensorrt_llm/_torch/modules/mamba/gdn_mixer.py @@ -20,7 +20,6 @@ _flashinfer_gdn_verify, fused_sigmoid_gating_delta_rule_update, ) -from tensorrt_llm._torch.pyexecutor.mamba_cache_manager import use_cpp_mamba_cache_manager from tensorrt_llm._utils import is_flashinfer_gdn_supported_arch from tensorrt_llm.mapping import Mapping @@ -1107,14 +1106,9 @@ def forward_core( batch_split_size = [num_prefills, num_decodes] has_initial_states = mamba_metadata.has_initial_states state_indices = mamba_metadata.state_indices[: num_prefills + num_decodes] - if use_cpp_mamba_cache_manager(): - conv_states = attn_metadata.kv_cache_manager.get_conv_states(self.layer_idx) - ssm_states = attn_metadata.kv_cache_manager.get_ssm_states(self.layer_idx) - layer_cache = None - else: - layer_cache = attn_metadata.kv_cache_manager.mamba_layer_cache(self.layer_idx) - conv_states = layer_cache.conv - ssm_states = layer_cache.temporal + layer_cache = attn_metadata.kv_cache_manager.mamba_layer_cache(self.layer_idx) + conv_states = layer_cache.conv + ssm_states = layer_cache.temporal state_indices_p, state_indices_d = torch.split(state_indices, batch_split_size) if num_prefills > 0: diff --git a/tensorrt_llm/_torch/modules/mamba/mamba2_mixer.py b/tensorrt_llm/_torch/modules/mamba/mamba2_mixer.py index 816c6bd1e27e..89adab83d5d9 100644 --- a/tensorrt_llm/_torch/modules/mamba/mamba2_mixer.py +++ b/tensorrt_llm/_torch/modules/mamba/mamba2_mixer.py @@ -24,8 +24,6 @@ from tensorrt_llm._torch.modules.mamba.mamba2_metadata import Mamba2Metadata from tensorrt_llm._torch.modules.multi_stream_utils import \ maybe_execute_in_parallel -from tensorrt_llm._torch.pyexecutor.mamba_cache_manager import \ - use_cpp_mamba_cache_manager from tensorrt_llm.logger import logger from tensorrt_llm.mapping import Mapping @@ -306,17 +304,10 @@ def forward( state_indices = mamba_metadata.state_indices[:num_prefills + num_decodes] - if use_cpp_mamba_cache_manager(): - conv_states = attn_metadata.kv_cache_manager.get_conv_states( - self.layer_idx) - ssm_states = attn_metadata.kv_cache_manager.get_ssm_states( - self.layer_idx) - layer_cache = None # Not used in C++ path - else: - layer_cache = attn_metadata.kv_cache_manager.mamba_layer_cache( - self.layer_idx) - conv_states = layer_cache.conv - ssm_states = layer_cache.temporal + layer_cache = attn_metadata.kv_cache_manager.mamba_layer_cache( + self.layer_idx) + conv_states = layer_cache.conv + ssm_states = layer_cache.temporal state_indices_p, state_indices_d = torch.split(state_indices, batch_split_size) diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 02ffb035705a..685bba132990 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -57,7 +57,6 @@ from .mamba_cache_manager import (BaseMambaCacheManager, CppMambaHybridCacheManager, MixedMambaHybridCacheManager, - use_cpp_mamba_cache_manager, use_py_mamba_cache_manager) from .model_engine import PyTorchModelEngine from .py_executor import PyExecutor @@ -93,13 +92,13 @@ def get_kv_cache_manager_cls( cache_transceiver_config: Optional[CacheTransceiverConfig] = None): """Resolve the concrete KV cache manager class for ``model_config``. - For hybrid mamba models the choice between ``Mixed`` ( TRTLLM_USE_CPP_MAMBA / TRTLLM_USE_PY_MAMBA) and - ``Cpp`` (unified pool with block reuse) is made here. Callers that don't - care about disagg can omit ``is_disagg`` and get the unified-pool default. + For hybrid mamba models the choice between ``Mixed`` (TRTLLM_USE_PY_MAMBA) + and ``Cpp`` (unified pool with block reuse) is made here. Callers that + don't care about disagg can omit ``is_disagg`` and get the unified-pool + default. Env-var overrides (agg mode only — disagg picks its inner impl via ``cache_transceiver_config.transceiver_runtime``): - * ``TRTLLM_USE_CPP_MAMBA=1`` — Mixed manager with CppMambaCacheManager. * ``TRTLLM_USE_PY_MAMBA=1`` — Mixed manager with PythonMambaCacheManager. """ config = model_config.pretrained_config @@ -135,10 +134,6 @@ def get_kv_cache_manager_cls( return MixedMambaHybridCacheManager if kv_cache_config.enable_block_reuse: return CppMambaHybridCacheManager - if use_cpp_mamba_cache_manager(): - logger.info( - "Using MixedMambaHybridCacheManager for hybrid mamba model") - return MixedMambaHybridCacheManager if (cache_transceiver_config is not None and cache_transceiver_config.transceiver_runtime == "PYTHON"): logger.info("Python transceiver detected; using " diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py index 8e6b47c8bff8..087862b07d5f 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py @@ -181,8 +181,7 @@ def __init__(self, self.kv_transfer_sender_future_timeout_ms = cache_transceiver_config.kv_transfer_sender_future_timeout_ms self.kv_transfer_poll_interval_ms = cache_transceiver_config.kv_transfer_poll_interval_ms - # Get RNN state manager and layer distribution if mamba_cache_manager is provided. - rnn_state_manager = None + # Get RNN layer distribution if mamba_cache_manager is provided. rnn_layer_num_per_pp_rank = [] if mamba_cache_manager is not None: if isinstance(mamba_cache_manager, CppMambaHybridCacheManager): @@ -191,21 +190,20 @@ def __init__(self, rnn_layer_num_per_pp_rank = dist.pp_allgather( mamba_cache_manager.local_num_mamba_layers) else: - rnn_state_manager = mamba_cache_manager._impl.mamba_impl - # Get the number of local RNN layers and allgather across PP ranks - rnn_local_layer_num = rnn_state_manager.get_num_local_layers() + # MixedMambaHybridCacheManager with PythonMambaCacheManager. rnn_layer_num_per_pp_rank = dist.pp_allgather( - rnn_local_layer_num) + len(mamba_cache_manager._impl.mamba_layer_offsets)) logger.info( f"RNN state transfer enabled: rnn_layer_num_per_pp={rnn_layer_num_per_pp_rank}" ) - self.impl = CacheTransceiverCpp( - kv_cache_manager.impl, total_num_kv_heads_per_layer, head_dim, - tokens_per_block, world_config, - pp_layer_num_per_pp_rank, dtype, attention_type, - cache_transceiver_config._to_pybind(), rnn_state_manager, - rnn_layer_num_per_pp_rank) + self.impl = CacheTransceiverCpp(kv_cache_manager.impl, + total_num_kv_heads_per_layer, head_dim, + tokens_per_block, world_config, + pp_layer_num_per_pp_rank, dtype, + attention_type, + cache_transceiver_config._to_pybind(), + rnn_layer_num_per_pp_rank) def respond_and_send_async(self, req: LlmRequest): return self.impl.respond_and_send_async(req) diff --git a/tensorrt_llm/_torch/pyexecutor/mamba_cache_manager.py b/tensorrt_llm/_torch/pyexecutor/mamba_cache_manager.py index d4d54d7df905..e02e7fca571e 100644 --- a/tensorrt_llm/_torch/pyexecutor/mamba_cache_manager.py +++ b/tensorrt_llm/_torch/pyexecutor/mamba_cache_manager.py @@ -23,8 +23,6 @@ import triton import triton.language as tl -import tensorrt_llm.bindings - if TYPE_CHECKING: from tensorrt_llm._torch.attention_backend.interface import AttentionMetadata from tensorrt_llm.llmapi.llm_args import DecodingBaseConfig @@ -35,17 +33,13 @@ BaseResourceManager, CacheTypeCpp, DataType, KVCacheManager, PoolConfiguration, get_pp_layers) from tensorrt_llm._torch.pyexecutor.scheduler import ScheduledRequests -from tensorrt_llm._utils import (nvtx_range, prefer_pinned, - torch_dtype_to_binding) +from tensorrt_llm._utils import nvtx_range, prefer_pinned from tensorrt_llm.bindings.internal.batch_manager import ( LinearAttentionMetadata, LinearCacheType) from tensorrt_llm.llmapi.llm_args import KvCacheConfig from tensorrt_llm.logger import logger from tensorrt_llm.mapping import Mapping -RnnStateManagerCpp = tensorrt_llm.bindings.internal.batch_manager.RnnStateManager -WorldConfig = tensorrt_llm.bindings.WorldConfig - GB = 1 << 30 # Replay kernels pad the token/window dimension to at least 16 for tensor-core @@ -145,21 +139,6 @@ def _mamba_rank_offset(mapping: Mapping) -> int: mapping.rank * 1_009) -def use_cpp_mamba_cache_manager() -> bool: - """Check if C++ MambaCacheManager should be used. - - Returns True if TRTLLM_USE_CPP_MAMBA='1' is set, False otherwise. - By default, PythonMambaCacheManager is used. - """ - cpp = os.environ.get('TRTLLM_USE_CPP_MAMBA', '0') == '1' - py = os.environ.get('TRTLLM_USE_PY_MAMBA', '0') == '1' - if cpp and py: - raise ValueError( - "TRTLLM_USE_CPP_MAMBA=1 and TRTLLM_USE_PY_MAMBA=1 are mutually " - "exclusive; unset one of them.") - return cpp - - def use_py_mamba_cache_manager() -> bool: """Check if PythonMambaCacheManager should be forced (agg mode override). @@ -170,13 +149,7 @@ def use_py_mamba_cache_manager() -> bool: CppMambaHybridCacheManager. Disagg mode is unaffected — it already picks PythonMambaCacheManager when transceiver_runtime='PYTHON'. """ - cpp = os.environ.get('TRTLLM_USE_CPP_MAMBA', '0') == '1' - py = os.environ.get('TRTLLM_USE_PY_MAMBA', '0') == '1' - if cpp and py: - raise ValueError( - "TRTLLM_USE_CPP_MAMBA=1 and TRTLLM_USE_PY_MAMBA=1 are mutually " - "exclusive; unset one of them.") - return py + return os.environ.get('TRTLLM_USE_PY_MAMBA', '0') == '1' class ReplayStateUpdateMetadata(NamedTuple): @@ -235,114 +208,6 @@ def mamba_layer_cache( ... -class CppMambaCacheManager(BaseResourceManager): - """Mamba state manager backed by the C++ RnnStateManager bindings. - - Manages only mamba states (conv + SSM). Used when TRTLLM_USE_CPP_MAMBA=1. - Supports disaggregated serving. - """ - - def __init__( - self, - d_state: int, - d_conv: int, - num_heads: int, - n_groups: int, - head_dim: int, - num_layers: int, - max_num_sequences: int, - mapping: Mapping, - dtype: torch.dtype, - ssm_cache_dtype: torch.dtype, - layer_mask: Optional[List[bool]] = None, - stream: Optional[torch.cuda.Stream] = None, - ) -> None: - self.mamba_ssm_cache_dtype = ssm_cache_dtype - - # get tp size - tp_size = mapping.tp_size if not mapping.enable_attention_dp else 1 - world_config = WorldConfig( - tensor_parallelism=tp_size, - pipeline_parallelism=mapping.pp_size, - rank=mapping.rank, - gpus_per_node=mapping.gpus_per_node, - ) - - dtype_binding = torch_dtype_to_binding(dtype) - ssm_cache_dtype_binding = torch_dtype_to_binding( - ssm_cache_dtype if ssm_cache_dtype is not None else dtype) - - self._stream = stream if stream is not None else torch.cuda.current_stream( - ) - - pp_layers, _ = get_pp_layers(num_layers, mapping, layer_mask=layer_mask) - - self.mamba_impl = RnnStateManagerCpp( - d_state=d_state, - d_conv=d_conv, - num_heads=num_heads, - n_groups=n_groups, - head_dim=head_dim, - max_batch_size=max_num_sequences, - world_config=world_config, - stream=self._stream.cuda_stream, - dtype=dtype_binding, - ssm_cache_dtype=ssm_cache_dtype_binding, - pp_layers=pp_layers, - num_layers=num_layers, - ) - self._max_num_sequences = max_num_sequences - - def get_max_resource_count(self) -> int: - # Return the maximum number of sequences that can be cached. - return self._max_num_sequences - - def get_needed_resource_to_completion(self, request: LlmRequest) -> int: - # For Mamba cache manager, we always need one slot per request. - return 1 - - def is_speculative(self) -> bool: - # C++ MambaCacheManager does not support speculative decoding - return False - - def prepare_resources(self, scheduled_batch: ScheduledRequests): - context_ids = [ - i.py_request_id for i in scheduled_batch.context_requests - ] - generation_ids = [ - i.py_request_id for i in scheduled_batch.generation_requests - ] - request_ids = context_ids + generation_ids - self.mamba_impl.allocate_cache_blocks(request_ids) - - def free_resources(self, request: LlmRequest): - self.mamba_impl.free_cache_block(request.py_request_id) - - def add_dummy_requests(self, request_ids: List[int], **kwargs): - # Allocate a permanent slot for every id, including CUDA-graph - # padding sentinels (matches PythonMambaCacheManager). Padding - # entries in get_state_indices then resolve via mCacheIndex to - # the sentinel's reserved slot and never alias a live request. - if request_ids: - self.mamba_impl.allocate_cache_blocks(request_ids) - - def get_state_indices(self, request_ids: List[int], - is_padding: List[bool]) -> List[int]: - return self.mamba_impl.get_state_indices(request_ids, is_padding) - - def get_conv_states(self, layer_idx: int) -> torch.Tensor: - return self.mamba_impl.get_conv_states(layer_idx) - - def get_ssm_states(self, layer_idx: int) -> torch.Tensor: - return self.mamba_impl.get_ssm_states(layer_idx) - - def get_mamba_ssm_cache_dtype(self) -> torch.dtype: - return self.mamba_ssm_cache_dtype - - def shutdown(self): - torch.cuda.empty_cache() - - class PythonMambaCacheManager(BaseResourceManager): """Pure-Python mamba state manager with speculative decoding support. @@ -976,7 +841,7 @@ def update_mamba_states(self, attn_metadata: "AttentionMetadata", class MambaCacheManager(BaseResourceManager, BaseMambaCacheManager): """Facade for standalone mamba state management (no KV cache). - Delegates to CppMambaCacheManager (when TRTLLM_USE_CPP_MAMBA=1) or PythonMambaCacheManager. + Delegates to PythonMambaCacheManager. """ def __init__( @@ -1000,51 +865,30 @@ def __init__( mamba_ssm_stochastic_rounding: bool = False, ) -> None: max_num_sequences = max_batch_size * mapping.pp_size - self._use_cpp = use_cpp_mamba_cache_manager() - - if self._use_cpp: - assert speculative_num_draft_tokens is None, \ - "speculative_num_draft_tokens is not supported in CppMambaCacheManager" - self._impl = CppMambaCacheManager( - d_state=d_state, - d_conv=d_conv, - num_heads=num_heads, - n_groups=n_groups, - head_dim=head_dim, - num_layers=num_layers, - max_num_sequences=max_num_sequences, - mapping=mapping, - dtype=dtype, - ssm_cache_dtype=ssm_cache_dtype, - layer_mask=layer_mask, - stream=stream, - ) - else: - self._impl = PythonMambaCacheManager( - d_state=d_state, - d_conv=d_conv, - num_heads=num_heads, - n_groups=n_groups, - head_dim=head_dim, - num_layers=num_layers, - max_batch_size=max_num_sequences, - spec_state_size=spec_state_size, - mapping=mapping, - dtype=dtype, - ssm_cache_dtype=ssm_cache_dtype, - layer_mask=layer_mask, - speculative_num_draft_tokens=speculative_num_draft_tokens, - model_type=model_type, - use_replay_state_update=use_replay_state_update, - mamba_ssm_stochastic_rounding=mamba_ssm_stochastic_rounding, - ) + + self._impl = PythonMambaCacheManager( + d_state=d_state, + d_conv=d_conv, + num_heads=num_heads, + n_groups=n_groups, + head_dim=head_dim, + num_layers=num_layers, + max_batch_size=max_num_sequences, + spec_state_size=spec_state_size, + mapping=mapping, + dtype=dtype, + ssm_cache_dtype=ssm_cache_dtype, + layer_mask=layer_mask, + speculative_num_draft_tokens=speculative_num_draft_tokens, + model_type=model_type, + use_replay_state_update=use_replay_state_update, + mamba_ssm_stochastic_rounding=mamba_ssm_stochastic_rounding, + ) def get_max_resource_count(self) -> int: return self._impl.get_max_resource_count() def filter_ctx_requests_by_capacity(self, context_requests: list) -> list: - if self._use_cpp: - return context_requests return self._impl.filter_ctx_requests_by_capacity(context_requests) def get_needed_resource_to_completion(self, request: LlmRequest) -> int: @@ -1068,12 +912,10 @@ def get_state_indices( @property def mamba_cache_free_blocks(self) -> List[int]: - assert not self._use_cpp, "mamba_cache_free_blocks is not supported in CppMambaCacheManager" return self._impl.mamba_cache_free_blocks @property def mamba_cache_index(self) -> Dict[int, int]: - assert not self._use_cpp, "mamba_cache_index is not supported in CppMambaCacheManager" return self._impl.mamba_cache_index def get_conv_states(self, layer_idx: int) -> torch.Tensor: @@ -1086,11 +928,8 @@ def get_mamba_ssm_cache_dtype(self) -> torch.dtype: return self._impl.get_mamba_ssm_cache_dtype() def get_mamba_ssm_rand_seed(self) -> Optional[torch.Tensor]: - """Delegate to the underlying Python manager. The C++ manager does - not allocate this buffer because it does not support speculative - decoding (and the SR-on-non-replay bug only fires under MTP).""" - getter = getattr(self._impl, 'get_mamba_ssm_rand_seed', None) - return getter() if getter is not None else None + """Delegate to the underlying Python manager.""" + return self._impl.get_mamba_ssm_rand_seed() @property def use_replay_state_update(self) -> bool: @@ -1098,45 +937,33 @@ def use_replay_state_update(self) -> bool: def get_replay_state_update_metadata( self) -> Optional[ReplayStateUpdateMetadata]: - get_metadata = getattr(self._impl, 'get_replay_state_update_metadata', - None) - if get_metadata is None: - return None - return get_metadata() + return self._impl.get_replay_state_update_metadata() def get_intermediate_ssm_states(self, layer_idx: int) -> Optional[torch.Tensor]: - assert not self._use_cpp, "get_intermediate_ssm_states is not supported in CppMambaCacheManager" return self._impl.get_intermediate_ssm_states(layer_idx) def get_intermediate_conv_states(self, layer_idx: int) -> Optional[torch.Tensor]: - assert not self._use_cpp, "get_intermediate_conv_states is not supported in CppMambaCacheManager" return self._impl.get_intermediate_conv_states(layer_idx) def get_replay_old_x(self, layer_idx: int) -> Optional[torch.Tensor]: - assert not self._use_cpp, "get_replay_old_x is not supported in CppMambaCacheManager" return self._impl.get_replay_old_x(layer_idx) def get_replay_old_B(self, layer_idx: int) -> Optional[torch.Tensor]: - assert not self._use_cpp, "get_replay_old_B is not supported in CppMambaCacheManager" return self._impl.get_replay_old_B(layer_idx) def get_replay_old_dt(self, layer_idx: int) -> Optional[torch.Tensor]: - assert not self._use_cpp, "get_replay_old_dt is not supported in CppMambaCacheManager" return self._impl.get_replay_old_dt(layer_idx) def get_replay_old_dA_cumsum(self, layer_idx: int) -> Optional[torch.Tensor]: - assert not self._use_cpp, "get_replay_old_dA_cumsum is not supported in CppMambaCacheManager" return self._impl.get_replay_old_dA_cumsum(layer_idx) def get_replay_cache_buf_idx(self) -> Optional[torch.Tensor]: - assert not self._use_cpp, "get_replay_cache_buf_idx is not supported in CppMambaCacheManager" return self._impl.get_replay_cache_buf_idx() def get_replay_prev_num_accepted_tokens(self) -> Optional[torch.Tensor]: - assert not self._use_cpp, "get_replay_prev_num_accepted_tokens is not supported in CppMambaCacheManager" return self._impl.get_replay_prev_num_accepted_tokens() def is_speculative(self) -> bool: @@ -1146,7 +973,6 @@ def mamba_layer_cache( self, layer_idx: int ) -> Union[PythonMambaCacheManager.State, PythonMambaCacheManager.SpeculativeState, None]: - assert not self._use_cpp, "mamba_layer_cache is not supported in CppMambaCacheManager" return self._impl.mamba_layer_cache(layer_idx) def shutdown(self): @@ -1159,10 +985,6 @@ def update_mamba_states(self, attn_metadata: "AttentionMetadata", # promotion is a clean no-op. if not self._impl.is_speculative(): return - # Belt-and-suspenders: C++ is non-speculative today so this is - # unreachable. Fires if C++ ever grows speculative support - # without also implementing the scatter there. - assert not self._use_cpp, "update_mamba_states is not supported in CppMambaCacheManager" self._impl.update_mamba_states(attn_metadata, num_accepted_tokens, state_indices) diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py index 797f2fd48666..695bbfa553ee 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py @@ -849,17 +849,15 @@ def drafting_loop_wrapper(model): # NOTE: TRTLLM_USE_PY_MAMBA is an agg-mode-only override and has # no effect in disagg. The disagg manager choice is driven solely # by transceiver_runtime: PYTHON => PythonMambaCacheManager, - # otherwise CppMambaCacheManager. Clear the var here so the - # mutual-exclusion check in use_cpp_mamba_cache_manager() does - # not fire after we force TRTLLM_USE_CPP_MAMBA=1 below. - if os.environ.pop("TRTLLM_USE_PY_MAMBA", "0") == "1": + # otherwise CppMambaHybridCacheManager (unified pool, default). + if os.environ.get("TRTLLM_USE_PY_MAMBA", "0") == "1": logger.warning( "TRTLLM_USE_PY_MAMBA is ignored in disaggregated serving; " "use cache_transceiver_config.transceiver_runtime='PYTHON' " "to select PythonMambaCacheManager.") else: logger.info("Disaggregated serving with hybrid model detected. " - "Enabling CppMambaHybridCacheManager.") + "Using CppMambaHybridCacheManager.") # Get draft config for one-engine speculative decoding if available draft_config = getattr(model_engine.model, 'draft_config', None) diff --git a/tests/integration/defs/accuracy/test_llm_api_pytorch.py b/tests/integration/defs/accuracy/test_llm_api_pytorch.py index 2e201782e786..475165015f0a 100644 --- a/tests/integration/defs/accuracy/test_llm_api_pytorch.py +++ b/tests/integration/defs/accuracy/test_llm_api_pytorch.py @@ -6764,8 +6764,8 @@ def test_nvfp4_marlin_multi_gpus(self, tp_size): ], ) def test_fp8_4gpus(self, attention_dp, use_cpp_mamba, monkeypatch): - monkeypatch.setenv("TRTLLM_USE_CPP_MAMBA", - "1" if use_cpp_mamba else "0") + monkeypatch.setenv("TRTLLM_USE_PY_MAMBA", + "1" if not use_cpp_mamba else "0") with LLM( f"{llm_models_root()}/NVIDIA-Nemotron-3-Super-120B-A12B-FP8", diff --git a/tests/unittest/_torch/executor/test_mamba_cache_manager.py b/tests/unittest/_torch/executor/test_mamba_cache_manager.py index 5f73a1f86881..30eaa2442424 100644 --- a/tests/unittest/_torch/executor/test_mamba_cache_manager.py +++ b/tests/unittest/_torch/executor/test_mamba_cache_manager.py @@ -5,7 +5,6 @@ import os from types import SimpleNamespace -from unittest.mock import MagicMock import pytest import torch @@ -14,7 +13,6 @@ from tensorrt_llm._torch.pyexecutor.llm_request import ATTENTION_DP_DUMMY_REQUEST_ID from tensorrt_llm._torch.pyexecutor.mamba_cache_manager import ( MIN_REPLAY_HISTORY_SIZE, - CppMambaCacheManager, CppMambaHybridCacheManager, PythonMambaCacheManager, _get_mamba_hybrid_pool_size, @@ -343,45 +341,6 @@ def test_non_mtp_pytorch_prepare_and_get_state_indices_flow(): ] -def test_cpp_add_dummy_requests_noop_on_empty_list(): - stub = SimpleNamespace(mamba_impl=MagicMock()) - CppMambaCacheManager.add_dummy_requests(stub, []) - stub.mamba_impl.allocate_cache_blocks.assert_not_called() - - -@skip_no_cuda -def test_cpp_get_state_indices_resolves_sentinel_to_reserved_slot(): - """End-to-end C++ path: add_dummy_requests + getStateIndices must - resolve the CUDA-graph sentinel to its reserved slot, distinct from - every live request's slot — guards the C++ mCacheIndex lookup, not - just the Python forwarder.""" - mgr = CppMambaCacheManager( - d_state=8, - d_conv=4, - num_heads=4, - n_groups=1, - head_dim=8, - num_layers=2, - max_num_sequences=8, - mapping=Mapping(world_size=1, tp_size=1, pp_size=1), - dtype=torch.float16, - ssm_cache_dtype=torch.float16, - ) - mgr.add_dummy_requests([100, 101, CUDA_GRAPH_DUMMY_REQUEST_ID]) - - request_ids = [100, 101, CUDA_GRAPH_DUMMY_REQUEST_ID] - is_padding = [False, False, True] - indices = mgr.get_state_indices(request_ids, is_padding) - - sentinel_slot = indices[2] - real_slots = {indices[0], indices[1]} - assert sentinel_slot not in real_slots, ( - f"sentinel slot {sentinel_slot} aliases a real request's slot {real_slots}" - ) - # Resolve again — reserved slot must be stable across calls. - assert mgr.get_state_indices(request_ids, is_padding) == indices - - # --------------------------------------------------------------------------- # CppMambaHybridCacheManager: recurrent-state snapshot pool sizing # diff --git a/tests/unittest/disaggregated/test_mamba_transfer.py b/tests/unittest/disaggregated/test_mamba_transfer.py index 32de0829850e..0cf0b3209899 100644 --- a/tests/unittest/disaggregated/test_mamba_transfer.py +++ b/tests/unittest/disaggregated/test_mamba_transfer.py @@ -329,8 +329,7 @@ def _read_actual(gen_managers, gen_request_ids) -> Dict: # --------------------------------------------------------------------------- # Main test logic # --------------------------------------------------------------------------- -def test_mamba_disagg_attention_dp_dummy_with_batch_size_one(monkeypatch): - monkeypatch.setenv("TRTLLM_USE_CPP_MAMBA", "0") +def test_mamba_disagg_attention_dp_dummy_with_batch_size_one(): mgr = _create_managers(1, max_batch_size=1, enable_attention_dp=True)[0] try: req = LlmRequest( diff --git a/tests/unittest/others/test_kv_cache_transceiver.py b/tests/unittest/others/test_kv_cache_transceiver.py index ccc059b1ab99..4dfb00f75e2d 100644 --- a/tests/unittest/others/test_kv_cache_transceiver.py +++ b/tests/unittest/others/test_kv_cache_transceiver.py @@ -717,9 +717,7 @@ def hybrid_dtypes(request): ], indirect=["hybrid_dtypes"], ) -def test_hybrid_cache_transceiver_single_process(backend, hybrid_dtypes, - monkeypatch): - monkeypatch.setenv("TRTLLM_USE_CPP_MAMBA", "1") +def test_hybrid_cache_transceiver_single_process(backend, hybrid_dtypes): mapping = Mapping(world_size=1, rank=0) kv_dtype, mamba_conv_dtype, mamba_ssm_dtype = hybrid_dtypes @@ -831,8 +829,7 @@ def transfers_done(): @pytest.mark.timeout(120) @pytest.mark.parametrize("backend", ["NIXL", "UCX"], ids=["NIXL", "UCX"]) -def test_hybrid_cache_transceiver_cancel_request(backend, monkeypatch): - monkeypatch.setenv("TRTLLM_USE_CPP_MAMBA", "1") +def test_hybrid_cache_transceiver_cancel_request(backend): mapping = Mapping(world_size=1, rank=0) dtype = DataType.HALF