diff --git a/cpp/include/tensorrt_llm/batch_manager/capacityScheduler.h b/cpp/include/tensorrt_llm/batch_manager/capacityScheduler.h index 2ea4f47ce4bc..0d11c25d14e2 100644 --- a/cpp/include/tensorrt_llm/batch_manager/capacityScheduler.h +++ b/cpp/include/tensorrt_llm/batch_manager/capacityScheduler.h @@ -87,7 +87,11 @@ class MaxRequestsScheduler : public BaseCapacityScheduler /// @brief Schedule requests using the MAX_UTILIZATION policy /// @details Try reserving resources to advance requests by one step, -/// may pause previously started requests. +/// may pause previously started requests. When a +/// ``crossKvCacheManager`` is supplied, requests in the +/// ``ENCODER_INIT`` state may be admitted for encoder compute +/// without consuming self- or cross-KV blocks; the later +/// ``CONTEXT_INIT`` decoder admission owns cross-pool budgeting. class MaxUtilizationScheduler : public BaseCapacityScheduler { public: @@ -96,8 +100,9 @@ class MaxUtilizationScheduler : public BaseCapacityScheduler LlmRequestState noScheduleAfterState = LlmRequestState::kGENERATION_COMPLETE); [[nodiscard]] std::tuple operator()( - kv_cache_manager::BaseKVCacheManager& kvCacheManager, OptionalRef peftCacheManager, - RequestList const& activeRequests) const; + kv_cache_manager::BaseKVCacheManager& kvCacheManager, + OptionalRef crossKvCacheManager, + OptionalRef peftCacheManager, RequestList const& activeRequests) const; private: SizeType32 mMaxNumRequests; @@ -106,6 +111,10 @@ class MaxUtilizationScheduler : public BaseCapacityScheduler }; /// @brief Schedule requests using the GUARANTEED_NO_EVICT policy +/// @details When a ``crossKvCacheManager`` is supplied, requests in the +/// ``ENCODER_INIT`` state may be admitted for encoder compute +/// without consuming self- or cross-KV blocks. The later +/// ``CONTEXT_INIT`` decoder admission owns cross-pool budgeting. class GuaranteedNoEvictScheduler : public BaseCapacityScheduler { public: @@ -158,7 +167,11 @@ class CapacityScheduler : public Algorithm * * @param kvCacheManager Required in MaxUtilizationScheduler (as a ref) and in GuaranteedNoEvictScheduler and * StaticBatchScheduler (as a const ref). - * @param crossKvCacheManager Optional used in GuaranteedNoEvictScheduler and StaticBatchScheduler. + * @param crossKvCacheManager Optional cross-attention KV cache manager. Used by + * MaxUtilizationScheduler (mutates: ``startScheduling`` / ``schedulingRemoveSequence``) + * and GuaranteedNoEvictScheduler / StaticBatchScheduler (read-only). Required for + * encoder-decoder admission. Encoder-init requests only require this pool + * to be configured; decoder context admission budgets blocks from it. * @param peftCacheManager Optional used in MaxUtilizationScheduler, GuaranteedNoEvictScheduler and * StaticBatchScheduler. * @param activeRequests @@ -168,7 +181,7 @@ class CapacityScheduler : public Algorithm [[nodiscard]] std::tuple operator()(RequestList const& activeRequests, OptionalRef kvCacheManager = std::nullopt, OptionalRef peftCacheManager = std::nullopt, - OptionalRef crossKvCacheManager = std::nullopt) const; + OptionalRef crossKvCacheManager = std::nullopt) const; /// @brief Sets the reorder policy to use AgentTreePolicy with the given configuration. /// @param agentPercentage The ratio of agent requests to schedule (0.0-1.0, -1.0 for random). diff --git a/cpp/include/tensorrt_llm/batch_manager/llmRequest.h b/cpp/include/tensorrt_llm/batch_manager/llmRequest.h index 886147a09c73..644813e381b7 100644 --- a/cpp/include/tensorrt_llm/batch_manager/llmRequest.h +++ b/cpp/include/tensorrt_llm/batch_manager/llmRequest.h @@ -665,9 +665,9 @@ class GenericLlmRequest return mEncoderUniqueTokens; } - /// @brief Get length of encoder input (could be tokens or features length) - /// @return An integer. - [[nodiscard]] SizeType32 getEncoderInputLen() const + /// @brief Get length of encoder input when present, without throwing for decoder-only requests. + /// @return Encoder input length, or nullopt when this request has no encoder side. + [[nodiscard]] std::optional tryGetEncoderInputLen() const { if (mEncoderInputFeatures.has_value()) { @@ -678,19 +678,45 @@ class GenericLlmRequest return getEncoderTokens().value()->size(); } - TLLM_THROW("GenericLlmRequest::getEncoderInputLen - Do not have encoder length!"); + return std::nullopt; } - /// @brief Get length of encoder output. Fall back to encoder input length if not present + /// @brief Get length of encoder input (could be tokens or features length) /// @return An integer. - [[nodiscard]] SizeType32 getEncoderOutputLen() const + [[nodiscard]] SizeType32 getEncoderInputLen() const + { + auto const encoderInputLen = tryGetEncoderInputLen(); + if (encoderInputLen.has_value()) + { + return encoderInputLen.value(); + } + + TLLM_THROW("GenericLlmRequest::getEncoderInputLen - Do not have encoder length!"); + } + + /// @brief Get length of encoder output when present, without throwing for decoder-only requests. + /// @return Encoder output length, or nullopt when this request has no encoder side. + [[nodiscard]] std::optional tryGetEncoderOutputLen() const { if (mEncoderOutputLength.has_value()) { return mEncoderOutputLength.value(); } - return getEncoderInputLen(); + return tryGetEncoderInputLen(); + } + + /// @brief Get length of encoder output, or throw if the request has no encoder side. + /// @return Explicit encoder output length, or encoder input length when the output length is not present. + [[nodiscard]] SizeType32 getEncoderOutputLen() const + { + auto const encoderOutputLen = tryGetEncoderOutputLen(); + if (encoderOutputLen.has_value()) + { + return encoderOutputLen.value(); + } + + TLLM_THROW("GenericLlmRequest::getEncoderInputLen - Do not have encoder length!"); } [[nodiscard]] std::optional>> getPositionIds() const diff --git a/cpp/include/tensorrt_llm/common/optionalRef.h b/cpp/include/tensorrt_llm/common/optionalRef.h index f55b377981d2..46723f1c697c 100644 --- a/cpp/include/tensorrt_llm/common/optionalRef.h +++ b/cpp/include/tensorrt_llm/common/optionalRef.h @@ -78,6 +78,13 @@ class OptionalRef { } + // Implicit conversion from OptionalRef to OptionalRef + template >> + OptionalRef(OptionalRef> const& other) + : opt(other ? std::optional>(std::ref(*other)) : std::nullopt) + { + } + T* operator->() const { return opt ? &(opt->get()) : nullptr; diff --git a/cpp/tensorrt_llm/batch_manager/capacityScheduler.cpp b/cpp/tensorrt_llm/batch_manager/capacityScheduler.cpp index 60b6cc14ebee..21a9e6d501d9 100644 --- a/cpp/tensorrt_llm/batch_manager/capacityScheduler.cpp +++ b/cpp/tensorrt_llm/batch_manager/capacityScheduler.cpp @@ -20,6 +20,7 @@ #include "tensorrt_llm/batch_manager/kvCacheManager.h" #include "tensorrt_llm/batch_manager/peftCacheManager.h" #include "tensorrt_llm/batch_manager/scheduledBlocksManager.h" +#include "tensorrt_llm/common/assert.h" #include "tensorrt_llm/common/logger.h" #include "tensorrt_llm/common/nvtxUtils.h" @@ -114,6 +115,19 @@ bool beneficialToSkip(std::optional const& return false; } +template +void checkRequiredCrossKvCacheManager( + LlmRequestState noScheduleUntilState, OptionalRef crossKvCacheManager) +{ + if (noScheduleUntilState != LlmRequestState::kENCODER_INIT) + { + return; + } + + TLLM_CHECK_WITH_INFO( + static_cast(crossKvCacheManager), "Encoder-decoder scheduling requires a cross_kv_cache_manager."); +} + } // namespace MaxRequestsScheduler::MaxRequestsScheduler( @@ -192,6 +206,8 @@ std::tuple GuaranteedNoEvictScheduler::impl( { RequestVector scheduledRequests; + checkRequiredCrossKvCacheManager(getNoScheduleUntilState(), crossKvCacheManager); + // Now check if we can add pending requests auto const maxPeftCachePages = peftCacheManager ? peftCacheManager->getMaxDevicePages() : std::numeric_limits::max(); @@ -287,6 +303,12 @@ std::tuple GuaranteedNoEvictScheduler::impl( // eliminating 2 redundant walks per request. bool const isFirstChunkContext = req->isContextInitState() && req->isFirstContextChunk() && !req->isDisaggGenerationInitState(); + // Encoder-init requests do not consume self- or cross-KV + // blocks. We still keep the cross reuse summary available for + // beneficial-to-skip so duplicate encoder inputs can be ordered + // consistently before their decoder-context admission budgets + // the cross pool. + bool const isEncoderInit = req->isEncoderInitState(); std::optional summary; std::optional crossSummary; if (isFirstChunkContext) @@ -305,9 +327,15 @@ std::tuple GuaranteedNoEvictScheduler::impl( crossSummary = crossKvCacheManager->analyzePrefixReuse(uniqueTokens, *req); } } - + else if (isEncoderInit && crossKvCacheManager && crossKvCacheManager->isEnableBlockReuse() + && !crossKvCacheManager->getBlockManager().isVariableWindow()) + { + // Encoder admission only needs the cross summary for reuse ordering. + auto uniqueTokens = *(req->getEncoderUniqueTokens().value()); + crossSummary = crossKvCacheManager->analyzePrefixReuse(uniqueTokens, *req); + } // Beneficial-to-skip check using the cached summary - if (!StaticBatchScheduling && skippingIsRelevant && isFirstChunkContext + if (!StaticBatchScheduling && skippingIsRelevant && (isFirstChunkContext || isEncoderInit) && beneficialToSkip( summary, crossSummary, newlyContributedContextBlocks, newlyContributedCrossContextBlocks)) { @@ -319,7 +347,23 @@ std::tuple GuaranteedNoEvictScheduler::impl( break; } - if (req->isContextInitState() || req->isDisaggGenerationInitState()) + if (isEncoderInit) + { + bool reqHasLora = req->getLoraTaskId().has_value(); + bool isNewTask = reqHasLora && !uniqTaskIds.count(req->getLoraTaskId().value()); + auto neededPeftPages = isNewTask && peftCacheManager ? peftCacheManager->determineNumPages(req) : 0; + + if (neededPeftPages <= availablePeftPages) + { + scheduledRequests.emplace_back(req); + availablePeftPages -= neededPeftPages; + if (isNewTask) + { + uniqTaskIds.insert(req->getLoraTaskId().value()); + } + } + } + else if (req->isContextInitState() || req->isDisaggGenerationInitState()) { // Check block availability using the cached summary when available. // enoughAvailableBlocks is check-only (no decrement) — safe if cross check fails. @@ -364,15 +408,23 @@ std::tuple GuaranteedNoEvictScheduler::impl( // the remote diff is easier to look at/rebase conflicts bool trySchedulingRequestMaxUtilization(std::shared_ptr const& req, SizeType32 maxNumRequests, RequestVector& scheduledRequests, kv_cache_manager::MaxUtilizationScheduledBlocksManager& blocksManager, + std::optional& crossBlocksManager, OptionalRef peftCacheManager, SizeType32& numScheduledPeftPages, std::unordered_set& seenTaskIds, std::optional const& cachedSummary); std::tuple MaxUtilizationScheduler::operator()( - kv_cache_manager::BaseKVCacheManager& kvCacheManager, OptionalRef peftCacheManager, - RequestList const& activeRequests) const + kv_cache_manager::BaseKVCacheManager& kvCacheManager, + OptionalRef crossKvCacheManager, + OptionalRef peftCacheManager, RequestList const& activeRequests) const { + checkRequiredCrossKvCacheManager(getNoScheduleUntilState(), crossKvCacheManager); + kvCacheManager.startScheduling(); + if (crossKvCacheManager) + { + crossKvCacheManager->startScheduling(); + } // The optimization of delaying requests won't work for variable window attention bool skippingIsRelevant = !kvCacheManager.getBlockManager().isVariableWindow(); @@ -380,6 +432,14 @@ std::tuple MaxUtilizationScheduler::operator()( // Keep track of number of requests and block needed for the scheduled requests auto scheduledBlocksManager = kv_cache_manager::MaxUtilizationScheduledBlocksManager(kvCacheManager, mTwoStepsLookAhead); + // Mirror the budget tracker for the cross pool when present. + // Encoder-init requests do not consume either tracker; decoder + // context/generation requests update both trackers in lockstep. + std::optional scheduledCrossBlocksManager; + if (crossKvCacheManager) + { + scheduledCrossBlocksManager.emplace(*crossKvCacheManager, mTwoStepsLookAhead); + } SizeType32 numScheduledPeftPages{0}; std::unordered_set seenTaskIds; @@ -387,7 +447,9 @@ std::tuple MaxUtilizationScheduler::operator()( auto [newlyContributedContextBlocks, newlyContributedCrossContextBlocks] = prefillWithChunkedContextsAlreadyExecuting(activeRequests, kvCacheManager); - // Find last active in case we need to evict + // Find last active in case we need to evict. Encoder-init requests are + // intentionally excluded here: they hold no started self- or cross-pool + // blocks, so pausing them would not free any KV budget. auto startedReqLambda = [this](std::shared_ptr const& req) { return (req->hasReachedState(getNoScheduleUntilState()) && !req->hasReachedState(getNoScheduleAfterState()) @@ -437,8 +499,9 @@ std::tuple MaxUtilizationScheduler::operator()( continue; } - bool const wasScheduled = trySchedulingRequestMaxUtilization(req, mMaxNumRequests, scheduledRequests, - scheduledBlocksManager, peftCacheManager, numScheduledPeftPages, seenTaskIds, summary); + bool const wasScheduled + = trySchedulingRequestMaxUtilization(req, mMaxNumRequests, scheduledRequests, scheduledBlocksManager, + scheduledCrossBlocksManager, peftCacheManager, numScheduledPeftPages, seenTaskIds, summary); if (wasScheduled) { TLLM_LOG_DEBUG("MaxUtilizationScheduler: request ID %lu -> start", req->mRequestId); @@ -455,6 +518,13 @@ std::tuple MaxUtilizationScheduler::operator()( // from the end of the vector and try again // Here we simulate freeing the kvCache blocks associated with that sequence kvCacheManager.schedulingRemoveSequence((*lastStartedReqIt)->mRequestId); + if (crossKvCacheManager) + { + // Mirror self-pool eviction on the cross pool so any cross + // blocks held by the paused request are released for reuse + // by other admissions in this iteration. + crossKvCacheManager->schedulingRemoveSequence((*lastStartedReqIt)->mRequestId); + } pausedRequests.emplace_back(*lastStartedReqIt); TLLM_LOG_INFO("MaxUtilizationScheduler: request ID %lu -> pause", (*lastStartedReqIt)->mRequestId); reqItEnd = std::next(lastStartedReqIt).base(); @@ -471,6 +541,7 @@ std::tuple MaxUtilizationScheduler::operator()( bool trySchedulingRequestMaxUtilization(std::shared_ptr const& req, SizeType32 maxNumRequests, RequestVector& scheduledRequests, kv_cache_manager::MaxUtilizationScheduledBlocksManager& blocksManager, + std::optional& crossBlocksManager, OptionalRef peftCacheManager, SizeType32& numScheduledPeftPages, std::unordered_set& seenTaskIds, std::optional const& cachedSummary) { @@ -482,16 +553,51 @@ bool trySchedulingRequestMaxUtilization(std::shared_ptr const& req, = (isNewTask && peftCacheManager) ? peftCacheManager->determineNumPages(req) : 0; TLLM_LOG_DEBUG( "MaxUtilizationScheduler: request ID %lu required peft pages: %i", req->mRequestId, numRequiredPeftPages); - // Use the cached summary when available to avoid a redundant tree walk - auto const scheduledBlocksIfFitsKvCache - = blocksManager.prepareNewNumberOfBlocksIfWeEndUpScheduling(*req, cachedSummary); bool fitsPeft = (peftCacheManager ? numRequiredPeftPages + numScheduledPeftPages <= peftCacheManager->getMaxDevicePages() : true); + if (req->isEncoderInitState()) + { + // Encoder admission does not reserve KV blocks. The scheduler + // entry point verifies the cross manager globally before encoder + // work can be admitted. + if (fitsPeft) + { + numScheduledPeftPages += numRequiredPeftPages; + scheduledRequests.emplace_back(req); + if (isNewTask) + { + seenTaskIds.insert(req->getLoraTaskId().value()); + } + return true; + } + return false; + } + + // Use the cached summary when available to avoid a redundant tree walk + auto const scheduledBlocksIfFitsKvCache + = blocksManager.prepareNewNumberOfBlocksIfWeEndUpScheduling(*req, cachedSummary); + // Context/generation requests must fit in both pools when a cross + // manager is present. Self-pool fit is checked first so that the + // budget probe is cheap when self is already saturated. + std::optional> crossScheduledIfFits; + if (crossBlocksManager) + { + crossScheduledIfFits = crossBlocksManager->prepareNewNumberOfBlocksIfWeEndUpScheduling(*req); + if (!crossScheduledIfFits) + { + return false; + } + } + if (scheduledBlocksIfFitsKvCache && fitsPeft) { blocksManager.updateScheduledBlocks(scheduledBlocksIfFitsKvCache.value()); + if (crossScheduledIfFits) + { + crossBlocksManager->updateScheduledBlocks(crossScheduledIfFits.value()); + } numScheduledPeftPages += numRequiredPeftPages; TLLM_LOG_DEBUG("MaxUtilizationScheduler: scheduled peft pages: %i", numRequiredPeftPages); scheduledRequests.emplace_back(req); @@ -546,7 +652,7 @@ void CapacityScheduler::setAgentTreeReorderPolicy( std::tuple CapacityScheduler::operator()(RequestList const& activeRequests, OptionalRef kvCacheManager, OptionalRef peftCacheManager, - OptionalRef crossKvCacheManager) const + OptionalRef crossKvCacheManager) const { NVTX3_SCOPED_RANGE(capacitySchedulerScheduling); @@ -566,7 +672,7 @@ std::tuple CapacityScheduler::opera else if constexpr (std::is_same_v, MaxUtilizationScheduler>) { std::tie(tmpFittingRequests, pausedRequests) - = scheduler(*kvCacheManager, peftCacheManager, requestsToSchedule); + = scheduler(*kvCacheManager, crossKvCacheManager, peftCacheManager, requestsToSchedule); } else if constexpr (std::is_same_v, GuaranteedNoEvictScheduler> || std::is_same_v, StaticBatchScheduler>) diff --git a/cpp/tensorrt_llm/common/attentionOp.cpp b/cpp/tensorrt_llm/common/attentionOp.cpp index fd66e5d9490d..0ea70da0d73b 100644 --- a/cpp/tensorrt_llm/common/attentionOp.cpp +++ b/cpp/tensorrt_llm/common/attentionOp.cpp @@ -1735,8 +1735,12 @@ int AttentionOp::enqueueContext(EnqueueContextParams const& params, cudaStrea preprocessingParams.qkv_bias = params.qkv_bias; preprocessingParams.tokens_info = decoder_params.tokensInfo; preprocessingParams.seq_lens = params.context_lengths; - // Indicate if chunked-context is used (i.e. q_seqlen > kv_seqlen). - preprocessingParams.cache_seq_lens = params.sequence_lengths; + // For self-attention, cache_seq_lens indicates whether chunked context is used + // (i.e. cache_seq_len > seq_len). + // For cross-attention, callers do not consistently use sequence_lengths as decoder length; use decoder + // context lengths so the encoder KV-cache write gate opens. + preprocessingParams.cache_seq_lens = isCrossAttention() ? params.context_lengths : params.sequence_lengths; + preprocessingParams.encoder_seq_lens = params.encoder_input_lengths; preprocessingParams.cu_seq_lens = contextCuQSeqlens; // Cross-attention only. diff --git a/cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp b/cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp index 086cdf7547c2..ee257713b85d 100644 --- a/cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp @@ -307,7 +307,27 @@ void initBindings(nb::module_& m) return std::optional(*encoderUniqueTokens.value()); } return std::optional(std::nullopt); - }); + }) + // Encoder-decoder accessors for PyExecutor encoder iteration. + // ``encoder_tokens`` returns the source-side tokens used to drive the + // encoder forward. ``encoder_output_len`` is the cross-KV capacity + // for the request (number of encoder hidden states the decoder + // cross-attention will read), which mirrors the C++ + // ``getEncoderOutputLen`` contract. + // ``try_get_encoder_output_len`` exposes the same value as an + // optional probe for decoder-only scheduler paths. + .def_prop_ro("encoder_tokens", + [](GenLlmReq& self) -> std::optional + { + auto const& encoderTokens = self.getEncoderTokens(); + if (encoderTokens.has_value() && encoderTokens.value()) + { + return std::optional(*encoderTokens.value()); + } + return std::nullopt; + }) + .def("try_get_encoder_output_len", &GenLlmReq::tryGetEncoderOutputLen) + .def_prop_ro("encoder_output_len", &GenLlmReq::getEncoderOutputLen); nb::class_(m, "LlmRequest", nb::dynamic_attr()) .def( diff --git a/cpp/tests/unit_tests/batch_manager/capacitySchedulerTest.cpp b/cpp/tests/unit_tests/batch_manager/capacitySchedulerTest.cpp index 9334d813e93e..267361cd73da 100644 --- a/cpp/tests/unit_tests/batch_manager/capacitySchedulerTest.cpp +++ b/cpp/tests/unit_tests/batch_manager/capacitySchedulerTest.cpp @@ -2268,3 +2268,123 @@ TEST_F(CapacitySchedulerTest, MaxUtilizationNoReuseWhenDisabled) // Both requests start at iteration 0 and finish together, so numIterations = maxNewTokens EXPECT_EQ(numIterations, maxNewTokens); } + +// ============================================================================ +// ENCODER_INIT admission tests for dual-pool capacity scheduling. +// ============================================================================ +// +// These tests exercise the C++ scheduler paths that admit requests in the +// LlmRequestState::kENCODER_INIT state. Encoder-init requests must not reserve +// blocks from either the self or cross KV cache; decoder CONTEXT_INIT admission +// owns that budgeting. They also must not be considered eviction victims by +// MaxUtilization. +// +// Unlike the legacy enc-dec tests above (which use prepRequestsForEncoderSkip +// to flip ENCODER_INIT → CONTEXT_INIT before the scheduler runs), the tests +// below construct the scheduler with no_schedule_until_state=kENCODER_INIT +// so the encoder phase reaches the policy code paths directly. + +namespace +{ +// Helper to create an encoder-decoder request that stays in the ENCODER_INIT +// state when it reaches the scheduler. +std::shared_ptr createEncoderInitRequest( + int32_t promptLen, int32_t maxNewTokens, int32_t encoderInputLen, uint64_t reqId) +{ + auto inputTokens = VecTokens(promptLen, 1); + auto encoderInputTokens = VecTokens(encoderInputLen, 1); + tensorrt_llm::executor::OutputConfig outConfig; + outConfig.excludeInputFromOutput = false; + outConfig.returnLogProbs = false; + outConfig.returnGenerationLogits = false; + outConfig.returnContextLogits = false; + outConfig.returnEncoderOutput = false; + bool streaming = false; + auto executorReq = tensorrt_llm::executor::Request( + inputTokens, maxNewTokens, streaming, tensorrt_llm::executor::SamplingConfig(), outConfig); + executorReq.setEncoderInputTokenIds(encoderInputTokens); + auto req = std::make_shared(reqId, executorReq); + // executor::Request ctor sets kENCODER_INIT when encoderInputTokenIds is present. + EXPECT_EQ(req->getState(), LlmRequestState::kENCODER_INIT); + return req; +} +} // namespace + +// GuaranteedNoEvict: a single encoder-init request is admitted without +// consuming self- or cross-pool blocks. +// Without a cross_kv_cache_manager, an encoder-init request cannot honour the +// dual-pool contract and must fail fast for both policies. +TEST_F(CapacitySchedulerTest, EncoderInitWithoutCrossManagerThrows) +{ + SizeType32 const maxNumRequests = 4; + SizeType32 const tokensPerBlock = 10; + SizeType32 const selfMaxTokens = 200; + SizeType32 const selfMaxTokensPerSeq = 100; + int32_t const encoderInputLen = 20; + + for (auto policy : {CapacitySchedulerPolicy::kGUARANTEED_NO_EVICT, CapacitySchedulerPolicy::kMAX_UTILIZATION}) + { + auto kvCacheManager = getKvCacheManager(maxNumRequests, tokensPerBlock, selfMaxTokens, selfMaxTokensPerSeq); + auto peftCacheManager = getPeftCacheManager(); + auto capacityScheduler = CapacityScheduler(maxNumRequests, policy, kvCacheManager != nullptr, + /*twoStepsLookAhead=*/false, LlmRequestState::kENCODER_INIT, LlmRequestState::kGENERATION_COMPLETE); + + RequestList activeRequests; + activeRequests.push_back( + createEncoderInitRequest(/*promptLen=*/10, /*maxNewTokens=*/40, encoderInputLen, /*id=*/1)); + + // No cross manager passed. + EXPECT_THROW((void) capacityScheduler( + activeRequests, kvCacheManager, peftCacheManager, /*crossKvCacheManager=*/std::nullopt), + tc::TllmException) + << "policy=" << static_cast(policy); + } +} + +// Cross pool pressure does not throttle encoder admission. The request will +// face the cross-pool budget on its later decoder CONTEXT_INIT iteration. +TEST_F(CapacitySchedulerTest, EncoderInitDoesNotConsumeCrossPool) +{ + SizeType32 const maxNumRequests = 4; + SizeType32 const tokensPerBlock = 10; + SizeType32 const selfMaxTokens = 400; + SizeType32 const selfMaxTokensPerSeq = 100; + // Cross pool fits one 20-token encoder sequence, but encoder admission + // should not consume it. + SizeType32 const crossMaxTokens = 20; + SizeType32 const crossMaxTokensPerSeq = 20; + int32_t const encoderInputLen = 20; + + for (auto policy : {CapacitySchedulerPolicy::kGUARANTEED_NO_EVICT, CapacitySchedulerPolicy::kMAX_UTILIZATION}) + { + auto kvCacheManager = getKvCacheManager(maxNumRequests, tokensPerBlock, selfMaxTokens, selfMaxTokensPerSeq); + auto crossKvCacheManager = getKvCacheManager(maxNumRequests, tokensPerBlock, crossMaxTokens, + crossMaxTokensPerSeq, /*sinkTokenLength=*/0, /*enableReuse=*/false, kv_cache_manager::CacheType::kCROSS); + auto peftCacheManager = getPeftCacheManager(); + auto capacityScheduler = CapacityScheduler(maxNumRequests, policy, kvCacheManager != nullptr, + /*twoStepsLookAhead=*/false, LlmRequestState::kENCODER_INIT, LlmRequestState::kGENERATION_COMPLETE); + + RequestList activeRequests; + activeRequests.push_back( + createEncoderInitRequest(/*promptLen=*/10, /*maxNewTokens=*/40, encoderInputLen, /*id=*/1)); + activeRequests.push_back( + createEncoderInitRequest(/*promptLen=*/10, /*maxNewTokens=*/40, encoderInputLen, /*id=*/2)); + + auto const selfFreeBefore = kvCacheManager->getNumFreeBlocks(); + auto const crossFreeBefore = crossKvCacheManager->getNumFreeBlocks(); + + auto [fittingRequests, fittingDisaggGenInitRequests, pausedRequests] + = capacityScheduler(activeRequests, kvCacheManager, peftCacheManager, crossKvCacheManager); + + EXPECT_EQ(fittingRequests.size(), 2u) << "policy=" << static_cast(policy); + EXPECT_EQ(fittingDisaggGenInitRequests.size(), 0u) << "policy=" << static_cast(policy); + EXPECT_EQ(pausedRequests.size(), 0u) << "policy=" << static_cast(policy); + EXPECT_EQ(fittingRequests.front()->mRequestId, 1u) << "policy=" << static_cast(policy); + EXPECT_EQ(fittingRequests.back()->mRequestId, 2u) << "policy=" << static_cast(policy); + + // Scheduling alone reserves blocks via in-memory bookkeeping only; the + // managers' free-block counters are unaffected by encoder admission. + EXPECT_EQ(kvCacheManager->getNumFreeBlocks(), selfFreeBefore) << "policy=" << static_cast(policy); + EXPECT_EQ(crossKvCacheManager->getNumFreeBlocks(), crossFreeBefore) << "policy=" << static_cast(policy); + } +} diff --git a/tensorrt_llm/_torch/attention_backend/interface.py b/tensorrt_llm/_torch/attention_backend/interface.py index 2856a1769fe8..5c6c5ee12582 100644 --- a/tensorrt_llm/_torch/attention_backend/interface.py +++ b/tensorrt_llm/_torch/attention_backend/interface.py @@ -431,6 +431,78 @@ def update_helix_param( Hook to be called when using helix parallelism. """ + def create_cross_metadata( + self, + encoder_seq_lens: torch.Tensor, + cross_kv_cache_manager: Union[KVCacheManager, KVCacheManagerV2, + None] = None, + *, + encoder_num_cached_tokens_per_seq: Optional[List[int]] = None, + ) -> "AttentionMetadata": + """Build a sub-metadata instance for cross-attention. + + The returned metadata shares Q-side fields (``seq_lens``, + ``request_ids``, ``num_contexts``) with ``self`` (the decoder + self-attention metadata) and overrides the K/V-side with the encoder + lengths so that ``returned.is_cross is True``. + + This is intended to be called by the runtime / unit tests once the + encoder lengths and cross-pool KV cache manager are known. The + returned object is a *new* metadata instance (not stored on + ``self.cross``); callers can attach it to ``self.cross`` if desired. + + Args: + encoder_seq_lens: Per-request encoder sequence length (CPU + int32 tensor). On the first decoder context step this is + the full encoder length; on generation steps it should be + ``0`` (no new K/V tokens to add to the cross pool — the + encoder K/V are already cached). + cross_kv_cache_manager: KV cache manager for the cross pool. + When ``None``, the returned metadata uses the stateless + (no-KV-cache) path (suitable for unit tests). + encoder_num_cached_tokens_per_seq: Per-request count of encoder + K/V tokens already present in the cross pool. ``None`` + defaults to 0 (context phase, nothing cached yet). + + Returns: + A new ``AttentionMetadata`` of the same subclass as ``self``, + with ``seq_lens_kv`` set to ``encoder_seq_lens`` so that + ``is_cross`` becomes ``True``. + """ + cross_md = copy.copy(self) + cross_md._saved_tensors = {} + if self.is_cuda_graph: + # Cross-attention has K/V lengths from the encoder, while + # self-attention has K/V lengths from the decoder. Keep their + # CUDA graph metadata buffers separate so preparing cross metadata + # cannot overwrite self-attention sequence lengths. + cross_md.cuda_graph_buffers = Buffers() + cross_md.kv_cache_manager = cross_kv_cache_manager + cross_md._seq_lens_kv = None + cross_md._seq_lens_kv_cuda = None + cross_md.cross = None + cross_md.seq_lens_kv = encoder_seq_lens + if encoder_num_cached_tokens_per_seq is not None: + from ..metadata import KVCacheParams + base_params = self.kv_cache_params + cross_md.kv_cache_params = KVCacheParams( + use_cache=base_params.use_cache if base_params is not None else + (cross_kv_cache_manager is not None), + num_cached_tokens_per_seq=list( + encoder_num_cached_tokens_per_seq), + block_ids_per_seq=base_params.block_ids_per_seq + if base_params is not None else None, + host_max_attention_window_sizes=base_params. + host_max_attention_window_sizes + if base_params is not None else None, + host_sink_token_length=base_params.host_sink_token_length + if base_params is not None else None, + num_extra_kv_tokens=base_params.num_extra_kv_tokens + if base_params is not None else 0, + ) + cross_md.__post_init__() + return cross_md + def update_for_spec_dec(self) -> None: """ Hook to be called during forward when using spec-dec one-model mode. diff --git a/tensorrt_llm/_torch/attention_backend/trtllm.py b/tensorrt_llm/_torch/attention_backend/trtllm.py index bf1268549778..0f009fc8f64a 100644 --- a/tensorrt_llm/_torch/attention_backend/trtllm.py +++ b/tensorrt_llm/_torch/attention_backend/trtllm.py @@ -1676,9 +1676,9 @@ def _run( update_kv_cache=forward_args.update_kv_cache, cross_kv=forward_args.cross_kv, relative_attention_bias=forward_args.relative_attention_bias, - relative_attention_max_distance=( - forward_args.relative_attention_max_distance), - + relative_attention_max_distance=forward_args. + relative_attention_max_distance, + position_embedding_type=self.position_embedding_type, # --- Module config (TrtllmAttention) --- rotary_inv_freq=self.rotary_inv_freq, rotary_cos_sin=self.rotary_cos_sin, @@ -1689,7 +1689,6 @@ def _run( head_size=self.head_dim, quant_mode=self.quant_mode, q_scaling=self.q_scaling, - position_embedding_type=self.position_embedding_type, rope_dim=self.rope_dim, rope_base=self.rope_base, rope_scale_type=self.rope_scale_type, diff --git a/tensorrt_llm/_torch/attention_backend/vanilla.py b/tensorrt_llm/_torch/attention_backend/vanilla.py index 7b0c138c77d3..b9648c171b1d 100644 --- a/tensorrt_llm/_torch/attention_backend/vanilla.py +++ b/tensorrt_llm/_torch/attention_backend/vanilla.py @@ -314,16 +314,23 @@ def no_kv_cache_forward( metadata: AttentionMetadata, *, attention_mask: AttentionMask = PredefinedAttentionMask.CAUSAL, - position_ids: Optional[torch.Tensor] = None) -> torch.Tensor: - """ - This function is used to perform attention without kv cache. + position_ids: Optional[torch.Tensor] = None, + **kwargs) -> torch.Tensor: + """Perform attention without kv cache. + + Supports both self-attention (Q and K/V have matching per-request + lengths) and cross-attention (Q-side lengths from + ``metadata.seq_lens``, K/V-side lengths from ``metadata.seq_lens_kv``, + i.e. ``metadata.is_cross is True``). + Args: - q (torch.Tensor): Query tensor with shape (seq_len, num_heads * head_dim) or (seq_len, (num_heads + 2 * num_kv_heads) * head_dim), - k (Optional[torch.Tensor]): Key tensor with shape (seq_len, num_heads * head_dim) or None, - v (Optional[torch.Tensor]): Value tensor with shape (seq_len, num_heads * head_dim) or None, + q: Query tensor, shape ``(seq_len_q, num_heads * head_dim)`` + or ``(seq_len_q, (num_heads + 2*num_kv_heads) * head_dim)``. + k: Key tensor, shape ``(seq_len_kv, num_kv_heads * head_dim)`` or + None (fused QKV input). + v: Value tensor, shape ``(seq_len_kv, num_kv_heads * head_dim)`` + or None (fused QKV input). """ - # lazy loading - from flash_attn.flash_attn_interface import flash_attn_varlen_func head_dim = q.shape[-1] is_fused_qkv = False if (k is None) or (v is None): @@ -345,15 +352,51 @@ def no_kv_cache_forward( assert q.dim() == 3 assert k.dim() == 3 assert v.dim() == 3 - seqlens_in_batch = metadata.seq_lens - assert seqlens_in_batch is not None, "seq_len can not be None for remove padding inputs attention!" - max_seqlen_in_batch = seqlens_in_batch.max().item() - cu_seqlens = F.pad( - torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), - (1, 0)).to(q.device) + seqlens_q = metadata.seq_lens + assert seqlens_q is not None, "seq_len can not be None for remove padding inputs attention!" + seqlens_kv = metadata.seq_lens_kv + # In cross-attention the K/V-side lengths differ from the Q-side + # lengths and must be tracked separately for cu_seqlens. + is_cross = metadata.is_cross + if is_fused_qkv and is_cross: + raise ValueError( + "Cross-attention with fused QKV input is not supported: pass " + "Q, K, V as separate tensors when metadata.is_cross is True.") + max_seqlen_q = int(seqlens_q.max().item()) + cu_seqlens_q = F.pad(torch.cumsum(seqlens_q, dim=0, dtype=torch.int32), + (1, 0)).to(q.device) + if is_cross: + assert seqlens_kv is not None, ( + "metadata.seq_lens_kv must be set for cross-attention " + "(no_kv_cache_forward). Got None.") + assert seqlens_kv.sum().item() == k.size(0), ( + "K tensor token count does not match metadata.seq_lens_kv: " + f"k.shape[0]={k.size(0)} vs sum(seq_lens_kv)=" + f"{seqlens_kv.sum().item()}.") + max_seqlen_k = int(seqlens_kv.max().item()) + cu_seqlens_k = F.pad( + torch.cumsum(seqlens_kv, dim=0, dtype=torch.int32), + (1, 0)).to(q.device) + else: + max_seqlen_k = max_seqlen_q + cu_seqlens_k = cu_seqlens_q + + # flash-attn only supports fp16/bf16; fall back to PyTorch SDPA for + # other dtypes (e.g. float32), mirroring the TRT backend's behaviour + # of disabling context_fmha for float32. + if q.dtype not in (torch.float16, torch.bfloat16): + return self._no_kv_cache_sdpa_fallback(q, k, v, num_heads, + num_kv_heads, head_dim, + seqlens_q, cu_seqlens_q, + max_seqlen_q, attention_mask, + seqlens_kv, cu_seqlens_k, + max_seqlen_k, is_cross) - max_seqlen_q = max_seqlen_k = max_seqlen_in_batch - cu_seqlens_q = cu_seqlens_k = cu_seqlens + from flash_attn.flash_attn_interface import flash_attn_varlen_func + + softmax_scale = None + if self.q_scaling is not None: + softmax_scale = 1 / (math.sqrt(head_dim) * self.q_scaling) attn_output_unpad = flash_attn_varlen_func( q, @@ -364,9 +407,9 @@ def no_kv_cache_forward( max_seqlen_q, max_seqlen_k, dropout_p=0.0, - softmax_scale=None, - causal=attention_mask == PredefinedAttentionMask.CAUSAL, - # window_size=(-1, -1), # -1 means infinite context window + softmax_scale=softmax_scale, + causal=attention_mask == PredefinedAttentionMask.CAUSAL + and not is_cross, alibi_slopes=None, deterministic=False, return_attn_probs=False, @@ -374,6 +417,65 @@ def no_kv_cache_forward( return attn_output_unpad.reshape(attn_output_unpad.size(0), -1) + def _no_kv_cache_sdpa_fallback(self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + num_heads: int, + num_kv_heads: int, + head_dim: int, + seqlens_q: torch.Tensor, + cu_seqlens_q: torch.Tensor, + max_seqlen_q: int, + attention_mask: AttentionMask, + seqlens_kv: Optional[torch.Tensor] = None, + cu_seqlens_k: Optional[torch.Tensor] = None, + max_seqlen_k: Optional[int] = None, + is_cross: bool = False) -> torch.Tensor: + """PyTorch SDPA fallback for dtypes not supported by flash-attn. + + When ``seqlens_kv`` / ``cu_seqlens_k`` are provided, K/V are sliced + independently of Q (cross-attention path). + """ + del max_seqlen_q, max_seqlen_k # only seqlens / cu_seqlens are used + is_causal = (attention_mask == PredefinedAttentionMask.CAUSAL) + num_kv_groups = num_heads // num_kv_heads + num_requests = seqlens_q.numel() + + if seqlens_kv is None or cu_seqlens_k is None: + seqlens_kv = seqlens_q + cu_seqlens_k = cu_seqlens_q + + outputs = [] + for i in range(num_requests): + start_q = cu_seqlens_q[i].item() + end_q = cu_seqlens_q[i + 1].item() + start_k = cu_seqlens_k[i].item() + end_k = cu_seqlens_k[i + 1].item() + q_s = q[start_q:end_q].transpose(0, 1).unsqueeze(0) + k_s = k[start_k:end_k].transpose(0, 1).unsqueeze(0) + v_s = v[start_k:end_k].transpose(0, 1).unsqueeze(0) + k_s = repeat_kv(k_s, num_kv_groups) + v_s = repeat_kv(v_s, num_kv_groups) + + qk_scale = None + if self.q_scaling is not None: + qk_scale = 1 / (math.sqrt(head_dim) * self.q_scaling) + + # SDPA's is_causal flag implies square attention. Cross-attention + # is never causal: the decoder Q attends to all encoder K/V tokens. + sdpa_is_causal = (is_causal and not is_cross + and (end_q - start_q) == (end_k - start_k)) + out = F.scaled_dot_product_attention(q_s, + k_s, + v_s, + is_causal=sdpa_is_causal, + scale=qk_scale) + outputs.append(out.squeeze(0).transpose(0, 1)) + + result = torch.cat(outputs, dim=0) + return result.reshape(result.size(0), -1) + def forward(self, q: torch.Tensor, k: Optional[torch.Tensor], diff --git a/tensorrt_llm/_torch/model_config.py b/tensorrt_llm/_torch/model_config.py index fa448d8876c2..d1a2fa5503b8 100644 --- a/tensorrt_llm/_torch/model_config.py +++ b/tensorrt_llm/_torch/model_config.py @@ -138,6 +138,7 @@ class ModelConfig(Generic[TConfig]): sparse_attention_config: Optional["SparseAttentionConfig"] = None is_generation: bool = True + is_encoder_decoder: bool = False max_num_tokens: int = 8192 max_seq_len: Optional[int] = None @@ -203,6 +204,10 @@ def __setattr__(self, key, value): super().__setattr__(key, value) def __post_init__(self): + if self.pretrained_config: + self.is_encoder_decoder = self.is_encoder_decoder_model( + self.pretrained_config) + if self.pretrained_config and hasattr(self.pretrained_config, "architectures"): self.is_generation = self.is_generation_model( @@ -270,6 +275,18 @@ def get_quant_config(self, name: Optional[str] = None) -> QuantConfig: raise ValueError(f'quant config of {name} is not found') + @staticmethod + def is_encoder_decoder_model(pretrained_config: Optional[TConfig]) -> bool: + if pretrained_config is None: + return False + text_config = pretrained_config + get_text_config = getattr(pretrained_config, "get_text_config", None) + if callable(get_text_config): + text_config = get_text_config() + elif hasattr(pretrained_config, "text_config"): + text_config = pretrained_config.text_config + return getattr(text_config, "is_encoder_decoder", False) + @staticmethod def is_generation_model(model_architectures: Optional[List[str]], mm_encoder_only: bool = False) -> bool: diff --git a/tensorrt_llm/_torch/models/__init__.py b/tensorrt_llm/_torch/models/__init__.py index 11760e684361..20e65ab1c119 100644 --- a/tensorrt_llm/_torch/models/__init__.py +++ b/tensorrt_llm/_torch/models/__init__.py @@ -7,6 +7,8 @@ from .modeling_afmoe import AfmoeForCausalLM from .modeling_auto import AutoModelForCausalLM +from .modeling_bart import (BartForConditionalGeneration, + MBartForConditionalGeneration) from .modeling_bert import BertForSequenceClassification from .modeling_clip import CLIPVisionModel from .modeling_cohere2 import Cohere2ForCausalLM @@ -49,6 +51,7 @@ from .modeling_starcoder2 import Starcoder2ForCausalLM from .modeling_step3p7 import Step3p7ForCausalLM from .modeling_step3p7vl import Step3p7VLForConditionalGeneration +from .modeling_t5 import T5ForConditionalGeneration from .modeling_utils import get_model_architecture from .modeling_vila import VilaModel @@ -56,6 +59,7 @@ __all__ = [ "AfmoeForCausalLM", "AutoModelForCausalLM", + "BartForConditionalGeneration", "BertForSequenceClassification", "CLIPVisionModel", "DeepseekV3ForCausalLM", @@ -86,6 +90,8 @@ "Qwen2MoeForCausalLM", "SiglipVisionModel", "Starcoder2ForCausalLM", + "T5ForConditionalGeneration", + "MBartForConditionalGeneration", "get_model_architecture", "VilaModel", "Qwen2VLModel", diff --git a/tensorrt_llm/_torch/models/modeling_bart.py b/tensorrt_llm/_torch/models/modeling_bart.py new file mode 100644 index 000000000000..7866140036c9 --- /dev/null +++ b/tensorrt_llm/_torch/models/modeling_bart.py @@ -0,0 +1,726 @@ +# 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"); +# 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. +"""PyTorch-flow BART / mBART encoder-decoder model for TensorRT-LLM. + +Covers ``BartForConditionalGeneration`` and ``MBartForConditionalGeneration``. + +Key differences from T5: + - LayerNorm instead of RMSNorm. + - Post-norm (residual → add → LayerNorm) instead of pre-norm. + - Learned absolute positional embeddings (not relative bias). + - GELU activation (not ReLU / gated). + - Bias in attention and MLP projections. + - Embedding scale = sqrt(d_model). +""" + +import math +from typing import Dict, Optional + +import torch +import torch.nn.functional as F +from torch import nn +from transformers import BartConfig + +from ..attention_backend import AttentionMetadata +from ..attention_backend.interface import PredefinedAttentionMask +from ..model_config import ModelConfig +from ..modules.attention import Attention +from ..modules.cross_attention import CrossAttention +from ..modules.embedding import Embedding, LMHead +from ..modules.layer_norm import LayerNorm +from ..modules.linear import TensorParallelMode +from ..modules.logits_processor import LogitsProcessor +from ..modules.mlp import MLP +from .modeling_utils import PostInitCaller, register_auto_model + +# --------------------------------------------------------------------------- +# Config helpers +# --------------------------------------------------------------------------- + + +def _bart_encoder_hidden_size(config: BartConfig) -> int: + return config.d_model + + +def _bart_decoder_hidden_size(config: BartConfig) -> int: + return config.d_model + + +def _bart_encoder_num_heads(config: BartConfig) -> int: + return config.encoder_attention_heads + + +def _bart_decoder_num_heads(config: BartConfig) -> int: + return config.decoder_attention_heads + + +def _bart_encoder_ffn_dim(config: BartConfig) -> int: + return config.encoder_ffn_dim + + +def _bart_decoder_ffn_dim(config: BartConfig) -> int: + return config.decoder_ffn_dim + + +def _bart_encoder_num_layers(config: BartConfig) -> int: + return config.encoder_layers + + +def _bart_decoder_num_layers(config: BartConfig) -> int: + return config.decoder_layers + + +def _bart_head_dim(config: BartConfig) -> int: + return config.d_model // config.encoder_attention_heads + + +def _packed_position_ids( + position_ids: Optional[torch.IntTensor], + hidden_states: torch.Tensor, +) -> Optional[torch.IntTensor]: + if position_ids is None: + return None + + position_ids = position_ids.reshape(-1) + if position_ids.numel() != hidden_states.shape[0]: + raise ValueError( + "BART packed position_ids must match hidden_states tokens: " + f"got {position_ids.numel()} positions for {hidden_states.shape[0]} tokens." + ) + return position_ids + + +# --------------------------------------------------------------------------- +# BART Attention +# --------------------------------------------------------------------------- + + +class BartSelfAttention(Attention): + """BART-style MHA with bias and no positional encoding in the kernel. + + BART uses learned positional embeddings added to the input before the + attention layer, so no RoPE or other in-kernel positional encoding is + needed. + """ + + def __init__( + self, + model_config: ModelConfig[BartConfig], + num_heads: int, + layer_idx: Optional[int] = None, + ): + config = model_config.pretrained_config + super().__init__( + hidden_size=config.d_model, + num_attention_heads=num_heads, + num_key_value_heads=num_heads, + max_position_embeddings=config.max_position_embeddings, + bias=True, + pos_embd_params=None, + layer_idx=layer_idx, + dtype=config.torch_dtype, + config=model_config, + ) + + def apply_rope(self, q, k, v, position_ids): + """BART uses learned pos embeddings, not RoPE — pass through.""" + return q, k, v + + +class BartCrossAttention(CrossAttention): + """BART-style cross-attention with bias.""" + + def __init__( + self, + model_config: ModelConfig[BartConfig], + layer_idx: Optional[int] = None, + ): + config = model_config.pretrained_config + num_heads = _bart_decoder_num_heads(config) + super().__init__( + hidden_size=config.d_model, + num_attention_heads=num_heads, + num_key_value_heads=num_heads, + encoder_hidden_size=config.d_model, + bias=True, + layer_idx=layer_idx, + dtype=config.torch_dtype, + config=model_config, + ) + + +# --------------------------------------------------------------------------- +# Encoder layer +# --------------------------------------------------------------------------- + + +class BartEncoderLayer(nn.Module): + """BART encoder layer: self-attention → add+LN → MLP → add+LN (post-norm).""" + + def __init__( + self, + model_config: ModelConfig[BartConfig], + layer_idx: int, + ): + super().__init__() + config = model_config.pretrained_config + hidden_size = config.d_model + ffn_dim = _bart_encoder_ffn_dim(config) + num_heads = _bart_encoder_num_heads(config) + + self.self_attn = BartSelfAttention(model_config, num_heads=num_heads, layer_idx=layer_idx) + + self.self_attn_layer_norm = LayerNorm( + hidden_size=hidden_size, + eps=1e-5, + dtype=config.torch_dtype, + has_bias=True, + ) + + self.mlp = MLP( + hidden_size=hidden_size, + intermediate_size=ffn_dim, + bias=True, + activation=F.gelu, + dtype=config.torch_dtype, + config=model_config, + layer_idx=layer_idx, + ) + + self.final_layer_norm = LayerNorm( + hidden_size=hidden_size, + eps=1e-5, + dtype=config.torch_dtype, + has_bias=True, + ) + + def forward( + self, + hidden_states: torch.Tensor, + attn_metadata: AttentionMetadata, + position_ids: Optional[torch.IntTensor] = None, + **kwargs, + ) -> torch.Tensor: + residual = hidden_states + hidden_states = self.self_attn( + position_ids=position_ids, + hidden_states=hidden_states, + attn_metadata=attn_metadata, + attention_mask=PredefinedAttentionMask.FULL, + ) + hidden_states = residual + hidden_states + hidden_states = self.self_attn_layer_norm(hidden_states) + + residual = hidden_states + hidden_states = self.mlp(hidden_states) + hidden_states = residual + hidden_states + hidden_states = self.final_layer_norm(hidden_states) + + return hidden_states + + +# --------------------------------------------------------------------------- +# Decoder layer +# --------------------------------------------------------------------------- + + +class BartDecoderLayer(nn.Module): + """BART decoder layer: self-attn → add+LN → cross-attn → add+LN → MLP → add+LN.""" + + def __init__( + self, + model_config: ModelConfig[BartConfig], + layer_idx: int, + ): + super().__init__() + config = model_config.pretrained_config + hidden_size = config.d_model + ffn_dim = _bart_decoder_ffn_dim(config) + num_heads = _bart_decoder_num_heads(config) + + self.self_attn = BartSelfAttention(model_config, num_heads=num_heads, layer_idx=layer_idx) + + self.self_attn_layer_norm = LayerNorm( + hidden_size=hidden_size, + eps=1e-5, + dtype=config.torch_dtype, + has_bias=True, + ) + + self.cross_attn = BartCrossAttention(model_config, layer_idx=layer_idx) + + self.cross_attn_layer_norm = LayerNorm( + hidden_size=hidden_size, + eps=1e-5, + dtype=config.torch_dtype, + has_bias=True, + ) + + self.mlp = MLP( + hidden_size=hidden_size, + intermediate_size=ffn_dim, + bias=True, + activation=F.gelu, + dtype=config.torch_dtype, + config=model_config, + layer_idx=layer_idx, + ) + + self.final_layer_norm = LayerNorm( + hidden_size=hidden_size, + eps=1e-5, + dtype=config.torch_dtype, + has_bias=True, + ) + + def forward( + self, + position_ids: torch.IntTensor, + hidden_states: torch.Tensor, + attn_metadata: AttentionMetadata, + encoder_hidden_states: Optional[torch.Tensor] = None, + cross_attn_metadata: Optional[AttentionMetadata] = None, + skip_cross_kv_projection: bool = False, + **kwargs, + ) -> torch.Tensor: + # Self-attention (post-norm) + residual = hidden_states + hidden_states = self.self_attn( + position_ids=position_ids, + hidden_states=hidden_states, + attn_metadata=attn_metadata, + attention_mask=PredefinedAttentionMask.CAUSAL, + ) + hidden_states = residual + hidden_states + hidden_states = self.self_attn_layer_norm(hidden_states) + + # Cross-attention (post-norm) + residual = hidden_states + hidden_states = self.cross_attn( + hidden_states=hidden_states, + encoder_hidden_states=encoder_hidden_states, + attn_metadata=attn_metadata, + cross_attn_metadata=cross_attn_metadata, + skip_cross_kv_projection=skip_cross_kv_projection, + ) + hidden_states = residual + hidden_states + hidden_states = self.cross_attn_layer_norm(hidden_states) + + # MLP (post-norm) + residual = hidden_states + hidden_states = self.mlp(hidden_states) + hidden_states = residual + hidden_states + hidden_states = self.final_layer_norm(hidden_states) + + return hidden_states + + +# --------------------------------------------------------------------------- +# Encoder / Decoder stacks +# --------------------------------------------------------------------------- + + +class BartEncoder(nn.Module): + """BART encoder: positional embedding + encoder layers.""" + + def __init__(self, model_config: ModelConfig[BartConfig]): + super().__init__() + config = model_config.pretrained_config + num_layers = _bart_encoder_num_layers(config) + + # HF BART uses offset=2 for the padding token, so the actual embedding + # table has max_position_embeddings + 2 entries. + self.embed_positions = Embedding( + config.max_position_embeddings + 2, + config.d_model, + dtype=config.torch_dtype, + ) + self.layernorm_embedding = LayerNorm( + hidden_size=config.d_model, + eps=1e-5, + dtype=config.torch_dtype, + has_bias=True, + ) + self.layers = nn.ModuleList( + [BartEncoderLayer(model_config, layer_idx=i) for i in range(num_layers)] + ) + + def forward( + self, + hidden_states: torch.Tensor, + attn_metadata: AttentionMetadata, + position_ids: Optional[torch.IntTensor] = None, + ) -> torch.Tensor: + position_ids = _packed_position_ids(position_ids, hidden_states) + if position_ids is not None: + hidden_states = hidden_states + self.embed_positions(position_ids) + hidden_states = self.layernorm_embedding(hidden_states) + + for layer in self.layers: + hidden_states = layer( + hidden_states=hidden_states, + attn_metadata=attn_metadata, + position_ids=position_ids, + ) + return hidden_states + + +class BartDecoder(nn.Module): + """BART decoder: positional embedding + decoder layers.""" + + def __init__(self, model_config: ModelConfig[BartConfig]): + super().__init__() + config = model_config.pretrained_config + num_layers = _bart_decoder_num_layers(config) + + self.embed_positions = Embedding( + config.max_position_embeddings + 2, + config.d_model, + dtype=config.torch_dtype, + ) + self.layernorm_embedding = LayerNorm( + hidden_size=config.d_model, + eps=1e-5, + dtype=config.torch_dtype, + has_bias=True, + ) + self.layers = nn.ModuleList( + [BartDecoderLayer(model_config, layer_idx=i) for i in range(num_layers)] + ) + + def forward( + self, + hidden_states: torch.Tensor, + attn_metadata: AttentionMetadata, + position_ids: Optional[torch.IntTensor] = None, + encoder_hidden_states: Optional[torch.Tensor] = None, + cross_attn_metadata: Optional[AttentionMetadata] = None, + skip_cross_kv_projection: bool = False, + ) -> torch.Tensor: + position_ids = _packed_position_ids(position_ids, hidden_states) + if position_ids is not None: + hidden_states = hidden_states + self.embed_positions(position_ids) + hidden_states = self.layernorm_embedding(hidden_states) + + for layer in self.layers: + hidden_states = layer( + position_ids=position_ids, + hidden_states=hidden_states, + attn_metadata=attn_metadata, + encoder_hidden_states=encoder_hidden_states, + cross_attn_metadata=cross_attn_metadata, + skip_cross_kv_projection=skip_cross_kv_projection, + ) + return hidden_states + + +# --------------------------------------------------------------------------- +# Top-level model +# --------------------------------------------------------------------------- + + +class BartModel(nn.Module): + """BART encoder-decoder body (no lm_head).""" + + def __init__(self, model_config: ModelConfig[BartConfig]): + super().__init__() + self.model_config = model_config + config = model_config.pretrained_config + + self.shared_embedding = Embedding( + config.vocab_size, + config.d_model, + dtype=config.torch_dtype, + mapping=model_config.mapping, + tensor_parallel_mode=TensorParallelMode.COLUMN, + gather_output=True, + ) + self.embed_scale = ( + math.sqrt(config.d_model) if getattr(config, "scale_embedding", False) else 1.0 + ) + # HF BART learned position embeddings reserve indices 0 and 1. + self.position_id_offset = 2 + + self.encoder = BartEncoder(model_config) + self.decoder = BartDecoder(model_config) + + def forward( + self, + attn_metadata: AttentionMetadata, + input_ids: Optional[torch.IntTensor] = None, + encoder_input_ids: Optional[torch.IntTensor] = None, + encoder_hidden_states: Optional[torch.Tensor] = None, + position_ids: Optional[torch.IntTensor] = None, + encoder_position_ids: Optional[torch.IntTensor] = None, + encoder_attn_metadata: Optional[AttentionMetadata] = None, + cross_attn_metadata: Optional[AttentionMetadata] = None, + skip_cross_kv_projection: bool = False, + inputs_embeds: Optional[torch.Tensor] = None, + **kwargs, + ) -> torch.Tensor: + if encoder_hidden_states is None and encoder_input_ids is not None: + assert encoder_attn_metadata is not None + encoder_embeds = self.shared_embedding(encoder_input_ids) * self.embed_scale + encoder_hidden_states = self.encoder( + hidden_states=encoder_embeds, + attn_metadata=encoder_attn_metadata, + position_ids=encoder_position_ids, + ) + + if inputs_embeds is None: + assert input_ids is not None + inputs_embeds = self.shared_embedding(input_ids) * self.embed_scale + + decoder_output = self.decoder( + hidden_states=inputs_embeds, + attn_metadata=attn_metadata, + position_ids=position_ids, + encoder_hidden_states=encoder_hidden_states, + cross_attn_metadata=cross_attn_metadata, + skip_cross_kv_projection=skip_cross_kv_projection, + ) + return decoder_output + + +@register_auto_model("BartForConditionalGeneration") +class BartForConditionalGeneration(nn.Module, metaclass=PostInitCaller): + """BART encoder-decoder model with LM head.""" + + def __init__(self, model_config: ModelConfig[BartConfig]): + super().__init__() + self.model_config = model_config + config = model_config.pretrained_config + + self.model = BartModel(model_config) + + self.lm_head = LMHead( + config.vocab_size, + config.d_model, + dtype=config.torch_dtype, + mapping=model_config.mapping, + tensor_parallel_mode=TensorParallelMode.COLUMN, + gather_output=True, + reduce_output=False, + ) + + if getattr(config, "tie_word_embeddings", False): + self.lm_head.weight = self.model.shared_embedding.weight + + self.logits_processor = LogitsProcessor() + + def __post_init__(self): + for _, module in self.named_modules(): + if callable(getattr(module, "create_weights", None)): + module.create_weights() + + def __pp_init__(self): + pass + + @property + def config(self): + return self.model_config.pretrained_config + + def forward( + self, + attn_metadata: AttentionMetadata, + input_ids: Optional[torch.IntTensor] = None, + position_ids: Optional[torch.IntTensor] = None, + encoder_input_ids: Optional[torch.IntTensor] = None, + encoder_hidden_states: Optional[torch.Tensor] = None, + encoder_position_ids: Optional[torch.IntTensor] = None, + encoder_attn_metadata: Optional[AttentionMetadata] = None, + cross_attn_metadata: Optional[AttentionMetadata] = None, + skip_cross_kv_projection: bool = False, + inputs_embeds: Optional[torch.Tensor] = None, + return_context_logits: bool = False, + **kwargs, + ) -> torch.Tensor: + hidden_states = self.model( + attn_metadata=attn_metadata, + input_ids=input_ids, + encoder_input_ids=encoder_input_ids, + encoder_hidden_states=encoder_hidden_states, + position_ids=position_ids, + encoder_position_ids=encoder_position_ids, + encoder_attn_metadata=encoder_attn_metadata, + cross_attn_metadata=cross_attn_metadata, + skip_cross_kv_projection=skip_cross_kv_projection, + inputs_embeds=inputs_embeds, + ) + + return self.logits_processor.forward( + hidden_states, + self.lm_head, + attn_metadata, + return_context_logits, + ) + + def infer_max_seq_len(self) -> int: + config = self.model_config.pretrained_config + return getattr(config, "max_position_embeddings", 1024) + + def load_weights(self, weights: Dict, **kwargs): + config = self.model_config.pretrained_config + tllm_weights = _convert_hf_bart_weights( + weights, config, dtype=self.model_config.torch_dtype + ) + + for name, module in self.named_modules(): + if len(list(module.parameters(recurse=False))) == 0: + continue + if name not in tllm_weights: + continue + w = tllm_weights[name] + if hasattr(module, "load_weights"): + module.load_weights(weights=w) + else: + for n, p in module.named_parameters(recurse=False): + if n in w[0]: + p.data.copy_(w[0][n][:]) + + +@register_auto_model("MBartForConditionalGeneration") +class MBartForConditionalGeneration(BartForConditionalGeneration): + """mBART reuses the BART architecture with the same weight schema.""" + + pass + + +def _convert_hf_bart_weights( + hf_weights: Dict[str, torch.Tensor], + config: BartConfig, + dtype: Optional[torch.dtype] = None, +) -> Dict: + """Map HuggingFace BART/mBART state_dict keys to TRT-LLM module-tree keys. + + Args: + hf_weights: HuggingFace model ``state_dict``. + config: HuggingFace ``BartConfig``. + dtype: Target precision. When specified, every weight tensor is cast + to this dtype before being returned — mirroring the legacy TRT path's + ``convert_weight_to_dtype(params, config.dtype)`` logic. + + HF BART weight layout (prefix ``model.``): + model.shared.weight + model.encoder.embed_positions.weight + model.encoder.layernorm_embedding.{weight,bias} + model.encoder.layers.{i}.self_attn.{q_proj,k_proj,v_proj,out_proj}.{weight,bias} + model.encoder.layers.{i}.self_attn_layer_norm.{weight,bias} + model.encoder.layers.{i}.fc1.{weight,bias} + model.encoder.layers.{i}.fc2.{weight,bias} + model.encoder.layers.{i}.final_layer_norm.{weight,bias} + model.decoder.embed_positions.weight + model.decoder.layernorm_embedding.{weight,bias} + model.decoder.layers.{i}.self_attn.{q_proj,k_proj,v_proj,out_proj}.{weight,bias} + model.decoder.layers.{i}.self_attn_layer_norm.{weight,bias} + model.decoder.layers.{i}.encoder_attn.{q_proj,k_proj,v_proj,out_proj}.{weight,bias} + model.decoder.layers.{i}.encoder_attn_layer_norm.{weight,bias} + model.decoder.layers.{i}.fc1.{weight,bias}, fc2.{weight,bias} + model.decoder.layers.{i}.final_layer_norm.{weight,bias} + lm_head.weight + """ + if dtype is not None: + hf_weights = {k: v.to(dtype) for k, v in hf_weights.items()} + + out: Dict[str, list] = {} + enc_layers = config.encoder_layers + dec_layers = config.decoder_layers + + # HF BartForConditionalGeneration uses "model." prefix; + # HF BartModel does not. Detect and normalise. + has_prefix = any(k.startswith("model.") for k in hf_weights) + p = "model." if has_prefix else "" + + def _get(key: str) -> torch.Tensor: + if key in hf_weights: + return hf_weights[key] + raise KeyError(f"Missing expected HF weight: {key}") + + def _maybe(key: str): + return hf_weights.get(key, None) + + def _wb(prefix: str) -> dict: + d = {"weight": _get(f"{prefix}.weight")} + b = _maybe(f"{prefix}.bias") + if b is not None: + d["bias"] = b + return d + + # Shared embedding + out["model.shared_embedding"] = [{"weight": _get(f"{p}shared.weight")}] + + # LM head + if "lm_head.weight" in hf_weights: + out["lm_head"] = [{"weight": _get("lm_head.weight")}] + + # Encoder positional embedding + out["model.encoder.embed_positions"] = [{"weight": _get(f"{p}encoder.embed_positions.weight")}] + out["model.encoder.layernorm_embedding"] = [_wb(f"{p}encoder.layernorm_embedding")] + + # Encoder layers + for i in range(enc_layers): + hpfx = f"{p}encoder.layers.{i}" + tgt = f"model.encoder.layers.{i}" + + # Self-attention (fused QKV) + out[f"{tgt}.self_attn.qkv_proj"] = [ + _wb(f"{hpfx}.self_attn.q_proj"), + _wb(f"{hpfx}.self_attn.k_proj"), + _wb(f"{hpfx}.self_attn.v_proj"), + ] + out[f"{tgt}.self_attn.o_proj"] = [_wb(f"{hpfx}.self_attn.out_proj")] + + out[f"{tgt}.self_attn_layer_norm"] = [_wb(f"{hpfx}.self_attn_layer_norm")] + + # MLP: BART uses fc1 (up_proj) and fc2 (down_proj) + out[f"{tgt}.mlp.up_proj"] = [_wb(f"{hpfx}.fc1")] + out[f"{tgt}.mlp.down_proj"] = [_wb(f"{hpfx}.fc2")] + + out[f"{tgt}.final_layer_norm"] = [_wb(f"{hpfx}.final_layer_norm")] + + # Decoder positional embedding + out["model.decoder.embed_positions"] = [{"weight": _get(f"{p}decoder.embed_positions.weight")}] + out["model.decoder.layernorm_embedding"] = [_wb(f"{p}decoder.layernorm_embedding")] + + # Decoder layers + for i in range(dec_layers): + hpfx = f"{p}decoder.layers.{i}" + tgt = f"model.decoder.layers.{i}" + + # Self-attention (fused QKV) + out[f"{tgt}.self_attn.qkv_proj"] = [ + _wb(f"{hpfx}.self_attn.q_proj"), + _wb(f"{hpfx}.self_attn.k_proj"), + _wb(f"{hpfx}.self_attn.v_proj"), + ] + out[f"{tgt}.self_attn.o_proj"] = [_wb(f"{hpfx}.self_attn.out_proj")] + + out[f"{tgt}.self_attn_layer_norm"] = [_wb(f"{hpfx}.self_attn_layer_norm")] + + # Cross-attention (separate projections) + out[f"{tgt}.cross_attn.q_proj"] = [_wb(f"{hpfx}.encoder_attn.q_proj")] + out[f"{tgt}.cross_attn.k_proj"] = [_wb(f"{hpfx}.encoder_attn.k_proj")] + out[f"{tgt}.cross_attn.v_proj"] = [_wb(f"{hpfx}.encoder_attn.v_proj")] + out[f"{tgt}.cross_attn.o_proj"] = [_wb(f"{hpfx}.encoder_attn.out_proj")] + + out[f"{tgt}.cross_attn_layer_norm"] = [_wb(f"{hpfx}.encoder_attn_layer_norm")] + + # MLP + out[f"{tgt}.mlp.up_proj"] = [_wb(f"{hpfx}.fc1")] + out[f"{tgt}.mlp.down_proj"] = [_wb(f"{hpfx}.fc2")] + + out[f"{tgt}.final_layer_norm"] = [_wb(f"{hpfx}.final_layer_norm")] + + return out diff --git a/tensorrt_llm/_torch/models/modeling_t5.py b/tensorrt_llm/_torch/models/modeling_t5.py new file mode 100644 index 000000000000..5dc5a5c0e389 --- /dev/null +++ b/tensorrt_llm/_torch/models/modeling_t5.py @@ -0,0 +1,1096 @@ +# 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"); +# 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. +"""PyTorch-flow T5 encoder-decoder model for TensorRT-LLM. + +Supports T5 (``T5ForConditionalGeneration``) and Flan-T5 (gated MLP variant). +mBART and BART share a separate ``modeling_bart.py`` file. + +Architecture: + Encoder: stack of self-attention (non-causal) layers with RMSNorm. + Decoder: stack of self-attention (causal) + cross-attention + MLP layers. + Top-level: encoder + decoder + lm_head. + +HF config normalization: + T5Config stores dims as ``d_model``, ``d_kv``, ``d_ff``, ``num_heads``, + ``num_layers``, ``num_decoder_layers``. The ``hidden_size`` / + ``num_hidden_layers`` / ``num_attention_heads`` aliases are available via + HF property accessors but ``num_key_value_heads`` and + ``intermediate_size`` are not — helper functions below extract them. +""" + +import math +from typing import Dict, Optional + +import torch +import torch.nn.functional as F +from torch import nn +from transformers import T5Config + +from tensorrt_llm.functional import PositionEmbeddingType + +from ..attention_backend import AttentionMetadata +from ..attention_backend.interface import PositionalEmbeddingParams, PredefinedAttentionMask +from ..model_config import ModelConfig +from ..modules.attention import Attention +from ..modules.cross_attention import CrossAttention +from ..modules.embedding import Embedding, LMHead +from ..modules.gated_mlp import GatedMLP +from ..modules.linear import TensorParallelMode +from ..modules.logits_processor import LogitsProcessor +from ..modules.mlp import MLP +from ..modules.rms_norm import RMSNorm +from .modeling_utils import PostInitCaller, register_auto_model + +# --------------------------------------------------------------------------- +# Config helpers +# --------------------------------------------------------------------------- + + +def _t5_num_kv_heads(config: T5Config) -> int: + """T5 uses MHA — KV heads == Q heads.""" + return config.num_heads + + +def _t5_intermediate_size(config: T5Config) -> int: + return config.d_ff + + +def _t5_is_gated_act(config: T5Config) -> bool: + return getattr(config, "is_gated_act", False) + + +def _t5_head_dim(config: T5Config) -> int: + return config.d_kv + + +def _t5_q_scaling(config: T5Config) -> float: + # TRT-LLM attention backends use 1 / (sqrt(head_dim) * q_scaling). + # T5 matches Hugging Face by leaving QK scores unscaled. + return 1.0 / math.sqrt(_t5_head_dim(config)) + + +def _t5_dense_act_fn(config: T5Config): + """Resolve the T5 MLP activation function from the HF config. + + Standard T5 uses ``relu``; Flan-T5 (``gated-gelu``) uses ``gelu_new``. + """ + + def _gelu_new(x: torch.Tensor) -> torch.Tensor: + return ( + 0.5 + * x + * (1.0 + torch.tanh(math.sqrt(2.0 / math.pi) * (x + 0.044715 * torch.pow(x, 3.0)))) + ) + + act_name = getattr(config, "dense_act_fn", None) or "relu" + _ACT_FN_MAP = { + "relu": F.relu, + "gelu": F.gelu, + "gelu_new": _gelu_new, + "silu": F.silu, + "swish": F.silu, + } + if act_name not in _ACT_FN_MAP: + raise ValueError( + f"Unsupported T5 dense_act_fn '{act_name}'. Supported: {list(_ACT_FN_MAP.keys())}" + ) + return _ACT_FN_MAP[act_name] + + +def _t5_gated_act_fn(config: T5Config): + act_fn = _t5_dense_act_fn(config) + + def gated_act_fn(hidden_states: torch.Tensor) -> torch.Tensor: + gate, up = hidden_states.chunk(2, dim=-1) + return act_fn(gate) * up + + return gated_act_fn + + +def _clamp_fp16_infs(hidden_states: torch.Tensor) -> torch.Tensor: + """Match Hugging Face T5's fp16 overflow guard after residual adds.""" + if hidden_states.dtype != torch.float16: + return hidden_states + + clamp_value = torch.where( + torch.isinf(hidden_states).any(), + torch.finfo(hidden_states.dtype).max - 1000, + torch.finfo(hidden_states.dtype).max, + ) + return torch.clamp(hidden_states, min=-clamp_value, max=clamp_value) + + +def _t5_encoder_num_layers(config: T5Config) -> int: + return config.num_layers + + +def _t5_decoder_num_layers(config: T5Config) -> int: + return getattr(config, "num_decoder_layers", None) or config.num_layers + + +# --------------------------------------------------------------------------- +# T5 Relative Position Bias +# --------------------------------------------------------------------------- + + +class T5RelativePositionBias(nn.Module): + """Learned relative position bias for T5 attention. + + Only instantiated on the first layer of each stack (encoder / decoder). + The computed bias is shared across all layers in the same stack. + """ + + def __init__( + self, + num_buckets: int, + num_heads: int, + max_distance: int, + is_decoder: bool, + dtype: Optional[torch.dtype] = None, + ): + super().__init__() + self.num_buckets = num_buckets + self.max_distance = max_distance + self.is_decoder = is_decoder + self.relative_attention_bias = nn.Embedding(num_buckets, num_heads, dtype=dtype) + + @staticmethod + def _relative_position_bucket( + relative_position: torch.Tensor, + bidirectional: bool = True, + num_buckets: int = 32, + max_distance: int = 128, + ) -> torch.Tensor: + relative_buckets = 0 + if bidirectional: + num_buckets //= 2 + relative_buckets += (relative_position > 0).to(torch.long) * num_buckets + relative_position = torch.abs(relative_position) + else: + relative_position = -torch.min(relative_position, torch.zeros_like(relative_position)) + + max_exact = num_buckets // 2 + is_small = relative_position < max_exact + + relative_position_if_large = max_exact + ( + torch.log(relative_position.float() / max_exact) + / math.log(max_distance / max_exact) + * (num_buckets - max_exact) + ).to(torch.long) + relative_position_if_large = torch.min( + relative_position_if_large, + torch.full_like(relative_position_if_large, num_buckets - 1), + ) + + relative_buckets += torch.where(is_small, relative_position, relative_position_if_large) + return relative_buckets + + def forward(self, query_length: int, key_length: int, device: torch.device) -> torch.Tensor: + """Return position bias of shape ``(1, num_heads, query_length, key_length)``.""" + context_position = torch.arange(query_length, dtype=torch.long, device=device)[:, None] + memory_position = torch.arange(key_length, dtype=torch.long, device=device)[None, :] + relative_position = memory_position - context_position + + bucket_ids = self._relative_position_bucket( + relative_position, + bidirectional=not self.is_decoder, + num_buckets=self.num_buckets, + max_distance=self.max_distance, + ) + values = self.relative_attention_bias(bucket_ids) + # (query_length, key_length, num_heads) → (1, num_heads, q, k) + return values.permute(2, 0, 1).unsqueeze(0) + + +# --------------------------------------------------------------------------- +# T5 Attention (self-attention with relative position bias support) +# --------------------------------------------------------------------------- + + +class T5Attention(Attention): + """T5-style multi-head self-attention. + + When ``position_bias`` is provided (from a ``T5RelativePositionBias`` + module living on layer 0), it is added to the QK^T scores before + softmax. Without a KV cache the module computes SDPA directly + (bypassing the VANILLA backend's ``flash_attn_varlen_func`` which + cannot accept an additive bias). With a KV cache it passes the learned + relative-attention table to the TRTLLM backend. + """ + + def __init__( + self, + model_config: ModelConfig[T5Config], + layer_idx: Optional[int] = None, + is_decoder: bool = True, + ): + config = model_config.pretrained_config + num_heads = config.num_heads + num_kv_heads = _t5_num_kv_heads(config) + hidden_size = config.d_model + + super().__init__( + hidden_size=hidden_size, + num_attention_heads=num_heads, + num_key_value_heads=num_kv_heads, + max_position_embeddings=512, + bias=False, + pos_embd_params=PositionalEmbeddingParams(type=PositionEmbeddingType.relative), + layer_idx=layer_idx, + dtype=config.torch_dtype, + config=model_config, + q_scaling=_t5_q_scaling(config), + head_dim=_t5_head_dim(config), + ) + self._is_decoder = is_decoder + self._head_dim = _t5_head_dim(config) + + def apply_rope(self, q, k, v, position_ids): + """T5 has no RoPE — pass through unchanged.""" + return q, k, v + + def _split_qkv( + self, + hidden_states: torch.Tensor, + num_tokens: int, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + qkv = self.qkv_proj(hidden_states) + q_size = self.num_heads * self._head_dim + kv_size = self.num_key_value_heads * self._head_dim + q, k, v = qkv[:num_tokens].split([q_size, kv_size, kv_size], dim=-1) + + q = q.view(-1, self.num_heads, self._head_dim) + k = k.view(-1, self.num_key_value_heads, self._head_dim) + v = v.view(-1, self.num_key_value_heads, self._head_dim) + return q, k, v + + @staticmethod + def _slice_position_bias( + position_bias: torch.Tensor, + query_length: int, + key_length: int, + ) -> torch.Tensor: + query_start = key_length - query_length + return position_bias[:, :, query_start:key_length, :key_length].squeeze(0) + + def _local_position_bias( + self, + position_bias: torch.Tensor, + hidden_states: torch.Tensor, + ) -> torch.Tensor: + if position_bias.shape[1] != self.num_heads: + head_start = self.tp_rank * self.num_heads + head_end = head_start + self.num_heads + if position_bias.shape[1] < head_end: + raise ValueError( + f"T5 position bias has {position_bias.shape[1]} heads, " + f"but rank {self.tp_rank} needs heads [{head_start}, {head_end})." + ) + position_bias = position_bias[:, head_start:head_end] + + return position_bias.to( + device=hidden_states.device, + dtype=hidden_states.dtype, + ).contiguous() + + def _local_relative_attention_bias( + self, + relative_attention_bias: torch.Tensor, + hidden_states: torch.Tensor, + ) -> torch.Tensor: + if relative_attention_bias.shape[0] != self.num_heads: + head_start = self.tp_rank * self.num_heads + head_end = head_start + self.num_heads + if relative_attention_bias.shape[0] < head_end: + raise ValueError( + f"T5 relative attention bias has {relative_attention_bias.shape[0]} heads, " + f"but rank {self.tp_rank} needs heads [{head_start}, {head_end})." + ) + relative_attention_bias = relative_attention_bias[head_start:head_end] + + return relative_attention_bias.to( + device=hidden_states.device, + dtype=hidden_states.dtype, + ).contiguous() + + def forward( + self, + position_ids: Optional[torch.IntTensor] = None, + hidden_states: Optional[torch.Tensor] = None, + attn_metadata: Optional[AttentionMetadata] = None, + attention_mask: Optional[PredefinedAttentionMask] = None, + position_bias: Optional[torch.Tensor] = None, + relative_attention_bias: Optional[torch.Tensor] = None, + relative_attention_max_distance: int = 0, + **kwargs, + ) -> torch.Tensor: + if position_bias is None and relative_attention_bias is None: + return super().forward( + position_ids=position_ids, + hidden_states=hidden_states, + attn_metadata=attn_metadata, + attention_mask=attention_mask, + **kwargs, + ) + + assert attn_metadata is not None + assert hidden_states is not None + if attn_metadata.kv_cache_manager is not None: + if relative_attention_bias is None: + raise ValueError("Cached T5 attention requires a relative attention bias table.") + relative_attention_bias = self._local_relative_attention_bias( + relative_attention_bias, + hidden_states, + ) + return super().forward( + position_ids=position_ids, + hidden_states=hidden_states, + attn_metadata=attn_metadata, + attention_mask=attention_mask, + relative_attention_bias=relative_attention_bias, + relative_attention_max_distance=relative_attention_max_distance, + **kwargs, + ) + + # Manual SDPA with additive position bias (no-KV-cache path). + assert position_bias is not None + position_bias = self._local_position_bias(position_bias, hidden_states) + num_tokens = attn_metadata.num_tokens + q, k, v = self._split_qkv(hidden_states, num_tokens) + + # Per-request SDPA with position bias applied to each request's scores. + seq_lens = attn_metadata.seq_lens + offset = 0 + outputs = [] + for seq_len in seq_lens: + sl = int(seq_len) + q_s = q[offset : offset + sl].transpose(0, 1) # (H, S, D) + k_s = k[offset : offset + sl].transpose(0, 1) + v_s = v[offset : offset + sl].transpose(0, 1) + + scores = torch.matmul(q_s, k_s.transpose(-2, -1)) + # position_bias: (1, H, qlen, klen) — slice to this request's lengths + scores = scores + self._slice_position_bias(position_bias, sl, sl) + + if self._is_decoder: + causal_mask = torch.triu( + torch.full((sl, sl), float("-inf"), device=scores.device, dtype=scores.dtype), + diagonal=1, + ) + scores = scores + causal_mask + + attn_weights = F.softmax(scores.float(), dim=-1).to(q.dtype) + out = torch.matmul(attn_weights, v_s) # (H, S, D) + outputs.append(out.transpose(0, 1)) # (S, H, D) + offset += sl + + attn_output = torch.cat(outputs, dim=0) # (T, H, D) + attn_output = attn_output.reshape(num_tokens, -1) + attn_output = self.o_proj(attn_output) + return attn_output + + +class T5CrossAttention(CrossAttention): + """T5-style cross-attention with the same sizing conventions.""" + + def __init__( + self, + model_config: ModelConfig[T5Config], + layer_idx: Optional[int] = None, + ): + config = model_config.pretrained_config + num_heads = config.num_heads + num_kv_heads = _t5_num_kv_heads(config) + hidden_size = config.d_model + + super().__init__( + hidden_size=hidden_size, + num_attention_heads=num_heads, + num_key_value_heads=num_kv_heads, + encoder_hidden_size=hidden_size, + bias=False, + layer_idx=layer_idx, + dtype=config.torch_dtype, + config=model_config, + q_scaling=_t5_q_scaling(config), + head_dim=_t5_head_dim(config), + ) + + +# --------------------------------------------------------------------------- +# Encoder layer +# --------------------------------------------------------------------------- + + +class T5EncoderLayer(nn.Module): + """T5 encoder layer: pre-norm self-attention + pre-norm MLP.""" + + def __init__( + self, + model_config: ModelConfig[T5Config], + layer_idx: int, + ): + super().__init__() + config = model_config.pretrained_config + hidden_size = config.d_model + intermediate_size = _t5_intermediate_size(config) + is_gated = _t5_is_gated_act(config) + + act_fn = _t5_gated_act_fn(config) if is_gated else _t5_dense_act_fn(config) + + self.self_attn = T5Attention(model_config, layer_idx=layer_idx, is_decoder=False) + + self.input_layernorm = RMSNorm( + hidden_size=hidden_size, + eps=config.layer_norm_epsilon, + dtype=config.torch_dtype, + ) + self.post_attention_layernorm = RMSNorm( + hidden_size=hidden_size, + eps=config.layer_norm_epsilon, + dtype=config.torch_dtype, + ) + + if is_gated: + self.mlp = GatedMLP( + hidden_size=hidden_size, + intermediate_size=intermediate_size, + bias=False, + activation=act_fn, + dtype=config.torch_dtype, + config=model_config, + layer_idx=layer_idx, + ) + else: + self.mlp = MLP( + hidden_size=hidden_size, + intermediate_size=intermediate_size, + bias=False, + activation=act_fn, + dtype=config.torch_dtype, + config=model_config, + layer_idx=layer_idx, + ) + + def forward( + self, + hidden_states: torch.Tensor, + attn_metadata: AttentionMetadata, + position_ids: Optional[torch.IntTensor] = None, + position_bias: Optional[torch.Tensor] = None, + **kwargs, + ) -> torch.Tensor: + residual = hidden_states + hidden_states = self.input_layernorm(hidden_states) + + hidden_states = self.self_attn( + position_ids=position_ids, + hidden_states=hidden_states, + attn_metadata=attn_metadata, + attention_mask=PredefinedAttentionMask.FULL, + position_bias=position_bias, + ) + hidden_states = residual + hidden_states + hidden_states = _clamp_fp16_infs(hidden_states) + + residual = hidden_states + hidden_states = self.post_attention_layernorm(hidden_states) + hidden_states = self.mlp(hidden_states) + hidden_states = residual + hidden_states + hidden_states = _clamp_fp16_infs(hidden_states) + + return hidden_states + + +# --------------------------------------------------------------------------- +# Decoder layer (self-attn + cross-attn + MLP) +# --------------------------------------------------------------------------- + + +class T5DecoderLayer(nn.Module): + """T5 decoder layer: pre-norm self-attention + pre-norm cross-attention + + pre-norm MLP.""" + + def __init__( + self, + model_config: ModelConfig[T5Config], + layer_idx: int, + ): + super().__init__() + config = model_config.pretrained_config + hidden_size = config.d_model + intermediate_size = _t5_intermediate_size(config) + is_gated = _t5_is_gated_act(config) + + act_fn = _t5_gated_act_fn(config) if is_gated else _t5_dense_act_fn(config) + + self.self_attn = T5Attention(model_config, layer_idx=layer_idx, is_decoder=True) + + self.cross_attn = T5CrossAttention(model_config, layer_idx=layer_idx) + + self.input_layernorm = RMSNorm( + hidden_size=hidden_size, + eps=config.layer_norm_epsilon, + dtype=config.torch_dtype, + ) + self.post_attention_layernorm = RMSNorm( + hidden_size=hidden_size, + eps=config.layer_norm_epsilon, + dtype=config.torch_dtype, + ) + self.cross_attn_layernorm = RMSNorm( + hidden_size=hidden_size, + eps=config.layer_norm_epsilon, + dtype=config.torch_dtype, + ) + + if is_gated: + self.mlp = GatedMLP( + hidden_size=hidden_size, + intermediate_size=intermediate_size, + bias=False, + activation=act_fn, + dtype=config.torch_dtype, + config=model_config, + layer_idx=layer_idx, + ) + else: + self.mlp = MLP( + hidden_size=hidden_size, + intermediate_size=intermediate_size, + bias=False, + activation=act_fn, + dtype=config.torch_dtype, + config=model_config, + layer_idx=layer_idx, + ) + + def forward( + self, + position_ids: torch.IntTensor, + hidden_states: torch.Tensor, + attn_metadata: AttentionMetadata, + encoder_hidden_states: Optional[torch.Tensor] = None, + cross_attn_metadata: Optional[AttentionMetadata] = None, + skip_cross_kv_projection: bool = False, + position_bias: Optional[torch.Tensor] = None, + relative_attention_bias: Optional[torch.Tensor] = None, + relative_attention_max_distance: int = 0, + **kwargs, + ) -> torch.Tensor: + # Self-attention (pre-norm) + residual = hidden_states + hidden_states = self.input_layernorm(hidden_states) + hidden_states = self.self_attn( + position_ids=position_ids, + hidden_states=hidden_states, + attn_metadata=attn_metadata, + attention_mask=PredefinedAttentionMask.CAUSAL, + position_bias=position_bias, + relative_attention_bias=relative_attention_bias, + relative_attention_max_distance=relative_attention_max_distance, + ) + hidden_states = residual + hidden_states + hidden_states = _clamp_fp16_infs(hidden_states) + + # Cross-attention (pre-norm) + residual = hidden_states + hidden_states = self.post_attention_layernorm(hidden_states) + hidden_states = self.cross_attn( + hidden_states=hidden_states, + encoder_hidden_states=encoder_hidden_states, + attn_metadata=attn_metadata, + cross_attn_metadata=cross_attn_metadata, + skip_cross_kv_projection=skip_cross_kv_projection, + ) + hidden_states = residual + hidden_states + hidden_states = _clamp_fp16_infs(hidden_states) + + # MLP (pre-norm) + residual = hidden_states + hidden_states = self.cross_attn_layernorm(hidden_states) + hidden_states = self.mlp(hidden_states) + hidden_states = residual + hidden_states + hidden_states = _clamp_fp16_infs(hidden_states) + + return hidden_states + + +# --------------------------------------------------------------------------- +# Encoder stack +# --------------------------------------------------------------------------- + + +class T5Encoder(nn.Module): + """T5 encoder: shared embedding → encoder layers → final RMSNorm.""" + + def __init__(self, model_config: ModelConfig[T5Config]): + super().__init__() + config = model_config.pretrained_config + num_layers = _t5_encoder_num_layers(config) + + self.relative_position_bias = T5RelativePositionBias( + num_buckets=config.relative_attention_num_buckets, + num_heads=config.num_heads, + max_distance=config.relative_attention_max_distance, + is_decoder=False, + dtype=config.torch_dtype, + ) + + self.layers = nn.ModuleList( + [T5EncoderLayer(model_config, layer_idx=i) for i in range(num_layers)] + ) + self.final_layernorm = RMSNorm( + hidden_size=config.d_model, + eps=config.layer_norm_epsilon, + dtype=config.torch_dtype, + ) + + def forward( + self, + hidden_states: torch.Tensor, + attn_metadata: AttentionMetadata, + position_ids: Optional[torch.IntTensor] = None, + ) -> torch.Tensor: + seq_len = hidden_states.shape[0] + position_bias = self.relative_position_bias(seq_len, seq_len, hidden_states.device) + + for layer in self.layers: + hidden_states = layer( + hidden_states=hidden_states, + attn_metadata=attn_metadata, + position_ids=position_ids, + position_bias=position_bias, + ) + hidden_states = self.final_layernorm(hidden_states) + return hidden_states + + +# --------------------------------------------------------------------------- +# Decoder stack +# --------------------------------------------------------------------------- + + +class T5Decoder(nn.Module): + """T5 decoder: decoder layers → final RMSNorm.""" + + def __init__(self, model_config: ModelConfig[T5Config]): + super().__init__() + config = model_config.pretrained_config + num_layers = _t5_decoder_num_layers(config) + + self.relative_position_bias = T5RelativePositionBias( + num_buckets=config.relative_attention_num_buckets, + num_heads=config.num_heads, + max_distance=config.relative_attention_max_distance, + is_decoder=True, + dtype=config.torch_dtype, + ) + + self.layers = nn.ModuleList( + [T5DecoderLayer(model_config, layer_idx=i) for i in range(num_layers)] + ) + self.final_layernorm = RMSNorm( + hidden_size=config.d_model, + eps=config.layer_norm_epsilon, + dtype=config.torch_dtype, + ) + + def forward( + self, + hidden_states: torch.Tensor, + attn_metadata: AttentionMetadata, + position_ids: Optional[torch.IntTensor] = None, + encoder_hidden_states: Optional[torch.Tensor] = None, + cross_attn_metadata: Optional[AttentionMetadata] = None, + skip_cross_kv_projection: bool = False, + ) -> torch.Tensor: + position_bias = None + relative_attention_bias = None + relative_attention_max_distance = 0 + if attn_metadata.kv_cache_manager is None: + seq_len = hidden_states.shape[0] + position_bias = self.relative_position_bias(seq_len, seq_len, hidden_states.device) + else: + relative_attention_bias = ( + self.relative_position_bias.relative_attention_bias.weight.transpose(0, 1) + ) + relative_attention_max_distance = self.relative_position_bias.max_distance + + for layer in self.layers: + hidden_states = layer( + position_ids=position_ids, + hidden_states=hidden_states, + attn_metadata=attn_metadata, + encoder_hidden_states=encoder_hidden_states, + cross_attn_metadata=cross_attn_metadata, + skip_cross_kv_projection=skip_cross_kv_projection, + position_bias=position_bias, + relative_attention_bias=relative_attention_bias, + relative_attention_max_distance=relative_attention_max_distance, + ) + hidden_states = self.final_layernorm(hidden_states) + return hidden_states + + +# --------------------------------------------------------------------------- +# Top-level model +# --------------------------------------------------------------------------- + + +class T5Model(nn.Module): + """Full T5 encoder-decoder model body (no lm_head). + + The shared embedding table is used for both encoder and decoder inputs + (T5 ties encoder/decoder embeddings by default). + """ + + def __init__(self, model_config: ModelConfig[T5Config]): + super().__init__() + self.model_config = model_config + config = model_config.pretrained_config + + self.shared_embedding = Embedding( + config.vocab_size, + config.d_model, + dtype=config.torch_dtype, + mapping=model_config.mapping, + tensor_parallel_mode=TensorParallelMode.COLUMN, + gather_output=True, + ) + + self.encoder = T5Encoder(model_config) + self.decoder = T5Decoder(model_config) + + def forward( + self, + attn_metadata: AttentionMetadata, + input_ids: Optional[torch.IntTensor] = None, + encoder_input_ids: Optional[torch.IntTensor] = None, + encoder_hidden_states: Optional[torch.Tensor] = None, + position_ids: Optional[torch.IntTensor] = None, + encoder_position_ids: Optional[torch.IntTensor] = None, + encoder_attn_metadata: Optional[AttentionMetadata] = None, + cross_attn_metadata: Optional[AttentionMetadata] = None, + skip_cross_kv_projection: bool = False, + inputs_embeds: Optional[torch.Tensor] = None, + **kwargs, + ) -> torch.Tensor: + """Forward the full encoder-decoder model. + + When ``encoder_hidden_states`` is already provided (from a previous + encoder pass cached by the runtime), skip the encoder entirely. + + Args: + attn_metadata: Decoder-side attention metadata. + input_ids: Decoder input token IDs. + encoder_input_ids: Encoder input token IDs. + encoder_hidden_states: Pre-computed encoder output. + position_ids: Decoder position IDs. + encoder_position_ids: Encoder position IDs. + encoder_attn_metadata: Encoder-side attention metadata. + cross_attn_metadata: Metadata for cross-attention layers. + skip_cross_kv_projection: If ``True``, skip K/V projection in + cross-attention (generation steps after the first context step). + inputs_embeds: Pre-computed decoder input embeddings. + """ + if encoder_hidden_states is None and encoder_input_ids is not None: + assert encoder_attn_metadata is not None + encoder_embeds = self.shared_embedding(encoder_input_ids) + encoder_hidden_states = self.encoder( + hidden_states=encoder_embeds, + attn_metadata=encoder_attn_metadata, + position_ids=encoder_position_ids, + ) + + if inputs_embeds is None: + assert input_ids is not None + inputs_embeds = self.shared_embedding(input_ids) + + decoder_output = self.decoder( + hidden_states=inputs_embeds, + attn_metadata=attn_metadata, + position_ids=position_ids, + encoder_hidden_states=encoder_hidden_states, + cross_attn_metadata=cross_attn_metadata, + skip_cross_kv_projection=skip_cross_kv_projection, + ) + return decoder_output + + +@register_auto_model("T5ForConditionalGeneration") +class T5ForConditionalGeneration(nn.Module, metaclass=PostInitCaller): + """T5 encoder-decoder model with LM head for conditional generation. + + Registered for the HF architecture name ``T5ForConditionalGeneration``. + """ + + def __init__(self, model_config: ModelConfig[T5Config]): + super().__init__() + self.model_config = model_config + config = model_config.pretrained_config + + self.model = T5Model(model_config) + + self.lm_head = LMHead( + config.vocab_size, + config.d_model, + dtype=config.torch_dtype, + mapping=model_config.mapping, + tensor_parallel_mode=TensorParallelMode.COLUMN, + gather_output=True, + reduce_output=False, + ) + + # T5 ties lm_head to shared embedding by default + if getattr(config, "tie_word_embeddings", True): + self.lm_head.weight = self.model.shared_embedding.weight + + self.logits_processor = LogitsProcessor() + + # T5 convention: scale logits by 1/sqrt(d_model) + self.rescale_before_lm_head = True + self.d_model = config.d_model + + def __post_init__(self): + for _, module in self.named_modules(): + if callable(getattr(module, "create_weights", None)): + module.create_weights() + + def __pp_init__(self): + pass + + @property + def config(self): + return self.model_config.pretrained_config + + def forward( + self, + attn_metadata: AttentionMetadata, + input_ids: Optional[torch.IntTensor] = None, + position_ids: Optional[torch.IntTensor] = None, + encoder_input_ids: Optional[torch.IntTensor] = None, + encoder_hidden_states: Optional[torch.Tensor] = None, + encoder_position_ids: Optional[torch.IntTensor] = None, + encoder_attn_metadata: Optional[AttentionMetadata] = None, + cross_attn_metadata: Optional[AttentionMetadata] = None, + skip_cross_kv_projection: bool = False, + inputs_embeds: Optional[torch.Tensor] = None, + return_context_logits: bool = False, + **kwargs, + ) -> torch.Tensor: + hidden_states = self.model( + attn_metadata=attn_metadata, + input_ids=input_ids, + encoder_input_ids=encoder_input_ids, + encoder_hidden_states=encoder_hidden_states, + position_ids=position_ids, + encoder_position_ids=encoder_position_ids, + encoder_attn_metadata=encoder_attn_metadata, + cross_attn_metadata=cross_attn_metadata, + skip_cross_kv_projection=skip_cross_kv_projection, + inputs_embeds=inputs_embeds, + ) + + if self.rescale_before_lm_head: + hidden_states = hidden_states * (self.d_model**-0.5) + + return self.logits_processor.forward( + hidden_states, + self.lm_head, + attn_metadata, + return_context_logits, + ) + + def infer_max_seq_len(self) -> int: + return 512 + + def load_weights(self, weights: Dict, **kwargs): + config = self.model_config.pretrained_config + tllm_weights = _convert_hf_t5_weights(weights, config, dtype=self.model_config.torch_dtype) + + if "lm_head.weight" in weights: + self.lm_head.weight = nn.Parameter(torch.empty_like(self.lm_head.weight)) + + for name, module in self.named_modules(): + if len(list(module.parameters(recurse=False))) == 0: + continue + if name not in tllm_weights: + continue + w = tllm_weights[name] + if hasattr(module, "load_weights"): + module.load_weights(weights=w) + else: + for n, p in module.named_parameters(recurse=False): + if n in w[0]: + p.data.copy_(w[0][n][:]) + + +def _convert_hf_t5_weights( + hf_weights: Dict[str, torch.Tensor], + config: T5Config, + dtype: Optional[torch.dtype] = None, +) -> Dict: + """Map HuggingFace T5 state_dict keys to TRT-LLM module-tree keys. + + Returns a dict keyed by TRT-LLM module path, where each value is a list of + weight dicts suitable for ``module.load_weights(weights=...)``. + + Args: + hf_weights: HuggingFace model ``state_dict``. + config: HuggingFace ``T5Config``. + dtype: Target precision. When specified, every weight tensor is cast + to this dtype before being returned — mirroring the legacy TRT path's + ``convert_weight_to_dtype(params, config.dtype)`` logic. + + HF T5 weight layout: + shared.weight + encoder.block.{i}.layer.0.SelfAttention.{q,k,v,o}.weight + encoder.block.{i}.layer.0.layer_norm.weight + encoder.block.{i}.layer.{1}.DenseReluDense.{wi,wo}.weight (non-gated) + encoder.block.{i}.layer.{1}.DenseReluDense.{wi_0,wi_1,wo}.weight (gated) + encoder.block.{i}.layer.{1}.layer_norm.weight + encoder.final_layer_norm.weight + decoder.block.{i}.layer.0.SelfAttention.{q,k,v,o}.weight + decoder.block.{i}.layer.1.EncDecAttention.{q,k,v,o}.weight + decoder.block.{i}.layer.{0,1,2}.layer_norm.weight + decoder.block.{i}.layer.2.DenseReluDense.{wi,wo|wi_0,wi_1,wo}.weight + decoder.final_layer_norm.weight + lm_head.weight + """ + if dtype is not None: + hf_weights = {k: v.to(dtype) for k, v in hf_weights.items()} + + out: Dict[str, list] = {} + is_gated = getattr(config, "is_gated_act", False) + enc_layers = config.num_layers + dec_layers = getattr(config, "num_decoder_layers", None) or config.num_layers + + def _get(key: str) -> torch.Tensor: + if key in hf_weights: + return hf_weights[key] + raise KeyError(f"Missing expected HF weight: {key}") + + # Shared embedding + out["model.shared_embedding"] = [{"weight": _get("shared.weight")}] + + # LM head + if "lm_head.weight" in hf_weights: + out["lm_head"] = [{"weight": _get("lm_head.weight")}] + + # Encoder + for i in range(enc_layers): + pfx = f"encoder.block.{i}" + tgt = f"model.encoder.layers.{i}" + + # Self-attention (fused QKV in TRT-LLM) + out[f"{tgt}.self_attn.qkv_proj"] = [ + {"weight": _get(f"{pfx}.layer.0.SelfAttention.q.weight")}, + {"weight": _get(f"{pfx}.layer.0.SelfAttention.k.weight")}, + {"weight": _get(f"{pfx}.layer.0.SelfAttention.v.weight")}, + ] + out[f"{tgt}.self_attn.o_proj"] = [{"weight": _get(f"{pfx}.layer.0.SelfAttention.o.weight")}] + + # Pre-attention layer norm + out[f"{tgt}.input_layernorm"] = [{"weight": _get(f"{pfx}.layer.0.layer_norm.weight")}] + + # MLP (layer.1 for encoder) + if is_gated: + out[f"{tgt}.mlp.gate_up_proj"] = [ + {"weight": _get(f"{pfx}.layer.1.DenseReluDense.wi_0.weight")}, + {"weight": _get(f"{pfx}.layer.1.DenseReluDense.wi_1.weight")}, + ] + else: + out[f"{tgt}.mlp.up_proj"] = [ + {"weight": _get(f"{pfx}.layer.1.DenseReluDense.wi.weight")} + ] + out[f"{tgt}.mlp.down_proj"] = [{"weight": _get(f"{pfx}.layer.1.DenseReluDense.wo.weight")}] + + # Post-attention (pre-MLP) layer norm + out[f"{tgt}.post_attention_layernorm"] = [ + {"weight": _get(f"{pfx}.layer.1.layer_norm.weight")} + ] + + # Encoder relative position bias (only layer 0 in HF) + rpb_key = "encoder.block.0.layer.0.SelfAttention.relative_attention_bias.weight" + if rpb_key in hf_weights: + out["model.encoder.relative_position_bias.relative_attention_bias"] = [ + {"weight": _get(rpb_key)} + ] + + # Encoder final layer norm + out["model.encoder.final_layernorm"] = [{"weight": _get("encoder.final_layer_norm.weight")}] + + # Decoder + for i in range(dec_layers): + pfx = f"decoder.block.{i}" + tgt = f"model.decoder.layers.{i}" + + # Self-attention (fused QKV) + out[f"{tgt}.self_attn.qkv_proj"] = [ + {"weight": _get(f"{pfx}.layer.0.SelfAttention.q.weight")}, + {"weight": _get(f"{pfx}.layer.0.SelfAttention.k.weight")}, + {"weight": _get(f"{pfx}.layer.0.SelfAttention.v.weight")}, + ] + out[f"{tgt}.self_attn.o_proj"] = [{"weight": _get(f"{pfx}.layer.0.SelfAttention.o.weight")}] + + # Self-attention layer norm + out[f"{tgt}.input_layernorm"] = [{"weight": _get(f"{pfx}.layer.0.layer_norm.weight")}] + + # Cross-attention (separate projections in CrossAttention module) + out[f"{tgt}.cross_attn.q_proj"] = [ + {"weight": _get(f"{pfx}.layer.1.EncDecAttention.q.weight")} + ] + out[f"{tgt}.cross_attn.k_proj"] = [ + {"weight": _get(f"{pfx}.layer.1.EncDecAttention.k.weight")} + ] + out[f"{tgt}.cross_attn.v_proj"] = [ + {"weight": _get(f"{pfx}.layer.1.EncDecAttention.v.weight")} + ] + out[f"{tgt}.cross_attn.o_proj"] = [ + {"weight": _get(f"{pfx}.layer.1.EncDecAttention.o.weight")} + ] + + # Cross-attention layer norm (post_attention_layernorm in T5DecoderLayer) + out[f"{tgt}.post_attention_layernorm"] = [ + {"weight": _get(f"{pfx}.layer.1.layer_norm.weight")} + ] + + # MLP (layer.2 for decoder) + if is_gated: + out[f"{tgt}.mlp.gate_up_proj"] = [ + {"weight": _get(f"{pfx}.layer.2.DenseReluDense.wi_0.weight")}, + {"weight": _get(f"{pfx}.layer.2.DenseReluDense.wi_1.weight")}, + ] + else: + out[f"{tgt}.mlp.up_proj"] = [ + {"weight": _get(f"{pfx}.layer.2.DenseReluDense.wi.weight")} + ] + out[f"{tgt}.mlp.down_proj"] = [{"weight": _get(f"{pfx}.layer.2.DenseReluDense.wo.weight")}] + + # Pre-MLP layer norm + out[f"{tgt}.cross_attn_layernorm"] = [{"weight": _get(f"{pfx}.layer.2.layer_norm.weight")}] + + # Decoder relative position bias (only layer 0 in HF) + rpb_key = "decoder.block.0.layer.0.SelfAttention.relative_attention_bias.weight" + if rpb_key in hf_weights: + out["model.decoder.relative_position_bias.relative_attention_bias"] = [ + {"weight": _get(rpb_key)} + ] + + # Decoder final layer norm + out["model.decoder.final_layernorm"] = [{"weight": _get("decoder.final_layer_norm.weight")}] + + return out diff --git a/tensorrt_llm/_torch/modules/attention.py b/tensorrt_llm/_torch/modules/attention.py index e23cc89a713b..b50273386373 100644 --- a/tensorrt_llm/_torch/modules/attention.py +++ b/tensorrt_llm/_torch/modules/attention.py @@ -102,6 +102,8 @@ def attn_custom_op_inplace( attention_window_size: Optional[int], attention_mask_data: Optional[torch.Tensor], attention_sinks: Optional[torch.Tensor], + relative_attention_bias: Optional[torch.Tensor], + relative_attention_max_distance: int, layer_idx: str, output: torch.Tensor, output_sf: Optional[torch.Tensor], @@ -112,18 +114,22 @@ def attn_custom_op_inplace( ) if attention_mask != CustomAttentionMask.CUSTOM else CustomAttentionMask( attention_mask) # NVFP4 output cannot be supported by torch compile for TRTLLM backend. - attn_layer._attn_impl(q, - k, - v, - metadata, - mask, - mrope_rotary_cos_sin, - mrope_position_deltas, - attention_window_size, - attention_mask_data, - output=output, - output_sf=output_sf, - attention_sinks=attention_sinks) + attn_layer._attn_impl( + q, + k, + v, + metadata, + mask, + mrope_rotary_cos_sin, + mrope_position_deltas, + attention_window_size, + attention_mask_data, + output=output, + output_sf=output_sf, + attention_sinks=attention_sinks, + relative_attention_bias=relative_attention_bias, + relative_attention_max_distance=relative_attention_max_distance, + ) def _helix_post_process( @@ -608,7 +614,10 @@ def __init__( self.num_heads, self.head_dim, self.num_key_value_heads, - pos_embd_params=self.pos_embd_params if self.rope_fusion else None, + pos_embd_params=(self.pos_embd_params if self.rope_fusion or + (self.pos_embd_params is not None + and not self.pos_embd_params.type.is_rope()) else + None), quant_config=self.quant_config, skip_create_weights_in_init=config.skip_create_weights_in_init, q_scaling=self.q_scaling, @@ -698,6 +707,8 @@ def _attn_impl( output: Optional[torch.Tensor] = None, output_sf: Optional[torch.Tensor] = None, attention_sinks: Optional[torch.Tensor] = None, + relative_attention_bias: Optional[torch.Tensor] = None, + relative_attention_max_distance: int = 0, has_lora: bool = False, multi_item_part_lens: Optional[list[list[int]]] = None, ): @@ -736,6 +747,9 @@ def _attn_impl( attention_mask_data=attention_mask_data, softmax_stats_tensor=softmax_stats, attention_sinks=attention_sinks, + relative_attention_bias=relative_attention_bias, + relative_attention_max_distance= + relative_attention_max_distance, multi_item_part_lens=multi_item_part_lens, )) if isinstance(attn_output, tuple): @@ -783,6 +797,8 @@ def _attn_impl( output=output[:num_tokens, :] if output is not None else None, output_sf=output_sf, attention_sinks=attention_sinks, + relative_attention_bias=relative_attention_bias, + relative_attention_max_distance=relative_attention_max_distance, multi_item_part_lens=multi_item_part_lens, )) if isinstance(attn_output, tuple): @@ -803,6 +819,8 @@ def forward_impl( attention_mask_data: Optional[torch.Tensor], mrope_config: Optional[dict], attention_sinks: Optional[torch.Tensor] = None, + relative_attention_bias: Optional[torch.Tensor] = None, + relative_attention_max_distance: int = 0, has_lora: bool = False, multi_item_part_lens: Optional[list[list[int]]] = None, ): @@ -836,6 +854,8 @@ def forward_impl( attention_window_size, attention_mask_data, attention_sinks, + relative_attention_bias, + relative_attention_max_distance, self.layer_idx_str, output, output_sf, @@ -852,6 +872,8 @@ def forward_impl( attention_window_size, attention_mask_data, attention_sinks=attention_sinks, + relative_attention_bias=relative_attention_bias, + relative_attention_max_distance=relative_attention_max_distance, has_lora=has_lora, multi_item_part_lens=multi_item_part_lens, ) @@ -872,6 +894,8 @@ def forward( attention_window_size: Optional[int] = None, attention_mask_data: Optional[torch.Tensor] = None, attention_sinks: Optional[torch.Tensor] = None, + relative_attention_bias: Optional[torch.Tensor] = None, + relative_attention_max_distance: int = 0, multi_item_part_lens: Optional[list[list[int]]] = None, **kwargs, ) -> torch.Tensor: @@ -956,6 +980,8 @@ def forward( assert self.attn_backend == "TRTLLM", ( f"Attention sinks are only supported with attn_backend='TRTLLM'. " f"Current backend: {self.attn_backend}.") + if relative_attention_bias is not None: + assert self.attn_backend == "TRTLLM", "Relative attention bias is only supported for TRTLLM backend." attn_output = self.forward_impl( q, @@ -967,6 +993,8 @@ def forward( attention_mask_data, mrope_config=mrope_config, attention_sinks=attention_sinks, + relative_attention_bias=relative_attention_bias, + relative_attention_max_distance=relative_attention_max_distance, has_lora=bool(lora_params), multi_item_part_lens=multi_item_part_lens, ) diff --git a/tensorrt_llm/_torch/modules/cross_attention.py b/tensorrt_llm/_torch/modules/cross_attention.py new file mode 100644 index 000000000000..ff0f440e245c --- /dev/null +++ b/tensorrt_llm/_torch/modules/cross_attention.py @@ -0,0 +1,249 @@ +# 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"); +# 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. +"""Cross-attention module for encoder-decoder models. + +Unlike self-attention, cross-attention uses Q from the decoder hidden states +and K/V from the encoder output (or from a cached cross-KV pool after the +first decoder context step). +""" + +from typing import Optional + +import torch +from torch import nn + +from ..attention_backend import AttentionMetadata +from ..attention_backend.interface import AttentionBackend, PredefinedAttentionMask +from ..attention_backend.utils import create_attention +from ..distributed import AllReduceParams +from ..model_config import ModelConfig +from .linear import Linear, TensorParallelMode + + +class CrossAttention(nn.Module): + """Cross-attention layer for encoder-decoder models. + + Computes attention where Q comes from decoder hidden states and K/V come + from encoder output. On the first decoder context step, K/V are projected + from encoder_hidden_states and written into the cross-KV cache pool. On + subsequent generation steps, K/V are read from the cache without + re-projection. + + The cross-attention sub-layer honors ``ModelConfig.attn_backend``: when + set to ``"TRTLLM"`` it dispatches through the production C++ attention op + on every supported architecture. Cross-attention currently uses the THOP + attention path because the ``trtllm_gen`` backend API does not yet carry + encoder K/V tensors. + + Encoder and decoder self-attention are unaffected and continue to use + whatever backend ``ModelConfig.attn_backend`` selects. + """ + + def __init__( + self, + *, + hidden_size: int, + num_attention_heads: int, + num_key_value_heads: int, + encoder_hidden_size: Optional[int] = None, + max_position_embeddings: int = 512, + bias: bool = False, + layer_idx: Optional[int] = None, + dtype: Optional[torch.dtype] = None, + dense_bias: Optional[bool] = None, + config: Optional[ModelConfig] = None, + q_scaling: float = 1.0, + head_dim: Optional[int] = None, + ): + super().__init__() + self.layer_idx = layer_idx + config = config or ModelConfig() + self.hidden_size = hidden_size + self.encoder_hidden_size = encoder_hidden_size or hidden_size + self.num_heads = num_attention_heads + if head_dim is not None: + self.head_dim = head_dim + else: + self.head_dim = getattr(config.pretrained_config, "head_dim", None) + if not isinstance(self.head_dim, int): + self.head_dim = self.hidden_size // self.num_heads + self.num_key_value_heads = num_key_value_heads + self.q_scaling = q_scaling + + if dense_bias is None: + dense_bias = bias + + self.mapping = config.mapping + tp_size = self.mapping.tp_size + if self.mapping.enable_attention_dp: + tp_size = 1 + + assert self.num_heads % tp_size == 0 + self.num_heads = self.num_heads // tp_size + self.num_key_value_heads = (self.num_key_value_heads + tp_size - 1) // tp_size + self.q_size = self.num_heads * self.head_dim + self.kv_size = self.num_key_value_heads * self.head_dim + + mapping = config.mapping + + self.q_proj = Linear( + self.hidden_size, + tp_size * self.q_size, + bias=bias, + dtype=dtype, + mapping=mapping, + tensor_parallel_mode=TensorParallelMode.COLUMN, + quant_config=config.get_quant_config(), + skip_create_weights_in_init=config.skip_create_weights_in_init, + ) + + self.k_proj = Linear( + self.encoder_hidden_size, + tp_size * self.kv_size, + bias=bias, + dtype=dtype, + mapping=mapping, + tensor_parallel_mode=TensorParallelMode.COLUMN, + quant_config=config.get_quant_config(), + skip_create_weights_in_init=config.skip_create_weights_in_init, + ) + + self.v_proj = Linear( + self.encoder_hidden_size, + tp_size * self.kv_size, + bias=bias, + dtype=dtype, + mapping=mapping, + tensor_parallel_mode=TensorParallelMode.COLUMN, + quant_config=config.get_quant_config(), + skip_create_weights_in_init=config.skip_create_weights_in_init, + ) + + self.o_proj = Linear( + tp_size * self.q_size, + self.hidden_size, + bias=dense_bias, + dtype=dtype, + mapping=mapping, + tensor_parallel_mode=TensorParallelMode.ROW, + quant_config=config.get_quant_config(), + skip_create_weights_in_init=config.skip_create_weights_in_init, + reduce_output=True, + ) + + # Cross-attention backend selection honors ``ModelConfig.attn_backend`` + # directly, mirroring the behavior of self-attention. + self.attn: AttentionBackend = create_attention( + config.attn_backend, + layer_idx, + self.num_heads, + self.head_dim, + self.num_key_value_heads, + q_scaling=self.q_scaling, + ) + + if not config.skip_create_weights_in_init: + self.create_weights() + + def create_weights(self): + self.q_proj.create_weights() + self.k_proj.create_weights() + self.v_proj.create_weights() + self.o_proj.create_weights() + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: Optional[torch.Tensor], + attn_metadata: AttentionMetadata, + cross_attn_metadata: AttentionMetadata, + skip_cross_kv_projection: bool = False, + all_reduce_params: Optional[AllReduceParams] = None, + **kwargs, + ) -> torch.Tensor: + """Forward pass for cross-attention. + + Args: + hidden_states: Decoder hidden states ``[num_tokens, hidden_size]``. + encoder_hidden_states: Encoder output. Required on the first + decoder context step (when ``skip_cross_kv_projection`` is + ``False``). ``None`` for generation steps. + attn_metadata: Decoder-side attention metadata (Q-side lengths). + cross_attn_metadata: Cross-attention metadata carrying encoder + K/V-side lengths, cross-pool block tables, etc. Must satisfy + ``cross_attn_metadata.is_cross is True`` (i.e. the K/V-side + ``seq_lens_kv`` differs from the Q-side ``seq_lens``). Always + required — build it via + ``attn_metadata.create_cross_metadata(encoder_seq_lens, + cross_kv_cache_manager)``. + skip_cross_kv_projection: When ``True``, K/V are read from the + cross-KV cache without re-projection (decoder generation + steps). When ``False``, K/V are projected from + ``encoder_hidden_states`` and written into the cache (first + decoder context step). + all_reduce_params: AllReduce parameters for TP output projection. + + Returns: + Output tensor ``[num_tokens, hidden_size]``. + """ + if cross_attn_metadata is None: + raise ValueError( + "cross_attn_metadata is required. Build it via " + "attn_metadata.create_cross_metadata(encoder_seq_lens, " + "cross_kv_cache_manager)." + ) + assert cross_attn_metadata.is_cross, ( + "cross_attn_metadata.is_cross must be True. Build it via " + "attn_metadata.create_cross_metadata(encoder_seq_lens, " + "cross_kv_cache_manager) so seq_lens_kv differs from " + "seq_lens." + ) + metadata = cross_attn_metadata + + q = self.q_proj(hidden_states) + + if not skip_cross_kv_projection: + assert encoder_hidden_states is not None, ( + "encoder_hidden_states is required when cross-KV projection " + "is not skipped (first decoder context step)." + ) + k = self.k_proj(encoder_hidden_states) + v = self.v_proj(encoder_hidden_states) + else: + # Generation step: skip projection, K/V are already in the + # cross-KV cache (written during the first decoder context step). + # The backend reads them from ``metadata.kv_cache_manager``. + assert metadata.kv_cache_manager is not None, ( + "skip_cross_kv_projection=True requires a populated " + "cross-KV cache manager on cross_attn_metadata." + ) + k = None + v = None + + num_tokens = attn_metadata.num_tokens + q = q[:num_tokens, :] + + attn_output = self.attn.forward( + q, + k, + v, + metadata, + attention_mask=PredefinedAttentionMask.FULL, + ) + if isinstance(attn_output, tuple): + attn_output = attn_output[0] + + attn_output = self.o_proj(attn_output, all_reduce_params=all_reduce_params) + return attn_output diff --git a/tensorrt_llm/_torch/modules/rms_norm.py b/tensorrt_llm/_torch/modules/rms_norm.py index 4a22bef2196d..8c9a744e3bdd 100644 --- a/tensorrt_llm/_torch/modules/rms_norm.py +++ b/tensorrt_llm/_torch/modules/rms_norm.py @@ -186,7 +186,8 @@ def _ensure_contiguous_with_dtype(t: torch.Tensor, key: str): gather=True, use_gemma=self.use_gemma, ) - elif IS_FLASHINFER_AVAILABLE: + elif IS_FLASHINFER_AVAILABLE and hidden_states.dtype in ( + torch.float16, torch.bfloat16): from ..custom_ops import (flashinfer_fused_add_rmsnorm, flashinfer_gemma_fused_add_rmsnorm, flashinfer_gemma_rmsnorm, diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index f69fc8488766..1fef2c239d37 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -39,7 +39,7 @@ from .guided_decoder import GuidedDecoder from .kv_cache_manager_v2 import KVCacheManagerV2 from .kv_cache_transceiver import AttentionTypeCpp, create_kv_cache_transceiver -from .llm_request import ExecutorResponse +from .llm_request import ExecutorResponse, LlmRequestState from .mamba_cache_manager import (BaseMambaCacheManager, CppMambaHybridCacheManager, MixedMambaHybridCacheManager, @@ -243,17 +243,23 @@ def __init__( self._draft_config = draft_config self._skip_est = skip_est - def _get_model_kv_cache_manager_cls(self, model_engine: PyTorchModelEngine): + def _get_model_kv_cache_manager_cls( + self, + model_engine: PyTorchModelEngine, + kv_cache_config_override: Optional[KvCacheConfig] = None, + ): + kv_cache_config = (kv_cache_config_override if kv_cache_config_override + is not None else self._kv_cache_config) config = model_engine.model.model_config.pretrained_config cls = get_kv_cache_manager_cls( model_engine.model.model_config, - self._kv_cache_config, + kv_cache_config, is_disagg=self._is_disagg, cache_transceiver_config=self._cache_transceiver_config) if cls == KVCacheManagerV2: if self._kv_connector_manager is not None or ( self._max_beam_width is not None and self._max_beam_width - > 1) or self._kv_cache_config.event_buffer_max_size > 0 or ( + > 1) or kv_cache_config.event_buffer_max_size > 0 or ( self._cache_transceiver_config is not None and self._cache_transceiver_config.backend is not None): # Per-layer head_dim models (e.g., Gemma4 hybrid) require V2's @@ -279,8 +285,9 @@ def _get_model_kv_cache_manager_cls(self, model_engine: PyTorchModelEngine): # the routing site so users see the warning where the decision is # actually made. if is_hybrid_linear(model_engine.model.model_config.pretrained_config) \ - and self._kv_cache_config.enable_block_reuse: - uses_v1_mamba_route = os.environ.get('TRTLLM_USE_CPP_MAMBA', '0') == '1' \ + and kv_cache_config.enable_block_reuse: + uses_v1_mamba_route = self._is_disagg \ + or os.environ.get('TRTLLM_USE_CPP_MAMBA', '0') == '1' \ or os.environ.get('TRTLLM_USE_PY_MAMBA', '0') == '1' \ or self._speculative_config is not None if uses_v1_mamba_route: @@ -290,32 +297,44 @@ def _get_model_kv_cache_manager_cls(self, model_engine: PyTorchModelEngine): ) return cls - def _per_manager_cache_cost(self, manager_cls, model_config, + def _per_manager_cache_cost(self, + manager_cls, + model_config, + kv_cache_config: Optional[KvCacheConfig] = None, **extra_kwargs) -> CacheCost: + kv_cache_config = (kv_cache_config if kv_cache_config is not None else + self._kv_cache_config) return CacheCost.from_raw( manager_cls.get_cache_size_per_token( model_config, self._mapping, tokens_per_block=self._tokens_per_block, max_batch_size=self._max_batch_size, - kv_cache_config=self._kv_cache_config, + kv_cache_config=kv_cache_config, **extra_kwargs)) - def _get_kv_size_per_token(self) -> CacheCost: + def _get_kv_size_per_token(self, + kv_cache_config: Optional[KvCacheConfig] = None + ) -> CacheCost: """Aggregate KV cost across target + (optional) draft as a CacheCost. ``max_batch_size`` and ``kv_cache_config`` are passed unconditionally; managers that don't need them ignore via ``**kwargs``. """ + kv_cache_config = (kv_cache_config if kv_cache_config is not None else + self._kv_cache_config) model_config = self._model_engine.model.model_config total = self._per_manager_cache_cost(self._kv_cache_manager_cls, - model_config) + model_config, kv_cache_config) + if self._is_encoder_decoder(): + total += CacheCost.from_raw(self._get_cross_kv_size_per_token()) if self._draft_model_engine is not None: draft_model_config = self._draft_model_engine.model.model_config draft_kv_cache_manager_cls = self._get_model_kv_cache_manager_cls( - self._draft_model_engine) + self._draft_model_engine, kv_cache_config) total += self._per_manager_cache_cost(draft_kv_cache_manager_cls, - draft_model_config) + draft_model_config, + kv_cache_config) elif self._should_create_separate_draft_kv_cache(): # One-model draft with separate KV cache layout. # Pass num_layers explicitly since the HF config may report a @@ -330,15 +349,17 @@ def _get_kv_size_per_token(self) -> CacheCost: # from target (e.g. hybrid target + plain transformer draft). draft_kv_cache_manager_cls = get_kv_cache_manager_cls( effective_draft_config, - self._kv_cache_config, + kv_cache_config, is_disagg=self._is_disagg) total += self._per_manager_cache_cost( - draft_kv_cache_manager_cls, effective_draft_config) + draft_kv_cache_manager_cls, effective_draft_config, + kv_cache_config) elif self._mapping.is_last_pp_rank(): # EAGLE3/MTP: draft layers only on last PP rank total += self._per_manager_cache_cost( self._kv_cache_manager_cls, effective_draft_config, + kv_cache_config, num_layers=self._get_num_draft_layers()) return total @@ -743,13 +764,17 @@ def configure_kv_cache_capacity(self, # ---------------------------handle max_gpu_total_bytes--------------------------------- def _create_kv_cache_manager( - self, - model_engine: PyTorchModelEngine, - estimating_kv_cache: bool = False) -> KVCacheManager: + self, + model_engine: PyTorchModelEngine, + estimating_kv_cache: bool = False, + kv_cache_config_override: Optional[KvCacheConfig] = None + ) -> KVCacheManager: mapping = self._mapping assert model_engine.model.model_config.is_generation, "Only construct KV cache for generation models." + kv_cache_config = (kv_cache_config_override if kv_cache_config_override + is not None else self._kv_cache_config) kv_cache_manager_cls = self._get_model_kv_cache_manager_cls( - model_engine) + model_engine, kv_cache_config) # When using separate draft KV cache in one-model speculative decoding, # use layer_mask to include only target layers. The draft layers should @@ -765,7 +790,7 @@ def _create_kv_cache_manager( model_engine=model_engine, kv_cache_manager_cls=kv_cache_manager_cls, mapping=mapping, - kv_cache_config=self._kv_cache_config, + kv_cache_config=kv_cache_config, tokens_per_block=self._tokens_per_block, max_seq_len=self._max_seq_len, max_batch_size=self._max_batch_size, @@ -876,17 +901,17 @@ def _create_one_model_draft_kv_cache_manager( # otherwise fall back to target model config for MTP). effective_draft_config = self._get_effective_draft_config() + draft_kv_config = (kv_cache_config_override if kv_cache_config_override + is not None else self._kv_cache_config) # Get the appropriate KV cache manager class for the draft model draft_kv_cache_manager_cls = get_kv_cache_manager_cls( - effective_draft_config, - self._kv_cache_config, - is_disagg=self._is_disagg) + effective_draft_config, draft_kv_config, is_disagg=self._is_disagg) # Use V2 if enabled and the base class is KVCacheManager if draft_kv_cache_manager_cls == KVCacheManagerV2: if self._kv_connector_manager is not None or ( self._max_beam_width is not None and self._max_beam_width - > 1) or self._kv_cache_config.event_buffer_max_size > 0 or ( + > 1) or draft_kv_config.event_buffer_max_size > 0 or ( self._cache_transceiver_config is not None and self._cache_transceiver_config.backend is not None): logger.warning( @@ -900,7 +925,6 @@ def _create_one_model_draft_kv_cache_manager( # the sparse_attention_config. Get it from effective_draft_config which # falls back to the target model's config for MTP mode. sparse_attn_config = effective_draft_config.sparse_attention_config - draft_kv_config = kv_cache_config_override if kv_cache_config_override is not None else self._kv_cache_config return _create_kv_cache_manager( model_engine=None, kv_cache_manager_cls=draft_kv_cache_manager_cls, @@ -926,11 +950,16 @@ def _create_one_model_draft_kv_cache_manager( ) def _get_target_and_draft_cache_costs( - self, ) -> Optional[tuple[CacheCost, CacheCost]]: + self, + kv_cache_config: Optional[KvCacheConfig] = None, + ) -> Optional[tuple[CacheCost, CacheCost]]: """Per-manager KV cache costs for target and draft layers.""" - total_kv = self._get_kv_size_per_token() + target_kv_cache_config = (kv_cache_config if kv_cache_config is not None + else self._kv_cache_config) + total_kv = self._get_kv_size_per_token(target_kv_cache_config) target_kv = self._per_manager_cache_cost( - self._kv_cache_manager_cls, self._model_engine.model.model_config) + self._kv_cache_manager_cls, self._model_engine.model.model_config, + target_kv_cache_config) # The draft contribution is whatever the aggregate has on top of the # target. Both pieces are CacheCost; subtraction is component-wise. draft_kv = CacheCost(slope=total_kv.slope - target_kv.slope, @@ -963,21 +992,20 @@ def _compute_draft_budget_shares( def _split_kv_cache_budget_for_draft( self, budget_attr: str, + target_kv_cache_config: Optional[KvCacheConfig] = None, draft_kv_cache_config: Optional[KvCacheConfig] = None, - ) -> Optional[KvCacheConfig]: + ) -> tuple[KvCacheConfig, Optional[KvCacheConfig]]: """Split a byte budget (attribute on ``KvCacheConfig``) between target and draft KV caches. - Splits the value of ``self._kv_cache_config.`` using the - affine target/draft cache costs, updates the target config in-place, - and merges the draft share into ``draft_kv_cache_config`` (cloning the - target config if needed). + Splits the value of ``target_kv_cache_config.`` using the + affine target/draft cache costs, then returns cloned target and draft + configs containing their respective shares. - Returns the (possibly newly created) draft config. The input - ``draft_kv_cache_config`` is returned unchanged when the split is not - applicable (the budget is not set, or the per-manager cache costs are - unavailable) — in those cases sharing ``self._kv_cache_config`` is - correct. + The input target config and the creator's base config are not mutated. + When the split is not applicable (the budget is not set, or the + per-manager cache costs are unavailable), the input configs are returned + unchanged. The affine fixed (intercept) cost models GPU-resident state (e.g. mamba SSM state). It is only charged against ``max_gpu_total_bytes``; for any @@ -993,13 +1021,17 @@ def _split_kv_cache_budget_for_draft( for non-GPU budgets remains so the draft never silently inherits the full budget and double-allocates it. """ - total_budget = getattr(self._kv_cache_config, budget_attr) or 0 + target_kv_cache_config = (target_kv_cache_config + if target_kv_cache_config is not None else + self._kv_cache_config) + total_budget = getattr(target_kv_cache_config, budget_attr) or 0 if total_budget <= 0: - return draft_kv_cache_config + return target_kv_cache_config, draft_kv_cache_config - cache_costs = self._get_target_and_draft_cache_costs() + cache_costs = self._get_target_and_draft_cache_costs( + target_kv_cache_config) if cache_costs is None: - return draft_kv_cache_config + return target_kv_cache_config, draft_kv_cache_config target_kv, draft_kv = cache_costs # The fixed (intercept) cost models GPU-resident state such as mamba SSM @@ -1040,9 +1072,11 @@ def _split_kv_cache_budget_for_draft( f"assigning the draft a zero {budget_attr} budget to avoid " f"double-allocating the full budget.") if draft_kv_cache_config is None: - draft_kv_cache_config = self._kv_cache_config.model_copy() + draft_kv_cache_config = target_kv_cache_config.model_copy() + else: + draft_kv_cache_config = draft_kv_cache_config.model_copy() setattr(draft_kv_cache_config, budget_attr, 0) - return draft_kv_cache_config + return target_kv_cache_config, draft_kv_cache_config target_budget, draft_budget = shares logger.info( @@ -1050,17 +1084,246 @@ def _split_kv_cache_budget_for_draft( f"target={target_budget / GB:.2f} GiB ({target_kv}), " f"draft={draft_budget / GB:.2f} GiB ({draft_kv})") - setattr(self._kv_cache_config, budget_attr, target_budget) + split_target_kv_cache_config = target_kv_cache_config.model_copy() + setattr(split_target_kv_cache_config, budget_attr, target_budget) if draft_kv_cache_config is None: - draft_kv_cache_config = self._kv_cache_config.model_copy() - setattr(draft_kv_cache_config, budget_attr, draft_budget) - return draft_kv_cache_config + split_draft_kv_cache_config = target_kv_cache_config.model_copy() + else: + split_draft_kv_cache_config = draft_kv_cache_config.model_copy() + setattr(split_draft_kv_cache_config, budget_attr, draft_budget) + return split_target_kv_cache_config, split_draft_kv_cache_config + + def _is_encoder_decoder(self) -> bool: + return self._model_engine.model.model_config.is_encoder_decoder + + @staticmethod + def _get_config_int_attr(config, names: tuple[str, ...]) -> Optional[int]: + for name in names: + value = getattr(config, name, None) + if isinstance(value, int): + return value + return None + + def _get_cross_kv_cache_layout( + self, + fallback_max_seq_len: Optional[int] = None + ) -> tuple[int, int, int, int]: + """Return decoder-layer count and encoder KV geometry for cross cache.""" + config = self._model_engine.model.model_config.pretrained_config + + num_layers = self._get_config_int_attr( + config, + ("num_decoder_layers", "decoder_layers", "num_hidden_layers", + "num_layers"), + ) + if num_layers is None: + raise ValueError( + "Unable to determine decoder layer count for cross KV cache.") - def _needs_gpu_kv_cache_budget_split(self) -> bool: + encoder_num_heads = self._get_config_int_attr( + config, + ("encoder_num_heads", "encoder_attention_heads", "num_heads", + "num_attention_heads"), + ) + if encoder_num_heads is None: + raise ValueError( + "Unable to determine encoder attention head count for cross KV cache." + ) + + num_kv_heads = self._get_config_int_attr( + config, + ("encoder_num_kv_heads", "encoder_num_key_value_heads", + "encoder_attention_heads", "encoder_num_heads", + "num_key_value_heads", "num_heads", "num_attention_heads"), + ) + if num_kv_heads is None: + num_kv_heads = encoder_num_heads + + encoder_hidden_size = self._get_config_int_attr( + config, ("encoder_hidden_size", "d_model", "hidden_size")) + if encoder_hidden_size is None: + raise ValueError( + "Unable to determine encoder hidden size for cross KV cache.") + + head_dim = self._get_config_int_attr( + config, + ("encoder_head_size", "encoder_head_dim", "d_kv"), + ) + if head_dim is None: + head_dim = encoder_hidden_size // encoder_num_heads + + max_seq_len = fallback_max_seq_len or self._max_seq_len + max_input_len = getattr(self._llm_args, "max_input_len", None) + if isinstance(max_input_len, int) and max_input_len > 0: + max_seq_len = max_input_len + encoder_limit = self._get_config_int_attr( + config, + ("max_encoder_input_len", "encoder_max_input_length", + "max_encoder_position_embeddings", + "encoder_max_position_embeddings", "max_position_embeddings", + "n_positions"), + ) + if encoder_limit is not None: + max_seq_len = min(max_seq_len, encoder_limit) + + return num_layers, num_kv_heads, head_dim, max_seq_len + + def _get_cross_kv_size_per_token(self) -> int: + """Estimate bytes/token for the encoder-decoder cross-attention pool.""" + from types import SimpleNamespace + + model_config = self._model_engine.model.model_config + config = model_config.pretrained_config + (num_layers, num_kv_heads, head_dim, + _) = self._get_cross_kv_cache_layout() + num_attention_heads = self._get_config_int_attr( + config, + ("encoder_num_heads", "encoder_attention_heads", "num_heads", + "num_attention_heads"), + ) + hidden_size = self._get_config_int_attr( + config, ("encoder_hidden_size", "d_model", "hidden_size")) + proxy_model_config = SimpleNamespace( + pretrained_config=SimpleNamespace( + num_key_value_heads=num_kv_heads, + num_attention_heads=num_attention_heads, + hidden_size=hidden_size, + head_dim=head_dim, + ), + quant_config=model_config.quant_config, + ) + return self._kv_cache_manager_cls.get_cache_size_per_token( + proxy_model_config, + self._mapping, + tokens_per_block=self._tokens_per_block, + num_layers=num_layers, + ) + + def _split_kv_cache_budget_for_cross( + self, + kv_cache_config: Optional[KvCacheConfig] = None, + ) -> tuple[KvCacheConfig, KvCacheConfig]: + """Split enc-dec KV cache budgets between self and cross pools. + + The cross manager must exist for every encoder-decoder runtime. During + both estimation and final construction, split the same memory-derived + budget sources used by the legacy TRT path: the free-memory fraction, + any explicit ``max_gpu_total_bytes`` override, and any explicit host + cache budget. ``max_tokens`` is a logical cap, not a memory split knob, + so it is intentionally left unchanged. The creator's base config is not + mutated. + """ + base_kv_cache_config = (kv_cache_config if kv_cache_config is not None + else self._kv_cache_config) + fraction = base_kv_cache_config.cross_kv_cache_fraction + if fraction is None: + raise ValueError("Encoder-decoder models require " + "cross_kv_cache_fraction to size the cross " + "KV cache pool.") + + self_kv_cache_config = base_kv_cache_config.model_copy() + cross_kv_cache_config = base_kv_cache_config.model_copy() + split_any_budget = False + + free_fraction = base_kv_cache_config.free_gpu_memory_fraction + if free_fraction is not None: + cross_fraction = free_fraction * fraction + self_fraction = free_fraction - cross_fraction + logger.info( + "Splitting encoder-decoder free GPU memory fraction: " + f"total={free_fraction:.3f}, self={self_fraction:.3f}, cross={cross_fraction:.3f}" + ) + self_kv_cache_config.free_gpu_memory_fraction = self_fraction + cross_kv_cache_config.free_gpu_memory_fraction = cross_fraction + split_any_budget = True + + total_budget = base_kv_cache_config.max_gpu_total_bytes + if total_budget is not None and total_budget > 0: + cross_budget = int(total_budget * fraction) + self_budget = total_budget - cross_budget + logger.info( + f"Splitting KV cache budget for encoder-decoder: " + f"total={total_budget / GB:.2f} GiB, " + f"self={self_budget / GB:.2f} GiB ({1 - fraction:.0%}), " + f"cross={cross_budget / GB:.2f} GiB ({fraction:.0%})") + self_kv_cache_config.max_gpu_total_bytes = self_budget + cross_kv_cache_config.max_gpu_total_bytes = cross_budget + split_any_budget = True + + host_cache_size = base_kv_cache_config.host_cache_size + if host_cache_size is not None and host_cache_size > 0: + cross_host_cache_size = int(host_cache_size * fraction) + self_host_cache_size = host_cache_size - cross_host_cache_size + logger.info( + f"Splitting KV cache host budget for encoder-decoder: " + f"total={host_cache_size / GB:.2f} GiB, " + f"self={self_host_cache_size / GB:.2f} GiB ({1 - fraction:.0%}), " + f"cross={cross_host_cache_size / GB:.2f} GiB ({fraction:.0%})") + self_kv_cache_config.host_cache_size = self_host_cache_size + cross_kv_cache_config.host_cache_size = cross_host_cache_size + split_any_budget = True + + if not split_any_budget: + raise ValueError("Unable to size the encoder-decoder cross KV " + "cache pool: neither free_gpu_memory_fraction nor " + "max_gpu_total_bytes nor host_cache_size is " + "available.") + + return self_kv_cache_config, cross_kv_cache_config + + def _create_cross_kv_cache_manager( + self, + cross_kv_cache_config: KvCacheConfig, + estimating_kv_cache: bool = False, + fallback_max_seq_len: Optional[int] = None, + ) -> KVCacheManager: + """Create a KV cache manager for the cross-attention pool. + + The cross pool stores encoder K/V projections that are written once + during the first decoder context step and read on every subsequent + decoder generation step. It uses ``CacheType.CROSS`` with decoder + layer count but encoder-side KV geometry. + + The manager class mirrors the self pool (``KVCacheManager`` for V1, + ``KVCacheManagerV2`` for V2) so that both pools share the same + runtime ABI and scheduler integration. V1 is the default and the + production target for encoder-decoder models. + """ + (num_layers, num_kv_heads, head_dim, + max_seq_len) = self._get_cross_kv_cache_layout(fallback_max_seq_len) + estimating_kv_cache = estimating_kv_cache and not self._skip_est + return _create_kv_cache_manager( + model_engine=self._model_engine, + kv_cache_manager_cls=self._kv_cache_manager_cls, + mapping=self._mapping, + kv_cache_config=cross_kv_cache_config, + tokens_per_block=self._tokens_per_block, + max_seq_len=max_seq_len, + max_batch_size=self._max_batch_size, + spec_config=None, + sparse_attn_config=None, + max_num_tokens=self._max_num_tokens, + max_beam_width=1, + kv_connector_manager=None, + estimating_kv_cache=estimating_kv_cache, + execution_stream=self._execution_stream, + num_layers=num_layers, + num_kv_heads=num_kv_heads, + head_dim=head_dim, + kv_cache_type=tensorrt_llm.bindings.internal.batch_manager. + CacheType.CROSS, + ) + + def _needs_gpu_kv_cache_budget_split( + self, + kv_cache_config: Optional[KvCacheConfig] = None, + ) -> bool: """Whether max_gpu_total_bytes must be split per manager.""" if issubclass(self._kv_cache_manager_cls, KVCacheManagerV2): return self._should_create_separate_draft_kv_cache() - return is_vswa_enabled(self._kv_cache_config) + kv_cache_config = (kv_cache_config if kv_cache_config is not None else + self._kv_cache_config) + return is_vswa_enabled(kv_cache_config) def build_managers(self, resources: Dict, @@ -1068,6 +1331,17 @@ def build_managers(self, """Construct KV caches for model and draft model (if applicable).""" if self._skip_est: self.configure_kv_cache_capacity() + original_max_seq_len = self._max_seq_len + + # For encoder-decoder models, split the self/cross budgets first so + # every enc-dec build creates a real cross pool. This must happen + # before any draft split so that the draft split operates on the + # already-reduced self-pool budget. + self_kv_cache_config = self._kv_cache_config + cross_kv_cache_config = None + if self._is_encoder_decoder(): + self_kv_cache_config, cross_kv_cache_config = self._split_kv_cache_budget_for_cross( + ) # Split combined KV cache budgets before creating managers. Skip during # estimation — estimation uses max_tokens-based logic and must not @@ -1079,26 +1353,36 @@ def build_managers(self, if not estimating_kv_cache and has_draft: # Used when each manager sizes pools from max_gpu_total_bytes (V2 # and V1 VSWA). V1 non-VSWA GPU uses shared max_tokens instead. - if self._needs_gpu_kv_cache_budget_split(): - draft_kv_cache_config = self._split_kv_cache_budget_for_draft( - "max_gpu_total_bytes", draft_kv_cache_config) + if self._needs_gpu_kv_cache_budget_split(self_kv_cache_config): + self_kv_cache_config, draft_kv_cache_config = ( + self._split_kv_cache_budget_for_draft( + "max_gpu_total_bytes", self_kv_cache_config, + draft_kv_cache_config)) # KVCacheManagerV2 does not support two-model draft budget splitting. v2_two_model = (issubclass(self._kv_cache_manager_cls, KVCacheManagerV2) and self._draft_model_engine is not None) if not v2_two_model: # Each manager sizes its host pool from host_cache_size directly. - draft_kv_cache_config = self._split_kv_cache_budget_for_draft( - "host_cache_size", draft_kv_cache_config) + self_kv_cache_config, draft_kv_cache_config = ( + self._split_kv_cache_budget_for_draft( + "host_cache_size", self_kv_cache_config, + draft_kv_cache_config)) kv_cache_manager = self._create_kv_cache_manager( - self._model_engine, estimating_kv_cache) + self._model_engine, + estimating_kv_cache, + kv_cache_config_override=self_kv_cache_config) - if not estimating_kv_cache and self._kv_connector_manager is not None and self._draft_model_engine is not None: + if (not estimating_kv_cache and self._kv_connector_manager is not None + and self._draft_model_engine is not None): raise NotImplementedError( "Connector manager is not supported for draft model.") draft_kv_cache_manager = None + draft_build_kv_cache_config = (draft_kv_cache_config + if draft_kv_cache_config is not None else + self_kv_cache_config) # Two-model speculative decoding: draft model has separate engine if self._draft_model_engine is not None: @@ -1106,31 +1390,31 @@ def build_managers(self, assert draft_kv_cache_config is None, ( "KVCacheManagerV2 does not support two-model speculative " "decoding with separate draft KV cache budget splitting.") - # For V1, apply the draft's split budgets temporarily. - if draft_kv_cache_config is not None: - saved_budget = self._kv_cache_config.max_gpu_total_bytes - saved_host = self._kv_cache_config.host_cache_size - self._kv_cache_config.max_gpu_total_bytes = ( - draft_kv_cache_config.max_gpu_total_bytes) - self._kv_cache_config.host_cache_size = ( - draft_kv_cache_config.host_cache_size) draft_kv_cache_manager = self._create_kv_cache_manager( - self._draft_model_engine, estimating_kv_cache) - if draft_kv_cache_config is not None: - self._kv_cache_config.max_gpu_total_bytes = saved_budget - self._kv_cache_config.host_cache_size = saved_host + self._draft_model_engine, + estimating_kv_cache, + kv_cache_config_override=draft_build_kv_cache_config) # One-model speculative decoding with different KV layouts elif self._should_create_separate_draft_kv_cache(): draft_kv_cache_manager = self._create_one_model_draft_kv_cache_manager( estimating_kv_cache, - kv_cache_config_override=draft_kv_cache_config) + kv_cache_config_override=draft_build_kv_cache_config) + + # Encoder-decoder cross-attention pool + cross_kv_cache_manager = None + if cross_kv_cache_config is not None: + cross_kv_cache_manager = self._create_cross_kv_cache_manager( + cross_kv_cache_config, estimating_kv_cache, + original_max_seq_len) resources[ResourceManagerType.KV_CACHE_MANAGER] = kv_cache_manager resources[ ResourceManagerType.DRAFT_KV_CACHE_MANAGER] = draft_kv_cache_manager + resources[ + ResourceManagerType.CROSS_KV_CACHE_MANAGER] = cross_kv_cache_manager def teardown_managers(self, resources: Dict) -> None: - """Clean up KV caches for model and draft model (if applicable).""" + """Clean up KV caches for model, draft model, and cross pool.""" resources[ResourceManagerType.KV_CACHE_MANAGER].shutdown() del resources[ResourceManagerType.KV_CACHE_MANAGER] draft_kv_cache_manager = resources[ @@ -1138,6 +1422,12 @@ def teardown_managers(self, resources: Dict) -> None: if draft_kv_cache_manager: draft_kv_cache_manager.shutdown() del resources[ResourceManagerType.DRAFT_KV_CACHE_MANAGER] + cross_kv_cache_manager = resources.get( + ResourceManagerType.CROSS_KV_CACHE_MANAGER) + if cross_kv_cache_manager is not None: + cross_kv_cache_manager.shutdown() + if ResourceManagerType.CROSS_KV_CACHE_MANAGER in resources: + del resources[ResourceManagerType.CROSS_KV_CACHE_MANAGER] def _build_per_layer_num_kv_heads( @@ -1194,6 +1484,9 @@ def _create_kv_cache_manager( is_draft: Optional[bool] = None, layer_mask: Optional[List[bool]] = None, num_layers: Optional[int] = None, + num_kv_heads: Optional[Union[int, List[int]]] = None, + head_dim: Optional[int] = None, + kv_cache_type=None, is_disagg: bool = False) -> KVCacheManager: """ Returns: @@ -1215,11 +1508,15 @@ def _create_kv_cache_manager( if is_draft is None: is_draft = model_engine.is_draft_model + if kv_cache_type is None: + kv_cache_type = tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF + hidden_size = config.hidden_size num_attention_heads = config.num_attention_heads - num_key_value_heads = getattr(config, 'num_key_value_heads', - num_attention_heads) - head_dim = getattr(config, "head_dim", None) + num_key_value_heads = num_kv_heads if num_kv_heads is not None else getattr( + config, 'num_key_value_heads', num_attention_heads) + if not isinstance(head_dim, int): + head_dim = getattr(config, "head_dim", None) if not isinstance(head_dim, int): head_dim = hidden_size // num_attention_heads @@ -1493,7 +1790,7 @@ def _create_kv_cache_manager( and kv_cache_manager_cls.__name__ == "KVCacheManager" else head_dim) kv_cache_manager = kv_cache_manager_cls( kv_cache_config, - tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF, + kv_cache_type, num_layers=num_hidden_layers, num_kv_heads=per_layer_num_kv_heads, head_dim=effective_head_dim, @@ -1720,18 +2017,32 @@ def create_py_executor_instance( resource_manager = ResourceManager(resources) - # Make sure the kv cache manager is always invoked last as it could + # Make sure the kv cache managers are always invoked last as they could # depend on the results of other resource managers. if kv_cache_manager is not None: resource_manager.resource_managers.move_to_end( ResourceManagerType.KV_CACHE_MANAGER, last=True) + cross_kv_cache_manager = resources.get( + ResourceManagerType.CROSS_KV_CACHE_MANAGER) + if cross_kv_cache_manager is not None: + resource_manager.resource_managers.move_to_end( + ResourceManagerType.CROSS_KV_CACHE_MANAGER, last=True) + # When scheduler_capacity == 1, attention dp dummy request will prevent the scheduling of DISAGG_GENERATION_INIT. # Enlarge scheduler capacity to avoid DISAGG_GENERATION_INIT stuck in the scheduler. scheduler_capacity = max_num_sequences if scheduler_capacity == 1 and mapping.enable_attention_dp and kv_cache_manager: scheduler_capacity += 1 + # For encoder-decoder models, requests start in ENCODER_INIT and the + # capacity scheduler must admit them already at that state so the + # encoder loop can run. Decoder-only deployments keep the default + # CONTEXT_INIT gating. + no_schedule_until_state = (LlmRequestState.ENCODER_INIT + if cross_kv_cache_manager is not None else + LlmRequestState.CONTEXT_INIT) + if isinstance(kv_cache_manager, KVCacheManagerV2): # V2: interleaved scheduler handles both capacity and budget draft_kv_cache_manager = resources.get( @@ -1749,6 +2060,8 @@ def create_py_executor_instance( if peft_cache_manager is not None else None, scheduler_capacity=scheduler_capacity, draft_kv_cache_manager=draft_kv_cache_manager, + cross_kv_cache_manager=cross_kv_cache_manager, + no_schedule_until_state=no_schedule_until_state, ) elif (scheduler_config is not None and scheduler_config.use_python_scheduler): @@ -1761,15 +2074,21 @@ def create_py_executor_instance( if peft_cache_manager is not None else None, scheduler_policy=scheduler_config.capacity_scheduler_policy, ctx_chunk_config=ctx_chunk_config, + cross_kv_cache_manager=cross_kv_cache_manager.impl + if cross_kv_cache_manager is not None else None, two_step_lookahead=mapping.has_pp(), - scheduler_capacity=scheduler_capacity) + scheduler_capacity=scheduler_capacity, + no_schedule_until_state=no_schedule_until_state) else: capacity_scheduler = BindCapacityScheduler( scheduler_capacity, kv_cache_manager.impl if kv_cache_manager is not None else None, peft_cache_manager.impl if peft_cache_manager is not None else None, scheduler_config.capacity_scheduler_policy, - two_step_lookahead=mapping.has_pp()) + cross_kv_cache_manager=cross_kv_cache_manager.impl + if cross_kv_cache_manager is not None else None, + two_step_lookahead=mapping.has_pp(), + no_schedule_until_state=no_schedule_until_state) mb_scheduler = BindMicroBatchScheduler(max_batch_size, max_num_tokens, ctx_chunk_config) diff --git a/tensorrt_llm/_torch/pyexecutor/llm_request.py b/tensorrt_llm/_torch/pyexecutor/llm_request.py index 3d2c0b8980b0..6f793607f539 100644 --- a/tensorrt_llm/_torch/pyexecutor/llm_request.py +++ b/tensorrt_llm/_torch/pyexecutor/llm_request.py @@ -292,6 +292,7 @@ class Diff: additional_generation_outputs_list: list[tuple[str, torch.Tensor]] = field( default_factory=list) + encoder_output: torch.Tensor | None = None def __init__(self, *, @@ -339,6 +340,7 @@ def __init__(self, name: [] for name in additional_outputs } if additional_outputs else None + self._encoder_output: Optional[torch.Tensor] = None self.diff = PyResult.Diff() def reset_diff(self): @@ -349,6 +351,8 @@ def get_diff(self) -> Diff: self.diff.context_logits_list[i] = context_logits.to("cpu") for i, generation_logits in enumerate(self.diff.generation_logits_list): self.diff.generation_logits_list[i] = generation_logits.to("cpu") + if self.diff.encoder_output is not None: + self.diff.encoder_output = self.diff.encoder_output.detach().cpu() return self.diff def apply_diff(self, diff: Diff): @@ -369,6 +373,8 @@ def apply_diff(self, diff: Diff): if diff.mrope_position_ids is not None: self._mrope_position_ids = diff.mrope_position_ids self._mrope_position_deltas = diff.mrope_position_deltas + if diff.encoder_output is not None: + self._encoder_output = diff.encoder_output if len(diff.additional_context_outputs_list) > 0: for name, additional_context_outputs in diff.additional_context_outputs_list: self._additional_context_outputs[name].append( @@ -431,6 +437,10 @@ def set_mrope_position( self.diff.mrope_position_ids = self._mrope_position_ids self.diff.mrope_position_deltas = self._mrope_position_deltas + def set_encoder_output(self, encoder_output: torch.Tensor): + self._encoder_output = encoder_output + self.diff.encoder_output = encoder_output + def transfer_remaining_device_logits(self): """Finalize any remaining generation logits transfers (for chunked mode)""" if self._generation_logits: @@ -560,6 +570,10 @@ def additional_generation_outputs(self) -> Dict[str, torch.Tensor] | None: output_list, dim=0) if len(output_list) > 1 else output_list[0] return outputs + @property + def encoder_output(self) -> torch.Tensor | None: + return self._encoder_output + class LlmResult: """LlmResult wraps `bindings.executor.Result` but detour some features to Python implementation""" @@ -567,7 +581,8 @@ class LlmResult: ('context_logits', 'generation_logits', 'log_probs', 'cum_log_probs', 'first_gen_log_probs', 'mm_embedding_handles', 'additional_context_outputs', 'additional_generation_outputs', - 'mrope_position_ids_handle', 'mrope_position_deltas_handle')) + 'encoder_output', 'mrope_position_ids_handle', + 'mrope_position_deltas_handle')) def __init__(self, result: Union[bytes, tensorrt_llm.bindings.executor.Result], @@ -659,6 +674,15 @@ def __init__( self.py_lora_path: str | None = kwargs.pop("py_lora_path", None) # Multimodal data self.py_multimodal_data = kwargs.pop("py_multimodal_data", None) + encoder_input_tokens = kwargs.get("encoder_input_tokens") + encoder_output_len = kwargs.get("encoder_output_len") + return_encoder_output = bool(kwargs.get("return_encoder_output", False)) + if return_encoder_output: + kwargs["return_encoder_output"] = False + if (llm_request is None and encoder_input_tokens is not None + and encoder_output_len is None): + encoder_output_len = len(encoder_input_tokens) + kwargs["encoder_output_len"] = encoder_output_len if llm_request is not None: super().__init__(llm_request) else: @@ -671,8 +695,16 @@ def __init__( return_perf_metrics=return_perf_metrics, stop_words_list=torch.tensor(stop_words_list, dtype=torch.int32) if stop_words_list else None, - **kwargs, - ) + **kwargs) + if encoder_output_len is not None and not hasattr( + self, "encoder_output_len"): + self.encoder_output_len = int(encoder_output_len) + if encoder_input_tokens is not None and not hasattr( + self, "encoder_tokens"): + encoder_tokens = (encoder_input_tokens.tolist() if hasattr( + encoder_input_tokens, "tolist") else list(encoder_input_tokens)) + self.encoder_tokens = encoder_tokens + self.py_return_encoder_output = return_encoder_output self.py_client_id = client_id self.py_request_id = self.request_id self.py_llm_request_type = self.llm_request_type @@ -704,6 +736,22 @@ def __init__( self.py_kv_transfer_start_time = None self.py_kv_transfer_timed_out = False + # Encoder-decoder runtime state. ``py_encoder_output`` holds the + # packed encoder hidden states produced by the encoder iteration as + # a GPU buffer for cross-attention projection and fallback paths. + # ``py_encoder_output_ready_event`` is + # recorded on the encoder stream when those hidden states become + # available; the scheduler queries it before admitting the request + # to a decoder context step. ``py_skip_cross_kv_projection`` controls + # whether the decoder's cross-attention projects K/V from + # ``encoder_output`` (False on the first context step, the only step + # that writes the cross pool) or reads cross-KV without projection + # (True on later decoder steps and chunks). All three are unused for + # decoder-only models. + self.py_encoder_output: Optional[torch.Tensor] = None + self.py_encoder_output_ready_event: Optional[torch.cuda.Event] = None + self.py_skip_cross_kv_projection: bool = False + # Performance timing info (step metrics, GPU events, context GPU timing) # Lazily created only when return_perf_metrics is enabled to avoid # overhead for every request. @@ -900,7 +948,10 @@ def create_child_request(self, child_id): if attr_name.startswith('py_'): attr_value = getattr(self, attr_name) setattr(py_request, attr_name, deepcopy(attr_value)) - elif attr_name in ['is_attention_dp_dummy', 'is_cuda_graph_dummy']: + elif attr_name in [ + 'is_attention_dp_dummy', 'is_cuda_graph_dummy', + 'encoder_tokens', 'encoder_output_len' + ]: setattr(py_request, attr_name, attr_value) # Rewrite specific attributes that should use child_request values. @@ -1104,8 +1155,9 @@ def executor_request_to_llm_request( guided_decoding_params=executor_request.guided_decoding_params, py_logits_post_processors=getattr(executor_request, "py_logits_post_processors", None), - encoder_input_tokens=None, - return_encoder_output=False, + encoder_input_tokens=executor_request.encoder_input_token_ids, + return_encoder_output=executor_request.output_config. + return_encoder_output, client_id=executor_request.client_id if executor_request.client_id is not None else req_id, priority=executor_request.priority, diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 152ed8c825fe..32f456d76407 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -993,6 +993,12 @@ def warmup(self, resource_manager: ResourceManager) -> None: # Reset the global cuda graph dummy requests in warmup. self.cuda_graph_runner.padding_dummy_requests = {} + if self._is_encoder_decoder_model(): + logger.info( + "Skipping warmup for encoder-decoder models; warmup dummy " + "requests do not carry encoder output state.") + return + if self.mapping.cp_size > 1: cp_type = self.mapping.cp_config.get("cp_type", None) if cp_type != CpType.HELIX: @@ -2204,6 +2210,131 @@ def _prepare_multimodal_indices(self, input_ids: list[int]): input_ids, vocab_size=vocab_size, mm_token_ids=mm_token_ids) return text_token_indices, mm_token_indices + def _is_encoder_decoder_model(self) -> bool: + return bool( + getattr(getattr(self.model, "model_config", None), + "is_encoder_decoder", False)) + + def _get_top_level_model(self) -> Any: + model = getattr(self.model, "_orig_mod", self.model) + top_level_model = getattr(model, "model", model) + return getattr(top_level_model, "_orig_mod", top_level_model) + + def _get_position_id_offset(self) -> int: + offset = getattr(self._get_top_level_model(), "position_id_offset", 0) + return 0 if offset is None else int(offset) + + def _apply_position_id_offset(self, position_ids: List[int]) -> List[int]: + offset = self._get_position_id_offset() + if offset == 0: + return position_ids + return [position_id + offset for position_id in position_ids] + + def _prepare_encoder_decoder_cross_attention_inputs( + self, + encoder_hidden_states: List[torch.Tensor], + encoder_seq_lens: List[int], + encoder_num_cached_tokens_per_seq: List[int], + attn_metadata: AttentionMetadata, + resource_manager: Optional[ResourceManager], + ) -> Dict[str, Any]: + if not encoder_seq_lens: + return {} + + if len(encoder_seq_lens) != attn_metadata.num_seqs: + raise RuntimeError( + "Cross-attention encoder lengths must align with decoder " + f"sequences: got {len(encoder_seq_lens)} encoder lengths for " + f"{attn_metadata.num_seqs} decoder sequences.") + + if resource_manager is None: + raise RuntimeError( + "Encoder-decoder decoder forward requires a resource manager " + "with a cross-KV cache manager.") + cross_kv_cache_manager = resource_manager.get_resource_manager( + ResourceManagerType.CROSS_KV_CACHE_MANAGER) + if cross_kv_cache_manager is None: + raise RuntimeError("Encoder-decoder decoder forward requires " + "ResourceManagerType.CROSS_KV_CACHE_MANAGER.") + + new_encoder_tokens = sum(encoder_seq_lens) + if encoder_hidden_states: + packed_encoder_hidden_states = ( + encoder_hidden_states[0] if len(encoder_hidden_states) == 1 else + torch.cat(encoder_hidden_states, dim=0)) + if packed_encoder_hidden_states.shape[0] != new_encoder_tokens: + raise RuntimeError( + "Packed encoder hidden states do not match cross-attention " + "metadata: got " + f"{packed_encoder_hidden_states.shape[0]} rows for " + f"{new_encoder_tokens} new encoder KV tokens.") + skip_cross_kv_projection = False + else: + if new_encoder_tokens != 0: + raise RuntimeError( + "Cross-attention metadata asks to project encoder K/V, " + "but no encoder hidden states were supplied.") + packed_encoder_hidden_states = None + skip_cross_kv_projection = True + + encoder_seq_lens_tensor = torch.tensor(encoder_seq_lens, + dtype=torch.int, + pin_memory=prefer_pinned()) + + def update_cross_metadata( + cross_attn_metadata: AttentionMetadata) -> AttentionMetadata: + base_params = attn_metadata.kv_cache_params + cross_attn_metadata.kv_cache_manager = cross_kv_cache_manager + cross_attn_metadata._seq_lens = attn_metadata.seq_lens + cross_attn_metadata._seq_lens_cuda = attn_metadata.seq_lens_cuda + cross_attn_metadata.cross = cross_attn_metadata + cross_attn_metadata.seq_lens_kv = encoder_seq_lens_tensor + if encoder_num_cached_tokens_per_seq is not None: + use_cache = (base_params.use_cache if base_params is not None + else (cross_kv_cache_manager is not None)) + block_ids_per_seq = (base_params.block_ids_per_seq + if base_params is not None else None) + host_max_attention_window_sizes = ( + base_params.host_max_attention_window_sizes + if base_params is not None else None) + host_sink_token_length = (base_params.host_sink_token_length + if base_params is not None else None) + num_extra_kv_tokens = (base_params.num_extra_kv_tokens + if base_params is not None else 0) + cross_attn_metadata.kv_cache_params = KVCacheParams( + use_cache=use_cache, + num_cached_tokens_per_seq=list( + encoder_num_cached_tokens_per_seq), + block_ids_per_seq=block_ids_per_seq, + host_max_attention_window_sizes= + host_max_attention_window_sizes, + host_sink_token_length=host_sink_token_length, + num_extra_kv_tokens=num_extra_kv_tokens, + ) + cross_attn_metadata.request_ids = attn_metadata.request_ids + cross_attn_metadata.prompt_lens = attn_metadata.prompt_lens + cross_attn_metadata.num_contexts = attn_metadata.num_contexts + return cross_attn_metadata + + if attn_metadata.is_cuda_graph and attn_metadata.has_cross_sub_metadata: + cross_attn_metadata = update_cross_metadata(attn_metadata.cross) + else: + cross_attn_metadata = attn_metadata.create_cross_metadata( + encoder_seq_lens=encoder_seq_lens_tensor, + cross_kv_cache_manager=cross_kv_cache_manager, + encoder_num_cached_tokens_per_seq= + encoder_num_cached_tokens_per_seq, + ) + if attn_metadata.is_cuda_graph: + attn_metadata.cross = cross_attn_metadata + cross_attn_metadata.prepare() + + return { + "encoder_hidden_states": packed_encoder_hidden_states, + "cross_attn_metadata": cross_attn_metadata, + "skip_cross_kv_projection": skip_cross_kv_projection, + } + def _can_use_incremental_update( self, scheduled_requests: ScheduledRequests, new_tokens_device: Optional[torch.Tensor], @@ -2743,6 +2874,11 @@ def _prepare_tp_inputs( # requests, whose outputs are discarded. mrope_dummy_seq_slot = self.max_num_tokens * self.mapping.pp_size num_accepted_draft_tokens = [] # per request + is_encoder_decoder = self._is_encoder_decoder_model() + cross_encoder_hidden_states: List[torch.Tensor] = [] + cross_encoder_seq_lens: List[int] = [ + ] # new encoder K/V tokens per decoder sequence + cross_encoder_cached_tokens_per_seq: List[int] = [] # if using tree decoding, we need to store the request type and accepted path for each request, # which will be used to update the hidden_states_read_indices. request_accepted_path = {} # per request @@ -2760,6 +2896,37 @@ def _prepare_tp_inputs( # (start_idx, end_idx, seq_slot) for first_draft requests first_draft_input_ids_positions = [] + def append_cross_attention_state(request: LlmRequest, + project_encoder_output: bool, + repeat: int = 1) -> None: + if not is_encoder_decoder: + return + + encoder_output_len = int(request.encoder_output_len) + if project_encoder_output: + encoder_output = getattr(request, "py_encoder_output", None) + if encoder_output is None: + raise RuntimeError( + "Decoder context request " + f"{request.py_request_id} has no encoder output. " + "The encoder iteration must populate " + "req.py_encoder_output before the first decoder " + "context step.") + if encoder_output.shape[0] != encoder_output_len: + raise RuntimeError( + "Decoder context request " + f"{request.py_request_id} encoder output length " + f"({encoder_output.shape[0]}) does not match " + f"encoder_output_len ({encoder_output_len}).") + cross_encoder_hidden_states.append(encoder_output) + cross_encoder_seq_lens.append(encoder_output_len) + cross_encoder_cached_tokens_per_seq.append(0) + return + + for _ in range(repeat): + cross_encoder_seq_lens.append(0) + cross_encoder_cached_tokens_per_seq.append(encoder_output_len) + for request in scheduled_requests.context_requests: request_ids.append(request.py_request_id) all_prompt_tokens = request.get_tokens(0) @@ -2797,6 +2964,10 @@ def _prepare_tp_inputs( past_seen_token_num = begin_compute num_cached_tokens_per_seq.append(past_seen_token_num) request.cached_tokens = num_cached_tokens_per_seq[-1] + append_cross_attention_state( + request, + project_encoder_output=not request.py_skip_cross_kv_projection + and not getattr(request, "is_dummy", False)) # Embed mask is required only for partial iterations (chunked # prefill or KV-cache reuse); full-prefill degrades gracefully. @@ -2992,6 +3163,8 @@ def _prepare_tp_inputs( else: prompt_lengths.append(request.py_prompt_len) + append_cross_attention_state(request, project_encoder_output=False) + for request in first_draft_requests: request_ids.append(request.py_request_id) all_prompt_tokens = request.get_tokens(0) @@ -3041,6 +3214,7 @@ def _prepare_tp_inputs( prompt_lengths.append(request.py_prompt_len) past_seen_token_num = begin_compute num_cached_tokens_per_seq.append(past_seen_token_num) + append_cross_attention_state(request, project_encoder_output=False) # update batch index request.py_batch_idx = request.py_seq_slot @@ -3171,6 +3345,9 @@ def _prepare_tp_inputs( strip_mm_data_for_generation(request.py_multimodal_data) request.py_batch_idx = request.py_seq_slot + append_cross_attention_state(request, + project_encoder_output=False, + repeat=beam_width) # Do not add a gen_request_seq_slot for CUDA graph dummy requests # to prevent access errors due to None values if not request.is_cuda_graph_dummy: @@ -3401,6 +3578,7 @@ def previous_seq_slots_device(): self.previous_pos_id_offsets_cuda *= 0 self.previous_kv_lens_offsets_cuda *= 0 + position_ids = self._apply_position_id_offset(position_ids) if self.use_mrope and mrope_position_ids: # Mixed batches may have only some requests with multimodal MRoPE # data. Seed the full (3,1,N) buffer from scalar position_ids @@ -3537,6 +3715,14 @@ def previous_seq_slots_device(): if hasattr(self.model.model_config.pretrained_config, 'chunk_size'): attn_metadata.mamba_chunk_size = self.model.model_config.pretrained_config.chunk_size attn_metadata.prepare() + cross_attention_inputs = ( + self._prepare_encoder_decoder_cross_attention_inputs( + cross_encoder_hidden_states, + cross_encoder_seq_lens, + cross_encoder_cached_tokens_per_seq, + attn_metadata, + resource_manager, + ) if is_encoder_decoder else {}) peft_cache_manager = resource_manager and resource_manager.get_resource_manager( ResourceManagerType.PEFT_CACHE_MANAGER) @@ -3578,6 +3764,7 @@ def previous_seq_slots_device(): "multimodal_params": multimodal_params_list, 'resource_manager': resource_manager, } + inputs.update(cross_attention_inputs) if self.use_mrope: if mrope_delta_write_seq_slots: @@ -3705,6 +3892,7 @@ def _prepare_tp_inputs_no_cache( pin_memory=prefer_pinned()) self.input_ids_cuda[:num_tokens].copy_(input_ids, non_blocking=True) + position_ids = self._apply_position_id_offset(position_ids) position_ids = torch.tensor(position_ids, dtype=torch.int, pin_memory=prefer_pinned()) @@ -4938,6 +5126,189 @@ def _forward_step_mm_encoder_only( return result + @nvtx_range("_prepare_tp_inputs_encoder") + def _prepare_tp_inputs_encoder( + self, + encoder_requests: List[LlmRequest], + resource_manager: Optional[ResourceManager] = None, + ): + """Pack encoder-side inputs for an encoder-decoder forward pass. + + Mirrors the no-cache path used by ``mm_encoder_only`` and the + legacy ``EncoderBuffers`` shape contract: ``encoder_input_ids`` + and ``encoder_position_ids`` are concatenated across requests + into a single ``[sum(encoder_output_len)]`` tensor, with one + non-causal :class:`AttentionMetadata` describing the packed + encoder batch. + + The encoder pass does not touch any KV-cache pool. The cross pool is + only written by the decoder's cross-attention on the first context + step. Self-pool blocks for the decoder are reserved on the next + scheduler iteration when the request transitions to ``CONTEXT_INIT``. + """ + if not encoder_requests: + raise ValueError( + "_prepare_tp_inputs_encoder called with no encoder requests") + + encoder_input_ids: List[int] = [] + encoder_position_ids: List[int] = [] + sequence_lengths: List[int] = [] + request_ids: List[int] = [] + + for request in encoder_requests: + tokens = request.encoder_tokens + if tokens is None: + raise ValueError( + f"Encoder request {request.py_request_id} has no " + "encoder_tokens; encoder_input_token_ids must be wired " + "through executor_request_to_llm_request.") + seq_len = len(tokens) + encoder_input_ids.extend(tokens) + encoder_position_ids.extend( + self._apply_position_id_offset(list(range(seq_len)))) + sequence_lengths.append(seq_len) + request_ids.append(request.py_request_id) + + num_tokens = len(encoder_input_ids) + assert num_tokens <= self.max_num_tokens, ( + f"encoder packed length ({num_tokens}) exceeds max_num_tokens " + f"({self.max_num_tokens})") + + # Build a fresh, no-cache attention metadata for the encoder + # pass. We do not reuse ``self.attn_metadata`` because that + # object is bound to the decoder's KV-cache manager. + encoder_attn_metadata = self.attn_backend.Metadata( + max_num_requests=self.batch_size, + max_num_tokens=self.max_num_tokens, + max_num_sequences=self.batch_size * self.max_beam_width, + kv_cache_manager=None, + mapping=self.mapping, + runtime_features=self.attn_runtime_features, + enable_flash_mla=self.model.model_config.enable_flash_mla, + enable_context_mla_with_cached_kv=False, + cache_indirection=None, + sparse_attention_config=self.sparse_attention_config, + num_heads_per_kv=1, + ) + assert isinstance( + encoder_attn_metadata, + (VanillaAttentionMetadata, TrtllmAttentionMetadata) + ), "Only vanilla and trtllm attention metadata are supported for the encoder pass" + + encoder_attn_metadata.seq_lens = torch.tensor( + sequence_lengths, + dtype=torch.int, + pin_memory=prefer_pinned(), + ) + encoder_attn_metadata.num_contexts = len(encoder_requests) + encoder_attn_metadata.max_seq_len = self.max_seq_len + encoder_attn_metadata.request_ids = request_ids + encoder_attn_metadata.prepare() + + encoder_input_ids_t = torch.tensor(encoder_input_ids, + dtype=torch.int, + pin_memory=prefer_pinned()) + encoder_position_ids_t = torch.tensor(encoder_position_ids, + dtype=torch.int, + pin_memory=prefer_pinned()) + + inputs = { + 'encoder_input_ids': + encoder_input_ids_t.to('cuda', non_blocking=True), + 'encoder_position_ids': + encoder_position_ids_t.to('cuda', non_blocking=True).unsqueeze(0), + 'encoder_attn_metadata': + encoder_attn_metadata, + 'encoder_seq_lens': + sequence_lengths, + 'resource_manager': + resource_manager, + } + return inputs + + @nvtx_range("_forward_step_encoder") + def _forward_step_encoder( + self, + inputs: Dict[str, Any], + ) -> torch.Tensor: + """Run the encoder stack and return packed encoder hidden states. + + Returns ``[sum(encoder_output_len), hidden_size]`` (matches the + ``EncoderBuffers`` shape contract from the legacy TRT path). + Slicing back into per-request hidden states is the executor's + responsibility — see :meth:`PyExecutor._scatter_encoder_output`. + """ + encoder = getattr(self.model, "encoder", None) + if encoder is None: + inner = getattr(self.model, "model", None) + encoder = getattr(inner, "encoder", + None) if inner is not None else None + if encoder is None: + raise AttributeError( + "Model does not expose an `encoder` submodule; encoder-decoder " + "models must define a top-level `encoder` (or `model.encoder`) " + "stack to participate in the encoder iteration.") + + # Encoder operates on packed token IDs. Models like T5 own the + # shared embedding on ``self.model`` rather than inside the + # encoder stack, so we go through the top-level model when + # available so the embedding is applied consistently with the + # decoder pass. + top_level_model = self._get_top_level_model() + embed = getattr(top_level_model, "shared_embedding", None) or getattr( + top_level_model, "embed_tokens", None) + encoder_input_ids = inputs['encoder_input_ids'] + if embed is not None: + hidden_states = embed(encoder_input_ids) + embed_scale = getattr(top_level_model, "embed_scale", None) + if embed_scale is not None: + hidden_states = hidden_states * embed_scale + else: + # Fall back to letting the encoder accept token ids directly. + hidden_states = encoder_input_ids + + encoder_attn_metadata = inputs['encoder_attn_metadata'] + position_ids = inputs.get('encoder_position_ids') + if position_ids is not None and position_ids.dim() == 2: + position_ids = position_ids.squeeze(0) + + encoder_hidden_states = encoder( + hidden_states=hidden_states, + attn_metadata=encoder_attn_metadata, + position_ids=position_ids, + ) + return encoder_hidden_states + + @nvtx_range("forward_encoder") + def forward_encoder( + self, + encoder_requests: List[LlmRequest], + resource_manager: Optional[ResourceManager] = None, + ) -> Tuple[torch.Tensor, List[int]]: + """Run the encoder stack for ``encoder_requests``. + + Returns a tuple ``(encoder_hidden_states, encoder_seq_lens)`` + where the hidden states tensor is shaped + ``[sum(encoder_seq_lens), hidden_size]`` (one packed batch). + The accompanying ``encoder_seq_lens`` list is in the same + ordering as ``encoder_requests``, so callers can split the + packed output 1:1. + + This entry point is the encoder-step analog of the legacy + ``TrtEncoderModel::forwardAsync`` (see §2.6/§2.7). The decoder + IFB step is unchanged and continues to flow through + :meth:`forward`. + """ + if not encoder_requests: + raise ValueError("forward_encoder called with no encoder requests") + + with torch.inference_mode(): + inputs = self._prepare_tp_inputs_encoder( + encoder_requests, resource_manager=resource_manager) + encoder_hidden_states = self._forward_step_encoder(inputs) + + return encoder_hidden_states, inputs['encoder_seq_lens'] + def _init_userbuffers(self, hidden_size): if self.mapping.tp_size <= 1 or self.mapping.pp_size > 1: return False diff --git a/tensorrt_llm/_torch/pyexecutor/model_loader.py b/tensorrt_llm/_torch/pyexecutor/model_loader.py index 1c9fc3566700..4e763cfd7613 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_loader.py +++ b/tensorrt_llm/_torch/pyexecutor/model_loader.py @@ -97,6 +97,28 @@ def validate_and_set_kv_cache_quant(model_config: ModelConfig, model_config.quant_config.kv_cache_quant_algo = mapped_pyt_quant +def validate_encoder_decoder_kv_cache_config(model_config: ModelConfig, + kv_cache_config) -> None: + """Validate encoder-decoder KV-cache requirements for the PyTorch runtime. + + Both V1 (``KVCacheManager``, default and production target) and V2 + (``KVCacheManagerV2``, additive secondary path) are supported for + encoder-decoder models. Both paths require ``cross_kv_cache_fraction`` + so the cross-attention pool can be sized. + """ + if model_config.is_encoder_decoder: + if kv_cache_config.cross_kv_cache_fraction is None: + raise ValueError( + "Encoder-decoder models require kv_cache_config.cross_kv_cache_fraction to be set." + ) + return + + if kv_cache_config.cross_kv_cache_fraction is not None: + raise ValueError( + "kv_cache_config.cross_kv_cache_fraction should only be set for encoder-decoder models." + ) + + def initialize_dummy_weights( model: torch.nn.Module, low: float = -1e-3, @@ -1021,6 +1043,8 @@ def _load_and_validate_config( f"{type(config.pretrained_config).__name__}: {e}. " f"AllReduce pre-allocation will be skipped.") + validate_encoder_decoder_kv_cache_config(config, + self.llm_args.kv_cache_config) validate_and_set_kv_cache_quant(config, self.llm_args.kv_cache_config.dtype) validate_and_set_mamba_ssm_cache_dtype( diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 1b624f72d88b..512fe9d0470c 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -8,7 +8,8 @@ from contextlib import contextmanager from enum import IntEnum from queue import Queue -from typing import Callable, Dict, Iterable, List, Optional, Tuple, Union +from typing import (TYPE_CHECKING, Callable, Dict, Iterable, List, Optional, + Tuple, Union) import torch @@ -82,6 +83,9 @@ create_waiting_queue) from .scheduler.adp_router import ADPRouter +if TYPE_CHECKING: + from ray.actor import ActorHandle + # Environment variable to specify iteration ranges for profiling start/stop. # Format: "start1-stop1,start2-stop2,..." or single iterations "iter1,iter2,..." PROFILE_START_STOP_ENV_VAR_NAME = "TLLM_PROFILE_START_STOP" @@ -337,15 +341,20 @@ def __init__( super(PyExecutor, self).__init__() self.device_id = torch.cuda.current_device() self.global_rank = dist.rank - # Store the execution stream for model forward operations. - # This stream is used for proper synchronization with KVCacheTransferManager. - # execution_stream can be provided by create_py_executor - # Create a new stream if none provided + # Store the execution stream for decoder/model forward operations. + # This stream is used for proper synchronization with + # KVCacheTransferManager. execution_stream can be provided by + # create_py_executor. Create a new stream if none provided. self.execution_stream = execution_stream if execution_stream is not None else torch.cuda.Stream( ) + # Encoder-decoder requests use a dedicated encoder stream so the + # encoder forward does not serialize the decoder forward when the + # two operate on disjoint request sets. Per-request CUDA events + # carry the encoder->decoder dependency to the eventual consumer. + self.encoder_stream = torch.cuda.Stream() logger.info( - f"[PyExecutor] execution_stream initialized: {self.execution_stream}. " - ) + f"[PyExecutor] execution_stream initialized: {self.execution_stream}; " + f"encoder_stream initialized: {self.encoder_stream}.") self.peft_cache_config = peft_cache_config @@ -666,6 +675,31 @@ def on_detected(): self._disagg_pp_termination_handler = DisaggPPTerminationHandler( self.dist, self._do_terminate_request) + # Encoder-decoder models execute the encoder and decoder in separate + # iterations. The encoder branch lives in ``_executor_loop`` only; + # ``_executor_loop_overlap`` has not been threaded yet. Reject + # pp_size > 1 for parity with the legacy TRT path (Encoder PP support + # is intentionally out of scope for this port). + is_encoder_decoder = bool( + getattr(getattr(self.model_engine.model, "model_config", None), + "is_encoder_decoder", False)) + if is_encoder_decoder: + if self.dist.pp_size > 1: + raise NotImplementedError( + "pp_size > 1 is not supported for encoder-decoder models " + "in the PyTorch flow; encoder send/recv hooks are out of " + "scope. Set pp_size=1 to run T5/BART/mBART.") + if not self.disable_overlap_scheduler: + raise NotImplementedError( + "Overlap scheduler is not yet wired for encoder-decoder " + "models. Set disable_overlap_scheduler=True for " + "encoder-decoder runs.") + if getattr(self.model_engine, "cuda_graph_config", + None) is not None: + raise NotImplementedError( + "CUDA graph is not supported for encoder-decoder models. " + "Disable cuda_graph_config for encoder-decoder runs.") + if self.dist.pp_size > 1: self.event_loop = self._executor_loop_pp # `TLLM_PP_ASYNC_BROADCAST_SAMPLE_STATE` controls whether to broadcast the sample state asynchronously. @@ -925,10 +959,9 @@ def __exit__(self, exc_type, exc_val, exc_tb): self.shutdown() def enqueue_requests( - self, - requests: List[ExecutorRequest], - result_wait_queue: "Optional[ray.actor.ActorHandle]" = None - ) -> List[int]: + self, + requests: List[ExecutorRequest], + result_wait_queue: "Optional[ActorHandle]" = None) -> List[int]: """ Enqueue new requests """ @@ -1069,7 +1102,7 @@ def enqueue_request( self, request: ExecutorRequest, query: Optional[List] = None, - result_wait_queue: "Optional[ray.actor.ActorHandle]" = None) -> int: + result_wait_queue: "Optional[ActorHandle]" = None) -> int: """ Enqueue a new request, query is only used in `StarAttention`. """ @@ -1368,7 +1401,10 @@ def get_req_stats(req: LlmRequest) -> RequestStats: req_stat.reused_blocks_per_request = req.reused_blocks req_stat.missed_blocks_per_request = req.missed_blocks req_stat.kv_cache_hit_rate_per_request = req.kv_cache_hit_rate - req_stat.scheduled = req in scheduled_requests.context_requests or req in scheduled_requests.generation_requests + req_stat.scheduled = (req in scheduled_requests.encoder_requests + or req in scheduled_requests.context_requests + or req + in scheduled_requests.generation_requests) if req.llm_request_type == LlmRequestType.LLMREQUEST_TYPE_CONTEXT_ONLY or req.llm_request_type == LlmRequestType.LLMREQUEST_TYPE_GENERATION_ONLY: req_stat.dis_serving_stats = DisServingRequestStats() req_stat.dis_serving_stats.kv_cache_transfer_ms = req.kv_cache_transfer_time_ms @@ -2079,7 +2115,8 @@ def _executor_loop_pp(self): logger.debug( f'iteration {self.iter_counter}, microbatch {microbatch_id}, ' f'has {len(self.active_requests)} active_requests, ' - f'scheduled {scheduled_batch.num_context_requests} context requests and ' + f'scheduled {scheduled_batch.num_encoder_requests} encoder requests, ' + f'{scheduled_batch.num_context_requests} context requests and ' f'{scheduled_batch.num_generation_requests} generation requests' ) @@ -2695,7 +2732,8 @@ def _prepare_and_schedule_batch(self): self.num_scheduled_requests = scheduled_batch.batch_size logger.debug( f'has {len(self.active_requests)} active_requests, ' - f'scheduled {scheduled_batch.num_context_requests} context requests and ' + f'scheduled {scheduled_batch.num_encoder_requests} encoder requests, ' + f'{scheduled_batch.num_context_requests} context requests and ' f'{scheduled_batch.num_generation_requests} generation requests') return scheduled_batch, iter_stats @@ -2873,6 +2911,15 @@ def _executor_loop(self): gpu_forward_end = None gpu_forward_events_from_perf_pool = False + # Run the encoder iteration first. After scatter the + # encoder requests transition to ``CONTEXT_INIT`` and are + # picked up by the next scheduler iteration as decoder + # context. The encoder pass is independent of the decoder + # ``can_queue`` gate, so an iteration with only encoder-init + # requests still makes forward progress. + if scheduled_batch.encoder_requests: + self._run_encoder_step(scheduled_batch.encoder_requests) + can_queue, _ = self._can_queue(scheduled_batch) if can_queue: @@ -4018,12 +4065,154 @@ def _schedule(self): self._revert_ctx_alloc(dropped) scheduled_requests = ScheduledRequests() + scheduled_requests.encoder_requests = scheduler_output.encoder_requests scheduled_requests.reset_context_requests(scheduled_context_requests) scheduled_requests.generation_requests = scheduler_output.generation_requests scheduled_requests.paused_requests = scheduler_output.paused_requests return scheduled_requests, scheduler_output.fitting_disagg_gen_init_requests, num_fitting + # --------------------------------------------------------------- + # Encoder-decoder support: encoder iteration in the executor loop. + # + # At a scheduling pass, the scheduler may admit encoder-init requests + # alongside decoder-context and generation requests. It returns them in + # disjoint buckets: + # + # * encoder requests (``LlmRequestState.ENCODER_INIT``), which run + # through ``ModelEngine.forward_encoder`` on this iteration. + # After scatter, they transition to ``CONTEXT_INIT`` and are + # re-admitted by the *next* iteration's scheduler pass for the + # decoder context step. + # + # * decoder-context requests (``CONTEXT_INIT`` and disagg-gen-init), + # which flow through the normal decoder IFB step. + # + # The invariant is that encoder and decoder context never share one + # micro-batch; this preserves the cross-KV lifecycle and the + # dual-pool budget. + # --------------------------------------------------------------- + @nvtx_range("_run_encoder_step") + def _run_encoder_step(self, encoder_requests: List[LlmRequest]) -> None: + """Drive one encoder iteration for ``encoder_requests``. + + Runs the encoder stack on the dedicated encoder stream, then + scatters the packed hidden states back onto the per-request + ``py_encoder_output`` field and transitions request state to + ``CONTEXT_INIT`` so the next scheduler pass picks them up as + decoder-context requests. A separate CUDA event is recorded for + each request on the encoder stream; the scheduler queries that + event before admitting the request to a decoder context step. + """ + if not encoder_requests: + return + + try: + self.encoder_stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(self.encoder_stream): + encoder_hidden_states, encoder_seq_lens = ( + self.model_engine.forward_encoder( + encoder_requests, + resource_manager=self.resource_manager, + )) + except Exception as e: + traceback.print_exc() + error_msg = str(e) + logger.error( + f"Encountered an error in encoder forward: {error_msg}") + self._handle_errors(error_msg, requests=encoder_requests) + return + + self._scatter_encoder_output(encoder_requests, encoder_hidden_states, + encoder_seq_lens) + for req in encoder_requests: + req.py_encoder_output_ready_event = torch.cuda.Event() + req.py_encoder_output_ready_event.record(self.encoder_stream) + # TODO(TRTLLM-12339): Honor return_encoder_output once the public + # LLM API shape for returned encoder hidden states is finalized. + + @nvtx_range("_scatter_encoder_output") + def _scatter_encoder_output( + self, + encoder_requests: List[LlmRequest], + encoder_hidden_states: torch.Tensor, + encoder_seq_lens: List[int], + ) -> None: + """Slice packed encoder hidden states into per-request tensors. + + Stores the slice for each request in ``req.py_encoder_output`` + as a temporary GPU buffer (consumed by the first decoder + context step), and transitions the request from + ``ENCODER_INIT`` to ``CONTEXT_INIT`` so the next scheduler + iteration admits it on the decoder side. + + ``py_skip_cross_kv_projection`` is initialized to ``False`` so + the *first* decoder context step projects K/V from + ``encoder_output`` and writes the cross-KV pool; the decoder + step flips it to ``True`` for later steps and chunks. + """ + if encoder_hidden_states is None: + raise RuntimeError( + "Encoder forward returned None hidden states; cannot " + "scatter encoder output to requests.") + + assert len(encoder_seq_lens) == len(encoder_requests), ( + "Encoder packed sequence lengths must match the number of " + "encoder requests") + assert encoder_hidden_states.shape[0] == sum(encoder_seq_lens), ( + "Encoder packed hidden states first dim must equal " + "sum(encoder_seq_lens)") + + offset = 0 + for req, seq_len in zip(encoder_requests, encoder_seq_lens): + req.py_encoder_output = encoder_hidden_states[offset:offset + + seq_len] + req.py_skip_cross_kv_projection = False + req.state = LlmRequestState.CONTEXT_INIT + offset += seq_len + + @nvtx_range("_attach_encoder_output_to_execution_stream") + def _attach_encoder_output_to_execution_stream( + self, scheduled_requests: ScheduledRequests) -> None: + """Hand encoder-produced tensors over to the execution stream. + + Per-request encoder output tensors are produced on the dedicated + ``encoder_stream`` and consumed by the decoder forward on + ``execution_stream``. Cross-stream correctness is guaranteed by + the scheduler: + ``drop_decoder_context_requests_waiting_for_encoder_output`` excludes + any ``CONTEXT_INIT`` request whose ``py_encoder_output_ready_event`` + has not completed, so by the time a request reaches this point the + encoder kernels for that request are already done. No + ``wait_event`` is therefore needed on the execution stream. + + Two pieces of bookkeeping remain that this helper performs: + + * ``record_stream`` is called on the encoder-output tensor so the + PyTorch caching allocator knows the storage is still in use on + the execution stream and must not be reused until the decoder + forward releases it. + * The spent ``py_encoder_output_ready_event`` is cleared so it + cannot be queried again on a later iteration. + """ + for req in scheduled_requests.context_requests: + ready_event = getattr(req, "py_encoder_output_ready_event", None) + if ready_event is None: + continue + + if req.py_encoder_output is not None: + req.py_encoder_output.record_stream(self.execution_stream) + req.py_encoder_output_ready_event = None + + def _mark_cross_kv_projection_consumed( + self, scheduled_requests: ScheduledRequests) -> None: + """Release temporary encoder outputs after decoder context consumes them.""" + for req in scheduled_requests.context_requests: + if getattr(req, "py_encoder_output", None) is None: + continue + req.py_encoder_output = None + req.py_skip_cross_kv_projection = True + @nvtx_range("_check_disagg_gen_transfer_status") def _check_disagg_gen_transfer_status(self): @@ -4536,11 +4725,13 @@ def forward(scheduled_requests, resource_manager, new_tensors_device, # Run model forward on the execution stream for proper synchronization # with KVCacheTransferManager's onboard/offload operations. self.execution_stream.wait_stream(torch.cuda.current_stream()) + self._attach_encoder_output_to_execution_stream(scheduled_requests) with torch.cuda.stream(self.execution_stream): outputs = forward(scheduled_requests, self.resource_manager, new_tensors_device, gather_context_logits, cache_indirection_buffer, num_accepted_tokens_device) + self._mark_cross_kv_projection_consumed(scheduled_requests) # Ensure the default stream waits for execution_stream to complete # before downstream operations use the outputs. diff --git a/tensorrt_llm/_torch/pyexecutor/resource_manager.py b/tensorrt_llm/_torch/pyexecutor/resource_manager.py index 137007e1efbc..183d5ab0c715 100644 --- a/tensorrt_llm/_torch/pyexecutor/resource_manager.py +++ b/tensorrt_llm/_torch/pyexecutor/resource_manager.py @@ -53,6 +53,7 @@ if TYPE_CHECKING: from tensorrt_llm._torch.attention_backend.interface import \ AttentionMetadata + from tensorrt_llm.llmapi.llm_args import DecodingBaseConfig BlocksPerWindow = Dict[int, Tuple[ int, @@ -78,6 +79,7 @@ class PoolConfiguration: class ResourceManagerType(enum.Enum): KV_CACHE_MANAGER = "KV_CACHE_MANAGER" DRAFT_KV_CACHE_MANAGER = "DRAFT_KV_CACHE_MANAGER" + CROSS_KV_CACHE_MANAGER = "CROSS_KV_CACHE_MANAGER" PEFT_CACHE_MANAGER = "PEFT_CACHE_MANAGER" SEQ_SLOT_MANAGER = "SEQ_SLOT_MANAGER" SPEC_RESOURCE_MANAGER = "SPEC_RESOURCE_MANAGER" @@ -668,49 +670,54 @@ def get_needed_resource_to_completion(self, request: LlmRequest) -> int: remaining_tokens / self.tokens_per_block) return need_blocks + def _context_seq_len(self, req: LlmRequest, is_cross: bool, + is_star_cp: bool) -> Optional[int]: + """Return the sequence length to pass to add_sequence_batch, or None to skip this request.""" + if is_cross: + if (getattr(req, "py_skip_cross_kv_projection", False) + or not req.is_first_context_chunk + or not self._kv_connector_should_add_sequence(req)): + return None + encoder_output_len = getattr(req, "encoder_output_len", None) + if encoder_output_len is None: + raise RuntimeError( + "Cross KV cache allocation requires " + f"encoder_output_len for request {req.py_request_id}.") + return int(encoder_output_len) + if is_star_cp: + if req.ctx_iters != 0: + return None + seq_len = sum(len(ctx_block) for ctx_block in req.ctx_blocks) + return seq_len + (len(req.query_id) if self.mapping.cp_rank + == self.mapping.cp_size - 1 else 0) + if not req.is_first_context_chunk or not self._kv_connector_should_add_sequence( + req): + return None + return req.prompt_len + def prepare_resources(self, scheduled_batch: ScheduledRequests): + # Cross/encoder K/V is allocated once and never grows; handle it on a + # dedicated path so the self-attention flow below stays unconditional. + if self.kv_cache_type == CacheTypeCpp.CROSS: + return self._prepare_cross_resources(scheduled_batch) + + is_star_cp = ('cp_type' in self.mapping.cp_config + and CpType.STAR == self.mapping.cp_config['cp_type']) with request_context(self.is_draft, scheduled_batch): # wait for all pending work to finish before launching offload/onboarding/partial copy self.impl.sync_transfer_manager_with_buffer_manager() - # Collect first-chunk requests eligible for add_sequence_batch. + # Collect first-chunk requests eligible for batch add_sequence_batch. # When block reuse is enabled, addSequenceBatch uses a two-phase # claim-then-onboard strategy that prevents host offloading from # evicting reusable blocks in the radix tree. - batch_request_infos = [] - batch_llm_requests = [] - batch_ctx_requests = [] - - # allocate KV Cache - is_star_cp = 'cp_type' in self.mapping.cp_config and CpType.STAR == self.mapping.cp_config[ - 'cp_type'] - - for req in scheduled_batch.context_requests: - req_beam_width = req.py_beam_width - if is_star_cp: - if req.ctx_iters == 0: - seq_len = sum( - len(ctx_block) for ctx_block in req.ctx_blocks) - prompt_len = seq_len + ( - len(req.query_id) if self.mapping.cp_rank - == self.mapping.cp_size - 1 else 0) - batch_request_infos.append( - (req.py_request_id, prompt_len, req_beam_width)) - batch_llm_requests.append(req) - batch_ctx_requests.append(req) - else: - if req.is_first_context_chunk and self._kv_connector_should_add_sequence( - req): - # Batch path: two-phase claim-then-onboard - batch_request_infos.append( - (req.py_request_id, req.prompt_len, req_beam_width)) - batch_llm_requests.append(req) - batch_ctx_requests.append(req) + batch_request_infos, batch_llm_requests = self._collect_context_sequences( + scheduled_batch, is_cross=False, is_star_cp=is_star_cp) if batch_request_infos: self.impl.add_sequence_batch(batch_request_infos, batch_llm_requests) - for req in batch_ctx_requests: + for req in batch_llm_requests: for _ in range(self.num_extra_kv_tokens): self.impl.add_token(req.py_request_id) for _ in range(get_draft_token_length(req)): @@ -750,6 +757,44 @@ def prepare_resources(self, scheduled_batch: ScheduledRequests): self.kv_connector_manager.build_scheduler_output( scheduled_batch, self) + def _collect_context_sequences(self, scheduled_batch: ScheduledRequests, + is_cross: bool, is_star_cp: bool): + """Build the (request_info, llm_request) lists for add_sequence_batch. + + Cross (encoder) sequences are sized from encoder_output_len with a beam + width of 1 (request-scoped); self-attention sequences use the request's + own beam width. + """ + batch_request_infos = [] + batch_llm_requests = [] + for req in scheduled_batch.context_requests: + seq_len = self._context_seq_len(req, is_cross, is_star_cp) + if seq_len is None: + continue + beam_width = 1 if is_cross else req.py_beam_width + batch_request_infos.append((req.py_request_id, seq_len, beam_width)) + batch_llm_requests.append(req) + return batch_request_infos, batch_llm_requests + + def _prepare_cross_resources(self, scheduled_batch: ScheduledRequests): + """Allocate cross (encoder) K/V blocks. + + Encoder K/V is written once at the first decoder context step and read + unchanged on every generation step, so it never grows: this skips the + decode-time token growth, draft-token reserve, and scheduler bookkeeping + that the self-attention path performs. + """ + with request_context(self.is_draft, scheduled_batch): + # wait for all pending work to finish before launching offload/onboarding/partial copy + self.impl.sync_transfer_manager_with_buffer_manager() + batch_request_infos, batch_llm_requests = self._collect_context_sequences( + scheduled_batch, is_cross=True, is_star_cp=False) + if batch_request_infos: + self.impl.add_sequence_batch(batch_request_infos, + batch_llm_requests) + # kernels wait for scheduled offload/onboard/partial copy work before launching + self.impl.refresh_blocks() + def extend_capacity_for_tokens(self, request: LlmRequest) -> None: """No-op for V1; interface kept consistent with V2.""" @@ -893,29 +938,41 @@ def update_resources(self, scheduled_batch: ScheduledRequests, attn_metadata: "AttentionMetadata" = None, kv_cache_dtype_byte_size: float = None): - # Rewind KV cache for requests with rejected draft tokens. - # Skip: - # - GENERATION_COMPLETE: finished requests - # - CONTEXT_INIT: requests whose state was reset after being paused with KV cache freed. - # With overlap scheduler, the scheduler pauses a request and frees KV cache at iteration N, - # while the previous batch (N-1) is still trying to update the KV cache after forward pass. - for request in scheduled_batch.generation_requests: - if request.state in (LlmRequestState.GENERATION_COMPLETE, - LlmRequestState.CONTEXT_INIT): - continue - if request.py_rewind_len > 0: - self.rewind_kv_cache(request, request.py_rewind_len) - # Symmetric companion to prepare_resources's reserve_slack - # add_token loop: when _kv_reserve_draft_tokens (e.g. dynamic - # tree's K*max_draft_len) exceeds the runtime draft length, - # those extra slots must also be rewound, otherwise the draft - # KV cache leaks reserve_slack tokens per generation iteration - # and eventually overflows mCacheBlockIndices. - runtime_draft_len = (request.py_rewind_len + - request.py_num_accepted_draft_tokens) - extra_rewind = self._kv_reserve_draft_tokens - runtime_draft_len - if extra_rewind > 0: - self.rewind_kv_cache(request, extra_rewind) + # Self-attention pools rewind rejected speculative tokens each step; + # cross/encoder K/V is immutable, so only the context-block commit below + # applies to it. + if self.kv_cache_type != CacheTypeCpp.CROSS: + if not self.is_draft: + from .kv_cache_manager_v2 import \ + _update_kv_cache_draft_token_location + + _update_kv_cache_draft_token_location(self, scheduled_batch, + attn_metadata, + kv_cache_dtype_byte_size) + + # Rewind KV cache for requests with rejected draft tokens. + # Skip: + # - GENERATION_COMPLETE: finished requests + # - CONTEXT_INIT: requests whose state was reset after being paused with KV cache freed. + # With overlap scheduler, the scheduler pauses a request and frees KV cache at iteration N, + # while the previous batch (N-1) is still trying to update the KV cache after forward pass. + for request in scheduled_batch.generation_requests: + if request.state in (LlmRequestState.GENERATION_COMPLETE, + LlmRequestState.CONTEXT_INIT): + continue + if request.py_rewind_len > 0: + self.rewind_kv_cache(request, request.py_rewind_len) + # Symmetric companion to prepare_resources's reserve_slack + # add_token loop: when _kv_reserve_draft_tokens (e.g. dynamic + # tree's K*max_draft_len) exceeds the runtime draft length, + # those extra slots must also be rewound, otherwise the draft + # KV cache leaks reserve_slack tokens per generation iteration + # and eventually overflows mCacheBlockIndices. + runtime_draft_len = (request.py_rewind_len + + request.py_num_accepted_draft_tokens) + extra_rewind = self._kv_reserve_draft_tokens - runtime_draft_len + if extra_rewind > 0: + self.rewind_kv_cache(request, extra_rewind) # For context requests, store completed context blocks for KV cache reuse. # We wait until context_remaining_length == 0 (all chunks processed) before @@ -2012,6 +2069,42 @@ def pin_blocks(self, request_id: int): def copy_batch_block_offsets(self, dst_tensor: torch.Tensor, request_ids: List[int], beam_width: int, num_context: int, num_seqs: int): + if self.kv_cache_type == CacheTypeCpp.CROSS and beam_width > 1: + # This branch is reached only via attribute aliasing, never a + # direct cross_kv_cache_manager.copy_batch_block_offsets(...) call: + # AttentionMetadata.create_cross_metadata() sets + # cross_md.kv_cache_manager = cross_kv_cache_manager + # (attention_backend/interface.py), and then + # TrtllmAttentionMetadata.prepare() calls + # self.kv_cache_manager.copy_batch_block_offsets(...) + # (attention_backend/trtllm.py), which dispatches here on the + # cross manager. + num_gen_requests = len(request_ids) - num_context + expected_num_seqs = num_context + num_gen_requests * beam_width + assert num_seqs == expected_num_seqs, ( + f"Cross KV cache block offsets expected {expected_num_seqs} " + f"decoder rows, got {num_seqs}.") + + # Cross KV is request-scoped: all decoder beams read the same + # encoder K/V blocks. Populate one host row per request, then + # expand generation rows across beams in the attention metadata + # tensor whose rows are decoder-sequence scoped. + self.impl.copy_batch_block_offsets(self.host_kv_cache_block_offsets, + request_ids, 1, 0) + for pool_idx in range(self.host_kv_cache_block_offsets.shape[0]): + if num_context > 0: + dst_tensor[pool_idx, :num_context].copy_( + self.host_kv_cache_block_offsets[ + pool_idx, :num_context], + non_blocking=True) + if num_gen_requests > 0: + gen_block_offsets = self.host_kv_cache_block_offsets[ + pool_idx, num_context:num_context + num_gen_requests] + dst_tensor[pool_idx, num_context:num_seqs].copy_( + gen_block_offsets.repeat_interleave(beam_width, dim=0), + non_blocking=True) + return + self.impl.copy_batch_block_offsets(self.host_kv_cache_block_offsets, request_ids[:num_context], 1, 0) self.impl.copy_batch_block_offsets(self.host_kv_cache_block_offsets, diff --git a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler.py b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler.py index 3716f2397635..ace36551f5c0 100644 --- a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler.py +++ b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler.py @@ -52,6 +52,7 @@ def _call_with_optional_summary( SchedulerOutput = namedtuple( "SchedulerOutput", [ + "encoder_requests", "context_requests", "generation_requests", "paused_requests", @@ -61,15 +62,68 @@ def _call_with_optional_summary( ) +def is_decoder_context_request_waiting_for_encoder_output(req: LlmRequest) -> bool: + """Return whether decoder-context scheduling is blocked on encoder output.""" + if not req.is_context_init_state: + return False + + ready_event = getattr(req, "py_encoder_output_ready_event", None) + return ready_event is not None and not ready_event.query() + + +def drop_decoder_context_requests_waiting_for_encoder_output( + active_requests: RequestList, +) -> RequestList: + """Drop ``CONTEXT_INIT`` requests whose encoder output is not ready yet.""" + filtered_requests: RequestList = [] + for req in active_requests: + if is_decoder_context_request_waiting_for_encoder_output(req): + logger.debug( + "Skipping context request %s until encoder output is ready.", + getattr(req, "py_request_id", req.request_id), + ) + continue + + filtered_requests.append(req) + + return filtered_requests + + +def split_encoder_from_decoder_context_requests( + requests: RequestList, +) -> tuple[RequestList, RequestList]: + """Split scheduled encoder-init requests from decoder-context requests.""" + encoder_requests: RequestList = [] + context_requests: RequestList = [] + for req in requests: + if req.is_encoder_init_state: + encoder_requests.append(req) + else: + context_requests.append(req) + return encoder_requests, context_requests + + +def _get_lora_task_id(req: LlmRequest): + # C++ uses std::optional comparison where nullopt < any_value, so + # requests without LoRA (nullopt) should come first. + lora_id = getattr(req, "lora_task_id", None) + if lora_id is None: + return (0, 0) + return (1, lora_id) + + class ScheduledRequests: """Scheduled requests separated into disjoint sets. The reason for the separation is that requests are handled differently in different phases. For example, + - encoder requests run on the encoder stack and never enter decoder forward. - context requests and generation requests execute different attention kernels. - only context requests that are at the last chunk and generation requests sample new tokens. """ + encoder_requests: RequestList + """Requests that are in the encoder phase.""" context_requests_chunking: RequestList """Requests that are in the middle of the context phase.""" context_requests_last_chunk: RequestList @@ -80,6 +134,7 @@ class ScheduledRequests: """Requests that are paused.""" def __init__(self): + self.encoder_requests: RequestList = [] self.context_requests_chunking: RequestList = [] self.context_requests_last_chunk: RequestList = [] self.generation_requests: RequestList = [] @@ -99,6 +154,10 @@ def can_run_cuda_graph(self) -> bool: def batch_size(self) -> int: return self.num_context_requests + len(self.generation_requests) + @property + def num_encoder_requests(self) -> int: + return len(self.encoder_requests) + @property def num_context_requests(self) -> int: return len(self.context_requests_chunking) + len(self.context_requests_last_chunk) @@ -114,6 +173,9 @@ def context_requests(self) -> RequestList: def all_requests(self) -> RequestList: return self.context_requests + self.generation_requests + def append_encoder_request(self, request: LlmRequest) -> None: + self.encoder_requests.append(request) + def append_context_request(self, request: LlmRequest) -> None: if request.is_last_context_chunk: self.context_requests_last_chunk.append(request) @@ -165,6 +227,7 @@ class SerializableSchedulerOutput: Need this class because LlmRequest is not serializable by pickle. """ + encoder_requests: list[int] # request ids of encoder requests context_requests_chunking: list[int] # request ids of context requests chunking context_requests_last_chunk: list[int] # request ids of context requests last chunk generation_requests: list[int] # request ids of generation requests @@ -182,6 +245,7 @@ def from_scheduler_result( num_fitting_requests: int, ) -> "SerializableSchedulerOutput": return cls( + encoder_requests=[req.request_id for req in scheduled_requests.encoder_requests], context_requests_chunking=[ req.request_id for req in scheduled_requests.context_requests_chunking ], @@ -201,6 +265,9 @@ def to_scheduler_result( ) -> tuple[ScheduledRequests, RequestList, int]: id_to_request = {req.request_id: req for req in active_requests} scheduled_requests = ScheduledRequests() + scheduled_requests.encoder_requests = [ + id_to_request[req_id] for req_id in self.encoder_requests + ] scheduled_requests.context_requests_chunking = [ id_to_request[req_id] for req_id in self.context_requests_chunking ] @@ -239,36 +306,55 @@ def __init__( kv_cache_manager, peft_cache_manager: tb_internal.batch_manager.PeftCacheManager | None, scheduler_policy: CapacitySchedulerPolicy = CapacitySchedulerPolicy.GUARANTEED_NO_EVICT, + *, + cross_kv_cache_manager=None, two_step_lookahead: bool = False, + no_schedule_until_state: LlmRequestState = LlmRequestState.CONTEXT_INIT, ): + """C++-bound capacity scheduler wrapper. + + ``cross_kv_cache_manager`` enables encoder-decoder dual-pool + scheduling. When provided, callers should also pass + ``no_schedule_until_state=LlmRequestState.ENCODER_INIT`` so the + scheduler admits requests already in ``ENCODER_INIT`` for the + encoder loop. The C++ ``CapacityScheduler`` already accepts a + cross manager in its ``__call__`` (legacy enc-dec relies on this); + the Python wrapper just widens its signature to expose it. + """ super(BindCapacityScheduler, self).__init__() self.kv_cache_manager = kv_cache_manager self.peft_cache_manager = peft_cache_manager + self.cross_kv_cache_manager = cross_kv_cache_manager self.impl = tb_internal.algorithms.CapacityScheduler( max_num_requests=max_num_requests, capacity_scheduler_policy=scheduler_policy._to_pybind(), has_kv_cache_manager=kv_cache_manager is not None, two_step_lookahead=two_step_lookahead, - no_schedule_until_state=LlmRequestState.CONTEXT_INIT, + no_schedule_until_state=no_schedule_until_state, no_schedule_after_state=LlmRequestState.GENERATION_COMPLETE, ) def schedule_request( self, active_requests: RequestList ) -> tuple[list[LlmRequest], list[LlmRequest], list[LlmRequest]]: - return self.impl(active_requests, self.kv_cache_manager, self.peft_cache_manager) + return self.impl( + active_requests, + self.kv_cache_manager, + self.peft_cache_manager, + self.cross_kv_cache_manager, + ) class MicroBatchScheduler(ABC): @abstractmethod def schedule( self, active_requests: RequestList, inflight_request_ids: set[int] - ) -> tuple[list[LlmRequest], list[LlmRequest]]: + ) -> tuple[list[LlmRequest], list[LlmRequest], list[LlmRequest]]: """ :param active_requests: list of active requests, up to maximum number of sequences :param inflight_request_ids: set of request ids that are inflight (of all micro batches) - :return: (contextRequests, generationRequests) + :return: (encoderRequests, contextRequests, generationRequests) """ # to be aligned with MicroBatchScheduler::scheduleRequests # in cpp/tensorrt_llm/batch_manager/microBatchScheduler.h @@ -296,10 +382,16 @@ def __init__( def schedule( self, active_requests: RequestList, inflight_request_ids: set[int] - ) -> tuple[list[LlmRequest], list[LlmRequest]]: - return self.impl( + ) -> tuple[list[LlmRequest], list[LlmRequest], list[LlmRequest]]: + encoder_or_context_requests, generation_requests = self.impl( active_requests, inflight_request_ids, self.max_batch_size, self.max_num_tokens ) + # Convert from binding type RequestVector to list[LlmRequest], + # so Python fields on LlmRequest won't be stripped away. + encoder_requests, context_requests = split_encoder_from_decoder_context_requests( + list(encoder_or_context_requests) + ) + return encoder_requests, context_requests, list(generation_requests) class SimpleScheduler(RequestScheduler): @@ -313,24 +405,25 @@ def __init__( def schedule_request( self, active_requests: RequestList, inflight_request_ids: set[int] ) -> SchedulerOutput: + active_requests = drop_decoder_context_requests_waiting_for_encoder_output(active_requests) fitting_requests, fitting_disagg_gen_init_requests, paused_requests = ( self.capacity_scheduler.schedule_request(active_requests) ) - context_requests, generation_requests = self.micro_batch_scheduler.schedule( - fitting_requests, inflight_request_ids + encoder_requests, context_requests, generation_requests = ( + self.micro_batch_scheduler.schedule(fitting_requests, inflight_request_ids) ) - # Convert from binding type RequestVector to list[LlmRequest], - # so Python fields on LlmRequest won't be stripped away return SchedulerOutput( - list(context_requests), - list(generation_requests), + encoder_requests, + context_requests, + generation_requests, list(paused_requests), list(fitting_disagg_gen_init_requests), len(fitting_requests), ) def can_schedule(self, requests: RequestList) -> bool: + requests = drop_decoder_context_requests_waiting_for_encoder_output(requests) fitting_requests, _, _ = self.capacity_scheduler.schedule_request(requests) return len(fitting_requests) == len(requests) @@ -395,6 +488,8 @@ def _can_be_scheduled(self, req: LlmRequest) -> bool: C++ reference: microBatchScheduler.cpp line 192-195 Optimized: use state_value property to avoid enum object creation """ + if is_decoder_context_request_waiting_for_encoder_output(req): + return False # Use state_value property (returns int directly, avoids enum object creation) state_value = req.state_value # Inline comparison: must have reached until_state but not after_state @@ -405,7 +500,8 @@ def _can_be_scheduled(self, req: LlmRequest) -> bool: def schedule( self, active_requests: RequestList, inflight_request_ids: set[int] - ) -> tuple[RequestList, RequestList]: + ) -> tuple[RequestList, RequestList, RequestList]: + encoder_requests: RequestList = [] context_requests: RequestList = [] generation_requests: RequestList = [] @@ -431,6 +527,9 @@ def schedule( if req.request_id in inflight_request_ids: continue + if is_decoder_context_request_waiting_for_encoder_output(req): + continue + # Skip if request cannot be scheduled yet or should no longer be scheduled, # manually inline the condition to reuse req.state_value if not ( @@ -455,7 +554,7 @@ def schedule( break logger.debug(f"encoder request scheduled: ID {req.request_id}") - context_requests.append(req) + encoder_requests.append(req) batch_num_tokens += req_num_tokens # --- B. Context Request Handling --- @@ -587,19 +686,20 @@ def schedule( # Sort requests for consistency with C++ # C++ reference: utils::sortRequests in inflightBatchingUtils.cpp + encoder_requests.sort(key=_get_lora_task_id) self._sort_requests(context_requests, generation_requests, not all_context_requests_fit) # Summary logs logger.debug( f"batchSize (num ctx/enc requests + num gen requests): " - f"{len(context_requests) + len(generation_requests)}" + f"{len(encoder_requests) + len(context_requests) + len(generation_requests)}" ) logger.debug( f"batchNumTokens (num ctx/enc input tokens + num gen input tokens) " f"/ maxNumTokens: {batch_num_tokens} / {max_num_tokens or 0}" ) - return context_requests, generation_requests + return encoder_requests, context_requests, generation_requests def _sort_requests( self, context_requests: RequestList, generation_requests: RequestList, chunks_present: bool @@ -613,29 +713,21 @@ def _sort_requests( 2. Sort all requests by lora task id for performance. """ - def get_lora_task_id(req: LlmRequest): - # C++ uses std::optional comparison where nullopt < any_value - # So requests without LoRA (nullopt) should come first - lora_id = getattr(req, "lora_task_id", None) - if lora_id is None: - return (0, 0) # (has_value=False, value=0) - comes first - return (1, lora_id) # (has_value=True, value) - sorted by value - if chunks_present: # Partition: non-last-chunk first, last-chunk at end not_last_chunk = [r for r in context_requests if not r.is_last_context_chunk] last_chunk = [r for r in context_requests if r.is_last_context_chunk] # Sort each group by lora_task_id - not_last_chunk.sort(key=get_lora_task_id) - last_chunk.sort(key=get_lora_task_id) + not_last_chunk.sort(key=_get_lora_task_id) + last_chunk.sort(key=_get_lora_task_id) # Rebuild the list in-place context_requests.clear() context_requests.extend(not_last_chunk) context_requests.extend(last_chunk) else: - context_requests.sort(key=get_lora_task_id) + context_requests.sort(key=_get_lora_task_id) - generation_requests.sort(key=get_lora_task_id) + generation_requests.sort(key=_get_lora_task_id) def _set_ctx_requests_chunk_size( self, @@ -901,6 +993,17 @@ class GuaranteedNoEvictPolicy(SchedulerPolicyBase): """ GuaranteedNoEvictScheduler: Reserve blocks for requests to complete without eviction. C++ reference: capacityScheduler.cpp:194-331 + + Encoder-decoder support: when ``cross_kv_cache_manager`` is configured + on the parent scheduler and ``no_schedule_until_state=ENCODER_INIT``, + encoder-init requests are considered in the same *scheduler pass* as + context/generation requests. This does not mean encoder and decoder + context execute in the same model iteration: encoder admission only + admits encoder compute and leaves both KV pools untouched until the + request transitions to ``CONTEXT_INIT`` on a later decoder-context + iteration. The in-pass classification preserves the legacy invariant + that encoder and decoder context never collide on the same self pool + budget. """ def __init__(self, static_batch: bool = False): @@ -939,7 +1042,10 @@ def schedule( pending_requests: RequestList = [] pending_dis_gen_init_requests: RequestList = [] - # First pass: process in-progress generation and classify requests + # First pass: process in-progress generation and classify requests. + # Encoder-init and context-init both fall into ``pending_requests`` + # and are budgeted in the second pass; they use distinct pool + # reservation rules but share ordering. for req in active_requests: if not scheduler._can_be_scheduled_with_disagg_exception(req): continue @@ -973,6 +1079,12 @@ def schedule( for requests in [pending_dis_gen_init_requests, pending_requests]: for req in requests: + if req.is_encoder_init_state and reserved_cross_blocks is None: + raise RuntimeError( + f"Encoder-init request {req.request_id} requires " + "a cross_kv_cache_manager." + ) + if ( not self.static_batch and skipping_is_relevant @@ -1000,7 +1112,26 @@ def schedule( cached_summary = summary_by_req.get(req_id) cached_cross_summary = cross_summary_by_req.get(req_id) - if req.is_context_init_state or req.is_disagg_generation_init_state: + if req.is_encoder_init_state: + # Encoder admission only admits encoder compute. + # KV block budgeting happens when the request is + # scheduled as decoder CONTEXT_INIT. Without a cross + # manager, the later decoder context cannot satisfy + # the dual-pool contract, so fail before running + # encoder work. + if has_peft: + lora_task_id, is_new_task, needed_peft_pages = ( + scheduler._get_peft_task_info(req, uniq_task_ids) + ) + if needed_peft_pages > available_peft_pages: + continue + available_peft_pages -= needed_peft_pages + if is_new_task: + uniq_task_ids.add(lora_task_id) + + scheduled_requests.append(req) + + elif req.is_context_init_state or req.is_disagg_generation_init_state: enough_blocks = reserved_blocks.enough_available_blocks( req, cached_summary=cached_summary ) @@ -1040,6 +1171,15 @@ class MaxUtilizationPolicy(SchedulerPolicyBase): """ MaxUtilizationScheduler: Maximize utilization, may pause started requests. C++ reference: capacityScheduler.cpp:341-425 + + Encoder-decoder support: encoder-init requests are considered in the + same *scheduler pass* as context/generation requests when + ``no_schedule_until_state=ENCODER_INIT`` and a + ``cross_kv_cache_manager`` is configured. Encoder admission only + schedules encoder compute; self- and cross-pool budgeting happens + when the request transitions to ``CONTEXT_INIT`` on a later + decoder-context iteration. Encoder requests are not eligible eviction + victims (they have no started KV blocks to free). """ def schedule( @@ -1052,6 +1192,12 @@ def schedule( scheduled_blocks_manager = MaxUtilizationScheduledBlocksManager( scheduler.kv_cache_manager, scheduler.two_step_lookahead ) + scheduled_cross_blocks_manager: Optional[MaxUtilizationScheduledBlocksManager] = None + if scheduler.cross_kv_cache_manager is not None: + scheduler.cross_kv_cache_manager.start_scheduling() + scheduled_cross_blocks_manager = MaxUtilizationScheduledBlocksManager( + scheduler.cross_kv_cache_manager, scheduler.two_step_lookahead + ) num_scheduled_peft_pages = 0 seen_task_ids: set[int] = set() @@ -1065,6 +1211,8 @@ def schedule( def is_started_request(req: LlmRequest) -> bool: if not scheduler._can_be_scheduled(req): return False + # Encoder-init requests have not allocated any self-pool blocks + # yet, so they are never started in the eviction sense. return ( req.is_context_init_state and not req.is_first_context_chunk ) or req.is_generation_in_progress_state @@ -1088,6 +1236,11 @@ def is_started_request(req: LlmRequest) -> bool: req_it += 1 continue + if req.is_encoder_init_state and scheduled_cross_blocks_manager is None: + raise RuntimeError( + f"Encoder-init request {req.request_id} requires a cross_kv_cache_manager." + ) + if skipping_is_relevant and scheduler._beneficial_to_skip( req, newly_contributed_context_blocks, @@ -1102,6 +1255,7 @@ def is_started_request(req: LlmRequest) -> bool: req, scheduled_requests, scheduled_blocks_manager, + scheduled_cross_blocks_manager, num_scheduled_peft_pages, seen_task_ids, cached_summary=summary_by_req.get(req.py_request_id), @@ -1120,6 +1274,10 @@ def is_started_request(req: LlmRequest) -> bool: if last_started_idx is not None: paused_req = requests_list[last_started_idx] scheduler.kv_cache_manager.scheduling_remove_sequence(paused_req.py_request_id) + if scheduler.cross_kv_cache_manager is not None: + scheduler.cross_kv_cache_manager.scheduling_remove_sequence( + paused_req.py_request_id + ) paused_requests.append(paused_req) logger.debug( f"MaxUtilizationScheduler: request ID {paused_req.request_id} -> pause" @@ -1136,6 +1294,7 @@ def _try_scheduling_request( req: LlmRequest, scheduled_requests: RequestList, scheduled_blocks_manager: "MaxUtilizationScheduledBlocksManager", + scheduled_cross_blocks_manager: Optional["MaxUtilizationScheduledBlocksManager"], num_scheduled_peft_pages: int, seen_task_ids: set[int], cached_summary: Optional[PrefixReuseSummary] = None, @@ -1143,11 +1302,25 @@ def _try_scheduling_request( if len(scheduled_requests) >= scheduler.max_num_requests: return False, num_scheduled_peft_pages - blocks_if_scheduled = scheduled_blocks_manager.prepare_blocks_if_schedulable( - req, cached_summary=cached_summary - ) - if blocks_if_scheduled is None: - return False, num_scheduled_peft_pages + # Encoder-init: no KV blocks are needed until the later decoder + # context admission. Still require the cross manager so a + # misconfigured enc-dec runtime fails before running encoder work. + if req.is_encoder_init_state: + blocks_if_scheduled = None + cross_blocks_if_scheduled = None + else: + blocks_if_scheduled = scheduled_blocks_manager.prepare_blocks_if_schedulable( + req, cached_summary=cached_summary + ) + if blocks_if_scheduled is None: + return False, num_scheduled_peft_pages + cross_blocks_if_scheduled: Optional[dict] = None + if scheduled_cross_blocks_manager is not None: + cross_blocks_if_scheduled = ( + scheduled_cross_blocks_manager.prepare_blocks_if_schedulable(req) + ) + if cross_blocks_if_scheduled is None: + return False, num_scheduled_peft_pages # PEFT check only when needed if scheduler.peft_cache_manager is not None: @@ -1168,7 +1341,10 @@ def _try_scheduling_request( if is_new_task: seen_task_ids.add(lora_task_id) - scheduled_blocks_manager.update_scheduled_blocks(blocks_if_scheduled) + if blocks_if_scheduled is not None: + scheduled_blocks_manager.update_scheduled_blocks(blocks_if_scheduled) + if scheduled_cross_blocks_manager is not None and cross_blocks_if_scheduled is not None: + scheduled_cross_blocks_manager.update_scheduled_blocks(cross_blocks_if_scheduled) scheduled_requests.append(req) return True, num_scheduled_peft_pages @@ -1372,6 +1548,8 @@ def _can_be_scheduled(self, req: LlmRequest) -> bool: but has not yet reached no_schedule_after_state. Optimized: use state_value property to avoid enum object creation """ + if is_decoder_context_request_waiting_for_encoder_output(req): + return False # Use state_value property (returns int directly, avoids enum object creation) state_value = req.state_value # Inline comparison: must have reached until_state but not after_state @@ -1594,6 +1772,7 @@ def __init__( cross_kv_cache_manager=None, two_step_lookahead: bool = False, scheduler_capacity: Optional[int] = None, + no_schedule_until_state: LlmRequestState = LlmRequestState.CONTEXT_INIT, ): # Use scheduler_capacity if provided, otherwise fall back to max_batch_size # scheduler_capacity may differ from max_batch_size (e.g., adjusted for attention_dp + disagg) @@ -1608,6 +1787,7 @@ def __init__( scheduler_policy=scheduler_policy, cross_kv_cache_manager=cross_kv_cache_manager, two_step_lookahead=two_step_lookahead, + no_schedule_until_state=no_schedule_until_state, ) # 2. Initialize Python MicroBatch Scheduler @@ -1631,22 +1811,25 @@ def __init__( max_batch_size=max_batch_size, max_num_tokens=max_num_tokens, ctx_chunk_config=py_chunk_config, + no_schedule_until_state=no_schedule_until_state, ) def schedule_request( self, active_requests: RequestList, inflight_request_ids: set[int] ) -> SchedulerOutput: + active_requests = drop_decoder_context_requests_waiting_for_encoder_output(active_requests) # Step 1: Capacity Check (Who fits in memory?) fitting_requests, fitting_disagg_gen_init, paused_requests = ( self.capacity_scheduler.schedule_request(active_requests) ) # Step 2: MicroBatch Check (Who fits in token budget? + Chunking) - context_requests, generation_requests = self.micro_batch_scheduler.schedule( - fitting_requests, inflight_request_ids + encoder_requests, context_requests, generation_requests = ( + self.micro_batch_scheduler.schedule(fitting_requests, inflight_request_ids) ) return SchedulerOutput( + encoder_requests=encoder_requests, context_requests=context_requests, generation_requests=generation_requests, paused_requests=paused_requests, @@ -1655,6 +1838,7 @@ def schedule_request( ) def can_schedule(self, requests: RequestList) -> bool: + requests = drop_decoder_context_requests_waiting_for_encoder_output(requests) # Dry run capacity check fitting, _, _ = self.capacity_scheduler.schedule_request(requests) return len(fitting) == len(requests) diff --git a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py index d00e7b3e1444..1575021bab5f 100644 --- a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py @@ -20,7 +20,13 @@ from tensorrt_llm.logger import logger from ..llm_request import LlmRequest, LlmRequestState, get_draft_token_length -from .scheduler import RequestList, RequestScheduler, SchedulerOutput +from .scheduler import ( + RequestList, + RequestScheduler, + SchedulerOutput, + _get_lora_task_id, + drop_decoder_context_requests_waiting_for_encoder_output, +) class ScheduleAction(enum.Enum): @@ -143,6 +149,7 @@ def __init__( no_schedule_until_state: LlmRequestState = LlmRequestState.CONTEXT_INIT, no_schedule_after_state: LlmRequestState = LlmRequestState.GENERATION_TO_COMPLETE, draft_kv_cache_manager=None, # KVCacheManagerV2 for MTP draft layers + cross_kv_cache_manager=None, # KVCacheManagerV2 for enc-dec cross-attn ): self.max_num_tokens = max_num_tokens self.max_num_requests = ( @@ -155,6 +162,7 @@ def __init__( ) self.kv_cache_manager = kv_cache_manager self.draft_kv_cache_manager = draft_kv_cache_manager + self.cross_kv_cache_manager = cross_kv_cache_manager if scheduler_policy != CapacitySchedulerPolicy.MAX_UTILIZATION: logger.warning( "KVCacheV2Scheduler only supports MAX_UTILIZATION for now, " @@ -172,10 +180,13 @@ def __init__( draft_mgr_name = ( type(draft_kv_cache_manager).__name__ if draft_kv_cache_manager is not None else "None" ) + cross_mgr_name = ( + type(cross_kv_cache_manager).__name__ if cross_kv_cache_manager is not None else "None" + ) logger.info( f"KVCacheV2Scheduler: tokens_per_block={self.tokens_per_block}, " f"max_num_tokens={max_num_tokens}, max_batch_size={max_batch_size}, " - f"draft_mgr={draft_mgr_name}" + f"draft_mgr={draft_mgr_name}, cross_mgr={cross_mgr_name}" ) if ctx_chunk_config is not None: self.chunking_enabled = True @@ -196,26 +207,35 @@ def __init__( def schedule_request( self, active_requests: RequestList, inflight_request_ids: set[int] ) -> SchedulerOutput: + active_requests = drop_decoder_context_requests_waiting_for_encoder_output(active_requests) # Main scheduling loop - (scheduled_ctx, scheduled_gen, evicted, disagg_candidates, has_chunking) = ( - self._schedule_loop(active_requests, inflight_request_ids) - ) + ( + scheduled_encoder, + scheduled_ctx, + scheduled_gen, + evicted, + disagg_candidates, + has_chunking, + ) = self._schedule_loop(active_requests, inflight_request_ids) # Sort by LoRA task ID + scheduled_encoder.sort(key=_get_lora_task_id) self._sort_requests(scheduled_ctx, scheduled_gen, has_chunking) return SchedulerOutput( + encoder_requests=scheduled_encoder, context_requests=scheduled_ctx, generation_requests=scheduled_gen, paused_requests=evicted, fitting_disagg_gen_init_requests=disagg_candidates, - num_fitting_requests=len(scheduled_ctx) + len(scheduled_gen), + num_fitting_requests=(len(scheduled_encoder) + len(scheduled_ctx) + len(scheduled_gen)), ) # ---- Main scheduling loop ---- def _schedule_loop(self, active_requests, inflight_request_ids): scheduled_ctx: RequestList = [] + scheduled_encoder: RequestList = [] scheduled_gen: RequestList = [] evicted: RequestList = [] disagg_candidates: RequestList = [] @@ -362,7 +382,7 @@ def _schedule_loop(self, active_requests, inflight_request_ids): action, tokens = self._try_schedule_encoder(req, budget) if action is ScheduleAction.STOP: break - scheduled_ctx.append(req) + scheduled_encoder.append(req) budget.commit(req, tokens, peft_pages) else: action, tokens, chunking_flag = self._try_schedule_context(req, budget) @@ -397,7 +417,14 @@ def _schedule_loop(self, active_requests, inflight_request_ids): f"kv_cache_config.max_tokens." ) - return scheduled_ctx, scheduled_gen, evicted, disagg_candidates, has_chunking + return ( + scheduled_encoder, + scheduled_ctx, + scheduled_gen, + evicted, + disagg_candidates, + has_chunking, + ) # ---- Per-type scheduling methods ---- @@ -406,20 +433,28 @@ def _try_schedule_encoder( ) -> tuple[ScheduleAction, int]: """Try to schedule an encoder request. + Encoder admission does not need KV blocks for decoder work. Decoder + context is the first step that reserves self- and cross-KV blocks and + writes K/V projections into the cross cache. + Returns ``(action, tokens)`` where *tokens* is meaningful only when *action* is ``SCHEDULED``. """ + # Encoder-decoder runtime requires a cross pool for the later decoder + # context step. If the runtime did not plumb one through, surface this + # loudly rather than silently routing to the self pool, which would + # corrupt the dual-pool contract. + if self.cross_kv_cache_manager is None: + raise RuntimeError( + f"Encoder-init request {req.py_request_id} requires a cross_kv_cache_manager." + ) + req_tokens = req.encoder_output_len if not budget.can_fit_tokens(req_tokens): return ScheduleAction.STOP, 0 assert self.max_context_length is None or req_tokens <= self.max_context_length, ( f"The number of encoder tokens ({req_tokens}) exceeds the limit value ({self.max_context_length})" ) - if not self.kv_cache_manager.prepare_context(req): - logger.debug(f"prepare_context failed for encoder request {req.py_request_id}") - return ScheduleAction.STOP, 0 - if not self.kv_cache_manager.resize_context(req, req_tokens): - return ScheduleAction.STOP, 0 return ScheduleAction.SCHEDULED, req_tokens def _try_schedule_disagg_gen_init( @@ -488,6 +523,11 @@ def _try_schedule_context_full( if not self.kv_cache_manager.resize_context(req, req_tokens): return ScheduleAction.SKIP, 0, False + cross_action = self._try_schedule_cross_context(req) + if cross_action is not ScheduleAction.SCHEDULED: + self._suspend_request(req) + return cross_action, 0, False + return ScheduleAction.SCHEDULED, req_tokens, False def _try_schedule_context_chunked( @@ -558,6 +598,12 @@ def _try_schedule_context_chunked( # draft tokens for last chunk. if not self.kv_cache_manager.resize_context(req, chunk_tokens): return ScheduleAction.SKIP, 0, False + + cross_action = self._try_schedule_cross_context(req) + if cross_action is not ScheduleAction.SCHEDULED: + self._suspend_request(req) + return cross_action, 0, False + chunking_flag = req.context_chunk_size < req.context_remaining_length return ScheduleAction.SCHEDULED, chunk_tokens, chunking_flag @@ -725,6 +771,98 @@ def _align_chunk_to_mm_block( return down_block_start - lo + @staticmethod + def _get_optional_encoder_output_len(req: LlmRequest) -> Optional[int]: + get_encoder_output_len = getattr(type(req), "try_get_encoder_output_len", None) + encoder_output_len = ( + get_encoder_output_len(req) + if get_encoder_output_len is not None + else getattr(req, "encoder_output_len", None) + ) + if encoder_output_len is None: + return None + return int(encoder_output_len) + + @classmethod + def _needs_cross_context_allocation(cls, req: LlmRequest) -> bool: + """Return whether decoder context must reserve cross-KV for *req*.""" + if cls._get_optional_encoder_output_len(req) is None: + return False + skip_projection = getattr(req, "py_skip_cross_kv_projection", False) + return not (isinstance(skip_projection, bool) and skip_projection) + + def _try_schedule_cross_context(self, req: LlmRequest) -> ScheduleAction: + """Reserve cross-KV blocks for the first decoder context step.""" + if not self._needs_cross_context_allocation(req): + return ScheduleAction.SCHEDULED + + if self.cross_kv_cache_manager is None: + logger.warning( + "Decoder context request %s requires cross-KV cache but " + "no cross_kv_cache_manager is configured. Skipping.", + req.py_request_id, + ) + return ScheduleAction.STOP + + req_tokens = self._get_optional_encoder_output_len(req) + if req_tokens is None: + return ScheduleAction.SCHEDULED + from ..kv_cache_manager_v2 import KVCacheManagerV2 + + if isinstance(self.cross_kv_cache_manager, KVCacheManagerV2): + if not self._try_schedule_cross_context_v2( + self.cross_kv_cache_manager, req, req_tokens + ): + return ScheduleAction.SKIP + return ScheduleAction.SCHEDULED + + if not self.cross_kv_cache_manager.prepare_context(req): + logger.debug( + "cross prepare_context failed for decoder context request %s", + req.py_request_id, + ) + return ScheduleAction.SKIP + if not self.cross_kv_cache_manager.resize_context(req, req_tokens): + return ScheduleAction.SKIP + return ScheduleAction.SCHEDULED + + @staticmethod + def _try_schedule_cross_context_v2( + cross_kv_cache_manager, req: LlmRequest, req_tokens: int + ) -> bool: + """Reserve V2 cross-KV without mutating decoder context position.""" + kv_cache = cross_kv_cache_manager.kv_cache_map.get(req.py_request_id) + if kv_cache is None: + if not req.is_first_context_chunk: + logger.debug( + "cross KV cache missing for non-first context chunk, request %s", + req.py_request_id, + ) + return False + input_tokens = ( + req.get_encoder_unique_tokens() + if cross_kv_cache_manager.enable_block_reuse + else None + ) + kv_cache = cross_kv_cache_manager._create_kv_cache( + req.py_request_id, req.lora_task_id, input_tokens + ) + kv_cache.cuda_stream = cross_kv_cache_manager._stream.cuda_stream + + if not cross_kv_cache_manager.enable_block_reuse: + kv_cache.stop_committing() + + if not cross_kv_cache_manager._resume_and_restore(req.py_request_id, kv_cache): + return False + + target_capacity = req_tokens + cross_kv_cache_manager.num_extra_kv_tokens + if not kv_cache.resize(max(kv_cache.capacity, target_capacity)): + if req.is_first_context_chunk: + kv_cache.suspend() + return False + + return True + def _try_schedule_generation( self, req: LlmRequest, @@ -856,10 +994,7 @@ def _try_evict_for_gen(self, req, requests_list, req_it, req_it_end, evicted): @staticmethod def _lora_key(req: LlmRequest): - lora_id = getattr(req, "lora_task_id", None) - if lora_id is None: - return (0, 0) - return (1, lora_id) + return _get_lora_task_id(req) def _sort_requests(self, context_requests, generation_requests, has_chunks): """Sort by LoRA task ID. Non-last chunks before last chunks.""" diff --git a/tensorrt_llm/executor/base_worker.py b/tensorrt_llm/executor/base_worker.py index 090742b38d19..35ed243d4c16 100644 --- a/tensorrt_llm/executor/base_worker.py +++ b/tensorrt_llm/executor/base_worker.py @@ -593,6 +593,7 @@ def _deduce_max_tokens(request: GenerationRequest, request.sampling_params.logits_processor, kv_cache_retention_config=request.kv_cache_retention_config, context_phase_params=context_phase_params, + encoder_input_token_ids=request.encoder_input_token_ids, type=request_type, disagg_request_id=disagg_request_id, cache_salt=request.cache_salt, diff --git a/tensorrt_llm/executor/executor.py b/tensorrt_llm/executor/executor.py index 5ea904531f22..fdd3dff44a48 100644 --- a/tensorrt_llm/executor/executor.py +++ b/tensorrt_llm/executor/executor.py @@ -136,6 +136,8 @@ def generate_async( scheduling_params: Optional[SchedulingParams] = None, cache_salt: Optional[str] = None, arrival_time: Optional[float] = None, + encoder_input_token_ids: Optional[Union[torch.Tensor, np.ndarray, + list]] = None, priority: float = DEFAULT_REQUEST_PRIORITY, ) -> GenerationResult: """Generate output for the given prompt token ids in the asynchronous mode. @@ -164,6 +166,7 @@ def generate_async( scheduling_params=scheduling_params, cache_salt=cache_salt, arrival_time=arrival_time, + encoder_input_token_ids=encoder_input_token_ids, priority=priority) result = self.submit(request) # release memory in time diff --git a/tensorrt_llm/executor/request.py b/tensorrt_llm/executor/request.py index 43ea11706e54..fdd6ba0627ce 100644 --- a/tensorrt_llm/executor/request.py +++ b/tensorrt_llm/executor/request.py @@ -109,6 +109,8 @@ def __init__( scheduling_params: Optional[SchedulingParams] = None, cache_salt: Optional[str] = None, arrival_time: Optional[float] = None, + encoder_input_token_ids: Optional[Union[torch.Tensor, np.ndarray, + list]] = None, priority: float = DEFAULT_REQUEST_PRIORITY, ): if isinstance(prompt_token_ids, list): @@ -152,11 +154,28 @@ def __init__( f"({self.MAX_CACHE_SALT_LEN}).") self.cache_salt = cache_salt self.arrival_time = arrival_time + self.encoder_input_token_ids = self._normalize_optional_token_ids( + encoder_input_token_ids, "encoder_input_token_ids") if not (0.0 <= priority <= 1.0): raise ValueError( f"priority must be a float in [0.0, 1.0], got {priority}") self.priority = priority + @staticmethod + def _normalize_optional_token_ids(token_ids: Optional[Union[torch.Tensor, + np.ndarray, + list]], + name: str) -> Optional[list]: + if token_ids is None: + return None + if isinstance(token_ids, list): + return token_ids + if isinstance(token_ids, (torch.Tensor, np.ndarray)): + return token_ids.tolist() + raise TypeError( + f"{name} ({token_ids}) should be an instance of torch.Tensor, np.ndarray or list" + ) + def set_id(self, id): assert self.id is None, f"Request ID is already set: {self.id}" self.id = id diff --git a/tensorrt_llm/llmapi/llm.py b/tensorrt_llm/llmapi/llm.py index d63e37f38193..2659d46449e2 100644 --- a/tensorrt_llm/llmapi/llm.py +++ b/tensorrt_llm/llmapi/llm.py @@ -39,7 +39,7 @@ create_input_processor_with_hash, maybe_compute_mm_embed_cumsum, prompt_inputs) from ..logger import logger -from ..sampling_params import SamplingParams +from ..sampling_params import LogitsProcessor, SamplingParams from ..scheduling_params import SchedulingParams from .llm_args import (TORCH_LLMARGS_EXPLICIT_DOCSTRING, TRT_LLMARGS_EXPLICIT_DOCSTRING, PeftCacheConfig, @@ -119,6 +119,82 @@ class EncoderOutput: prompt: Optional[str] = None +class _BartForcedTokensLogitsProcessor(LogitsProcessor): + """Apply BART forced BOS/EOS tokens from Hugging Face generation config.""" + + _DECODER_PROMPT_LEN = 1 + + def __init__( + self, + *, + forced_bos_token_id: Optional[int], + forced_eos_token_id: Optional[int], + max_tokens: int, + ) -> None: + self.forced_bos_token_id = forced_bos_token_id + self.forced_eos_token_id = forced_eos_token_id + self.max_tokens = max_tokens + + def __call__( + self, + req_id: int, + logits: torch.Tensor, + token_ids: List[List[int]], + stream_ptr: Optional[int], + client_id: Optional[int], + ) -> None: + del req_id, client_id + if stream_ptr is None: + self._apply(token_ids, logits) + return + with torch.cuda.stream(torch.cuda.ExternalStream(stream_ptr)): + self._apply(token_ids, logits) + + def _apply(self, token_ids: List[List[int]], logits: torch.Tensor) -> None: + for beam_idx, beam_token_ids in enumerate(token_ids): + forced_token_id = self._forced_token_id(beam_token_ids) + if forced_token_id is not None: + self._force_token(logits, beam_idx, len(token_ids), + forced_token_id) + + def _forced_token_id(self, token_ids: List[int]) -> Optional[int]: + generated_len = max(len(token_ids) - self._DECODER_PROMPT_LEN, 0) + if generated_len == 0: + return self.forced_bos_token_id + if (self.max_tokens > 0 and generated_len == self.max_tokens - 1): + return self.forced_eos_token_id + return None + + @staticmethod + def _force_token(logits: torch.Tensor, beam_idx: int, beam_count: int, + token_id: int) -> None: + if token_id < 0 or token_id >= logits.shape[-1]: + raise ValueError( + f"Forced BART token id {token_id} is outside the logits " + f"vocabulary dimension {logits.shape[-1]}") + + target = logits + if logits.dim() > 1 and logits.shape[0] == beam_count: + target = logits[beam_idx] + target[:] = float("-inf") + target[..., token_id] = 0 + + +def _contains_bart_forced_tokens_logits_processor(processor: Any) -> bool: + if isinstance(processor, _BartForcedTokensLogitsProcessor): + return True + if isinstance(processor, list): + return any( + _contains_bart_forced_tokens_logits_processor(item) + for item in processor) + processors = getattr(processor, "processors", None) + if isinstance(processors, list): + return any( + _contains_bart_forced_tokens_logits_processor(item) + for item in processors) + return False + + TRT_LLM_DOCSTRING = TRT_LLMARGS_EXPLICIT_DOCSTRING + """ Attributes: @@ -147,6 +223,7 @@ class PreprocessedInputs: prompt_token_ids: List[int] query_token_ids: Optional[List[int]] = None multimodal_params: Optional[MultimodalParams] = None + encoder_input_token_ids: Optional[List[int]] = None class BaseLLM: @@ -326,6 +403,76 @@ def disaggregated_params(self) -> dict: ) if self._executor else {} return self._disaggregated_params + @staticmethod + def _is_token_id_list(value: Any) -> bool: + return isinstance(value, list) and all( + isinstance(token, int) for token in value) + + @classmethod + def _is_unbatched_inputs(cls, inputs: Any) -> bool: + if inputs is None: + return False + if isinstance(inputs, str) or isinstance(inputs, dict): + return True + if cls._is_token_id_list(inputs): + return True + return False + + @classmethod + def _is_unbatched_optional_inputs(cls, *values: Any) -> bool: + for value in values: + if value is None: + continue + return cls._is_unbatched_inputs(value) + return True + + @classmethod + def _item_at(cls, + maybe_batched: Any, + pos: int, + *, + token_ids_are_scalar: bool = False) -> Any: + if maybe_batched is None: + return None + if token_ids_are_scalar and cls._is_token_id_list(maybe_batched): + return maybe_batched + if isinstance(maybe_batched, list): + return maybe_batched[pos] + return maybe_batched + + @staticmethod + def _copy_prompt_inputs(inputs: PromptInputs) -> PromptInputs: + if isinstance(inputs, dict): + return dict(inputs) + return inputs + + @classmethod + def _normalize_token_ids(cls, token_ids: Any, name: str) -> List[int]: + if cls._is_token_id_list(token_ids): + return list(token_ids) + if hasattr(token_ids, "tolist"): + normalized = token_ids.tolist() + if cls._is_token_id_list(normalized): + return normalized + raise TypeError(f"{name} must be a list of token ids.") + + def _is_encoder_decoder_model(self) -> bool: + return bool(getattr(self._hf_model_config, "is_encoder_decoder", False)) + + def _get_decoder_start_token_id(self) -> int: + configs = [ + self._generation_config, + self._hf_model_config, + getattr(self._hf_model_config, "text_config", None), + ] + for attr_name in ("decoder_start_token_id", "bos_token_id"): + for config in configs: + token_id = getattr(config, attr_name, None) + if token_id is not None: + return int(token_id) + raise ValueError( + "decoder_start_token_id is required for encoder-decoder models.") + def generate( self, inputs: Union[PromptInputs, Sequence[PromptInputs]], @@ -370,21 +517,24 @@ def generate( Returns: Union[tensorrt_llm.llmapi.RequestOutput, List[tensorrt_llm.llmapi.RequestOutput]]: The output data of the completion request to the LLM. """ - unbatched = not isinstance(inputs, list) - if not unbatched: + unbatched = self._is_unbatched_optional_inputs(inputs) + if inputs is not None and not unbatched: if isinstance(inputs[0], int): unbatched = True - if unbatched: + if unbatched and inputs is not None: inputs = [inputs] - inputs = [prompt_inputs(i) for i in inputs] + if inputs is None: + request_inputs_list = [None] + else: + request_inputs_list = [prompt_inputs(i) for i in inputs] if isinstance(priority, list): - if len(priority) != len(inputs): + if len(priority) != len(request_inputs_list): raise ValueError( f"priority list length ({len(priority)}) does not match " - f"number of prompts ({len(inputs)})") + f"number of prompts ({len(request_inputs_list)})") for p in priority: if not (0.0 <= p <= 1.0): raise ValueError( @@ -394,25 +544,19 @@ def generate( raise ValueError( f"priority must be a float in [0.0, 1.0], got {priority}") - def _item_at(maybe_batched: Union[Any, Sequence[Any]], pos: int) -> Any: - if isinstance(maybe_batched, list): - return maybe_batched[pos] - else: - return maybe_batched - futures = [] - for i, request_inputs in enumerate(inputs): + for i, request_input in enumerate(request_inputs_list): future = self.generate_async( - request_inputs, - sampling_params=_item_at(sampling_params, i), - lora_request=_item_at(lora_request, i), - prompt_adapter_request=_item_at(prompt_adapter_request, i), - kv_cache_retention_config=_item_at(kv_cache_retention_config, - i), - disaggregated_params=_item_at(disaggregated_params, i), - scheduling_params=_item_at(scheduling_params, i), - cache_salt=_item_at(cache_salt, i), - priority=_item_at(priority, i), + request_input, + sampling_params=self._item_at(sampling_params, i), + lora_request=self._item_at(lora_request, i), + prompt_adapter_request=self._item_at(prompt_adapter_request, i), + kv_cache_retention_config=self._item_at( + kv_cache_retention_config, i), + disaggregated_params=self._item_at(disaggregated_params, i), + scheduling_params=self._item_at(scheduling_params, i), + cache_salt=self._item_at(cache_salt, i), + priority=self._item_at(priority, i), streaming=False, ) futures.append(future) @@ -464,7 +608,6 @@ def generate_async( Returns: tensorrt_llm.llmapi.RequestOutput: The output data of the completion request to the LLM. """ - if self._encode_only: raise RuntimeError( "generate_async() is not available when encode_only=True. " @@ -490,9 +633,19 @@ def generate_async( prompt = None query_token_ids = inputs.query_token_ids multimodal_params = inputs.multimodal_params + preprocessed_encoder_input_token_ids = inputs.encoder_input_token_ids + if preprocessed_encoder_input_token_ids is not None: + preprocessed_encoder_input_token_ids = self._normalize_token_ids( + preprocessed_encoder_input_token_ids, + "inputs.encoder_input_token_ids") + encoder_input_token_ids = preprocessed_encoder_input_token_ids else: - prompt_token_ids, prompt, query_token_ids, multimodal_params = ( - self._preprocess(inputs, sampling_params, disaggregated_params)) + (prompt_token_ids, prompt, query_token_ids, multimodal_params, + encoder_input_token_ids) = self._preprocess( + inputs, + sampling_params, + disaggregated_params, + ) arrival_time = steady_clock_now( ) if self.args.return_perf_metrics else None @@ -520,6 +673,7 @@ def generate_async( scheduling_params=scheduling_params, cache_salt=cache_salt, arrival_time=arrival_time, + encoder_input_token_ids=encoder_input_token_ids, priority=priority, ) @@ -532,19 +686,41 @@ def generate_async( def _preprocess( self, - inputs: PromptInputs, + inputs: Optional[PromptInputs], sampling_params: SamplingParams, disaggregated_params: Optional[DisaggregatedParams] = None, ) -> Tuple[List[int], Optional[str], Optional[List[int]], - Optional[MultimodalParams]]: + Optional[MultimodalParams], Optional[List[int]]]: """Preprocess raw prompts into token IDs and multimodal params. This is the CPU-heavy portion of generate_async (tokenization, multimodal processing, hash computation). Returns: - `(prompt_token_ids, prompt, query_token_ids, multimodal_params)` + `(prompt_token_ids, prompt, query_token_ids, multimodal_params, encoder_input_token_ids)` """ + if isinstance(inputs, dict): + inputs = self._copy_prompt_inputs(inputs) + if "encoder_inputs" in inputs: + raise ValueError( + "encoder_inputs is not supported. Pass encoder input as " + "inputs.") + if "encoder_input_token_ids" in inputs: + raise ValueError( + "encoder_input_token_ids is not supported. Pass encoder " + "token IDs as inputs.") + if "decoder_input_token_ids" in inputs: + raise ValueError( + "decoder_input_token_ids is not supported. Pass decoder " + "token IDs as inputs.") + + if inputs is None or (isinstance(inputs, dict) + and "prompt" not in inputs + and "prompt_token_ids" not in inputs): + raise TypeError( + f"The inputs must be type str or list of int, but got {type(inputs)}" + ) + inputs = prompt_inputs(inputs) if "multi_item_part_lens" in inputs: @@ -728,7 +904,13 @@ def _preprocess( f"The inputs must be type str or list of int, but got {type(inputs)}" ) - return prompt_token_ids, prompt, query_token_ids, multimodal_params + normalized_encoder_input_token_ids = None + if self._is_encoder_decoder_model(): + normalized_encoder_input_token_ids = prompt_token_ids + prompt_token_ids = [self._get_decoder_start_token_id()] + + return (prompt_token_ids, prompt, query_token_ids, multimodal_params, + normalized_encoder_input_token_ids) @set_api_status("prototype") def preprocess( @@ -750,13 +932,18 @@ def preprocess( passed directly to :meth:`generate_async` as `inputs`. """ sampling_params = self._prepare_sampling_params(sampling_params) - prompt_token_ids, _prompt, query_token_ids, multimodal_params = ( - self._preprocess(inputs, sampling_params, disaggregated_params)) + (prompt_token_ids, _prompt, query_token_ids, multimodal_params, + encoder_input_token_ids) = self._preprocess( + inputs, + sampling_params, + disaggregated_params, + ) return PreprocessedInputs( prompt_token_ids=prompt_token_ids, query_token_ids=query_token_ids, multimodal_params=multimodal_params, + encoder_input_token_ids=encoder_input_token_ids, ) @set_api_status("prototype") @@ -1057,6 +1244,7 @@ def _prepare_sampling_params( ) sampling_params._setup(self.tokenizer, self._hf_model_config, self._generation_config) + self._add_bart_forced_tokens_logits_processor(sampling_params) add_thinking_budget_logits_processor( sampling_params, reasoning_parser=self.args.reasoning_parser, @@ -1082,6 +1270,38 @@ def _prepare_sampling_params( sampling_params.return_perf_metrics = sampling_params.return_perf_metrics or self.args.return_perf_metrics return sampling_params + def _add_bart_forced_tokens_logits_processor( + self, sampling_params: SamplingParams) -> None: + if self.args.backend != "pytorch": + return + if getattr(self._hf_model_config, "model_type", None) != "bart": + return + if self._generation_config is None: + return + + forced_bos_token_id = getattr(self._generation_config, + "forced_bos_token_id", None) + forced_eos_token_id = getattr(self._generation_config, + "forced_eos_token_id", None) + if forced_bos_token_id is None and forced_eos_token_id is None: + return + + existing = sampling_params.logits_processor + if _contains_bart_forced_tokens_logits_processor(existing): + return + + processor = _BartForcedTokensLogitsProcessor( + forced_bos_token_id=forced_bos_token_id, + forced_eos_token_id=forced_eos_token_id, + max_tokens=sampling_params.max_tokens, + ) + if existing is None: + sampling_params.logits_processor = processor + elif isinstance(existing, list): + existing.append(processor) + else: + sampling_params.logits_processor = [existing, processor] + def _check_arguments(self, prompt_len: int, query_len: int, sampling_params: SamplingParams, is_gen_only: bool) -> None: diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 6425cb7d2f2d..6f06d3b68895 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -2796,7 +2796,7 @@ class KvCacheConfig(StrictBaseModel, PybindMirror): cross_kv_cache_fraction: Optional[float] = Field( default=None, description= - "The fraction of the KV Cache memory should be reserved for cross attention. If set to p, self attention will use 1-p of KV Cache memory and cross attention will use p of KV Cache memory. Default is 50%. Should only be set when using encoder-decoder model." + "The fraction of the KV Cache memory should be reserved for cross attention. If set to p, self attention will use 1-p of KV Cache memory and cross attention will use p of KV Cache memory. Defaults to None (unset); must be set when using an encoder-decoder model and must not be set otherwise." ) secondary_offload_min_priority: Optional[int] = Field( default=None, @@ -2935,6 +2935,17 @@ def validate_free_gpu_memory_fraction(cls, v: float): ) return v + @field_validator('cross_kv_cache_fraction') + @classmethod + def validate_cross_kv_cache_fraction(cls, v: Optional[float]): + if v is None: + return v + if not 0 <= v <= 1: + raise ValueError( + "kv_cache_config.cross_kv_cache_fraction must be a float between 0 and 1" + ) + return v + @field_validator('dtype') @classmethod def validate_dtype(cls, v: str): diff --git a/tests/integration/defs/llmapi/test_llm_api_pytorch_bart.py b/tests/integration/defs/llmapi/test_llm_api_pytorch_bart.py new file mode 100644 index 000000000000..3ae0e3652a3d --- /dev/null +++ b/tests/integration/defs/llmapi/test_llm_api_pytorch_bart.py @@ -0,0 +1,385 @@ +# 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. + +from pathlib import Path + +import pytest +from transformers import AutoTokenizer + +from tensorrt_llm.llmapi import ( + LLM, + CudaGraphConfig, + KvCacheConfig, + RequestOutput, + SamplingParams, + SchedulerConfig, +) + +from ..conftest import llm_models_root + +_SOURCE_TEXT = ( + "Summarize: NVIDIA builds fast inference software for large language models. " + "TensorRT-LLM supports encoder-decoder models such as BART and T5." +) +_MIXED_ENCODER_SOURCE_TEXTS = [ + _SOURCE_TEXT, + ( + "Summarize: The city opened a new public library on Monday. Residents said " + "the library has quiet rooms, computer access, and a large children section." + ), +] +_MODEL_NAME = "bart-large-cnn" +_MAX_NEW_TOKENS = 8 +_MAX_SEQUENCE_LENGTH = 128 +_MAX_KV_TOKENS = 384 +_MIN_GPU_MEMORY_MB = 16_000 +_FREE_GPU_MEMORY_FRACTION = 0.2 +_CROSS_KV_CACHE_FRACTION = 0.5 +_EXPECTED_GREEDY_OUTPUT_TOKEN_IDS = [0, 565, 35354, 13963, 12, 6006, 448, 2] +_EXPECTED_TEXT_FRAGMENT = "TensorRT" +_MIXED_ENCODER_EXPECTED_TEXT_FRAGMENTS = [ + _EXPECTED_TEXT_FRAGMENT, + "library", +] + + +def _test_case( + torch_dtype: str, + use_kv_cache_manager_v2: bool, + enable_cuda_graph: bool, + num_beams: int, + num_return_sequences: int, + exact_match: bool, + feature_id: str, +): + expected_output_token_ids = [_EXPECTED_GREEDY_OUTPUT_TOKEN_IDS] if num_beams == 1 else None + assert not exact_match or expected_output_token_ids is not None + + return pytest.param( + expected_output_token_ids, + torch_dtype, + use_kv_cache_manager_v2, + enable_cuda_graph, + num_beams, + num_return_sequences, + exact_match, + id=f"{feature_id}-{_MODEL_NAME}", + ) + + +_TEST_CASES = [ + _test_case( + torch_dtype="bfloat16", + use_kv_cache_manager_v2=False, + enable_cuda_graph=False, + num_beams=1, + num_return_sequences=1, + exact_match=True, + feature_id="bf16-kv-v1-cuda-graph-off-greedy", + ), + _test_case( + torch_dtype="float16", + use_kv_cache_manager_v2=False, + enable_cuda_graph=False, + num_beams=1, + num_return_sequences=1, + exact_match=False, + feature_id="fp16-kv-v1-cuda-graph-off-greedy", + ), + _test_case( + torch_dtype="bfloat16", + use_kv_cache_manager_v2=False, + enable_cuda_graph=False, + num_beams=2, + num_return_sequences=2, + exact_match=False, + feature_id="bf16-kv-v1-cuda-graph-off-beam2", + ), + _test_case( + torch_dtype="bfloat16", + use_kv_cache_manager_v2=True, + enable_cuda_graph=False, + num_beams=1, + num_return_sequences=1, + exact_match=True, + feature_id="bf16-kv-v2-cuda-graph-off-greedy", + ), +] + + +def _mixed_batch_test_case( + torch_dtype: str, + use_kv_cache_manager_v2: bool, + num_beams: int, + num_return_sequences: int, + feature_id: str, +): + return pytest.param( + torch_dtype, + use_kv_cache_manager_v2, + num_beams, + num_return_sequences, + id=f"{feature_id}-{_MODEL_NAME}", + ) + + +_MIXED_BATCH_TEST_CASES = [ + _mixed_batch_test_case( + torch_dtype="bfloat16", + use_kv_cache_manager_v2=False, + num_beams=1, + num_return_sequences=1, + feature_id="bf16-kv-v1-cuda-graph-off-greedy-batch2", + ), + _mixed_batch_test_case( + torch_dtype="bfloat16", + use_kv_cache_manager_v2=True, + num_beams=1, + num_return_sequences=1, + feature_id="bf16-kv-v2-cuda-graph-off-greedy-batch2", + ), +] + +pytestmark = [ + pytest.mark.skip_less_device(1), + pytest.mark.skip_less_device_memory(_MIN_GPU_MEMORY_MB), + pytest.mark.threadleak(enabled=False), +] + + +def _get_bart_model_path() -> str: + try: + models_root = Path(llm_models_root()) + except AssertionError as exc: + pytest.skip(str(exc)) + + model_path = models_root / _MODEL_NAME + if not model_path.exists(): + pytest.skip(f"{_MODEL_NAME} is not available under {models_root}") + return str(model_path) + + +def _sampling_params(num_beams: int, num_return_sequences: int) -> SamplingParams: + if num_beams == 1: + assert num_return_sequences == 1 + return SamplingParams( + max_tokens=_MAX_NEW_TOKENS, + temperature=0.0, + ) + + return SamplingParams( + best_of=num_beams, + max_tokens=_MAX_NEW_TOKENS, + n=num_return_sequences, + temperature=0.0, + use_beam_search=True, + ) + + +def _cuda_graph_config( + enabled: bool, + batch_sizes: list[int] | None = None, +) -> CudaGraphConfig | None: + return CudaGraphConfig(batch_sizes=batch_sizes or [1]) if enabled else None + + +def _assert_bart_response( + response: RequestOutput, + num_return_sequences: int, +) -> list[list[int]]: + assert response.finished + + assert len(response.outputs) == num_return_sequences + token_ids_by_output = [] + for output in response.outputs: + assert output.token_ids is not None + assert 0 < len(output.token_ids) <= _MAX_NEW_TOKENS + token_ids_by_output.append(output.token_ids) + return token_ids_by_output + + +def _print_generated_text( + tokenizer, case_id: str, label: str, token_ids_by_output: list[list[int]] +) -> None: + for output_idx, token_ids in enumerate(token_ids_by_output): + text = tokenizer.decode(token_ids, skip_special_tokens=True) + print(f"{case_id} {label}[{output_idx}]: {text!r} token_ids={token_ids}") + + +def _assert_expected_generation( + tokenizer, + token_ids_by_output: list[list[int]], + exact_match: bool, + expected_token_ids_by_output: list[list[int]] | None, + expected_text_fragment: str | None = _EXPECTED_TEXT_FRAGMENT, +) -> None: + decoded_text_by_output = [ + tokenizer.decode(token_ids, skip_special_tokens=True) for token_ids in token_ids_by_output + ] + assert all(decoded_text_by_output) + if expected_token_ids_by_output is None: + if expected_text_fragment is not None: + assert all(expected_text_fragment in text for text in decoded_text_by_output) + else: + assert token_ids_by_output[0] == expected_token_ids_by_output[0] + if len(token_ids_by_output) > 1: + assert len({tuple(token_ids) for token_ids in token_ids_by_output}) == len( + token_ids_by_output + ) + if not exact_match: + return + + assert expected_token_ids_by_output is not None + assert token_ids_by_output == expected_token_ids_by_output + + +@pytest.mark.parametrize( + "expected_output_token_ids_by_output,torch_dtype,use_kv_cache_manager_v2," + "enable_cuda_graph,num_beams,num_return_sequences,exact_match", + _TEST_CASES, +) +def test_bart_pytorch_generate_encoder_decoder_end_to_end( + monkeypatch: pytest.MonkeyPatch, + expected_output_token_ids_by_output: list[list[int]] | None, + torch_dtype: str, + use_kv_cache_manager_v2: bool, + enable_cuda_graph: bool, + num_beams: int, + num_return_sequences: int, + exact_match: bool, +) -> None: + monkeypatch.setenv("TLLM_WORKER_USE_SINGLE_PROCESS", "1") + monkeypatch.setenv("TRTLLM_SKIP_KV_CACHE_ESTIMATION", "1") + + model_path = _get_bart_model_path() + tokenizer = AutoTokenizer.from_pretrained(model_path) + case_id = ( + f"model={_MODEL_NAME}, dtype={torch_dtype}, kv_v2={use_kv_cache_manager_v2}, " + f"cuda_graph={enable_cuda_graph}, beams={num_beams}, returns={num_return_sequences}" + ) + sampling_params = _sampling_params(num_beams, num_return_sequences) + + with LLM( + model_path, + backend="pytorch", + attn_backend="TRTLLM", + cuda_graph_config=_cuda_graph_config(enable_cuda_graph), + disable_overlap_scheduler=True, + dtype=torch_dtype, + enable_chunked_prefill=False, + kv_cache_config=KvCacheConfig( + enable_block_reuse=False, + max_tokens=_MAX_KV_TOKENS, + free_gpu_memory_fraction=_FREE_GPU_MEMORY_FRACTION, + cross_kv_cache_fraction=_CROSS_KV_CACHE_FRACTION, + use_kv_cache_manager_v2=use_kv_cache_manager_v2, + ), + max_batch_size=1, + max_beam_width=num_beams, + max_input_len=_MAX_SEQUENCE_LENGTH, + max_num_tokens=_MAX_SEQUENCE_LENGTH, + max_seq_len=_MAX_SEQUENCE_LENGTH, + model_kwargs={"torch_dtype": torch_dtype}, + scheduler_config=SchedulerConfig(use_python_scheduler=True), + ) as llm: + response = llm.generate( + _SOURCE_TEXT, + sampling_params=sampling_params, + use_tqdm=False, + ) + token_ids = _assert_bart_response( + response, + num_return_sequences=num_return_sequences, + ) + _print_generated_text(tokenizer, case_id, "output", token_ids) + _assert_expected_generation( + tokenizer, + token_ids, + exact_match, + expected_output_token_ids_by_output, + ) + + +@pytest.mark.parametrize( + "torch_dtype,use_kv_cache_manager_v2,num_beams,num_return_sequences", + _MIXED_BATCH_TEST_CASES, +) +def test_bart_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch( + monkeypatch: pytest.MonkeyPatch, + torch_dtype: str, + use_kv_cache_manager_v2: bool, + num_beams: int, + num_return_sequences: int, +) -> None: + monkeypatch.setenv("TLLM_WORKER_USE_SINGLE_PROCESS", "1") + monkeypatch.setenv("TRTLLM_SKIP_KV_CACHE_ESTIMATION", "1") + + model_path = _get_bart_model_path() + tokenizer = AutoTokenizer.from_pretrained(model_path) + sampling_params = _sampling_params(num_beams, num_return_sequences) + case_id = ( + f"model={_MODEL_NAME}, dtype={torch_dtype}, kv_v2={use_kv_cache_manager_v2}, " + f"cuda_graph=False, beams={num_beams}, returns={num_return_sequences}, " + "mixed_encoder_lengths=True, batch_size=2" + ) + with LLM( + model_path, + backend="pytorch", + attn_backend="TRTLLM", + cuda_graph_config=None, + disable_overlap_scheduler=True, + dtype=torch_dtype, + enable_chunked_prefill=False, + kv_cache_config=KvCacheConfig( + enable_block_reuse=False, + max_tokens=_MAX_KV_TOKENS, + free_gpu_memory_fraction=_FREE_GPU_MEMORY_FRACTION, + cross_kv_cache_fraction=_CROSS_KV_CACHE_FRACTION, + use_kv_cache_manager_v2=use_kv_cache_manager_v2, + ), + max_batch_size=len(_MIXED_ENCODER_SOURCE_TEXTS), + max_beam_width=num_beams, + max_input_len=_MAX_SEQUENCE_LENGTH, + max_num_tokens=_MAX_SEQUENCE_LENGTH, + max_seq_len=_MAX_SEQUENCE_LENGTH, + model_kwargs={"torch_dtype": torch_dtype}, + scheduler_config=SchedulerConfig(use_python_scheduler=True), + ) as llm: + responses = llm.generate( + _MIXED_ENCODER_SOURCE_TEXTS, + sampling_params=sampling_params, + use_tqdm=False, + ) + + assert len(responses) == len(_MIXED_ENCODER_SOURCE_TEXTS) + + for request_idx, response in enumerate(responses): + token_ids = _assert_bart_response( + response, + num_return_sequences=num_return_sequences, + ) + _print_generated_text( + tokenizer, + f"{case_id}, request={request_idx}", + "output", + token_ids, + ) + _assert_expected_generation( + tokenizer, + token_ids, + exact_match=False, + expected_token_ids_by_output=None, + expected_text_fragment=_MIXED_ENCODER_EXPECTED_TEXT_FRAGMENTS[request_idx], + ) diff --git a/tests/integration/defs/llmapi/test_llm_api_pytorch_t5.py b/tests/integration/defs/llmapi/test_llm_api_pytorch_t5.py new file mode 100644 index 000000000000..c42bb0a40057 --- /dev/null +++ b/tests/integration/defs/llmapi/test_llm_api_pytorch_t5.py @@ -0,0 +1,678 @@ +# 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. + +from pathlib import Path + +import pytest +from transformers import AutoTokenizer + +from tensorrt_llm.llmapi import ( + LLM, + CudaGraphConfig, + KvCacheConfig, + RequestOutput, + SamplingParams, + SchedulerConfig, +) + +from ..conftest import llm_models_root + +_SOURCE_TEXT = "translate English to German: The house is wonderful." +_MIXED_ENCODER_SOURCE_TEXTS = [ + _SOURCE_TEXT, + "translate English to German: The book is on the table.", +] +_MAX_NEW_TOKENS = 4 +_MAX_SEQUENCE_LENGTH = 64 +_MAX_KV_TOKENS = 256 +_MIN_GPU_MEMORY_MB = 16_000 +_FLAN_T5_XXL_MIN_GPU_MEMORY_MB = 80_000 +_FREE_GPU_MEMORY_FRACTION = 0.2 +_CROSS_KV_CACHE_FRACTION = 0.5 +_EXPECTED_TRANSLATION_FRAGMENT = "Haus" +_EXPECTED_OUTPUT_TOKEN_IDS_BY_MODEL = { + "t5-small": [644, 4598, 229, 19250], + "t5-base": [644, 4598, 229, 19250], + "t5-large": [644, 4598, 229, 19250], + "flan-t5-small": [644, 4598, 229, 9685], + "byt5-small": [258, 35, 119, 114], +} +# Known HF references for returned beam hypotheses. The tests exact-match greedy +# outputs and the best beam when a reference is available; lower-ranked BF16 +# alternatives can differ on very close scores, so beam tests also assert that +# all requested outputs are present, non-empty, and distinct. +_HF_BEAM_OUTPUT_TOKEN_IDS_BY_MODEL_AND_BEAMS = { + ("t5-small", 2): [ + [644, 4598, 229, 19250], + [644, 4598, 229, 3], + ], + ("t5-base", 2): [ + [644, 4598, 229, 19250], + [644, 4598, 229, 3], + ], + ("t5-large", 2): [ + [644, 4598, 229, 19250], + [644, 4598, 229, 3], + ], + ("flan-t5-small", 2): [ + [644, 4598, 229, 9685], + [644, 4598, 229, 19250], + ], +} +_MIXED_ENCODER_OUTPUT_TOKEN_IDS_BY_MODEL_AND_BEAMS = { + ("t5-small", 1): [ + [[644, 4598, 229, 19250]], + [[644, 4675, 4186, 219]], + ], + ("t5-small", 2): [ + _HF_BEAM_OUTPUT_TOKEN_IDS_BY_MODEL_AND_BEAMS[("t5-small", 2)], + [ + [644, 4675, 229, 219], + [644, 4675, 4186, 219], + ], + ], + ("flan-t5-small", 2): [ + _HF_BEAM_OUTPUT_TOKEN_IDS_BY_MODEL_AND_BEAMS[("flan-t5-small", 2)], + [ + [316, 4675, 229, 219], + [316, 4675, 229, 256], + ], + ], +} +_MIXED_ENCODER_EXPECTED_TEXT_FRAGMENTS_BY_MODEL = { + "t5-small": [_EXPECTED_TRANSLATION_FRAGMENT, "Buch"], + "flan-t5-small": [_EXPECTED_TRANSLATION_FRAGMENT, "Buch"], +} + + +def _test_case( + model_name: str, + torch_dtype: str, + use_kv_cache_manager_v2: bool, + enable_cuda_graph: bool, + num_beams: int, + num_return_sequences: int, + exact_match: bool, + feature_id: str, + marks=(), +): + if num_beams == 1: + expected_output_token_ids = ( + [_EXPECTED_OUTPUT_TOKEN_IDS_BY_MODEL[model_name]] + if model_name in _EXPECTED_OUTPUT_TOKEN_IDS_BY_MODEL + else None + ) + elif num_return_sequences == num_beams: + expected_output_token_ids = _HF_BEAM_OUTPUT_TOKEN_IDS_BY_MODEL_AND_BEAMS.get( + (model_name, num_beams) + ) + else: + expected_output_token_ids = ( + [_EXPECTED_OUTPUT_TOKEN_IDS_BY_MODEL[model_name]] + if model_name in _EXPECTED_OUTPUT_TOKEN_IDS_BY_MODEL + else None + ) + + assert not exact_match or expected_output_token_ids is not None + + return pytest.param( + model_name, + expected_output_token_ids, + torch_dtype, + use_kv_cache_manager_v2, + enable_cuda_graph, + num_beams, + num_return_sequences, + exact_match, + id=f"{feature_id}-{model_name}", + marks=marks, + ) + + +_TEST_CASES = [ + # Primary coverage: v1 cache manager and beam search. + _test_case( + model_name="t5-small", + torch_dtype="bfloat16", + use_kv_cache_manager_v2=False, + enable_cuda_graph=False, + num_beams=2, + num_return_sequences=2, + exact_match=False, + feature_id="bf16-kv-v1-cuda-graph-off-beam2", + ), + _test_case( + model_name="flan-t5-small", + torch_dtype="bfloat16", + use_kv_cache_manager_v2=False, + enable_cuda_graph=False, + num_beams=2, + num_return_sequences=2, + exact_match=False, + feature_id="bf16-kv-v1-cuda-graph-off-beam2", + ), + _test_case( + model_name="t5-base", + torch_dtype="bfloat16", + use_kv_cache_manager_v2=False, + enable_cuda_graph=False, + num_beams=2, + num_return_sequences=2, + exact_match=False, + feature_id="bf16-kv-v1-cuda-graph-off-beam2", + ), + _test_case( + model_name="t5-large", + torch_dtype="bfloat16", + use_kv_cache_manager_v2=False, + enable_cuda_graph=False, + num_beams=2, + num_return_sequences=2, + exact_match=False, + feature_id="bf16-kv-v1-cuda-graph-off-beam2", + ), + _test_case( + model_name="flan-t5-base", + torch_dtype="bfloat16", + use_kv_cache_manager_v2=False, + enable_cuda_graph=False, + num_beams=2, + num_return_sequences=2, + exact_match=False, + feature_id="bf16-kv-v1-cuda-graph-off-beam2", + ), + _test_case( + model_name="flan-t5-large", + torch_dtype="bfloat16", + use_kv_cache_manager_v2=False, + enable_cuda_graph=False, + num_beams=2, + num_return_sequences=2, + exact_match=False, + feature_id="bf16-kv-v1-cuda-graph-off-beam2", + ), + _test_case( + model_name="flan-t5-xl", + torch_dtype="bfloat16", + use_kv_cache_manager_v2=False, + enable_cuda_graph=False, + num_beams=2, + num_return_sequences=2, + exact_match=False, + feature_id="bf16-kv-v1-cuda-graph-off-beam2", + ), + _test_case( + model_name="flan-t5-xxl", + torch_dtype="bfloat16", + use_kv_cache_manager_v2=False, + enable_cuda_graph=False, + num_beams=2, + num_return_sequences=2, + exact_match=False, + feature_id="bf16-kv-v1-cuda-graph-off-beam2", + marks=pytest.mark.skip_less_device_memory(_FLAN_T5_XXL_MIN_GPU_MEMORY_MB), + ), + # Non-CUDA-graph smoke for the same v1 beam path. + _test_case( + model_name="t5-small", + torch_dtype="bfloat16", + use_kv_cache_manager_v2=False, + enable_cuda_graph=False, + num_beams=2, + num_return_sequences=2, + exact_match=False, + feature_id="bf16-kv-v1-cuda-graph-off-beam2", + ), + # Greedy smoke for the priority v1 path. + _test_case( + model_name="t5-small", + torch_dtype="bfloat16", + use_kv_cache_manager_v2=False, + enable_cuda_graph=False, + num_beams=1, + num_return_sequences=1, + exact_match=True, + feature_id="bf16-kv-v1-cuda-graph-off-greedy", + ), + # Precision coverage for beam search. KVCacheManagerV2 currently requires + # max_beam_width == 1, so beam-search precision coverage uses v1. + _test_case( + model_name="t5-small", + torch_dtype="float16", + use_kv_cache_manager_v2=False, + enable_cuda_graph=False, + num_beams=2, + num_return_sequences=2, + exact_match=False, + feature_id="fp16-kv-v1-cuda-graph-off-beam2", + ), + _test_case( + model_name="t5-small", + torch_dtype="float32", + use_kv_cache_manager_v2=False, + enable_cuda_graph=False, + num_beams=2, + num_return_sequences=2, + exact_match=False, + feature_id="fp32-kv-v1-cuda-graph-off-beam2", + ), + _test_case( + model_name="flan-t5-small", + torch_dtype="float16", + use_kv_cache_manager_v2=False, + enable_cuda_graph=False, + num_beams=2, + num_return_sequences=2, + exact_match=False, + feature_id="fp16-kv-v1-cuda-graph-off-beam2", + ), + _test_case( + model_name="flan-t5-small", + torch_dtype="float32", + use_kv_cache_manager_v2=False, + enable_cuda_graph=False, + num_beams=2, + num_return_sequences=2, + exact_match=False, + feature_id="fp32-kv-v1-cuda-graph-off-beam2", + ), + # Precision coverage for v2 on its supported greedy path. + _test_case( + model_name="t5-small", + torch_dtype="bfloat16", + use_kv_cache_manager_v2=True, + enable_cuda_graph=False, + num_beams=1, + num_return_sequences=1, + exact_match=True, + feature_id="bf16-kv-v2-cuda-graph-off-greedy", + ), + _test_case( + model_name="t5-small", + torch_dtype="float16", + use_kv_cache_manager_v2=True, + enable_cuda_graph=False, + num_beams=1, + num_return_sequences=1, + exact_match=True, + feature_id="fp16-kv-v2-cuda-graph-off-greedy", + ), + _test_case( + model_name="t5-small", + torch_dtype="float32", + use_kv_cache_manager_v2=True, + enable_cuda_graph=False, + num_beams=1, + num_return_sequences=1, + exact_match=True, + feature_id="fp32-kv-v2-cuda-graph-off-greedy", + ), + _test_case( + model_name="flan-t5-small", + torch_dtype="bfloat16", + use_kv_cache_manager_v2=True, + enable_cuda_graph=False, + num_beams=1, + num_return_sequences=1, + exact_match=True, + feature_id="bf16-kv-v2-cuda-graph-off-greedy", + ), + _test_case( + model_name="flan-t5-small", + torch_dtype="float16", + use_kv_cache_manager_v2=True, + enable_cuda_graph=False, + num_beams=1, + num_return_sequences=1, + exact_match=True, + feature_id="fp16-kv-v2-cuda-graph-off-greedy", + ), + _test_case( + model_name="flan-t5-small", + torch_dtype="float32", + use_kv_cache_manager_v2=True, + enable_cuda_graph=False, + num_beams=1, + num_return_sequences=1, + exact_match=True, + feature_id="fp32-kv-v2-cuda-graph-off-greedy", + ), + # ByT5 sanity coverage keeps the known-stable expected output path. + _test_case( + model_name="byt5-small", + torch_dtype="bfloat16", + use_kv_cache_manager_v2=True, + enable_cuda_graph=False, + num_beams=1, + num_return_sequences=1, + exact_match=True, + feature_id="bf16-kv-v2-cuda-graph-off-greedy", + ), +] + + +def _mixed_batch_test_case( + model_name: str, + torch_dtype: str, + use_kv_cache_manager_v2: bool, + num_beams: int, + num_return_sequences: int, + exact_match: bool, + feature_id: str, + marks=(), +): + expected_output_token_ids_by_request = ( + _MIXED_ENCODER_OUTPUT_TOKEN_IDS_BY_MODEL_AND_BEAMS.get((model_name, num_beams)) + if exact_match or num_beams > 1 + else None + ) + assert not exact_match or expected_output_token_ids_by_request is not None + + return pytest.param( + model_name, + expected_output_token_ids_by_request, + torch_dtype, + use_kv_cache_manager_v2, + num_beams, + num_return_sequences, + exact_match, + id=f"{feature_id}-{model_name}", + marks=marks, + ) + + +_MIXED_BATCH_TEST_CASES = [ + _mixed_batch_test_case( + model_name="t5-small", + torch_dtype="bfloat16", + use_kv_cache_manager_v2=False, + num_beams=2, + num_return_sequences=2, + exact_match=False, + feature_id="bf16-kv-v1-cuda-graph-off-beam2-batch2", + ), + _mixed_batch_test_case( + model_name="flan-t5-small", + torch_dtype="bfloat16", + use_kv_cache_manager_v2=False, + num_beams=2, + num_return_sequences=2, + exact_match=False, + feature_id="bf16-kv-v1-cuda-graph-off-beam2-batch2", + ), + _mixed_batch_test_case( + model_name="t5-small", + torch_dtype="bfloat16", + use_kv_cache_manager_v2=False, + num_beams=1, + num_return_sequences=1, + exact_match=True, + feature_id="bf16-kv-v1-cuda-graph-off-greedy-batch2", + ), + _mixed_batch_test_case( + model_name="t5-small", + torch_dtype="bfloat16", + use_kv_cache_manager_v2=True, + num_beams=1, + num_return_sequences=1, + exact_match=True, + feature_id="bf16-kv-v2-cuda-graph-off-greedy-batch2", + ), +] + +pytestmark = [ + pytest.mark.skip_less_device(1), + pytest.mark.skip_less_device_memory(_MIN_GPU_MEMORY_MB), + pytest.mark.threadleak(enabled=False), +] + + +def _get_t5_model_path(model_name: str) -> str: + try: + models_root = Path(llm_models_root()) + except AssertionError as exc: + pytest.skip(str(exc)) + + model_path = models_root / model_name + if not model_path.exists(): + pytest.skip(f"{model_name} is not available under {models_root}") + return str(model_path) + + +def _sampling_params(num_beams: int, num_return_sequences: int) -> SamplingParams: + if num_beams == 1: + assert num_return_sequences == 1 + return SamplingParams( + max_tokens=_MAX_NEW_TOKENS, + temperature=0.0, + ) + + return SamplingParams( + best_of=num_beams, + max_tokens=_MAX_NEW_TOKENS, + n=num_return_sequences, + temperature=0.0, + use_beam_search=True, + ) + + +def _cuda_graph_config( + enabled: bool, + batch_sizes: list[int] | None = None, +) -> CudaGraphConfig | None: + return CudaGraphConfig(batch_sizes=batch_sizes or [1]) if enabled else None + + +def _assert_t5_response( + response: RequestOutput, + num_return_sequences: int, +) -> list[list[int]]: + assert response.finished + + assert len(response.outputs) == num_return_sequences + token_ids_by_output = [] + for output in response.outputs: + assert output.token_ids is not None + assert 0 < len(output.token_ids) <= _MAX_NEW_TOKENS + token_ids_by_output.append(output.token_ids) + return token_ids_by_output + + +def _print_generated_text( + tokenizer, case_id: str, label: str, token_ids_by_output: list[list[int]] +) -> None: + for output_idx, token_ids in enumerate(token_ids_by_output): + text = tokenizer.decode(token_ids, skip_special_tokens=True) + print(f"{case_id} {label}[{output_idx}]: {text!r} token_ids={token_ids}") + + +def _assert_expected_generation( + tokenizer, + token_ids_by_output: list[list[int]], + exact_match: bool, + expected_token_ids_by_output: list[list[int]] | None, + expected_text_fragment: str = _EXPECTED_TRANSLATION_FRAGMENT, +) -> None: + decoded_text_by_output = [ + tokenizer.decode(token_ids, skip_special_tokens=True) for token_ids in token_ids_by_output + ] + assert all(decoded_text_by_output) + if expected_token_ids_by_output is None: + assert all(expected_text_fragment in text for text in decoded_text_by_output) + else: + assert token_ids_by_output[0] == expected_token_ids_by_output[0] + if len(token_ids_by_output) > 1: + assert len({tuple(token_ids) for token_ids in token_ids_by_output}) == len( + token_ids_by_output + ) + if not exact_match: + return + + assert expected_token_ids_by_output is not None + assert token_ids_by_output == expected_token_ids_by_output + + +@pytest.mark.parametrize( + "model_name,expected_output_token_ids_by_output,torch_dtype,use_kv_cache_manager_v2," + "enable_cuda_graph,num_beams,num_return_sequences,exact_match", + _TEST_CASES, +) +def test_t5_pytorch_generate_encoder_decoder_end_to_end( + monkeypatch: pytest.MonkeyPatch, + model_name: str, + expected_output_token_ids_by_output: list[list[int]] | None, + torch_dtype: str, + use_kv_cache_manager_v2: bool, + enable_cuda_graph: bool, + num_beams: int, + num_return_sequences: int, + exact_match: bool, +) -> None: + monkeypatch.setenv("TLLM_WORKER_USE_SINGLE_PROCESS", "1") + monkeypatch.setenv("TRTLLM_SKIP_KV_CACHE_ESTIMATION", "1") + + model_path = _get_t5_model_path(model_name) + tokenizer = AutoTokenizer.from_pretrained(model_path) + case_id = ( + f"model={model_name}, dtype={torch_dtype}, kv_v2={use_kv_cache_manager_v2}, " + f"cuda_graph={enable_cuda_graph}, beams={num_beams}, returns={num_return_sequences}" + ) + sampling_params = _sampling_params(num_beams, num_return_sequences) + + with LLM( + model_path, + backend="pytorch", + attn_backend="TRTLLM", + cuda_graph_config=_cuda_graph_config(enable_cuda_graph), + disable_overlap_scheduler=True, + dtype=torch_dtype, + enable_chunked_prefill=False, + kv_cache_config=KvCacheConfig( + enable_block_reuse=False, + max_tokens=_MAX_KV_TOKENS, + free_gpu_memory_fraction=_FREE_GPU_MEMORY_FRACTION, + cross_kv_cache_fraction=_CROSS_KV_CACHE_FRACTION, + use_kv_cache_manager_v2=use_kv_cache_manager_v2, + ), + max_batch_size=1, + max_beam_width=num_beams, + max_input_len=_MAX_SEQUENCE_LENGTH, + max_num_tokens=_MAX_SEQUENCE_LENGTH, + max_seq_len=_MAX_SEQUENCE_LENGTH, + model_kwargs={"torch_dtype": torch_dtype}, + scheduler_config=SchedulerConfig(use_python_scheduler=True), + ) as llm: + response = llm.generate( + _SOURCE_TEXT, + sampling_params=sampling_params, + use_tqdm=False, + ) + token_ids = _assert_t5_response( + response, + num_return_sequences=num_return_sequences, + ) + _print_generated_text(tokenizer, case_id, "output", token_ids) + _assert_expected_generation( + tokenizer, + token_ids, + exact_match, + expected_output_token_ids_by_output, + ) + + +@pytest.mark.parametrize( + "model_name,expected_output_token_ids_by_request,torch_dtype,use_kv_cache_manager_v2," + "num_beams,num_return_sequences,exact_match", + _MIXED_BATCH_TEST_CASES, +) +def test_t5_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch( + monkeypatch: pytest.MonkeyPatch, + model_name: str, + expected_output_token_ids_by_request: list[list[list[int]] | None] | None, + torch_dtype: str, + use_kv_cache_manager_v2: bool, + num_beams: int, + num_return_sequences: int, + exact_match: bool, +) -> None: + monkeypatch.setenv("TLLM_WORKER_USE_SINGLE_PROCESS", "1") + monkeypatch.setenv("TRTLLM_SKIP_KV_CACHE_ESTIMATION", "1") + + model_path = _get_t5_model_path(model_name) + tokenizer = AutoTokenizer.from_pretrained(model_path) + sampling_params = _sampling_params(num_beams, num_return_sequences) + case_id = ( + f"model={model_name}, dtype={torch_dtype}, kv_v2={use_kv_cache_manager_v2}, " + f"cuda_graph=False, beams={num_beams}, returns={num_return_sequences}, " + "mixed_encoder_lengths=True, batch_size=2" + ) + with LLM( + model_path, + backend="pytorch", + attn_backend="TRTLLM", + cuda_graph_config=None, + disable_overlap_scheduler=True, + dtype=torch_dtype, + enable_chunked_prefill=False, + kv_cache_config=KvCacheConfig( + enable_block_reuse=False, + max_tokens=_MAX_KV_TOKENS, + free_gpu_memory_fraction=_FREE_GPU_MEMORY_FRACTION, + cross_kv_cache_fraction=_CROSS_KV_CACHE_FRACTION, + use_kv_cache_manager_v2=use_kv_cache_manager_v2, + ), + max_batch_size=len(_MIXED_ENCODER_SOURCE_TEXTS), + max_beam_width=num_beams, + max_input_len=_MAX_SEQUENCE_LENGTH, + max_num_tokens=_MAX_SEQUENCE_LENGTH, + max_seq_len=_MAX_SEQUENCE_LENGTH, + model_kwargs={"torch_dtype": torch_dtype}, + scheduler_config=SchedulerConfig(use_python_scheduler=True), + ) as llm: + responses = llm.generate( + _MIXED_ENCODER_SOURCE_TEXTS, + sampling_params=sampling_params, + use_tqdm=False, + ) + + assert len(responses) == len(_MIXED_ENCODER_SOURCE_TEXTS) + + for request_idx, response in enumerate(responses): + expected_token_ids = ( + None + if expected_output_token_ids_by_request is None + else expected_output_token_ids_by_request[request_idx] + ) + expected_text_fragment = _MIXED_ENCODER_EXPECTED_TEXT_FRAGMENTS_BY_MODEL[model_name][ + request_idx + ] + + token_ids = _assert_t5_response( + response, + num_return_sequences=num_return_sequences, + ) + _print_generated_text( + tokenizer, + f"{case_id}, request={request_idx}", + "output", + token_ids, + ) + _assert_expected_generation( + tokenizer, + token_ids, + exact_match=exact_match, + expected_token_ids_by_output=expected_token_ids, + expected_text_fragment=expected_text_fragment, + ) diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index e9571ad9724a..43f513bbe6aa 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -81,6 +81,37 @@ l0_b200: - accuracy/test_llm_api_pytorch.py::TestQwen3_5_9B::test_bf16[mtp_on] - accuracy/test_llm_api_pytorch.py::TestQwen3_5_9B::test_bf16[mtp_off] - disaggregated/test_workers.py::test_workers_kv_cache_aware_router_eviction[TinyLlama-1.1B-Chat-v1.0] # nvbugs 5300551 + - llmapi/test_llm_api_pytorch_bart.py::test_bart_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-greedy-bart-large-cnn] + - llmapi/test_llm_api_pytorch_bart.py::test_bart_pytorch_generate_encoder_decoder_end_to_end[fp16-kv-v1-cuda-graph-off-greedy-bart-large-cnn] + - llmapi/test_llm_api_pytorch_bart.py::test_bart_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-bart-large-cnn] + - llmapi/test_llm_api_pytorch_bart.py::test_bart_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v2-cuda-graph-off-greedy-bart-large-cnn] + - llmapi/test_llm_api_pytorch_bart.py::test_bart_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch[bf16-kv-v1-cuda-graph-off-greedy-batch2-bart-large-cnn] + - llmapi/test_llm_api_pytorch_bart.py::test_bart_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch[bf16-kv-v2-cuda-graph-off-greedy-batch2-bart-large-cnn] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-t5-small0] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-flan-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-t5-base] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-t5-large] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-flan-t5-base] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-flan-t5-large] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-flan-t5-xl] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-flan-t5-xxl] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-t5-small1] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-greedy-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[fp16-kv-v1-cuda-graph-off-beam2-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[fp32-kv-v1-cuda-graph-off-beam2-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[fp16-kv-v1-cuda-graph-off-beam2-flan-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[fp32-kv-v1-cuda-graph-off-beam2-flan-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v2-cuda-graph-off-greedy-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[fp16-kv-v2-cuda-graph-off-greedy-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[fp32-kv-v2-cuda-graph-off-greedy-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v2-cuda-graph-off-greedy-flan-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[fp16-kv-v2-cuda-graph-off-greedy-flan-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[fp32-kv-v2-cuda-graph-off-greedy-flan-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v2-cuda-graph-off-greedy-byt5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch[bf16-kv-v1-cuda-graph-off-beam2-batch2-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch[bf16-kv-v1-cuda-graph-off-beam2-batch2-flan-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch[bf16-kv-v1-cuda-graph-off-greedy-batch2-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch[bf16-kv-v2-cuda-graph-off-greedy-batch2-t5-small] - test_e2e.py::test_ptp_quickstart_advanced[Llama3.1-8B-NVFP4-nvfp4-quantized/Meta-Llama-3.1-8B] - test_e2e.py::test_ptp_quickstart_advanced[Llama3.1-8B-FP8-llama-3.1-model/Llama-3.1-8B-Instruct-FP8] - test_e2e.py::test_ptp_quickstart_advanced_mtp[DeepSeek-V3-Lite-BF16-DeepSeek-V3-Lite/bf16] diff --git a/tests/integration/test_lists/test-db/l0_h100.yml b/tests/integration/test_lists/test-db/l0_h100.yml index c12e322ed194..67ea2ecd7edd 100644 --- a/tests/integration/test_lists/test-db/l0_h100.yml +++ b/tests/integration/test_lists/test-db/l0_h100.yml @@ -166,6 +166,37 @@ l0_h100: - disaggregated/test_disaggregated_single_gpu.py::test_disaggregated_logits[True-TinyLlama-1.1B-Chat-v1.0] - disaggregated/test_disaggregated_single_gpu.py::test_disaggregated_logprobs[False-TinyLlama-1.1B-Chat-v1.0] - disaggregated/test_disaggregated_single_gpu.py::test_disaggregated_logprobs[True-TinyLlama-1.1B-Chat-v1.0] + - llmapi/test_llm_api_pytorch_bart.py::test_bart_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-greedy-bart-large-cnn] + - llmapi/test_llm_api_pytorch_bart.py::test_bart_pytorch_generate_encoder_decoder_end_to_end[fp16-kv-v1-cuda-graph-off-greedy-bart-large-cnn] + - llmapi/test_llm_api_pytorch_bart.py::test_bart_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-bart-large-cnn] + - llmapi/test_llm_api_pytorch_bart.py::test_bart_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v2-cuda-graph-off-greedy-bart-large-cnn] + - llmapi/test_llm_api_pytorch_bart.py::test_bart_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch[bf16-kv-v1-cuda-graph-off-greedy-batch2-bart-large-cnn] + - llmapi/test_llm_api_pytorch_bart.py::test_bart_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch[bf16-kv-v2-cuda-graph-off-greedy-batch2-bart-large-cnn] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-t5-small0] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-flan-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-t5-base] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-t5-large] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-flan-t5-base] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-flan-t5-large] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-flan-t5-xl] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-flan-t5-xxl] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-beam2-t5-small1] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-off-greedy-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[fp16-kv-v1-cuda-graph-off-beam2-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[fp32-kv-v1-cuda-graph-off-beam2-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[fp16-kv-v1-cuda-graph-off-beam2-flan-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[fp32-kv-v1-cuda-graph-off-beam2-flan-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v2-cuda-graph-off-greedy-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[fp16-kv-v2-cuda-graph-off-greedy-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[fp32-kv-v2-cuda-graph-off-greedy-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v2-cuda-graph-off-greedy-flan-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[fp16-kv-v2-cuda-graph-off-greedy-flan-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[fp32-kv-v2-cuda-graph-off-greedy-flan-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v2-cuda-graph-off-greedy-byt5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch[bf16-kv-v1-cuda-graph-off-beam2-batch2-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch[bf16-kv-v1-cuda-graph-off-beam2-batch2-flan-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch[bf16-kv-v1-cuda-graph-off-greedy-batch2-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch[bf16-kv-v2-cuda-graph-off-greedy-batch2-t5-small] - test_e2e.py::test_trtllm_bench_iteration_log[PyTorch-streaming-meta-llama/Llama-3.1-8B-llama-3.1-model/Meta-Llama-3.1-8B] - test_e2e.py::test_trtllm_bench_iteration_log[PyTorch-non-streaming-meta-llama/Llama-3.1-8B-llama-3.1-model/Meta-Llama-3.1-8B] - test_e2e.py::test_trtllm_bench_request_rate_and_concurrency[enable_concurrency-] diff --git a/tests/unittest/_torch/attention/test_vanilla_attention.py b/tests/unittest/_torch/attention/test_vanilla_attention.py index 44a70622747e..8f2b0a3296fd 100644 --- a/tests/unittest/_torch/attention/test_vanilla_attention.py +++ b/tests/unittest/_torch/attention/test_vanilla_attention.py @@ -1,10 +1,14 @@ import unittest +from unittest.mock import patch import torch +import torch.nn.functional as F import tensorrt_llm from tensorrt_llm._torch.attention_backend import (VanillaAttention, VanillaAttentionMetadata) +from tensorrt_llm._torch.attention_backend.interface import \ + PredefinedAttentionMask from tensorrt_llm._torch.metadata import KVCacheParams from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm.bindings.executor import KvCacheConfig @@ -13,6 +17,56 @@ class TestVanillaAttention(unittest.TestCase): + def test_sdpa_fallback_uses_metadata_cross_flag_for_causal_mask(self): + vanilla_attn = VanillaAttention(layer_idx=0, + num_heads=1, + head_dim=1, + num_kv_heads=1) + q = torch.ones(2, 1, 1) + k = torch.ones(2, 1, 1) + v = torch.ones(2, 1, 1) + seqlens = torch.tensor([2], dtype=torch.int32) + cu_seqlens = torch.tensor([0, 2], dtype=torch.int32) + observed_is_causal = [] + + def fake_sdpa(q_s, k_s, v_s, *, is_causal, scale): + del k_s, v_s, scale + observed_is_causal.append(is_causal) + return torch.zeros_like(q_s) + + with patch.object(F, "scaled_dot_product_attention", fake_sdpa): + vanilla_attn._no_kv_cache_sdpa_fallback( + q, + k, + v, + num_heads=1, + num_kv_heads=1, + head_dim=1, + seqlens_q=seqlens, + cu_seqlens_q=cu_seqlens, + max_seqlen_q=2, + attention_mask=PredefinedAttentionMask.CAUSAL, + seqlens_kv=seqlens.clone(), + cu_seqlens_k=cu_seqlens.clone(), + max_seqlen_k=2, + is_cross=True, + ) + vanilla_attn._no_kv_cache_sdpa_fallback( + q, + k, + v, + num_heads=1, + num_kv_heads=1, + head_dim=1, + seqlens_q=seqlens, + cu_seqlens_q=cu_seqlens, + max_seqlen_q=2, + attention_mask=PredefinedAttentionMask.CAUSAL, + is_cross=False, + ) + + self.assertEqual(observed_is_causal, [False, True]) + def test_vanilla_attention(self): num_heads = 32 num_kv_heads = 8 diff --git a/tests/unittest/_torch/executor/test_chunked_logits.py b/tests/unittest/_torch/executor/test_chunked_logits.py index 8d474e8d209b..fd7e33901fd9 100644 --- a/tests/unittest/_torch/executor/test_chunked_logits.py +++ b/tests/unittest/_torch/executor/test_chunked_logits.py @@ -131,6 +131,21 @@ def test_transfer_remaining_device_logits(self, sample_logits): # Should not raise errors + @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") + def test_get_diff_moves_encoder_output_to_cpu(self): + """Test get_diff normalizes encoder_output for cross-rank sync.""" + result = PyResult(prompt_len=5, max_new_tokens=10) + encoder_output = torch.arange(6, dtype=torch.float32, + device="cuda").reshape(2, 3) + + result.set_encoder_output(encoder_output) + diff = result.get_diff() + + assert diff.encoder_output is not None + assert diff.encoder_output.device.type == "cpu" + assert result.encoder_output is encoder_output + torch.testing.assert_close(diff.encoder_output, encoder_output.cpu()) + class TestGetLatestLogitsUnexcluded: """Tests for PyResult.get_latest_logits_unexcluded""" diff --git a/tests/unittest/_torch/executor/test_dual_pool_kv_cache.py b/tests/unittest/_torch/executor/test_dual_pool_kv_cache.py new file mode 100644 index 000000000000..1d5c00617c6e --- /dev/null +++ b/tests/unittest/_torch/executor/test_dual_pool_kv_cache.py @@ -0,0 +1,892 @@ +# 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"); +# 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. +"""Tests for dual-pool KV cache construction (enc-dec Steps 4 and 5). + +Validates budget splitting, ResourceManagerType.CROSS_KV_CACHE_MANAGER +registration, and the cross pool wiring for both the V1 ``KVCacheManager`` +(default and production target) and the V2 ``KVCacheManagerV2`` +(additive secondary path) scheduler integrations. +""" + +from types import SimpleNamespace +from unittest.mock import Mock, patch + +import pytest + +from tensorrt_llm._torch.pyexecutor._util import KvCacheCreator +from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import KVCacheManagerV2 +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager, ResourceManagerType +from tensorrt_llm.llmapi.llm_args import CapacitySchedulerPolicy + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_mock_kv_cache_config( + cross_kv_cache_fraction=None, + max_gpu_total_bytes=None, + use_kv_cache_manager_v2=True, + max_tokens=None, + free_gpu_memory_fraction=0.9, + host_cache_size=None, +): + """Create a mock KvCacheConfig with the fields KvCacheCreator needs.""" + config = Mock() + config.cross_kv_cache_fraction = cross_kv_cache_fraction + config.max_gpu_total_bytes = max_gpu_total_bytes + config.use_kv_cache_manager_v2 = use_kv_cache_manager_v2 + config.max_tokens = max_tokens + config.free_gpu_memory_fraction = free_gpu_memory_fraction + config.host_cache_size = host_cache_size + config.max_attention_window = None + config.event_buffer_max_size = 0 + + def model_copy(): + c = Mock() + c.cross_kv_cache_fraction = config.cross_kv_cache_fraction + c.max_gpu_total_bytes = config.max_gpu_total_bytes + c.use_kv_cache_manager_v2 = config.use_kv_cache_manager_v2 + c.max_tokens = config.max_tokens + c.free_gpu_memory_fraction = config.free_gpu_memory_fraction + c.host_cache_size = config.host_cache_size + c.max_attention_window = config.max_attention_window + c.event_buffer_max_size = config.event_buffer_max_size + return c + + config.model_copy = model_copy + return config + + +def _make_mock_model_config( + is_encoder_decoder=False, + is_generation=True, + **pretrained_overrides, +): + """Minimal mock ModelConfig for KvCacheCreator.""" + model_config = Mock() + model_config.is_encoder_decoder = is_encoder_decoder + model_config.is_generation = is_generation + model_config.sparse_attention_config = None + + pretrained = Mock() + pretrained.num_hidden_layers = 6 + pretrained.num_attention_heads = 8 + pretrained.num_key_value_heads = 8 + pretrained.hidden_size = 512 + pretrained.head_dim = 64 + pretrained.vocab_size = 32000 + pretrained.quantization = Mock() + pretrained.quantization.quant_algo = None + pretrained.quantization.kv_cache_quant_algo = None + for key, value in pretrained_overrides.items(): + setattr(pretrained, key, value) + if "encoder_attention_heads" not in pretrained_overrides: + pretrained.encoder_attention_heads = pretrained.num_attention_heads + if "decoder_attention_heads" not in pretrained_overrides: + pretrained.decoder_attention_heads = pretrained.num_attention_heads + if "encoder_layers" not in pretrained_overrides: + pretrained.encoder_layers = pretrained.num_hidden_layers + if "decoder_layers" not in pretrained_overrides: + pretrained.decoder_layers = pretrained.num_hidden_layers + if "d_model" not in pretrained_overrides: + pretrained.d_model = pretrained.hidden_size + if "max_position_embeddings" not in pretrained_overrides: + pretrained.max_position_embeddings = 1024 + model_config.pretrained_config = pretrained + model_config.quant_config = None + return model_config + + +def _make_mock_model_engine(model_config): + """Minimal mock PyTorchModelEngine.""" + engine = Mock() + engine.model.model_config = model_config + engine.dtype = "bfloat16" + engine.is_draft_model = False + engine.kv_cache_manager_key = ResourceManagerType.KV_CACHE_MANAGER + return engine + + +def _make_creator(kv_cache_config, model_config=None, is_enc_dec=False, manager_cls=None): + """Create a KvCacheCreator with minimal mocking. + + ``manager_cls`` selects the KV cache manager class the creator binds to. + Defaults to ``KVCacheManagerV2`` when ``kv_cache_config.use_kv_cache_manager_v2`` + is True, otherwise the V1 ``KVCacheManager``. Tests can override + explicitly via ``manager_cls`` to exercise either path independently. + """ + if model_config is None: + model_config = _make_mock_model_config(is_encoder_decoder=is_enc_dec) + model_engine = _make_mock_model_engine(model_config) + + if manager_cls is None: + manager_cls = ( + KVCacheManagerV2 + if getattr(kv_cache_config, "use_kv_cache_manager_v2", True) + else KVCacheManager + ) + + with patch( + "tensorrt_llm._torch.pyexecutor._util.get_kv_cache_manager_cls", + return_value=manager_cls, + ): + creator = KvCacheCreator.__new__(KvCacheCreator) + creator._model_engine = model_engine + creator._draft_model_engine = None + creator._mapping = Mock() + creator._mapping.enable_attention_dp = False + creator._mapping.tp_size = 1 + creator._mapping.pp_size = 1 + creator._mapping.cp_config = {} + creator._mapping.is_last_pp_rank.return_value = True + creator._kv_cache_config = kv_cache_config + creator._max_kv_tokens_in = kv_cache_config.max_tokens + creator._max_num_tokens = 4096 + creator._max_beam_width = 1 + creator._kv_connector_manager = None + creator._llm_args = Mock() + creator._llm_args.extra_resource_managers = {} + creator._cache_transceiver_config = None + creator._speculative_config = None + creator._sparse_attention_config = None + creator._tokens_per_block = 64 + creator._max_seq_len = 2048 + creator._max_batch_size = 8 + creator._net_max_seq_len = 2048 + creator._dummy_reqs = None + creator._profiling_stage_data = None + creator._kv_cache_manager_cls = manager_cls + creator._execution_stream = None + creator._draft_config = None + creator._skip_est = True + return creator + + +# --------------------------------------------------------------------------- +# Tests: _split_kv_cache_budget_for_cross +# --------------------------------------------------------------------------- + + +class TestSplitKvCacheBudgetForCross: + """Test the budget splitting method directly.""" + + def test_split_50_50(self): + total = 10 * (1 << 30) # 10 GiB + config = _make_mock_kv_cache_config( + cross_kv_cache_fraction=0.5, + max_gpu_total_bytes=total, + free_gpu_memory_fraction=0.8, + ) + + creator = _make_creator(config, is_enc_dec=True) + self_config, cross_config = creator._split_kv_cache_budget_for_cross() + + assert self_config is not config + assert cross_config.max_gpu_total_bytes == total // 2 + assert self_config.max_gpu_total_bytes == total - total // 2 + assert cross_config.free_gpu_memory_fraction == pytest.approx(0.4) + assert self_config.free_gpu_memory_fraction == pytest.approx(0.4) + assert config.max_gpu_total_bytes == total + assert config.free_gpu_memory_fraction == pytest.approx(0.8) + + def test_split_30_70(self): + total = 10 * (1 << 30) + config = _make_mock_kv_cache_config( + cross_kv_cache_fraction=0.3, + max_gpu_total_bytes=total, + free_gpu_memory_fraction=0.8, + ) + + creator = _make_creator(config, is_enc_dec=True) + self_config, cross_config = creator._split_kv_cache_budget_for_cross() + + expected_cross = int(total * 0.3) + expected_self = total - expected_cross + assert cross_config.max_gpu_total_bytes == expected_cross + assert self_config.max_gpu_total_bytes == expected_self + assert cross_config.free_gpu_memory_fraction == pytest.approx(0.24) + assert self_config.free_gpu_memory_fraction == pytest.approx(0.56) + assert config.max_gpu_total_bytes == total + assert config.free_gpu_memory_fraction == pytest.approx(0.8) + + def test_no_split_when_fraction_is_none(self): + total = 10 * (1 << 30) + config = _make_mock_kv_cache_config(cross_kv_cache_fraction=None, max_gpu_total_bytes=total) + + creator = _make_creator(config, is_enc_dec=True) + with pytest.raises(ValueError, match="cross_kv_cache_fraction"): + creator._split_kv_cache_budget_for_cross() + + def test_split_free_fraction_when_budget_is_none(self): + config = _make_mock_kv_cache_config( + cross_kv_cache_fraction=0.5, + max_gpu_total_bytes=None, + max_tokens=1000, + free_gpu_memory_fraction=0.8, + ) + + creator = _make_creator(config, is_enc_dec=True) + self_config, cross_config = creator._split_kv_cache_budget_for_cross() + + assert cross_config.max_tokens == 1000 + assert self_config.max_tokens == 1000 + assert cross_config.free_gpu_memory_fraction == pytest.approx(0.4) + assert self_config.free_gpu_memory_fraction == pytest.approx(0.4) + assert config.free_gpu_memory_fraction == pytest.approx(0.8) + + def test_split_free_fraction_when_budget_is_zero(self): + config = _make_mock_kv_cache_config( + cross_kv_cache_fraction=0.5, + max_gpu_total_bytes=0, + max_tokens=1000, + free_gpu_memory_fraction=0.8, + ) + + creator = _make_creator(config, is_enc_dec=True) + self_config, cross_config = creator._split_kv_cache_budget_for_cross() + + assert cross_config.max_tokens == 1000 + assert self_config.max_tokens == 1000 + assert cross_config.free_gpu_memory_fraction == pytest.approx(0.4) + assert self_config.free_gpu_memory_fraction == pytest.approx(0.4) + assert config.free_gpu_memory_fraction == pytest.approx(0.8) + + def test_raises_when_no_budget_source_exists(self): + config = _make_mock_kv_cache_config( + cross_kv_cache_fraction=0.5, + max_gpu_total_bytes=0, + max_tokens=None, + free_gpu_memory_fraction=None, + ) + + creator = _make_creator(config, is_enc_dec=True) + with pytest.raises(ValueError, match="Unable to size"): + creator._split_kv_cache_budget_for_cross() + + def test_is_encoder_decoder_helper(self): + dec_config = _make_mock_model_config(is_encoder_decoder=False) + dec_creator = _make_creator(_make_mock_kv_cache_config(), model_config=dec_config) + assert not dec_creator._is_encoder_decoder() + + enc_dec_config = _make_mock_model_config(is_encoder_decoder=True) + enc_dec_creator = _make_creator(_make_mock_kv_cache_config(), model_config=enc_dec_config) + assert enc_dec_creator._is_encoder_decoder() + + def test_budgets_sum_to_total(self): + """Self + cross budgets always sum to the original total.""" + total = 7 * (1 << 30) + 123 # non-round number + config = _make_mock_kv_cache_config(cross_kv_cache_fraction=0.4, max_gpu_total_bytes=total) + + creator = _make_creator(config, is_enc_dec=True) + self_config, cross_config = creator._split_kv_cache_budget_for_cross() + + assert (self_config.max_gpu_total_bytes + cross_config.max_gpu_total_bytes) == total + assert config.max_gpu_total_bytes == total + + def test_host_cache_budget_is_split_without_mutating_base_config(self): + """Self + cross host cache budgets sum to the original host budget.""" + total_host = 7 * (1 << 30) + 123 + config = _make_mock_kv_cache_config( + cross_kv_cache_fraction=0.4, + max_gpu_total_bytes=8 * (1 << 30), + host_cache_size=total_host, + ) + + creator = _make_creator(config, is_enc_dec=True) + self_config, cross_config = creator._split_kv_cache_budget_for_cross() + + expected_cross_host = int(total_host * 0.4) + expected_self_host = total_host - expected_cross_host + assert cross_config.host_cache_size == expected_cross_host + assert self_config.host_cache_size == expected_self_host + assert (self_config.host_cache_size + cross_config.host_cache_size) == total_host + assert config.host_cache_size == total_host + + def test_host_cache_budget_counts_as_split_budget_source(self): + total_host = 4 * (1 << 30) + config = _make_mock_kv_cache_config( + cross_kv_cache_fraction=0.25, + max_gpu_total_bytes=None, + free_gpu_memory_fraction=None, + host_cache_size=total_host, + ) + + creator = _make_creator(config, is_enc_dec=True) + self_config, cross_config = creator._split_kv_cache_budget_for_cross() + + assert cross_config.host_cache_size == total_host // 4 + assert self_config.host_cache_size == total_host - total_host // 4 + assert config.host_cache_size == total_host + + +# --------------------------------------------------------------------------- +# Tests: ResourceManagerType enum +# --------------------------------------------------------------------------- + + +class TestResourceManagerType: + """Verify CROSS_KV_CACHE_MANAGER exists in the enum.""" + + def test_cross_kv_cache_manager_in_enum(self): + assert ResourceManagerType.CROSS_KV_CACHE_MANAGER.value == "CROSS_KV_CACHE_MANAGER" + + +# --------------------------------------------------------------------------- +# Tests: Cross-pool geometry and build_managers coverage +# --------------------------------------------------------------------------- + + +class TestCrossKvCacheConstruction: + """Exercise the Steps 4 and 5 construction path beyond helper math.""" + + @pytest.mark.parametrize("use_kv_cache_manager_v2", [False, True]) + def test_create_cross_kv_cache_manager_uses_encoder_geometry(self, use_kv_cache_manager_v2): + expected_cls = KVCacheManagerV2 if use_kv_cache_manager_v2 else KVCacheManager + config = _make_mock_kv_cache_config( + cross_kv_cache_fraction=0.5, + max_gpu_total_bytes=8 * (1 << 30), + use_kv_cache_manager_v2=use_kv_cache_manager_v2, + ) + model_config = _make_mock_model_config( + is_encoder_decoder=True, + num_hidden_layers=10, + num_attention_heads=16, + num_key_value_heads=16, + hidden_size=768, + head_dim=48, + encoder_layers=8, + decoder_layers=10, + encoder_attention_heads=12, + d_model=768, + max_position_embeddings=1024, + ) + creator = _make_creator(config, model_config=model_config, manager_cls=expected_cls) + cross_cfg = config.model_copy() + + with patch( + "tensorrt_llm._torch.pyexecutor._util._create_kv_cache_manager", + return_value=Mock(), + ) as create_mock: + creator._create_cross_kv_cache_manager(cross_cfg) + + kwargs = create_mock.call_args.kwargs + # Cross pool must use the same manager class as the self pool so + # both pools share the same runtime ABI. V1 is the default and + # production target; V2 is an additive secondary path. + assert kwargs["kv_cache_manager_cls"] is expected_cls + assert kwargs["num_layers"] == 10 + assert kwargs["num_kv_heads"] == 12 + assert kwargs["head_dim"] == 64 + assert kwargs["max_seq_len"] == 1024 + + import tensorrt_llm + + assert kwargs["kv_cache_type"] == ( + tensorrt_llm.bindings.internal.batch_manager.CacheType.CROSS + ) + + def test_cross_layout_uses_max_input_len_for_encoder_capacity(self): + config = _make_mock_kv_cache_config( + cross_kv_cache_fraction=0.5, + max_gpu_total_bytes=8 * (1 << 30), + ) + model_config = _make_mock_model_config( + is_encoder_decoder=True, + max_position_embeddings=4096, + ) + creator = _make_creator(config, model_config=model_config) + creator._llm_args.max_input_len = 1536 + creator._max_seq_len = 864 + + _, _, _, max_seq_len = creator._get_cross_kv_cache_layout(fallback_max_seq_len=2048) + + assert max_seq_len == 1536 + + def test_build_managers_cross_pool_ignores_mutated_self_max_seq_len(self): + config = _make_mock_kv_cache_config( + cross_kv_cache_fraction=0.5, + max_gpu_total_bytes=8 * (1 << 30), + ) + model_config = _make_mock_model_config( + is_encoder_decoder=True, + max_position_embeddings=4096, + ) + creator = _make_creator(config, model_config=model_config) + creator._llm_args.max_input_len = None + creator._max_seq_len = 2048 + creator.configure_kv_cache_capacity = Mock() + creator._should_create_separate_draft_kv_cache = Mock(return_value=False) + + def create_self_manager(*_args, **_kwargs): + creator._max_seq_len = 864 + manager = Mock() + manager.max_seq_len = 864 + return manager + + captured_cross_max_seq_lens = [] + + def create_cross_manager(*_args, **kwargs): + captured_cross_max_seq_lens.append(kwargs["max_seq_len"]) + manager = Mock() + manager.max_seq_len = kwargs["max_seq_len"] + return manager + + creator._create_kv_cache_manager = Mock(side_effect=create_self_manager) + with patch( + "tensorrt_llm._torch.pyexecutor._util._create_kv_cache_manager", + side_effect=create_cross_manager, + ): + creator.build_managers({}, estimating_kv_cache=False) + + assert creator._max_seq_len == 864 + assert captured_cross_max_seq_lens == [2048] + + def test_get_kv_size_per_token_includes_cross_pool_for_enc_dec(self): + config = _make_mock_kv_cache_config( + cross_kv_cache_fraction=0.5, max_gpu_total_bytes=8 * (1 << 30) + ) + model_config = _make_mock_model_config( + is_encoder_decoder=True, + num_hidden_layers=10, + num_attention_heads=16, + num_key_value_heads=16, + hidden_size=768, + head_dim=48, + encoder_attention_heads=12, + decoder_layers=10, + d_model=768, + ) + creator = _make_creator(config, model_config=model_config) + + with patch.object( + creator._kv_cache_manager_cls, + "get_cache_size_per_token", + side_effect=[100, 40], + ) as get_size_mock: + kv_size = creator._get_kv_size_per_token() + + assert kv_size.slope == 140 + assert kv_size.intercept == 0 + assert get_size_mock.call_count == 2 + + cross_call = get_size_mock.call_args_list[1] + proxy_model_config = cross_call.args[0] + assert proxy_model_config.pretrained_config.num_key_value_heads == 12 + assert proxy_model_config.pretrained_config.num_attention_heads == 12 + assert proxy_model_config.pretrained_config.head_dim == 64 + assert cross_call.kwargs["num_layers"] == 10 + + @pytest.mark.parametrize("use_kv_cache_manager_v2", [False, True]) + def test_build_managers_registers_cross_pool_for_enc_dec(self, use_kv_cache_manager_v2): + creator = _make_creator( + _make_mock_kv_cache_config( + cross_kv_cache_fraction=0.5, + max_gpu_total_bytes=8 * (1 << 30), + use_kv_cache_manager_v2=use_kv_cache_manager_v2, + ), + is_enc_dec=True, + ) + creator.configure_kv_cache_capacity = Mock() + creator._should_create_separate_draft_kv_cache = Mock(return_value=False) + creator._split_kv_cache_budget_for_cross = Mock(return_value=(Mock(), Mock())) + creator._create_kv_cache_manager = Mock(return_value=Mock()) + creator._create_cross_kv_cache_manager = Mock(return_value=Mock()) + + resources = {} + creator.build_managers(resources, estimating_kv_cache=False) + + creator._create_cross_kv_cache_manager.assert_called_once() + + @pytest.mark.parametrize("use_kv_cache_manager_v2", [False, True]) + def test_build_managers_registers_cross_pool_for_enc_dec_estimation( + self, use_kv_cache_manager_v2 + ): + creator = _make_creator( + _make_mock_kv_cache_config( + cross_kv_cache_fraction=0.5, + max_gpu_total_bytes=0, + max_tokens=1024, + free_gpu_memory_fraction=0.8, + use_kv_cache_manager_v2=use_kv_cache_manager_v2, + ), + is_enc_dec=True, + ) + creator.configure_kv_cache_capacity = Mock() + creator._should_create_separate_draft_kv_cache = Mock(return_value=False) + creator._split_kv_cache_budget_for_cross = Mock(return_value=(Mock(), Mock())) + creator._create_kv_cache_manager = Mock(return_value=Mock()) + creator._create_cross_kv_cache_manager = Mock(return_value=Mock()) + + resources = {} + creator.build_managers(resources, estimating_kv_cache=True) + + creator._create_cross_kv_cache_manager.assert_called_once() + + @pytest.mark.parametrize("use_kv_cache_manager_v2", [False, True]) + def test_build_managers_uses_split_cross_budget_without_mutating_base_config( + self, use_kv_cache_manager_v2 + ): + total_budget = 10 * (1 << 30) + creator = _make_creator( + _make_mock_kv_cache_config( + cross_kv_cache_fraction=0.5, + max_gpu_total_bytes=total_budget, + free_gpu_memory_fraction=0.9, + use_kv_cache_manager_v2=use_kv_cache_manager_v2, + ), + is_enc_dec=True, + ) + creator.configure_kv_cache_capacity = Mock() + creator._should_create_separate_draft_kv_cache = Mock(return_value=False) + + self_budgets = [] + cross_budgets = [] + + def create_self_manager(*_args, **kwargs): + self_cfg = kwargs["kv_cache_config_override"] + self_budgets.append( + ( + self_cfg.free_gpu_memory_fraction, + self_cfg.max_gpu_total_bytes, + ) + ) + return Mock() + + def create_cross_manager(cross_cfg, *_args, **_kwargs): + cross_budgets.append( + ( + cross_cfg.free_gpu_memory_fraction, + cross_cfg.max_gpu_total_bytes, + ) + ) + return Mock() + + creator._create_kv_cache_manager = Mock(side_effect=create_self_manager) + creator._create_cross_kv_cache_manager = Mock(side_effect=create_cross_manager) + + resources = {} + creator.build_managers(resources, estimating_kv_cache=True) + + assert creator._kv_cache_config.free_gpu_memory_fraction == pytest.approx(0.9) + assert creator._kv_cache_config.max_gpu_total_bytes == total_budget + + creator.build_managers(resources, estimating_kv_cache=False) + + expected_split = total_budget // 2 + assert self_budgets == [ + (pytest.approx(0.45), expected_split), + (pytest.approx(0.45), expected_split), + ] + assert cross_budgets == [ + (pytest.approx(0.45), expected_split), + (pytest.approx(0.45), expected_split), + ] + + def test_build_managers_skips_cross_pool_for_decoder_only(self): + creator = _make_creator( + _make_mock_kv_cache_config( + cross_kv_cache_fraction=None, + max_gpu_total_bytes=8 * (1 << 30), + ), + is_enc_dec=False, + ) + creator.configure_kv_cache_capacity = Mock() + creator._should_create_separate_draft_kv_cache = Mock(return_value=False) + creator._split_kv_cache_budget_for_cross = Mock() + creator._create_kv_cache_manager = Mock(return_value=Mock()) + creator._create_cross_kv_cache_manager = Mock() + + resources = {} + creator.build_managers(resources, estimating_kv_cache=False) + + creator._split_kv_cache_budget_for_cross.assert_not_called() + creator._create_cross_kv_cache_manager.assert_not_called() + assert resources[ResourceManagerType.CROSS_KV_CACHE_MANAGER] is None + + +# --------------------------------------------------------------------------- +# Tests: KVCacheV2Scheduler cross_kv_cache_manager parameter +# --------------------------------------------------------------------------- + + +class TestKVCacheV2SchedulerCrossParam: + """KVCacheV2Scheduler should accept and store cross_kv_cache_manager.""" + + def _make_mock_kv_mgr(self, tokens_per_block=64): + mgr = Mock(spec=KVCacheManagerV2) + mgr.tokens_per_block = tokens_per_block + return mgr + + def test_default_cross_is_none(self): + from tensorrt_llm._torch.pyexecutor.scheduler.scheduler_v2 import KVCacheV2Scheduler + + kv_mgr = self._make_mock_kv_mgr() + scheduler = KVCacheV2Scheduler( + max_batch_size=8, + max_num_tokens=4096, + kv_cache_manager=kv_mgr, + scheduler_policy=CapacitySchedulerPolicy.MAX_UTILIZATION, + ) + assert scheduler.cross_kv_cache_manager is None + + def test_cross_kv_cache_manager_is_stored(self): + from tensorrt_llm._torch.pyexecutor.scheduler.scheduler_v2 import KVCacheV2Scheduler + + kv_mgr = self._make_mock_kv_mgr() + cross_mgr = self._make_mock_kv_mgr() + scheduler = KVCacheV2Scheduler( + max_batch_size=8, + max_num_tokens=4096, + kv_cache_manager=kv_mgr, + scheduler_policy=CapacitySchedulerPolicy.MAX_UTILIZATION, + cross_kv_cache_manager=cross_mgr, + ) + assert scheduler.cross_kv_cache_manager is cross_mgr + + def test_factory_forwards_encoder_init_until_state_for_cross_pool(self): + """The executor factory must widen V2 scheduling to ENCODER_INIT. + + Without this, V2 enc-dec requests are filtered by the default + CONTEXT_INIT state gate before the encoder loop can see them. + """ + from tensorrt_llm._torch.pyexecutor._util import create_py_executor_instance + from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequestState + + kv_mgr = Mock() + kv_mgr.tokens_per_block = 64 + cross_mgr = Mock() + resources = { + ResourceManagerType.KV_CACHE_MANAGER: kv_mgr, + ResourceManagerType.CROSS_KV_CACHE_MANAGER: cross_mgr, + ResourceManagerType.DRAFT_KV_CACHE_MANAGER: None, + } + mapping = SimpleNamespace( + pp_size=1, + enable_attention_dp=False, + has_pp=lambda: False, + ) + model_engine = SimpleNamespace( + spec_config=None, + model=SimpleNamespace( + model_config=SimpleNamespace( + pretrained_config=SimpleNamespace( + kv_lora_rank=None, + qk_rope_head_dim=None, + ), + ), + ), + ) + llm_args = SimpleNamespace( + extra_resource_managers={}, + disable_overlap_scheduler=True, + enable_early_first_token_response=False, + kv_cache_config=SimpleNamespace(enable_kv_pool_rebalance=False), + ) + + with ( + patch( + "tensorrt_llm._torch.pyexecutor._util.KVCacheManagerV2", + new=Mock, + ), + patch( + "tensorrt_llm._torch.pyexecutor._util.KVCacheV2Scheduler", + ) as scheduler_cls, + patch( + "tensorrt_llm._torch.pyexecutor._util.create_kv_cache_transceiver", + return_value=None, + ), + patch( + "tensorrt_llm._torch.pyexecutor._util.PyExecutor", + ), + ): + scheduler_cls.return_value = Mock() + create_py_executor_instance( + dist=Mock(), + resources=resources, + mapping=mapping, + llm_args=llm_args, + ctx_chunk_config=None, + model_engine=model_engine, + start_worker=False, + sampler=Mock(), + drafter=None, + max_seq_len=128, + max_batch_size=8, + max_beam_width=1, + max_num_tokens=4096, + ) + + kwargs = scheduler_cls.call_args.kwargs + assert kwargs["cross_kv_cache_manager"] is cross_mgr + assert kwargs["no_schedule_until_state"] == LlmRequestState.ENCODER_INIT + + +# --------------------------------------------------------------------------- +# Tests: V1 scheduler cross_kv_cache_manager wiring. +# --------------------------------------------------------------------------- + + +class TestBindCapacitySchedulerCrossParam: + """C++-bound V1 ``BindCapacityScheduler`` exposes cross-KV wiring. + + The C++ ``CapacityScheduler`` already accepts a cross manager. The Python + wrapper forwards the cross pool and the ENCODER_INIT gating. + """ + + def test_default_cross_is_none_and_default_until_state(self): + from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequestState + from tensorrt_llm._torch.pyexecutor.scheduler.scheduler import BindCapacityScheduler + + with patch( + "tensorrt_llm._torch.pyexecutor.scheduler.scheduler.tb_internal.algorithms.CapacityScheduler" + ) as cap_cls: + cap_cls.return_value = Mock() + scheduler = BindCapacityScheduler( + max_num_requests=8, + kv_cache_manager=Mock(), + peft_cache_manager=None, + scheduler_policy=CapacitySchedulerPolicy.MAX_UTILIZATION, + ) + + assert scheduler.cross_kv_cache_manager is None + kwargs = cap_cls.call_args.kwargs + assert kwargs["no_schedule_until_state"] == LlmRequestState.CONTEXT_INIT + + def test_cross_kv_cache_manager_and_until_state_are_forwarded(self): + from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequestState + from tensorrt_llm._torch.pyexecutor.scheduler.scheduler import BindCapacityScheduler + + cross_mgr = Mock() + kv_mgr = Mock() + with patch( + "tensorrt_llm._torch.pyexecutor.scheduler.scheduler.tb_internal.algorithms.CapacityScheduler" + ) as cap_cls: + impl = Mock() + cap_cls.return_value = impl + scheduler = BindCapacityScheduler( + max_num_requests=8, + kv_cache_manager=kv_mgr, + peft_cache_manager=None, + scheduler_policy=CapacitySchedulerPolicy.MAX_UTILIZATION, + cross_kv_cache_manager=cross_mgr, + no_schedule_until_state=LlmRequestState.ENCODER_INIT, + ) + + # Construction forwarded the gating to the C++ binding. + ctor_kwargs = cap_cls.call_args.kwargs + assert ctor_kwargs["no_schedule_until_state"] == LlmRequestState.ENCODER_INIT + + # schedule_request must forward the cross manager to the C++ + # __call__ so the dual-pool scheduling logic activates. + impl.return_value = ([], [], []) + scheduler.schedule_request([]) + impl.assert_called_once_with([], kv_mgr, None, cross_mgr) + + +class TestSimpleUnifiedSchedulerCrossParam: + """V1 Python ``SimpleUnifiedScheduler`` exposes cross-KV wiring.""" + + def test_cross_kv_cache_manager_and_until_state_are_forwarded(self): + from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequestState + from tensorrt_llm._torch.pyexecutor.scheduler.scheduler import SimpleUnifiedScheduler + + kv_mgr = Mock() + kv_mgr.is_variable_window = False + kv_mgr.enable_block_reuse = False + cross_mgr = Mock() + cross_mgr.is_variable_window = False + cross_mgr.enable_block_reuse = False + + scheduler = SimpleUnifiedScheduler( + max_batch_size=8, + max_num_tokens=4096, + kv_cache_manager=kv_mgr, + peft_cache_manager=None, + scheduler_policy=CapacitySchedulerPolicy.MAX_UTILIZATION, + cross_kv_cache_manager=cross_mgr, + no_schedule_until_state=LlmRequestState.ENCODER_INIT, + ) + + assert scheduler.capacity_scheduler.cross_kv_cache_manager is cross_mgr + assert scheduler.capacity_scheduler.no_schedule_until_state == LlmRequestState.ENCODER_INIT + assert ( + scheduler.micro_batch_scheduler.no_schedule_until_state == LlmRequestState.ENCODER_INIT + ) + + +# --------------------------------------------------------------------------- +# Tests: V1 dual-pool smoke test. +# --------------------------------------------------------------------------- + + +class TestV1DualPoolSmoke: + """Smoke test exercising V1 dual-pool construction. + + Constructs both pools as V1 ``KVCacheManager`` instances with + ``CacheType.SELF`` / ``CacheType.CROSS`` (via mocked + ``_create_kv_cache_manager``) and verifies that ``build_managers`` + wires both pools into the resource map for the V1 production path. + + Running an actual encoder + decoder context iteration requires GPUs + and a full model engine; that lives in the integration suite. Here + we verify the V1 construction wiring with mocks consistent with the + rest of this file. + """ + + def test_build_managers_uses_v1_kv_cache_manager_for_both_pools(self): + kv_cache_config = _make_mock_kv_cache_config( + cross_kv_cache_fraction=0.5, + max_gpu_total_bytes=8 * (1 << 30), + use_kv_cache_manager_v2=False, + ) + creator = _make_creator(kv_cache_config, is_enc_dec=True, manager_cls=KVCacheManager) + creator.configure_kv_cache_capacity = Mock() + creator._should_create_separate_draft_kv_cache = Mock(return_value=False) + creator._split_kv_cache_budget_for_cross = Mock(return_value=(Mock(), Mock())) + + # Both _create_kv_cache_manager (self pool) and + # _create_cross_kv_cache_manager are exercised through the + # underlying free-function _create_kv_cache_manager so we can + # assert the manager_cls and CacheType for each call. + import tensorrt_llm + + cache_type_self = tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF + cache_type_cross = tensorrt_llm.bindings.internal.batch_manager.CacheType.CROSS + + # Stub the self-pool path (_create_kv_cache_manager method) to + # avoid invoking the heavyweight free function. + self_mgr = Mock(spec=KVCacheManager) + self_mgr.kv_cache_type = cache_type_self + creator._create_kv_cache_manager = Mock(return_value=self_mgr) + + cross_mgr = Mock(spec=KVCacheManager) + cross_mgr.kv_cache_type = cache_type_cross + with patch( + "tensorrt_llm._torch.pyexecutor._util._create_kv_cache_manager", + return_value=cross_mgr, + ) as create_mock: + resources = {} + creator.build_managers(resources, estimating_kv_cache=False) + + # Self pool: registered as KV_CACHE_MANAGER. + assert resources[ResourceManagerType.KV_CACHE_MANAGER] is self_mgr + + # Cross pool: registered as CROSS_KV_CACHE_MANAGER and built + # with the V1 KVCacheManager class + CacheType.CROSS. + assert resources[ResourceManagerType.CROSS_KV_CACHE_MANAGER] is cross_mgr + cross_kwargs = create_mock.call_args.kwargs + assert cross_kwargs["kv_cache_manager_cls"] is KVCacheManager + assert cross_kwargs["kv_cache_type"] == cache_type_cross diff --git a/tests/unittest/_torch/executor/test_kv_cache_budget_split.py b/tests/unittest/_torch/executor/test_kv_cache_budget_split.py index 3c0cd8f3b491..8aa9ad22196d 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_budget_split.py +++ b/tests/unittest/_torch/executor/test_kv_cache_budget_split.py @@ -66,27 +66,39 @@ class TestSplitGpuBudgetForDraft: def test_gpu_budget_split_proportionally(self): total_gpu = 10 * GB c = _make_creator( - max_gpu_total_bytes=total_gpu, total_kv_per_token=100, target_kv_per_token=80 + max_gpu_total_bytes=total_gpu, + total_kv_per_token=100, + target_kv_per_token=80, ) - draft_config = c._split_kv_cache_budget_for_draft("max_gpu_total_bytes") + target_config, draft_config = c._split_kv_cache_budget_for_draft("max_gpu_total_bytes") assert draft_config is not None - assert c._kv_cache_config.max_gpu_total_bytes == 8 * GB + assert target_config.max_gpu_total_bytes == 8 * GB assert draft_config.max_gpu_total_bytes == 2 * GB + assert target_config.host_cache_size is None + assert c._kv_cache_config.max_gpu_total_bytes == total_gpu assert c._kv_cache_config.host_cache_size is None def test_returns_none_when_no_gpu_budget(self): c = _make_creator(max_gpu_total_bytes=0) - assert c._split_kv_cache_budget_for_draft("max_gpu_total_bytes") is None + target_config, draft_config = c._split_kv_cache_budget_for_draft("max_gpu_total_bytes") + + assert target_config is c._kv_cache_config + assert draft_config is None def test_returns_none_when_draft_kv_zero(self): c = _make_creator( - max_gpu_total_bytes=10 * GB, total_kv_per_token=100, target_kv_per_token=100 + max_gpu_total_bytes=10 * GB, + total_kv_per_token=100, + target_kv_per_token=100, ) - assert c._split_kv_cache_budget_for_draft("max_gpu_total_bytes") is None + target_config, draft_config = c._split_kv_cache_budget_for_draft("max_gpu_total_bytes") + + assert target_config is c._kv_cache_config + assert draft_config is None class TestSplitHostCacheBudgetForDraft: @@ -100,11 +112,13 @@ def test_host_budget_split_proportionally(self): target_kv_per_token=80, ) - draft_config = c._split_kv_cache_budget_for_draft("host_cache_size") + target_config, draft_config = c._split_kv_cache_budget_for_draft("host_cache_size") assert draft_config is not None - assert c._kv_cache_config.host_cache_size == 16 * GB + assert target_config.host_cache_size == 16 * GB assert draft_config.host_cache_size == 4 * GB + assert target_config.max_gpu_total_bytes == total_gpu + assert c._kv_cache_config.host_cache_size == total_host assert c._kv_cache_config.max_gpu_total_bytes == total_gpu def test_host_budget_not_doubled(self): @@ -117,10 +131,10 @@ def test_host_budget_not_doubled(self): target_kv_per_token=80, ) - draft_config = c._split_kv_cache_budget_for_draft("host_cache_size") + target_config, draft_config = c._split_kv_cache_budget_for_draft("host_cache_size") assert draft_config is not None - assert (c._kv_cache_config.host_cache_size + draft_config.host_cache_size) == total_host + assert (target_config.host_cache_size + draft_config.host_cache_size) == total_host def test_host_split_without_gpu_budget_uses_slope_ratio(self): """V1 non-VSWA: host split must not depend on max_gpu_total_bytes.""" @@ -132,10 +146,10 @@ def test_host_split_without_gpu_budget_uses_slope_ratio(self): target_kv_per_token=80, ) - draft_config = c._split_kv_cache_budget_for_draft("host_cache_size") + target_config, draft_config = c._split_kv_cache_budget_for_draft("host_cache_size") assert draft_config is not None - assert c._kv_cache_config.host_cache_size == 16 * GB + assert target_config.host_cache_size == 16 * GB assert draft_config.host_cache_size == 4 * GB def test_host_split_merges_into_existing_draft_config(self): @@ -148,13 +162,18 @@ def test_host_split_merges_into_existing_draft_config(self): target_kv_per_token=80, ) - draft_config = c._split_kv_cache_budget_for_draft("max_gpu_total_bytes") - draft_config = c._split_kv_cache_budget_for_draft("host_cache_size", draft_config) + target_config, draft_config = c._split_kv_cache_budget_for_draft("max_gpu_total_bytes") + target_config, draft_config = c._split_kv_cache_budget_for_draft( + "host_cache_size", target_config, draft_config + ) + assert draft_config is not None assert draft_config.max_gpu_total_bytes == 2 * GB - assert c._kv_cache_config.max_gpu_total_bytes == 8 * GB + assert target_config.max_gpu_total_bytes == 8 * GB assert draft_config.host_cache_size == 4 * GB - assert c._kv_cache_config.host_cache_size == 16 * GB + assert target_config.host_cache_size == 16 * GB + assert c._kv_cache_config.max_gpu_total_bytes == total_gpu + assert c._kv_cache_config.host_cache_size == total_host def test_host_split_after_gpu_split_is_unaffected_by_target_only_gpu_budget(self): """Regression: host split used to read max_gpu_total_bytes (already @@ -170,10 +189,13 @@ def test_host_split_after_gpu_split_is_unaffected_by_target_only_gpu_budget(self target_kv_per_token=80, ) - draft_config = c._split_kv_cache_budget_for_draft("max_gpu_total_bytes") - draft_config = c._split_kv_cache_budget_for_draft("host_cache_size", draft_config) + target_config, draft_config = c._split_kv_cache_budget_for_draft("max_gpu_total_bytes") + target_config, draft_config = c._split_kv_cache_budget_for_draft( + "host_cache_size", target_config, draft_config + ) - assert c._kv_cache_config.host_cache_size == 16 * GB + assert draft_config is not None + assert target_config.host_cache_size == 16 * GB assert draft_config.host_cache_size == 4 * GB def test_no_host_cache_leaves_none(self): @@ -184,7 +206,10 @@ def test_no_host_cache_leaves_none(self): target_kv_per_token=80, ) - assert c._split_kv_cache_budget_for_draft("host_cache_size") is None + target_config, draft_config = c._split_kv_cache_budget_for_draft("host_cache_size") + + assert target_config is c._kv_cache_config + assert draft_config is None def test_zero_host_cache_unchanged(self): c = _make_creator( @@ -194,25 +219,28 @@ def test_zero_host_cache_unchanged(self): target_kv_per_token=80, ) - assert c._split_kv_cache_budget_for_draft("host_cache_size") is None + target_config, draft_config = c._split_kv_cache_budget_for_draft("host_cache_size") + + assert target_config is c._kv_cache_config + assert draft_config is None @pytest.mark.parametrize("target_frac", [0.5, 0.75, 0.9, 0.95]) def test_various_ratios(self, target_frac): - total_gpu = 10 * GB total_host = 20 * GB total_kv = 1000 target_kv = int(total_kv * target_frac) c = _make_creator( - max_gpu_total_bytes=total_gpu, + max_gpu_total_bytes=10 * GB, host_cache_size=total_host, total_kv_per_token=total_kv, target_kv_per_token=target_kv, ) - draft_config = c._split_kv_cache_budget_for_draft("host_cache_size") + target_config, draft_config = c._split_kv_cache_budget_for_draft("host_cache_size") - assert (c._kv_cache_config.host_cache_size + draft_config.host_cache_size) == total_host + assert draft_config is not None + assert (target_config.host_cache_size + draft_config.host_cache_size) == total_host def test_budgets_sum_to_original_with_gpu_and_host(self): total_gpu = 15 * GB @@ -224,20 +252,20 @@ def test_budgets_sum_to_original_with_gpu_and_host(self): target_kv_per_token=700, ) - draft_config = c._split_kv_cache_budget_for_draft("max_gpu_total_bytes") - draft_config = c._split_kv_cache_budget_for_draft("host_cache_size", draft_config) + target_config, draft_config = c._split_kv_cache_budget_for_draft("max_gpu_total_bytes") + target_config, draft_config = c._split_kv_cache_budget_for_draft( + "host_cache_size", target_config, draft_config + ) - assert ( - c._kv_cache_config.max_gpu_total_bytes + draft_config.max_gpu_total_bytes - ) == total_gpu - assert (c._kv_cache_config.host_cache_size + draft_config.host_cache_size) == total_host + assert draft_config is not None + assert (target_config.max_gpu_total_bytes + draft_config.max_gpu_total_bytes) == total_gpu + assert (target_config.host_cache_size + draft_config.host_cache_size) == total_host + assert c._kv_cache_config.max_gpu_total_bytes == total_gpu + assert c._kv_cache_config.host_cache_size == total_host class TestHostSplitIgnoresGpuFixedCost: - """The fixed (intercept) cost models GPU-resident state (e.g. mamba SSM - state) and is not charged against host offload memory. The host split must - therefore stay proportional to the per-token (slope) cost even when the - GPU-resident fixed cost dwarfs the host budget.""" + """The fixed cost models GPU-resident state and is not host memory.""" def test_host_split_proportional_despite_large_intercept(self): total_host = 10 * GB @@ -246,14 +274,13 @@ def test_host_split_proportional_despite_large_intercept(self): host_cache_size=total_host, total_kv_per_token=100, target_kv_per_token=80, - total_kv_intercept=50 * GB, # huge GPU fixed cost, irrelevant to host + total_kv_intercept=50 * GB, ) - draft_config = c._split_kv_cache_budget_for_draft("host_cache_size") + target_config, draft_config = c._split_kv_cache_budget_for_draft("host_cache_size") - # Intercept ignored for host -> proportional on slope (draft 20/100). assert draft_config is not None - assert c._kv_cache_config.host_cache_size == 8 * GB + assert target_config.host_cache_size == 8 * GB assert draft_config.host_cache_size == 2 * GB def test_host_split_sums_to_original_despite_large_intercept(self): @@ -266,14 +293,14 @@ def test_host_split_sums_to_original_despite_large_intercept(self): total_kv_intercept=100 * GB, ) - draft_config = c._split_kv_cache_budget_for_draft("host_cache_size") + target_config, draft_config = c._split_kv_cache_budget_for_draft("host_cache_size") - assert (c._kv_cache_config.host_cache_size + draft_config.host_cache_size) == total_host + assert draft_config is not None + assert (target_config.host_cache_size + draft_config.host_cache_size) == total_host class TestGpuSplitChargesFixedCost: - """``max_gpu_total_bytes`` is where the GPU-resident fixed cost lives, so it - is charged the intercept and fails fast when the budget can't fit it.""" + """``max_gpu_total_bytes`` carries the GPU-resident fixed cost.""" def test_gpu_split_subtracts_intercept(self): total_gpu = 10 * GB @@ -281,44 +308,40 @@ def test_gpu_split_subtracts_intercept(self): max_gpu_total_bytes=total_gpu, total_kv_per_token=100, target_kv_per_token=80, - total_kv_intercept=5 * GB, # draft intercept = 5 - 0 = 5 GB + total_kv_intercept=5 * GB, target_kv_intercept=0, ) - draft_config = c._split_kv_cache_budget_for_draft("max_gpu_total_bytes") + target_config, draft_config = c._split_kv_cache_budget_for_draft("max_gpu_total_bytes") - # slope_budget = 10 - 5 = 5 GB; draft slope share = 5 * 20/100 = 1 GB; - # draft_budget = draft_intercept (5) + 1 = 6 GB; target = 4 GB. + assert draft_config is not None assert draft_config.max_gpu_total_bytes == 6 * GB - assert c._kv_cache_config.max_gpu_total_bytes == 4 * GB + assert target_config.max_gpu_total_bytes == 4 * GB def test_gpu_split_infeasible_raises(self): - """A GPU budget too small for the combined fixed cost is fatal (the run - would OOM), so the split must fail fast instead of degrading.""" - total_gpu = 1 * GB + """A GPU budget too small for fixed cost must fail fast.""" c = _make_creator( - max_gpu_total_bytes=total_gpu, + max_gpu_total_bytes=1 * GB, total_kv_per_token=100, target_kv_per_token=80, - total_kv_intercept=2 * GB, # fixed cost exceeds the gpu budget + total_kv_intercept=2 * GB, ) with pytest.raises(ValueError, match="GPU budget"): c._split_kv_cache_budget_for_draft("max_gpu_total_bytes") def test_gpu_raise_does_not_block_subsequent_host_split(self): - """GPU split raises on infeasible budget, but a host-only split with the - same large intercept still succeeds proportionally.""" total_host = 10 * GB c = _make_creator( - max_gpu_total_bytes=0, # no gpu split attempted + max_gpu_total_bytes=0, host_cache_size=total_host, total_kv_per_token=100, target_kv_per_token=80, total_kv_intercept=2 * GB, ) - draft_config = c._split_kv_cache_budget_for_draft("host_cache_size") + target_config, draft_config = c._split_kv_cache_budget_for_draft("host_cache_size") - assert c._kv_cache_config.host_cache_size == 8 * GB + assert draft_config is not None + assert target_config.host_cache_size == 8 * GB assert draft_config.host_cache_size == 2 * GB diff --git a/tests/unittest/_torch/executor/test_kv_cache_v2_scheduler.py b/tests/unittest/_torch/executor/test_kv_cache_v2_scheduler.py index d1168837b619..be9446523727 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_v2_scheduler.py +++ b/tests/unittest/_torch/executor/test_kv_cache_v2_scheduler.py @@ -54,6 +54,7 @@ def make_gen_request( req.is_context_init_state = False req.is_generation_in_progress_state = True req.is_first_context_chunk = is_first_context_chunk + req.py_encoder_output_ready_event = None return req @@ -84,6 +85,8 @@ def make_ctx_request( req.is_context_init_state = True req.is_generation_in_progress_state = False req.encoder_output_len = encoder_output_len + req.py_encoder_output_ready_event = None + req.py_skip_cross_kv_projection = False return req @@ -178,6 +181,7 @@ def make_scheduler( scheduler_capacity=None, no_schedule_until_state=None, no_schedule_after_state=None, + cross_kv_cache_manager=None, ): """Create KVCacheV2Scheduler, patching isinstance check for mock mgr.""" from tensorrt_llm._torch.pyexecutor.scheduler.scheduler_v2 import KVCacheV2Scheduler @@ -191,6 +195,8 @@ def make_scheduler( kwargs["no_schedule_until_state"] = no_schedule_until_state if no_schedule_after_state is not None: kwargs["no_schedule_after_state"] = no_schedule_after_state + if cross_kv_cache_manager is not None: + kwargs["cross_kv_cache_manager"] = cross_kv_cache_manager return KVCacheV2Scheduler( max_batch_size=max_batch_size, max_num_tokens=max_num_tokens, @@ -203,16 +209,31 @@ def make_scheduler( ) -def make_encoder_scheduler(kv_cache_manager, **kwargs): +def make_encoder_scheduler(kv_cache_manager, cross_kv_cache_manager=None, **kwargs): """Scheduler with state range widened to include ENCODER_INIT (matches - C++ trtEncoderModel pattern).""" + C++ trtEncoderModel pattern). + + Encoder-decoder runtime requires a cross_kv_cache_manager for the later + decoder-context cross-KV step. By default we wire a fresh mock cross + manager that succeeds; tests that exercise misconfiguration pass an + explicit ``None`` through ``make_scheduler`` directly. + """ + if cross_kv_cache_manager is None: + cross_kv_cache_manager = make_kv_cache_manager() return make_scheduler( kv_cache_manager, no_schedule_until_state=LlmRequestState.ENCODER_INIT, + cross_kv_cache_manager=cross_kv_cache_manager, **kwargs, ) +def make_not_ready_event(): + event = Mock() + event.query.return_value = False + return event + + def ids(reqs): return [r.request_id for r in reqs] @@ -288,7 +309,7 @@ def test_encoder_budget_exhausted(self): sched = make_encoder_scheduler(mgr, max_num_tokens=100) reqs = [make_encoder_request(0, encoder_output_len=200)] out = sched.schedule_request(reqs, set()) - assert len(out.context_requests) == 0 + assert len(out.encoder_requests) == 0 def test_gen_with_draft_tokens(self): mgr = make_kv_cache_manager() @@ -624,24 +645,6 @@ def test_chunked_fail_then_gen(self): assert ids(out.generation_requests) == [1] -class TestKVCacheFailuresEncoder: - """Encoder KV failures.""" - - def test_encoder_prepare_fails(self): - mgr = make_kv_cache_manager(prepare_context_fn=lambda req: False) - sched = make_encoder_scheduler(mgr, max_num_tokens=1000) - reqs = [make_encoder_request(0, encoder_output_len=100)] - out = sched.schedule_request(reqs, set()) - assert len(out.context_requests) == 0 - - def test_encoder_resize_fails(self): - mgr = make_kv_cache_manager(resize_context_fn=lambda req, n: False) - sched = make_encoder_scheduler(mgr, max_num_tokens=1000) - reqs = [make_encoder_request(0, encoder_output_len=100)] - out = sched.schedule_request(reqs, set()) - assert len(out.context_requests) == 0 - - # =========================================================================== # Eviction (MAX_UTILIZATION) # =========================================================================== @@ -914,7 +917,7 @@ def test_encoder_peft_check(self): sched = make_encoder_scheduler(mgr, peft_cache_manager=peft) reqs = [make_encoder_request(0, encoder_output_len=100, lora_task_id=1)] out = sched.schedule_request(reqs, set()) - assert len(out.context_requests) == 0 + assert len(out.encoder_requests) == 0 def test_mixed_peft_gen_claims_reduce_ctx(self): """[gen(task1), ctx(task2)], pages for 2 total → both ok.""" @@ -972,14 +975,15 @@ def test_encoder_scheduled(self): sched = make_encoder_scheduler(mgr, max_num_tokens=200) reqs = [make_encoder_request(0, encoder_output_len=100)] out = sched.schedule_request(reqs, set()) - assert ids(out.context_requests) == [0] + assert ids(out.encoder_requests) == [0] + assert ids(out.context_requests) == [] def test_encoder_budget_overflow(self): mgr = make_kv_cache_manager() sched = make_encoder_scheduler(mgr, max_num_tokens=100) reqs = [make_encoder_request(0, encoder_output_len=200)] out = sched.schedule_request(reqs, set()) - assert len(out.context_requests) == 0 + assert len(out.encoder_requests) == 0 def test_encoder_exceeds_budget(self): """encoder_output_len > max_num_tokens → break (not scheduled).""" @@ -987,7 +991,7 @@ def test_encoder_exceeds_budget(self): sched = make_encoder_scheduler(mgr, max_num_tokens=1000) reqs = [make_encoder_request(0, encoder_output_len=2000)] out = sched.schedule_request(reqs, set()) - assert len(out.context_requests) == 0 + assert len(out.encoder_requests) == 0 def test_encoder_plus_gen(self): mgr = make_kv_cache_manager() @@ -997,23 +1001,10 @@ def test_encoder_plus_gen(self): make_gen_request(1), ] out = sched.schedule_request(reqs, set()) - assert ids(out.context_requests) == [0] + assert ids(out.encoder_requests) == [0] + assert ids(out.context_requests) == [] assert ids(out.generation_requests) == [1] - def test_encoder_prepare_fails(self): - mgr = make_kv_cache_manager(prepare_context_fn=lambda req: False) - sched = make_encoder_scheduler(mgr, max_num_tokens=1000) - reqs = [make_encoder_request(0, encoder_output_len=100)] - out = sched.schedule_request(reqs, set()) - assert len(out.context_requests) == 0 - - def test_encoder_resize_fails(self): - mgr = make_kv_cache_manager(resize_context_fn=lambda req, n: False) - sched = make_encoder_scheduler(mgr, max_num_tokens=1000) - reqs = [make_encoder_request(0, encoder_output_len=100)] - out = sched.schedule_request(reqs, set()) - assert len(out.context_requests) == 0 - def test_multiple_encoders(self): mgr = make_kv_cache_manager() sched = make_encoder_scheduler(mgr, max_num_tokens=100) @@ -1022,7 +1013,8 @@ def test_multiple_encoders(self): make_encoder_request(1, encoder_output_len=50), ] out = sched.schedule_request(reqs, set()) - assert ids(out.context_requests) == [0, 1] + assert ids(out.encoder_requests) == [0, 1] + assert ids(out.context_requests) == [] def test_encoder_counts_toward_batch(self): """Gen wins phase 1; encoder scheduled in phase 2 still occupies a batch slot.""" @@ -1036,9 +1028,111 @@ def test_encoder_counts_toward_batch(self): out = sched.schedule_request(reqs, set()) # gen(1) scheduled first (phase 1), encoder(0) fills remaining slot (phase 2) assert ids(out.generation_requests) == [1] - assert ids(out.context_requests) == [0] + assert ids(out.encoder_requests) == [0] + assert ids(out.context_requests) == [] # encoder(2) excluded — encoder(0) counted toward batch + def test_encoder_does_not_touch_kv_pools(self): + """Encoder admission must not touch either KV pool. + + This guards the dual-pool contract: both self- and cross-pool + allocation are decoder-context responsibilities. + """ + self_mgr = make_kv_cache_manager() + cross_mgr = make_kv_cache_manager() + sched = make_encoder_scheduler( + self_mgr, cross_kv_cache_manager=cross_mgr, max_num_tokens=1000 + ) + req = make_encoder_request(0, encoder_output_len=100) + out = sched.schedule_request([req], set()) + assert ids(out.encoder_requests) == [0] + assert ids(out.context_requests) == [] + # Self pool stays untouched. + self_mgr.prepare_context.assert_not_called() + self_mgr.resize_context.assert_not_called() + self_mgr.try_allocate_generation.assert_not_called() + cross_mgr.prepare_context.assert_not_called() + cross_mgr.resize_context.assert_not_called() + cross_mgr.try_allocate_generation.assert_not_called() + + def test_encoder_without_cross_manager_raises(self): + """No cross_kv_cache_manager -> encoder request cannot be admitted. + + The dual-pool contract requires a cross manager. Without one + the scheduler errors rather than silently routing to the + self pool (which would corrupt self-pool sizing). + """ + from tensorrt_llm._torch.pyexecutor.scheduler.scheduler_v2 import KVCacheV2Scheduler + + mgr = make_kv_cache_manager() + # Build a scheduler with ENCODER_INIT gating but no cross manager. + with patch( + "tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2.KVCacheManagerV2", + new=type(mgr), + ): + sched = KVCacheV2Scheduler( + max_batch_size=8, + max_num_tokens=1000, + kv_cache_manager=mgr, + scheduler_policy=CapacitySchedulerPolicy.MAX_UTILIZATION, + no_schedule_until_state=LlmRequestState.ENCODER_INIT, + ) + reqs = [make_encoder_request(0, encoder_output_len=100)] + with pytest.raises(RuntimeError, match="requires a cross_kv_cache_manager"): + sched.schedule_request(reqs, set()) + # Self pool must not be touched. + mgr.prepare_context.assert_not_called() + + def test_encoder_then_context_defers_cross_pool_to_context(self): + """Cross-pool allocation is deferred from ENCODER_INIT to CONTEXT_INIT.""" + self_mgr = make_kv_cache_manager() + cross_mgr = make_kv_cache_manager() + sched = make_encoder_scheduler( + self_mgr, cross_kv_cache_manager=cross_mgr, max_num_tokens=1000 + ) + + # Iteration 1: ENCODER_INIT → encoder compute admission. + enc_req = make_encoder_request(0, encoder_output_len=80) + out1 = sched.schedule_request([enc_req], set()) + assert ids(out1.encoder_requests) == [0] + assert ids(out1.context_requests) == [] + self_mgr.prepare_context.assert_not_called() + self_mgr.resize_context.assert_not_called() + cross_mgr.prepare_context.assert_not_called() + cross_mgr.resize_context.assert_not_called() + + # Iteration 2: CONTEXT_INIT (post-encoder transition) → both pools. + ctx_req = make_ctx_request(0, context_remaining_length=50, encoder_output_len=80) + out2 = sched.schedule_request([ctx_req], set()) + assert ids(out2.context_requests) == [0] + self_mgr.prepare_context.assert_called_once_with(ctx_req) + self_mgr.resize_context.assert_called_once_with(ctx_req, 50) + cross_mgr.prepare_context.assert_called_once_with(ctx_req) + cross_mgr.resize_context.assert_called_once_with(ctx_req, 80) + + def test_later_context_chunk_reuses_cross_pool_without_resizing(self): + """Later decoder chunks read existing cross-KV without reallocation.""" + self_mgr = make_kv_cache_manager() + cross_mgr = make_kv_cache_manager() + sched = make_encoder_scheduler( + self_mgr, cross_kv_cache_manager=cross_mgr, max_num_tokens=1000 + ) + ctx_req = make_ctx_request( + 0, + context_remaining_length=50, + is_first_context_chunk=False, + encoder_output_len=80, + ) + ctx_req.py_skip_cross_kv_projection = True + + out = sched.schedule_request([ctx_req], set()) + + assert ids(out.context_requests) == [0] + self_mgr.prepare_context.assert_called_once_with(ctx_req) + self_mgr.resize_context.assert_called_once_with(ctx_req, 50) + cross_mgr.prepare_context.assert_not_called() + cross_mgr.resize_context.assert_not_called() + # =========================================================================== # Disaggregated Serving @@ -1589,6 +1683,18 @@ def test_context_init_passes(self): out = sched.schedule_request(reqs, set()) assert len(out.context_requests) == 1 + def test_context_init_waiting_on_encoder_event_is_filtered(self): + mgr = make_kv_cache_manager() + sched = make_scheduler(mgr, max_num_tokens=100) + blocked_ctx = make_ctx_request(0, context_remaining_length=10) + blocked_ctx.py_encoder_output_ready_event = make_not_ready_event() + gen_req = make_gen_request(1) + + out = sched.schedule_request([blocked_ctx, gen_req], set()) + + assert ids(out.context_requests) == [] + assert ids(out.generation_requests) == [1] + def test_gen_in_progress_passes(self): """GEN_IN_PROGRESS (13) is in [10,14) range → not filtered.""" mgr = make_kv_cache_manager() @@ -1691,6 +1797,7 @@ def test_output_fields_correct(self): make_disagg_request(2), ] out = sched.schedule_request(reqs, set()) + assert len(out.encoder_requests) == 0 assert len(out.context_requests) == 1 assert len(out.generation_requests) == 1 assert len(out.fitting_disagg_gen_init_requests) == 1 @@ -1790,7 +1897,8 @@ def test_encoder_then_gen(self): make_gen_request(1), ] out = sched.schedule_request(reqs, set()) - assert ids(out.context_requests) == [0] + assert ids(out.encoder_requests) == [0] + assert ids(out.context_requests) == [] assert ids(out.generation_requests) == [1] def test_ctx_fail_budget_preserved(self): @@ -1900,7 +2008,8 @@ def test_single_request_each_type(self): # Single encoder (needs widened state range) enc_sched = make_encoder_scheduler(mgr, max_num_tokens=1000) out = enc_sched.schedule_request([make_encoder_request(2, 50)], set()) - assert len(out.context_requests) == 1 + assert len(out.encoder_requests) == 1 + assert len(out.context_requests) == 0 def test_all_inflight(self): mgr = make_kv_cache_manager() diff --git a/tests/unittest/_torch/executor/test_py_scheduler.py b/tests/unittest/_torch/executor/test_py_scheduler.py index bb7ae2448ac5..b2d60ed823e0 100644 --- a/tests/unittest/_torch/executor/test_py_scheduler.py +++ b/tests/unittest/_torch/executor/test_py_scheduler.py @@ -25,6 +25,9 @@ from dataclasses import dataclass, field from typing import List, Optional +from unittest.mock import Mock + +import pytest from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest, LlmRequestState, SamplingConfig from tensorrt_llm._torch.pyexecutor.scheduler.scheduler import ( @@ -32,7 +35,9 @@ ContextChunkingConfig, PyCapacityScheduler, PyMicroBatchScheduler, + SimpleScheduler, SimpleUnifiedScheduler, + drop_decoder_context_requests_waiting_for_encoder_output, ) from tensorrt_llm.llmapi.llm_args import CapacitySchedulerPolicy @@ -78,12 +83,14 @@ def make_context_request( beam_width: int = 1, draft_tokens_len: int = 0, context_position: int = 0, + encoder_output_len: int = 0, ) -> LlmRequest: req = _make_request( request_id=request_id, prompt_len=prompt_len, beam_width=beam_width, draft_tokens_len=draft_tokens_len, + encoder_output_len=encoder_output_len, state=LlmRequestState.CONTEXT_INIT, ) if context_position > 0: @@ -160,9 +167,13 @@ def get_kv_cache_stats(self) -> MockKVCacheStats: ) def get_remaining_blocks_to_completion(self, req, window_size: int) -> int: + if req.is_encoder_init_state: + return 0 return self._blocks_per_request def get_needed_blocks_one_step(self, req, two_step_lookahead: bool, window_size: int) -> int: + if req.is_encoder_init_state: + return 0 return self._blocks_per_request def scheduling_has_free_blocks(self, total: int, window_size: int) -> bool: @@ -223,7 +234,7 @@ def test_simple_context_only(self): make_context_request(1, prompt_len=10), make_context_request(2, prompt_len=10), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) assert len(ctx) == 2 assert len(gen) == 0 assert ctx[0].request_id == 0 @@ -237,7 +248,7 @@ def test_simple_generation_only(self): make_generation_request(1), make_generation_request(2), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) assert len(ctx) == 0 assert len(gen) == 2 assert gen[0].request_id == 0 @@ -255,7 +266,7 @@ def test_context_generation_overlap(self): make_context_request(2, prompt_len=10), make_generation_request(3), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) assert len(ctx) == 2 assert len(gen) == 2 assert {r.request_id for r in ctx} == {0, 2} @@ -272,7 +283,7 @@ def test_max_num_tokens_limits_context(self): make_context_request(0, prompt_len=10), make_context_request(1, prompt_len=10), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) # Only 1 fits within token budget assert len(ctx) == 1 assert ctx[0].request_id == 0 @@ -288,7 +299,7 @@ def test_max_num_tokens_allows_gen_after_context(self): make_generation_request(1, beam_width=1), make_generation_request(2, beam_width=1), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) # context: 10 tokens, gen1: 1 token, gen2: 1 token => total 12 assert len(ctx) == 1 assert len(gen) == 2 @@ -301,7 +312,7 @@ def test_max_batch_size_limits_total(self): make_generation_request(1), make_generation_request(2), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) # batch_size=2: should schedule context_0 + gen_1 assert len(ctx) + len(gen) == 2 @@ -317,7 +328,7 @@ def test_beam_width_1(self): make_generation_request(2, beam_width=1), make_generation_request(3, beam_width=1), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) # context: 10, gen: 1+1 = 12 total. Can't fit gen_3 (would be 13). assert len(ctx) == 1 assert len(gen) == 2 @@ -333,7 +344,7 @@ def test_beam_width_4(self): make_generation_request(1, beam_width=4), make_generation_request(2, beam_width=4), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) # context: 10, gen1: 4 = 14. gen2: +4 = 18 > 15. assert len(ctx) == 1 assert len(gen) == 1 @@ -349,7 +360,7 @@ def test_beam_width_mismatch_skipped(self): make_generation_request(1, beam_width=4), make_generation_request(2, beam_width=1), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) # gen_0 sets beam_width=1, gen_1 is skipped (beam_width=4), gen_2 fits assert len(gen) == 2 assert gen[0].request_id == 0 @@ -366,7 +377,7 @@ def test_draft_tokens_count_toward_budget(self): make_context_request(0, prompt_len=10, draft_tokens_len=3), make_generation_request(1, draft_tokens_len=2), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) # context: 10+3=13, gen: 1+2=3, total=16 > 15 => only context fits assert len(ctx) == 1 assert len(gen) == 0 @@ -382,7 +393,7 @@ def test_gen_draft_tokens(self): make_generation_request(1, beam_width=1, draft_tokens_len=3), make_generation_request(2, beam_width=1, draft_tokens_len=3), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) # Each gen costs 1+3=4. Two fit (8), three don't (12 > 10). assert len(gen) == 2 @@ -394,7 +405,7 @@ def test_inflight_requests_excluded(self): make_context_request(1, prompt_len=10), make_generation_request(2), ] - ctx, gen = scheduler.schedule(requests, {0, 2}) + _enc, ctx, gen = scheduler.schedule(requests, {0, 2}) # Only request 1 is not in flight assert len(ctx) == 1 assert ctx[0].request_id == 1 @@ -408,7 +419,7 @@ def test_completed_requests_filtered(self): make_completed_request(1), make_generation_request(2), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) # Completed request 1 is filtered by state gating assert len(ctx) == 1 assert len(gen) == 1 @@ -429,7 +440,7 @@ def test_simple_no_overlap(self): make_context_request(2, prompt_len=10), make_context_request(3, prompt_len=10), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) assert len(ctx) == 2 assert len(gen) == 0 assert ctx[0].request_id == 0 @@ -443,7 +454,7 @@ def test_simple_no_overlap(self): make_context_request(2, prompt_len=10), make_context_request(3, prompt_len=10), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) assert len(gen) == 2 assert gen[0].request_id == 0 assert gen[1].request_id == 1 @@ -455,7 +466,7 @@ def test_simple_no_overlap(self): make_context_request(2, prompt_len=10), make_context_request(3, prompt_len=10), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) assert len(ctx) == 2 assert ctx[0].request_id == 2 assert ctx[1].request_id == 3 @@ -475,7 +486,7 @@ def test_simple_no_overlap_max_num_tokens(self): # C++: Req 0: (0,1,2,3,4), Req 1: () r0 = make_context_request(0, prompt_len=12) r1 = make_context_request(1, prompt_len=12) - ctx, gen = scheduler.schedule([r0, r1], set()) + _enc, ctx, gen = scheduler.schedule([r0, r1], set()) assert len(ctx) >= 1 # First request gets a chunk within budget req0 = next(r for r in ctx if r.request_id == 0) @@ -499,7 +510,7 @@ def test_simple_no_overlap_max_context_length(self): # Two requests with promptLen=10 fit within maxContextLength=12 r0 = make_context_request(0, prompt_len=10) r1 = make_context_request(1, prompt_len=10) - ctx, gen = scheduler.schedule([r0, r1], set()) + _enc, ctx, gen = scheduler.schedule([r0, r1], set()) assert len(ctx) == 2 # Each chunk should be at most max_context_length for r in ctx: @@ -507,7 +518,7 @@ def test_simple_no_overlap_max_context_length(self): # Request with promptLen=17 needs chunking (17 > 12) r3 = make_context_request(3, prompt_len=17) - ctx2, gen2 = scheduler.schedule([r3], set()) + _enc2, ctx2, _ = scheduler.schedule([r3], set()) assert len(ctx2) == 1 assert ctx2[0].context_chunk_size <= 12 @@ -527,17 +538,17 @@ def test_simple_with_overlap(self): requests = [make_context_request(i, prompt_len=10) for i in range(4)] # Step 1: slot 0 — req0, req1 scheduled - ctx0, _ = scheduler.schedule(requests, set()) + _enc0, ctx0, _ = scheduler.schedule(requests, set()) assert {r.request_id for r in ctx0} == {0, 1} slot0_inflight = {r.request_id for r in ctx0} # Step 2: slot 1 — req0/req1 still inflight, req2/req3 scheduled - ctx1, _ = scheduler.schedule(requests, slot0_inflight) + _enc1, ctx1, _ = scheduler.schedule(requests, slot0_inflight) assert {r.request_id for r in ctx1} == {2, 3} slot1_inflight = {r.request_id for r in ctx1} # Step 3: slot 0 freed (inflight = slot1 only) — req0/req1 scheduled again - ctx2, _ = scheduler.schedule(requests, slot1_inflight) + _enc2, ctx2, _ = scheduler.schedule(requests, slot1_inflight) assert {r.request_id for r in ctx2} == {0, 1} def test_gen_draft_tokens_max_num_tokens(self): @@ -550,7 +561,7 @@ def test_gen_draft_tokens_max_num_tokens(self): """ scheduler = PyMicroBatchScheduler(max_batch_size=64, max_num_tokens=128) requests = [make_generation_request(i, beam_width=1, draft_tokens_len=63) for i in range(4)] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) # Each request costs 1 + 63 = 64 tokens; 2 fit (128 = budget), 3 don't (192 > 128). assert len(gen) == 2 assert gen[0].request_id == 0 @@ -558,6 +569,75 @@ def test_gen_draft_tokens_max_num_tokens(self): assert len(ctx) == 0 +class TestEncoderOutputReadinessFiltering: + def test_drop_decoder_context_requests_waiting_for_encoder_output(self): + ready_ctx = make_context_request(1) + ready_ctx.py_encoder_output_ready_event = Mock() + ready_ctx.py_encoder_output_ready_event.query.return_value = True + + blocked_ctx = make_context_request(2) + blocked_ctx.py_encoder_output_ready_event = Mock() + blocked_ctx.py_encoder_output_ready_event.query.return_value = False + + gen_req = make_generation_request(3) + + filtered = drop_decoder_context_requests_waiting_for_encoder_output( + [ready_ctx, blocked_ctx, gen_req] + ) + + assert [req.request_id for req in filtered] == [1, 3] + + def test_simple_unified_scheduler_skips_unready_context_request(self): + scheduler = SimpleUnifiedScheduler( + max_batch_size=8, + max_num_tokens=128, + kv_cache_manager=MockKVCacheManager(), + peft_cache_manager=None, + scheduler_policy=CapacitySchedulerPolicy.GUARANTEED_NO_EVICT, + ) + ready_ctx = make_context_request(1) + blocked_ctx = make_context_request(2) + blocked_ctx.py_encoder_output_ready_event = Mock() + blocked_ctx.py_encoder_output_ready_event.query.return_value = False + gen_req = make_generation_request(3) + + out = scheduler.schedule_request([ready_ctx, blocked_ctx, gen_req], set()) + + assert [req.request_id for req in out.context_requests] == [1] + assert [req.request_id for req in out.generation_requests] == [3] + + def test_py_micro_batch_scheduler_skips_unready_context_request(self): + scheduler = PyMicroBatchScheduler(max_batch_size=8, max_num_tokens=128) + ready_ctx = make_context_request(1) + blocked_ctx = make_context_request(2) + blocked_ctx.py_encoder_output_ready_event = Mock() + blocked_ctx.py_encoder_output_ready_event.query.return_value = False + gen_req = make_generation_request(3) + + _enc, ctx, gen = scheduler.schedule([ready_ctx, blocked_ctx, gen_req], set()) + + assert [req.request_id for req in ctx] == [1] + assert [req.request_id for req in gen] == [3] + + def test_simple_scheduler_prefilters_before_capacity(self): + capacity_scheduler = Mock() + capacity_scheduler.schedule_request.return_value = ([], [], []) + micro_batch_scheduler = Mock() + micro_batch_scheduler.schedule.return_value = ([], [], []) + scheduler = SimpleScheduler(capacity_scheduler, micro_batch_scheduler) + + ready_ctx = make_context_request(1) + blocked_ctx = make_context_request(2) + blocked_ctx.py_encoder_output_ready_event = Mock() + blocked_ctx.py_encoder_output_ready_event.query.return_value = False + gen_req = make_generation_request(3) + + scheduler.schedule_request([ready_ctx, blocked_ctx, gen_req], set()) + + filtered_requests = capacity_scheduler.schedule_request.call_args.args[0] + assert [req.request_id for req in filtered_requests] == [1, 3] + + # ############################################################################ # # Part 2: Context Chunking Tests @@ -586,7 +666,7 @@ def test_equal_progress_basic(self): make_context_request(0, prompt_len=20), make_context_request(1, prompt_len=20), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) assert len(ctx) == 2 # Each should get ~5 tokens (equal progress, unit=5, total=10) total_chunk = sum(r.context_chunk_size for r in ctx) @@ -606,7 +686,7 @@ def test_equal_progress_uneven_remaining(self): make_context_request(0, prompt_len=3), # Only 3 tokens remaining make_context_request(1, prompt_len=20), # Lots remaining ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) assert len(ctx) == 2 # Look up by request_id since sort reorders (not-last-chunk first) req0 = next(r for r in ctx if r.request_id == 0) @@ -628,7 +708,7 @@ def test_fcfs_basic(self): make_context_request(0, prompt_len=20), make_context_request(1, prompt_len=20), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) # FCFS: request 0 gets up to budget, request 1 gets remainder assert len(ctx) >= 1 # First request should get more tokens @@ -645,7 +725,7 @@ def test_fcfs_fills_first_request(self): make_context_request(0, prompt_len=10), make_context_request(1, prompt_len=20), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) assert len(ctx) == 2 # Look up by request_id since sort reorders (not-last-chunk first) req0 = next(r for r in ctx if r.request_id == 0) @@ -669,7 +749,7 @@ def test_chunk_with_generation(self): make_context_request(1, prompt_len=20), make_context_request(2, prompt_len=20), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) assert len(gen) == 1 # Remaining budget for context: 15 - 1 = 14 total_ctx_tokens = sum(r.context_chunk_size for r in ctx) @@ -688,7 +768,7 @@ def test_chunk_size_zero_not_scheduled(self): make_context_request(0, prompt_len=20), make_context_request(1, prompt_len=20), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) # With budget 5, at most one request gets chunk_size=5, the other might get 0 for r in ctx: assert r.context_chunk_size > 0 @@ -705,7 +785,7 @@ def test_chunking_with_max_context_length(self): requests = [ make_context_request(0, prompt_len=20), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) assert len(ctx) == 1 # max_context_length = max_num_tokens = 12, so chunk <= 12 assert ctx[0].context_chunk_size <= 12 @@ -721,7 +801,7 @@ def test_continued_chunking(self): ) req = make_context_request(0, prompt_len=20, context_position=10) # remaining = 20 - 10 = 10 - ctx, gen = scheduler.schedule([req], set()) + _enc, ctx, gen = scheduler.schedule([req], set()) assert len(ctx) == 1 assert ctx[0].context_chunk_size <= 10 # remaining context @@ -738,7 +818,7 @@ def test_last_chunk_allows_draft_tokens(self): # prompt_len=8, so chunk_size will be 8. Unit=10, remainder=2. # Draft tokens=2 fits in remainder. req = make_context_request(0, prompt_len=8, draft_tokens_len=2) - ctx, gen = scheduler.schedule([req], set()) + _enc, ctx, gen = scheduler.schedule([req], set()) assert len(ctx) == 1 assert req.is_last_context_chunk @@ -753,7 +833,7 @@ def test_draft_tokens_discarded_when_no_space(self): ) # prompt_len=5, chunk_size=5, unit=5, remainder=0. Draft=3 won't fit. req = make_context_request(0, prompt_len=5, draft_tokens_len=3) - ctx, gen = scheduler.schedule([req], set()) + _enc, ctx, gen = scheduler.schedule([req], set()) assert len(ctx) == 1 def test_chunked_context_draft_tokens_max_num_tokens(self): @@ -772,7 +852,7 @@ def test_chunked_context_draft_tokens_max_num_tokens(self): ctx_chunk_config=config, ) requests = [make_context_request(i, prompt_len=2041, draft_tokens_len=8) for i in range(4)] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) assert len(ctx) == 4 for req in ctx: assert req.num_draft_tokens == 7 @@ -798,7 +878,7 @@ def test_chunked_context_draft_tokens_max_context_length(self): make_context_request(0, prompt_len=6, draft_tokens_len=5), make_context_request(1, prompt_len=6, draft_tokens_len=5), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) assert len(ctx) == 2 for req in ctx: assert req.num_draft_tokens == 4 @@ -809,7 +889,7 @@ def test_no_chunking_context_fits(self): max_batch_size=4, max_num_tokens=20, ctx_chunk_config=None ) req = make_context_request(0, prompt_len=15) - ctx, gen = scheduler.schedule([req], set()) + _enc, ctx, gen = scheduler.schedule([req], set()) assert len(ctx) == 1 def test_no_chunking_context_exceeds_budget(self): @@ -824,7 +904,7 @@ def test_no_chunking_context_exceeds_budget(self): make_context_request(0, prompt_len=8), make_context_request(1, prompt_len=8), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) # First request (8) fits (8 <= 10). Second (8+8=16 > 10) breaks the loop. assert len(ctx) == 1 assert ctx[0].request_id == 0 @@ -838,7 +918,7 @@ def test_sort_by_lora_task_id(self): r0 = _make_request(0, state=LlmRequestState.GENERATION_IN_PROGRESS, lora_task_id=5) r1 = _make_request(1, state=LlmRequestState.GENERATION_IN_PROGRESS) r2 = _make_request(2, state=LlmRequestState.GENERATION_IN_PROGRESS, lora_task_id=3) - ctx, gen = scheduler.schedule([r0, r1, r2], set()) + _enc, ctx, gen = scheduler.schedule([r0, r1, r2], set()) # None < any value, so order should be: r1(None), r2(3), r0(5) assert gen[0].request_id == 1 assert gen[1].request_id == 2 @@ -879,7 +959,7 @@ def test_reusable_tokens_reduce_compute_budget(self): req0.estimated_reusable_tokens = 15 req1.estimated_reusable_tokens = 15 - ctx, gen = scheduler.schedule([req0, req1], set()) + _enc, ctx, gen = scheduler.schedule([req0, req1], set()) assert len(ctx) == 2 def test_reusable_tokens_zero_has_no_effect(self): @@ -898,7 +978,7 @@ def test_reusable_tokens_zero_has_no_effect(self): req0.estimated_reusable_tokens = 0 req1.estimated_reusable_tokens = 0 - ctx, gen = scheduler.schedule([req0, req1], set()) + _enc, ctx, gen = scheduler.schedule([req0, req1], set()) assert len(ctx) == 1 assert ctx[0].request_id == 0 @@ -919,7 +999,7 @@ def test_reusable_tokens_chunked_context_fcfs_full_context_fits(self): req = make_context_request(0, prompt_len=20) req.estimated_reusable_tokens = 10 - ctx, gen = scheduler.schedule([req], set()) + _enc, ctx, gen = scheduler.schedule([req], set()) assert len(ctx) == 1 # Full context fits — chunk_size equals the full prompt length. assert ctx[0].context_chunk_size == 20 @@ -943,7 +1023,7 @@ def test_reusable_tokens_only_on_first_chunk(self): req0.estimated_reusable_tokens = 20 assert not req0.is_first_context_chunk - ctx, gen = scheduler.schedule([req0], set()) + _enc, ctx, gen = scheduler.schedule([req0], set()) assert len(ctx) == 1 # Remaining tokens = 30 - 10 = 20; reuse ignored → compute = 20 # Budget = 30, 20 <= 30, so it fits. @@ -953,7 +1033,7 @@ def test_reusable_tokens_only_on_first_chunk(self): scheduler2 = PyMicroBatchScheduler( max_batch_size=4, max_num_tokens=30, ctx_chunk_config=None ) - ctx2, _ = scheduler2.schedule([req0, req1], set()) + _enc2, ctx2, _ = scheduler2.schedule([req0, req1], set()) req0_again = make_context_request(0, prompt_len=30, context_position=10) req0_again.estimated_reusable_tokens = 20 # Confirm: fresh req0 (non-first chunk) + req1 → only req0 fits @@ -967,7 +1047,7 @@ def test_reusable_tokens_only_on_first_chunk(self): scheduler3 = PyMicroBatchScheduler( max_batch_size=4, max_num_tokens=30, ctx_chunk_config=None ) - ctx3, _ = scheduler3.schedule([req2, req3], set()) + _enc3, ctx3, _ = scheduler3.schedule([req2, req3], set()) # req2 compute = max(1, 30-20) = 10; req3 compute = 20; total = 30 → both fit. assert len(ctx3) == 2 @@ -987,7 +1067,7 @@ def test_reusable_tokens_no_chunking_min_cost_is_one(self): req0.estimated_reusable_tokens = 15 # exceeds prompt length req1 = make_context_request(1, prompt_len=10) - ctx, gen = scheduler.schedule([req0, req1], set()) + _enc, ctx, gen = scheduler.schedule([req0, req1], set()) # req0 compute = 1; req1 compute = 10; 1 + 10 = 11 > 10 → only req0 fits. assert len(ctx) == 1 assert ctx[0].request_id == 0 @@ -1021,7 +1101,7 @@ def test_reusable_tokens_fcfs_over_budget_multi_request(self): req1.estimated_reusable_tokens = 8 req2.estimated_reusable_tokens = 8 - ctx, gen = scheduler.schedule([req0, req1, req2], set()) + _enc, ctx, gen = scheduler.schedule([req0, req1, req2], set()) # Note: ctx is sorted — partially-chunked requests come before full-context ones. # Look up by request_id rather than by position. @@ -1067,7 +1147,7 @@ def test_reusable_tokens_equal_progress(self): req0.estimated_reusable_tokens = 10 req1.estimated_reusable_tokens = 0 - ctx, gen = scheduler.schedule([req0, req1], set()) + _enc, ctx, gen = scheduler.schedule([req0, req1], set()) chunks = {r.request_id: r.context_chunk_size for r in ctx} assert len(ctx) == 2, "Both requests should be scheduled" @@ -1616,7 +1696,7 @@ def test_full_scheduler_path(self): max_batch_size=4, max_num_tokens=100, ctx_chunk_config=config ) req = make_context_request(0, prompt_len=30) - ctx, gen = scheduler.schedule([req], set()) + _enc, ctx, gen = scheduler.schedule([req], set()) # Despite budget=100 >> prompt=30, FORCE_CHUNK limits chunk to unit_size=10. assert len(ctx) == 1 assert ctx[0].context_chunk_size == 10 @@ -1636,7 +1716,7 @@ def test_full_scheduler_multiple_requests(self): make_context_request(1, prompt_len=15), make_context_request(2, prompt_len=5), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) assert len(ctx) == 3 # Find by request_id since sorting may reorder. chunks = {r.request_id: r.context_chunk_size for r in ctx} @@ -1658,7 +1738,7 @@ def test_full_scheduler_with_generation(self): make_generation_request(0), # costs 1 token make_context_request(1, prompt_len=30), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) assert len(gen) == 1 assert len(ctx) == 1 # Budget remaining = 15 - 1 (gen) = 14; chunk = min(30, 10) = 10 @@ -1697,7 +1777,7 @@ def test_draft_tokens_greater_than_chunk_size(self): make_context_request(2, prompt_len=3, draft_tokens_len=17), ] - ctx, gen = scheduler.schedule(requests, set()) + _enc, ctx, gen = scheduler.schedule(requests, set()) assert len(ctx) == 3 req0 = next(r for r in ctx if r.request_id == 0) @@ -1744,7 +1824,7 @@ def test_mixed_batch_zero_draft_does_not_consume_speculative_budget(self): make_context_request(1, prompt_len=8, draft_tokens_len=0), make_context_request(2, prompt_len=3, draft_tokens_len=13), ] - ctx, _ = scheduler.schedule(requests, set()) + _enc, ctx, _ = scheduler.schedule(requests, set()) assert len(ctx) == 3 r0 = next(r for r in ctx if r.request_id == 0) r1 = next(r for r in ctx if r.request_id == 1) @@ -1775,7 +1855,7 @@ def test_short_draft_request_charges_only_kept_drafts(self) -> None: make_context_request(1, prompt_len=3, draft_tokens_len=13), make_context_request(2, prompt_len=3, draft_tokens_len=13), ] - ctx, _ = scheduler.schedule(requests, set()) + _enc, ctx, _ = scheduler.schedule(requests, set()) assert len(ctx) == 3 r0 = next(r for r in ctx if r.request_id == 0) @@ -1806,7 +1886,7 @@ def test_no_draft_tokens_bypasses_fit_draft(self): ) # Requests with zero draft tokens, large enough to exhaust budget. requests = [make_context_request(i, prompt_len=24, draft_tokens_len=0) for i in range(4)] - ctx, _ = scheduler.schedule(requests, set()) + _enc, ctx, _ = scheduler.schedule(requests, set()) assert ctx, "expected at least one context request to be scheduled" # All scheduled requests must still report 0 draft tokens — no discard happened. for r in ctx: @@ -2287,6 +2367,7 @@ def test_full_pipeline_output_structure(self): make_generation_request(1), ] output = scheduler.schedule_request(requests, set()) + assert hasattr(output, "encoder_requests") assert hasattr(output, "context_requests") assert hasattr(output, "generation_requests") assert hasattr(output, "paused_requests") @@ -2294,9 +2375,33 @@ def test_full_pipeline_output_structure(self): assert hasattr(output, "num_fitting_requests") assert len(output.context_requests) == 1 assert len(output.generation_requests) == 1 + assert len(output.encoder_requests) == 0 assert output.context_requests[0].request_id == 0 assert output.generation_requests[0].request_id == 1 + def test_full_pipeline_separates_encoder_requests(self): + """Encoder admission should not be exposed as decoder context.""" + kv = MockKVCacheManager(num_free_blocks=100, blocks_per_request=5) + cross_kv = MockKVCacheManager(num_free_blocks=100, blocks_per_request=5) + scheduler = SimpleUnifiedScheduler( + max_batch_size=4, + max_num_tokens=100, + kv_cache_manager=kv, + peft_cache_manager=None, + scheduler_policy=CapacitySchedulerPolicy.GUARANTEED_NO_EVICT, + cross_kv_cache_manager=cross_kv, + no_schedule_until_state=LlmRequestState.ENCODER_INIT, + ) + requests = [ + make_encoder_request(0, encoder_output_len=10), + make_context_request(1, prompt_len=10), + ] + + output = scheduler.schedule_request(requests, set()) + + assert [req.request_id for req in output.encoder_requests] == [0] + assert [req.request_id for req in output.context_requests] == [1] + def test_paused_requests_propagated(self): """Paused requests from capacity scheduler appear in output.""" kv = MockKVCacheManager(num_free_blocks=100, blocks_per_request=5) @@ -2339,10 +2444,8 @@ def test_should_fit_with_cross_blocks(self): cross_kv_cache_manager=cross_kv, scheduler_policy=CapacitySchedulerPolicy.GUARANTEED_NO_EVICT, ) - r0 = make_context_request(0, prompt_len=10) - r0.encoder_output_len = 10 - r1 = make_context_request(1, prompt_len=10) - r1.encoder_output_len = 10 + r0 = make_context_request(0, prompt_len=10, encoder_output_len=10) + r1 = make_context_request(1, prompt_len=10, encoder_output_len=10) fitting, disagg, paused = scheduler.schedule_request([r0, r1]) assert len(fitting) == 2 @@ -2356,14 +2459,161 @@ def test_doesnt_fit_with_cross_blocks(self): cross_kv_cache_manager=cross_kv, scheduler_policy=CapacitySchedulerPolicy.GUARANTEED_NO_EVICT, ) - r0 = make_context_request(0, prompt_len=10) - r0.encoder_output_len = 10 - r1 = make_context_request(1, prompt_len=10) - r1.encoder_output_len = 10 + r0 = make_context_request(0, prompt_len=10, encoder_output_len=10) + r1 = make_context_request(1, prompt_len=10, encoder_output_len=10) fitting, disagg, paused = scheduler.schedule_request([r0, r1]) assert len(fitting) == 1 +class TestPyCapacitySchedulerEncoderInit: + """ + V1 capacity scheduler ``ENCODER_INIT`` admission across policies. + + Encoder admission schedules encoder compute but does not reserve self- or + cross-KV blocks; the later decoder ``CONTEXT_INIT`` admission owns that + budgeting. Tests below cover both ``GuaranteedNoEvictPolicy`` and + ``MaxUtilizationPolicy``, plus the safety fallback when no cross + manager is configured. + + All tests below widen ``no_schedule_until_state=ENCODER_INIT`` so + that encoder-init requests pass the state gate (the default + ``CONTEXT_INIT`` rejects them, matching legacy decoder-only setups). + """ + + def _make_scheduler(self, kv, cross_kv, policy, max_num_requests=4): + return PyCapacityScheduler( + max_num_requests=max_num_requests, + kv_cache_manager=kv, + cross_kv_cache_manager=cross_kv, + scheduler_policy=policy, + no_schedule_until_state=LlmRequestState.ENCODER_INIT, + ) + + def test_guaranteed_no_evict_admits_encoder_with_cross_pool(self): + kv = MockKVCacheManager(num_free_blocks=100, blocks_per_request=5) + cross_kv = MockKVCacheManager(num_free_blocks=4, blocks_per_request=2) + scheduler = self._make_scheduler(kv, cross_kv, CapacitySchedulerPolicy.GUARANTEED_NO_EVICT) + # Two encoder-init requests admit even though the cross pool would + # only fit two decoder-context allocations. Cross budget is checked + # later, when each request reaches CONTEXT_INIT. + requests = [make_encoder_request(0, encoder_output_len=10)] + requests.append(make_encoder_request(1, encoder_output_len=10)) + fitting, disagg, paused = scheduler.schedule_request(requests) + assert {r.request_id for r in fitting} == {0, 1} + + def test_guaranteed_no_evict_encoder_does_not_consume_cross_pool(self): + """Cross pool pressure does not throttle encoder admission.""" + kv = MockKVCacheManager(num_free_blocks=100, blocks_per_request=5) + + class CrossPoolThatWouldRejectEncoder(MockKVCacheManager): + def get_remaining_blocks_to_completion(self, req, window_size: int) -> int: + if req.is_encoder_init_state: + return self._blocks_per_request + return super().get_remaining_blocks_to_completion(req, window_size) + + # Only 1 cross block free. If encoder admission reserved cross blocks, + # this would reject each encoder request. + cross_kv = CrossPoolThatWouldRejectEncoder(num_free_blocks=1, blocks_per_request=2) + scheduler = self._make_scheduler(kv, cross_kv, CapacitySchedulerPolicy.GUARANTEED_NO_EVICT) + requests = [ + make_encoder_request(0, encoder_output_len=10), + make_encoder_request(1, encoder_output_len=10), + ] + fitting, disagg, paused = scheduler.schedule_request(requests) + assert {r.request_id for r in fitting} == {0, 1} + + def test_guaranteed_no_evict_raises_encoder_without_cross_pool(self): + """No cross manager -> misconfigured enc-dec request is a hard error.""" + kv = MockKVCacheManager(num_free_blocks=100, blocks_per_request=5) + scheduler = self._make_scheduler(kv, None, CapacitySchedulerPolicy.GUARANTEED_NO_EVICT) + requests = [make_encoder_request(0, encoder_output_len=10)] + with pytest.raises(RuntimeError, match="requires a cross_kv_cache_manager"): + scheduler.schedule_request(requests) + + def test_guaranteed_no_evict_encoder_does_not_consume_self_pool(self): + """Self pool stays available for decoder context even when encoders + are admitted. + + Two encoders + one decoder context all admit because: + - Encoders do not reserve from either KV pool. + - The decoder context has the entire self pool to itself. + """ + # Self pool: enough for one decoder context (5 blocks). + kv = MockKVCacheManager(num_free_blocks=5, blocks_per_request=5) + cross_kv = MockKVCacheManager(num_free_blocks=10, blocks_per_request=2) + scheduler = self._make_scheduler(kv, cross_kv, CapacitySchedulerPolicy.GUARANTEED_NO_EVICT) + requests = [ + make_encoder_request(0, encoder_output_len=10), + make_encoder_request(1, encoder_output_len=10), + make_context_request(2, prompt_len=10), + ] + fitting, disagg, paused = scheduler.schedule_request(requests) + # All three admitted: self pool isn't dented by encoders. + assert {r.request_id for r in fitting} == {0, 1, 2} + + def test_max_utilization_admits_encoder_with_cross_pool(self): + kv = MockKVCacheManager(num_free_blocks=100, blocks_per_request=5) + cross_kv = MockKVCacheManager(num_free_blocks=10, blocks_per_request=2) + scheduler = self._make_scheduler(kv, cross_kv, CapacitySchedulerPolicy.MAX_UTILIZATION) + requests = [ + make_encoder_request(0, encoder_output_len=10), + make_generation_request(1), + ] + fitting, disagg, paused = scheduler.schedule_request(requests) + assert {r.request_id for r in fitting} == {0, 1} + assert len(paused) == 0 + + def test_max_utilization_encoder_does_not_consume_cross_pool(self): + """Cross pool pressure does not throttle MaxUtilization encoder admission.""" + kv = MockKVCacheManager(num_free_blocks=100, blocks_per_request=5) + + class CrossPoolThatWouldRejectEncoder(MockKVCacheManager): + def get_needed_blocks_one_step( + self, req, two_step_lookahead: bool, window_size: int + ) -> int: + if req.is_encoder_init_state: + return self._blocks_per_request + return super().get_needed_blocks_one_step(req, two_step_lookahead, window_size) + + # Only 1 cross block free. If encoder admission reserved cross blocks, + # this would reject each encoder request. + cross_kv = CrossPoolThatWouldRejectEncoder(num_free_blocks=1, blocks_per_request=2) + scheduler = self._make_scheduler(kv, cross_kv, CapacitySchedulerPolicy.MAX_UTILIZATION) + requests = [ + make_encoder_request(0, encoder_output_len=10), + make_encoder_request(1, encoder_output_len=10), + ] + fitting, disagg, paused = scheduler.schedule_request(requests) + assert {r.request_id for r in fitting} == {0, 1} + assert len(paused) == 0 + + def test_max_utilization_raises_encoder_without_cross_pool(self): + """MaxUtilization without a cross manager fails enc-dec admission.""" + kv = MockKVCacheManager(num_free_blocks=100, blocks_per_request=5) + scheduler = self._make_scheduler(kv, None, CapacitySchedulerPolicy.MAX_UTILIZATION) + requests = [make_encoder_request(0, encoder_output_len=10)] + with pytest.raises(RuntimeError, match="requires a cross_kv_cache_manager"): + scheduler.schedule_request(requests) + + def test_max_utilization_encoder_not_evictable_victim(self): + """Encoder-init has no started self-pool blocks → never an eviction + victim. When a decoder context request can't fit, MaxUtilization + skips the encoder while looking for a victim and gives up.""" + kv = MockKVCacheManager(num_free_blocks=3, blocks_per_request=5) + cross_kv = MockKVCacheManager(num_free_blocks=10, blocks_per_request=2) + scheduler = self._make_scheduler(kv, cross_kv, CapacitySchedulerPolicy.MAX_UTILIZATION) + # First-chunk context can't fit (3 < 5), encoder ahead of it + # is not a valid eviction victim (no started self blocks). + requests = [ + make_encoder_request(0, encoder_output_len=10), + make_context_request(1), + ] + fitting, disagg, paused = scheduler.schedule_request(requests) + # Encoder admits cleanly; context can't fit (no victim available). + assert 0 in {r.request_id for r in fitting} + assert 1 not in {r.request_id for r in fitting} + + class TestPyCapacitySchedulerPriority: """ Tests for priority-related scheduling behavior. diff --git a/tests/unittest/_torch/executor/test_scheduler_serializable_output.py b/tests/unittest/_torch/executor/test_scheduler_serializable_output.py index fcce6d88ea9e..b775d9105bec 100644 --- a/tests/unittest/_torch/executor/test_scheduler_serializable_output.py +++ b/tests/unittest/_torch/executor/test_scheduler_serializable_output.py @@ -24,6 +24,7 @@ def test_serializable_scheduler_output_round_trip(): # Create scheduler result: scheduled_requests, fitting_disagg_gen_init_requests, num_fitting_requests scheduled_requests = ScheduledRequests() + scheduled_requests.encoder_requests = [request_pool[7]] scheduled_requests.context_requests_last_chunk = [request_pool[1], request_pool[2]] scheduled_requests.generation_requests = [request_pool[3]] scheduled_requests.paused_requests = [request_pool[4]] @@ -47,6 +48,9 @@ def test_serializable_scheduler_output_round_trip(): # Verify the restored scheduler result is correct assert restored_num_fitting == num_fitting_requests + assert _request_ids(restored_schedule.encoder_requests) == _request_ids( + scheduled_requests.encoder_requests + ) assert _request_ids(restored_schedule.context_requests_chunking) == _request_ids( scheduled_requests.context_requests_chunking ) diff --git a/tests/unittest/_torch/test_model_config.py b/tests/unittest/_torch/test_model_config.py index ba879df12c0d..c1ddb8d8c3fb 100644 --- a/tests/unittest/_torch/test_model_config.py +++ b/tests/unittest/_torch/test_model_config.py @@ -4,7 +4,10 @@ import torch from tensorrt_llm._torch.model_config import ModelConfig -from tensorrt_llm._torch.pyexecutor.model_loader import validate_and_set_kv_cache_quant +from tensorrt_llm._torch.pyexecutor.model_loader import ( + validate_and_set_kv_cache_quant, + validate_encoder_decoder_kv_cache_config, +) from tensorrt_llm.mapping import Mapping from tensorrt_llm.models.modeling_utils import QuantAlgo, QuantConfig @@ -16,6 +19,7 @@ def make_pretrained_config( head_dim: int | None = None, num_hidden_layers: int = 1, vocab_size: int = 3000, + is_encoder_decoder: bool = False, ): # A minimal config object that provides the attributes used by # ModelConfig.get_bindings_model_config(). @@ -32,6 +36,7 @@ def make_pretrained_config( num_hidden_layers=num_hidden_layers, vocab_size=vocab_size, torch_dtype=torch.float16, + is_encoder_decoder=is_encoder_decoder, ) @@ -100,6 +105,15 @@ def _make_model_config_with_kv_quant(kv_cache_quant_algo): return ModelConfig(quant_config=QuantConfig(kv_cache_quant_algo=kv_cache_quant_algo)) +def _make_kv_cache_config( + *, use_kv_cache_manager_v2: bool = False, cross_kv_cache_fraction: float | None = None +): + return types.SimpleNamespace( + use_kv_cache_manager_v2=use_kv_cache_manager_v2, + cross_kv_cache_fraction=cross_kv_cache_fraction, + ) + + def test_validate_and_set_kv_cache_quant_auto_uses_checkpoint(): model_config = _make_model_config_with_kv_quant(QuantAlgo.FP8) validate_and_set_kv_cache_quant(model_config, "auto") @@ -116,3 +130,76 @@ def test_validate_and_set_kv_cache_quant_rejects_invalid_dtype(): model_config = _make_model_config_with_kv_quant(QuantAlgo.FP8) with pytest.raises(ValueError, match="Accepted types are"): validate_and_set_kv_cache_quant(model_config, "invalid_dtype") + + +def test_model_config_sets_is_encoder_decoder_from_pretrained_config(): + model_config = ModelConfig( + pretrained_config=make_pretrained_config( + head_dim=4, + is_encoder_decoder=True, + ) + ) + + assert model_config.is_encoder_decoder is True + + +def test_validate_encoder_decoder_kv_cache_config_accepts_v1_enc_dec(): + """V1 KVCacheManager is the default and production target for enc-dec models. + + Both V1 and V2 are supported as long as ``cross_kv_cache_fraction`` is set. + """ + model_config = ModelConfig( + pretrained_config=make_pretrained_config( + head_dim=4, + is_encoder_decoder=True, + ) + ) + + validate_encoder_decoder_kv_cache_config( + model_config, + _make_kv_cache_config(use_kv_cache_manager_v2=False, cross_kv_cache_fraction=0.5), + ) + + +def test_validate_encoder_decoder_kv_cache_config_requires_cross_fraction(): + model_config = ModelConfig( + pretrained_config=make_pretrained_config( + head_dim=4, + is_encoder_decoder=True, + ) + ) + + with pytest.raises(ValueError, match="cross_kv_cache_fraction to be set"): + validate_encoder_decoder_kv_cache_config( + model_config, + _make_kv_cache_config(use_kv_cache_manager_v2=True), + ) + + +def test_validate_encoder_decoder_kv_cache_config_rejects_cross_fraction_for_decoder_only(): + model_config = ModelConfig( + pretrained_config=make_pretrained_config( + head_dim=4, + is_encoder_decoder=False, + ) + ) + + with pytest.raises(ValueError, match="should only be set for encoder-decoder models"): + validate_encoder_decoder_kv_cache_config( + model_config, + _make_kv_cache_config(cross_kv_cache_fraction=0.5), + ) + + +def test_validate_encoder_decoder_kv_cache_config_accepts_v2_enc_dec(): + model_config = ModelConfig( + pretrained_config=make_pretrained_config( + head_dim=4, + is_encoder_decoder=True, + ) + ) + + validate_encoder_decoder_kv_cache_config( + model_config, + _make_kv_cache_config(use_kv_cache_manager_v2=True, cross_kv_cache_fraction=0.5), + )