diff --git a/cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp b/cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp index f1c57ec064ac..02355e9f7507 100644 --- a/cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp +++ b/cpp/tensorrt_llm/nanobind/batch_manager/bindings.cpp @@ -32,6 +32,7 @@ #include "tensorrt_llm/runtime/torchView.h" #include +#include #include #include #include @@ -40,6 +41,7 @@ #include #include #include +#include #include #include @@ -577,6 +579,153 @@ void initBindings(nb::module_& m) "Add new tokens to multiple LLM requests. The tokens vector should contain tokens for beam beam_idx of all " "requests in order."); + m.def( + "prepare_encoder_decoder_inputs", + [](std::vector> const& contextRequests, + std::vector> const& generationRequests, at::Tensor const& inputIds, + at::Tensor const& positionIds, at::Tensor const& sequenceLengths, at::Tensor const& promptLengths, + at::Tensor const& cachedTokenLengths, at::Tensor const& kvLengths, at::Tensor const& encoderKvLengths, + at::Tensor const& previousBatchIndices, SizeType32 positionIdOffset) + { + auto checkIntBuffer = [](at::Tensor const& tensor, char const* name) + { + TLLM_CHECK_WITH_INFO(tensor.device().is_cpu(), "%s must be a CPU tensor", name); + TLLM_CHECK_WITH_INFO(tensor.scalar_type() == at::kInt, "%s must have torch.int32 dtype", name); + TLLM_CHECK_WITH_INFO(tensor.is_contiguous(), "%s must be contiguous", name); + }; + checkIntBuffer(inputIds, "input_ids"); + checkIntBuffer(positionIds, "position_ids"); + checkIntBuffer(sequenceLengths, "sequence_lengths"); + checkIntBuffer(promptLengths, "prompt_lengths"); + checkIntBuffer(cachedTokenLengths, "cached_token_lengths"); + checkIntBuffer(kvLengths, "kv_lengths"); + checkIntBuffer(encoderKvLengths, "encoder_kv_lengths"); + checkIntBuffer(previousBatchIndices, "previous_batch_indices"); + + auto const numSequences = contextRequests.size() + generationRequests.size(); + TLLM_CHECK_WITH_INFO(sequenceLengths.numel() >= static_cast(numSequences), + "sequence_lengths capacity is smaller than the batch"); + TLLM_CHECK_WITH_INFO(promptLengths.numel() >= static_cast(numSequences), + "prompt_lengths capacity is smaller than the batch"); + TLLM_CHECK_WITH_INFO(cachedTokenLengths.numel() >= static_cast(numSequences), + "cached_token_lengths capacity is smaller than the batch"); + TLLM_CHECK_WITH_INFO(kvLengths.numel() >= static_cast(numSequences), + "kv_lengths capacity is smaller than the batch"); + TLLM_CHECK_WITH_INFO(encoderKvLengths.numel() >= static_cast(numSequences), + "encoder_kv_lengths capacity is smaller than the batch"); + TLLM_CHECK_WITH_INFO(previousBatchIndices.numel() >= static_cast(generationRequests.size()), + "previous_batch_indices capacity is smaller than the generation batch"); + + auto* inputIdsPtr = inputIds.data_ptr(); + auto* positionIdsPtr = positionIds.data_ptr(); + auto* sequenceLengthsPtr = sequenceLengths.data_ptr(); + auto* promptLengthsPtr = promptLengths.data_ptr(); + auto* cachedTokenLengthsPtr = cachedTokenLengths.data_ptr(); + auto* kvLengthsPtr = kvLengths.data_ptr(); + auto* encoderKvLengthsPtr = encoderKvLengths.data_ptr(); + auto* previousBatchIndicesPtr = previousBatchIndices.data_ptr(); + + std::vector requestIds; + std::vector encoderSequenceLengths; + std::vector encoderCachedTokenLengths; + requestIds.reserve(numSequences); + encoderSequenceLengths.reserve(numSequences); + encoderCachedTokenLengths.reserve(numSequences); + + SizeType32 numTokens{0}; + SizeType32 numContextTokens{0}; + SizeType32 numPreviousBatchRequests{0}; + SizeType32 cachedKvTokens{0}; + SizeType32 contextKvTokens{0}; + SizeType32 generationKvTokens{0}; + SizeType32 maxKvLength{0}; + SizeType32 contextEncoderKvTokens{0}; + SizeType32 generationEncoderKvTokens{0}; + SizeType32 maxEncoderKvLength{0}; + for (auto const& request : contextRequests) + { + auto const sequenceIdx = requestIds.size(); + auto const begin = request->getContextCurrentPosition(); + auto const chunkSize = request->getContextChunkSize(); + auto const& tokens = request->getTokens(0); + TLLM_CHECK_WITH_INFO(begin + chunkSize <= static_cast(tokens.size()), + "Context chunk exceeds the request token count"); + TLLM_CHECK_WITH_INFO(inputIds.numel() >= static_cast(numTokens + chunkSize), + "input_ids capacity is smaller than the packed context"); + TLLM_CHECK_WITH_INFO(positionIds.numel() >= static_cast(numTokens + chunkSize), + "position_ids capacity is smaller than the packed context"); + + std::copy_n(tokens.data() + begin, chunkSize, inputIdsPtr + numTokens); + std::iota(positionIdsPtr + numTokens, positionIdsPtr + numTokens + chunkSize, begin + positionIdOffset); + sequenceLengthsPtr[sequenceIdx] = chunkSize; + promptLengthsPtr[sequenceIdx] = chunkSize; + cachedTokenLengthsPtr[sequenceIdx] = begin; + auto const kvLength = begin + chunkSize; + kvLengthsPtr[sequenceIdx] = kvLength; + auto const encoderKvLength = request->getEncoderOutputLen(); + encoderKvLengthsPtr[sequenceIdx] = encoderKvLength; + cachedKvTokens += begin; + contextKvTokens += kvLength; + maxKvLength = std::max(maxKvLength, kvLength); + contextEncoderKvTokens += encoderKvLength; + maxEncoderKvLength = std::max(maxEncoderKvLength, encoderKvLength); + numTokens += chunkSize; + numContextTokens += chunkSize; + + requestIds.push_back(request->mRequestId); + encoderSequenceLengths.push_back(encoderKvLength); + encoderCachedTokenLengths.push_back(0); + } + + bool sawDummyRequest{false}; + for (auto const& request : generationRequests) + { + auto const sequenceIdx = requestIds.size(); + auto const isDummy = request->isDummyRequest(); + sawDummyRequest = sawDummyRequest || isDummy; + TLLM_CHECK_WITH_INFO( + isDummy || !sawDummyRequest, "CUDA graph dummy requests must follow real generation requests"); + TLLM_CHECK_WITH_INFO( + positionIds.numel() > numTokens, "position_ids capacity is smaller than the packed batch"); + + auto const pastSeenTokens = request->getMaxBeamNumTokens() - (isDummy ? 1 : 0); + positionIdsPtr[numTokens] = pastSeenTokens + positionIdOffset; + sequenceLengthsPtr[sequenceIdx] = 1; + promptLengthsPtr[sequenceIdx] = request->mPromptLen; + cachedTokenLengthsPtr[sequenceIdx] = pastSeenTokens; + auto const kvLength = pastSeenTokens + 1; + kvLengthsPtr[sequenceIdx] = kvLength; + auto const encoderKvLength = request->getEncoderOutputLen(); + encoderKvLengthsPtr[sequenceIdx] = encoderKvLength; + cachedKvTokens += pastSeenTokens; + generationKvTokens += kvLength; + maxKvLength = std::max(maxKvLength, kvLength); + generationEncoderKvTokens += encoderKvLength; + maxEncoderKvLength = std::max(maxEncoderKvLength, encoderKvLength); + ++numTokens; + + if (!isDummy) + { + TLLM_CHECK_WITH_INFO( + request->mSeqSlot.has_value(), "A real generation request must have a sequence slot"); + previousBatchIndicesPtr[numPreviousBatchRequests++] = request->mSeqSlot.value(); + } + + requestIds.push_back(request->mRequestId); + encoderSequenceLengths.push_back(0); + encoderCachedTokenLengths.push_back(encoderKvLength); + } + + return std::make_tuple(requestIds, encoderSequenceLengths, encoderCachedTokenLengths, numTokens, + numContextTokens, numPreviousBatchRequests, cachedKvTokens, contextKvTokens, generationKvTokens, + maxKvLength, contextEncoderKvTokens, generationEncoderKvTokens, maxEncoderKvLength); + }, + nb::arg("context_requests"), nb::arg("generation_requests"), nb::arg("input_ids"), nb::arg("position_ids"), + nb::arg("sequence_lengths"), nb::arg("prompt_lengths"), nb::arg("cached_token_lengths"), nb::arg("kv_lengths"), + nb::arg("encoder_kv_lengths"), nb::arg("previous_batch_indices"), nb::arg("position_id_offset") = 0, + nb::call_guard(), + "Prepare the persistent CPU input buffers for a simple encoder-decoder batch."); + m.def( "make_decoding_batch_input", [](tb::DecoderInputBuffers& decoderInputBuffers, runtime::decoder::DecoderState& decoderState, diff --git a/docs/source/models/encoder-decoder.md b/docs/source/models/encoder-decoder.md index 5675eed4d742..8d09f6816282 100644 --- a/docs/source/models/encoder-decoder.md +++ b/docs/source/models/encoder-decoder.md @@ -45,7 +45,7 @@ The following table describes the supported and recommended configurations. | Beam search | Yes with V1 | Configure `max_beam_width` when constructing `LLM`, then set `use_beam_search=True` in `SamplingParams`. | | Attention backend | `TRTLLM` | Use this backend for encoder-decoder models. It is required when `tensor_parallel_size > 1`. | | Decoder CUDA graphs | Yes, except in FP32 | `CudaGraphConfig` captures decoder work. V1 supports greedy and beam search; V2 supports its single-beam path. FP32 encoder-decoder models decline capture at engine init and log a warning instead of failing. | -| Encoder CUDA graphs | No | `EncodeCudaGraphConfig` is disabled for encoder-decoder models. The encoder runs eagerly. | +| Encoder CUDA graphs | Yes | Set `encoder_cuda_graph_config=EncodeCudaGraphConfig(...)` and `encoder_max_batch_size`. Usually set `encoder_max_batch_size` lower than `max_batch_size`. The `TRTLLM` attention backend is required. | | Overlap scheduler | Yes | Enabled by default. V1 supports greedy decoding and beam search; V2 remains limited to `max_beam_width=1`. | | Tensor parallelism | Yes | Use `tensor_parallel_size > 1` with `attn_backend="TRTLLM"`. Attention head counts must be divisible by the TP size. | | Pipeline parallelism | No | Keep `pipeline_parallel_size=1`. | @@ -397,12 +397,12 @@ to return only the best hypothesis. Beam search expands decoder-side cache and compute requirements. Include this expansion when sizing the self-attention KV pool and CUDA graph batch sizes. -## Enable decoder CUDA graphs +## Enable encoder and decoder CUDA graphs -Pass `CudaGraphConfig` to capture and replay decoder iterations: +Configure the decoder and encoder graph grids separately: ```python -from tensorrt_llm.llmapi import CudaGraphConfig +from tensorrt_llm.llmapi import CudaGraphConfig, EncodeCudaGraphConfig llm = LLM( @@ -410,10 +410,19 @@ llm = LLM( backend="pytorch", attn_backend="TRTLLM", max_batch_size=8, + encoder_max_batch_size=2, + encoder_max_num_tokens=2048, cuda_graph_config=CudaGraphConfig( max_batch_size=8, enable_padding=True, ), + encoder_cuda_graph_config=EncodeCudaGraphConfig( + batch_sizes=[1, 2], + num_tokens=[128, 256, 512, 1024, 2048], + seq_lens=[128, 256, 512, 1024], + enable_padding=True, + ), + enable_encoder_decoder_mixed_cuda_graph=True, kv_cache_config=KvCacheConfig( free_gpu_memory_fraction=0.8, cross_kv_cache_fraction=0.5, @@ -421,14 +430,31 @@ llm = LLM( ) ``` -This configuration captures decoder work only; the encoder continues to run -eagerly. With beam search, graph batch sizes must cover the active decoder -sequences after beam expansion. Padding lets nearby runtime batch sizes reuse a -captured graph. - -Do not use `EncodeCudaGraphConfig` for an encoder-decoder model. The runtime -warns and disables it. Piecewise CUDA graphs through `TorchCompileConfig` are -also unsupported for this model type. +`cuda_graph_config` controls decoder and mixed decoder graphs. +`encoder_cuda_graph_config` controls encoder-forward graph buckets for batch +size, total packed tokens, and maximum sequence length. The +`encoder_max_batch_size` value is the hard encoder capacity and admission +limit. With beam search, decoder graph batch sizes must cover the active +decoder sequences after beam expansion. + +`max_batch_size` controls the total decoder concurrency, while +`encoder_max_batch_size` controls encoder microbatch admission. For better +performance, tune `encoder_max_batch_size`, `encoder_max_num_tokens`, and the +encoder CUDA graph buckets together for the production workload. Start with +`encoder_max_batch_size` smaller than `max_batch_size`, such as 2 versus 8, +then adjust the limits and capture buckets based on benchmark results. + +`enable_encoder_decoder_mixed_cuda_graph` is primarily a performance option. It +reduces CPU launch overhead for decoder iterations that mix newly admitted +context requests with ongoing generation requests. The option defaults to +`True`, but becomes effective only when the encoder and decoder graph +configurations produce usable capture shapes. Set it to `False` to disable +mixed graphs while retaining the separate encoder and decoder CUDA graphs. + +Passing `EncodeCudaGraphConfig` through `cuda_graph_config` remains unsupported +for encoder-decoder models; pass it through `encoder_cuda_graph_config` +instead. Piecewise CUDA graphs through `TorchCompileConfig` are also +unsupported for this model type. ## Control the overlap scheduler @@ -491,10 +517,22 @@ attn_backend: TRTLLM dtype: bfloat16 disable_overlap_scheduler: false enable_chunked_prefill: false +max_batch_size: 8 +encoder_max_batch_size: 2 +encoder_max_num_tokens: 1024 max_beam_width: 1 max_input_len: 512 max_num_tokens: 2048 max_seq_len: 512 +cuda_graph_config: + max_batch_size: 8 + enable_padding: true +encoder_cuda_graph_config: + batch_sizes: [1, 2] + num_tokens: [128, 256, 512, 1024] + seq_lens: [128, 256, 512] + enable_padding: true +enable_encoder_decoder_mixed_cuda_graph: true kv_cache_config: enable_block_reuse: false free_gpu_memory_fraction: 0.8 @@ -509,7 +547,6 @@ Start the server: ```bash trtllm-serve google/flan-t5-small \ --backend pytorch \ - --max_batch_size 4 \ --config enc-dec-config.yaml ``` @@ -549,68 +586,36 @@ Use these guidelines as a starting point: - Set `max_seq_len` to at least the larger of the maximum encoder input length and maximum decoded sequence length. The current encoder-decoder runtime uses this value while sizing both phases. -- Set `max_num_tokens` high enough for all encoder tokens admitted together and - for the active decoder tokens. This is especially important for mixed-length - batches. -- Increase `max_batch_size` for more concurrent requests. Beam width multiplies - the number of active decoder sequences but not the number of source requests. +- Set `max_num_tokens` high enough for the active decoder tokens. +- Set `encoder_max_num_tokens` high enough for all encoder tokens in one + encoder microbatch. This is especially important for mixed-length batches. +- Increase `max_batch_size` for more concurrent requests. Start with a smaller + `encoder_max_batch_size`, such as 2 when `max_batch_size=8`, to bound encoder + memory and admission cost without reducing decoder concurrency. Beam width + multiplies the number of active decoder sequences but not the number of + source requests. - Tune `free_gpu_memory_fraction` first, then tune `cross_kv_cache_fraction` based on whether the cross-attention or self-attention pool is exhausted. ## Performance -The following benchmarks compare the PyTorch backend with the legacy TensorRT -encoder-decoder path for large-batch inference. The measurements use BF16 on -one H100 80 GB GPU with greedy decoding, an output limit of 128 tokens, and -mixed encoder input lengths from 260 to 440 tokens. The Flan-T5-XL results are -the average of ten timed runs after three warmup runs. The BART results are the -average of 20 timed runs after five warmup runs. Executed-token throughput -includes the terminal EOS token when a sequence emits it. - -The PyTorch configuration uses the `TRTLLM` attention backend, the overlap -scheduler, the Python scheduler, decoder CUDA graphs with padding, KV cache -manager V1, `max_input_len=512`, `max_seq_len=1024`, and -`max_num_tokens=65536`. Block reuse and chunked prefill are disabled. The KV -cache uses `free_gpu_memory_fraction=0.3` and -`cross_kv_cache_fraction=0.5`. - -The legacy TensorRT configuration uses separate BF16 encoder and decoder -engines built for batch size 128 and beam width 1. The encoder supports 512 -input tokens and 65,536 tokens per iteration; the decoder supports a sequence -length of 129. The benchmark runs these engines through `ModelRunnerCpp` with -greedy `top_k=1` decoding and the same KV cache fractions. For BART, the legacy -TensorRT benchmark starts the decoder with token IDs `[2, 0]` and generates at -most 127 more tokens. The PyTorch LLM API applies the same decoder prefix -internally and counts token ID 0 as the first output token; customers do not -need to provide the decoder prefix. Both paths use token ID 2 as EOS and stop -when the model generates it naturally. If a sequence reaches the output limit, -it retains the model-selected final token and reports a length stop instead of -forcing EOS. This setup also lets beam search begin after the shared decoder -prefix without a per-step Python logits processor. - -### Flan-T5-XL - -For Flan-T5-XL, the PyTorch backend performs on par with the legacy TensorRT -path, with slightly lower latency and higher executed-token throughput across -the tested batch sizes. - -| Batch size | Legacy TensorRT latency | PyTorch latency | PyTorch latency improvement over legacy TensorRT | Legacy TensorRT executed tokens/s | PyTorch executed tokens/s | -| ---: | ---: | ---: | ---: | ---: | ---: | -| 32 | 727.6 ms | 706.1 ms | 3.0% | 3,153 | 3,312 | -| 64 | 1,225.0 ms | 1,136.7 ms | 7.2% | 3,863 | 4,184 | -| 128 | 2,056.8 ms | 1,999.3 ms | 2.8% | 4,601 | 4,768 | - -### BART-large-CNN - -For BART-large-CNN, the PyTorch backend has 21.9% to 36.0% higher latency than -the legacy TensorRT path across the tested batch sizes. - -| Batch size | Legacy TensorRT latency | PyTorch latency | PyTorch latency difference | Legacy TensorRT executed tokens/s | PyTorch executed tokens/s | -| ---: | ---: | ---: | ---: | ---: | ---: | -| 32 | 229.9 ms | 280.2 ms | 21.9% slower | 12,611 | 10,662 | -| 64 | 252.2 ms | 343.0 ms | 36.0% slower | 22,007 | 16,209 | -| 128 | 352.3 ms | 472.7 ms | 34.2% slower | 31,544 | 23,518 | +Configure encoder, decoder, and mixed decoder CUDA graphs for the expected +serving workload. With representative capture buckets, the PyTorch backend can +outperform the legacy TensorRT encoder-decoder path while avoiding the engine +build and checkpoint conversion steps. + +For example, a BF16 FLAN-T5 Large serving benchmark on one H100 80 GB GPU used +encoder CUDA graphs, padded decoder CUDA graphs, and mixed encoder-decoder CUDA +graphs. Compared with the legacy TensorRT path, the PyTorch backend delivered +65.8%, 12.6%, and 11.8% higher request throughput at concurrencies 8, 32, and +64, respectively. P99 latency was 51.6%, 31.0%, and 33.5% lower. + +Follow [Enable encoder and decoder CUDA graphs](#enable-encoder-and-decoder-cuda-graphs) +and choose capture buckets that cover the batch sizes, packed encoder token +counts, and sequence lengths expected in production. Capture grids that omit +common runtime shapes fall back to eager execution and can lose these +performance benefits. Performance depends on the model, request distribution, decoding settings, and GPU configuration. Benchmark with a representative workload before deployment. @@ -644,11 +649,12 @@ Check all of the following: - `pipeline_parallel_size=1` and `context_parallel_size=1`. - `enable_attention_dp=False`. -### CUDA graphs do not capture the encoder +### Encoder CUDA graphs fall back to eager execution -This is expected. `CudaGraphConfig` accelerates decoder iterations only. The -encoder path runs eagerly, and `EncodeCudaGraphConfig` is disabled for -encoder-decoder models. +Check that `encoder_cuda_graph_config` and `encoder_max_batch_size` are set, +that the encoder graph buckets cover the request shape, and that +`attn_backend="TRTLLM"`. Unsupported shapes and attention backends fall back to +eager encoder execution. ### Output quality differs from the Hugging Face example diff --git a/docs/source/models/supported-models.md b/docs/source/models/supported-models.md index 4f4db5525b60..f93cf2eac866 100644 --- a/docs/source/models/supported-models.md +++ b/docs/source/models/supported-models.md @@ -6,6 +6,7 @@ The following is a table of supported models for the PyTorch backend: | Architecture | Model | HuggingFace Example | | ------------------------------------ | ---------------------------------- | -------------------------------------------- | | `AfmoeForCausalLM` | Arcee Foundation MoE (Trinity) | `arcee-ai/Trinity-Mini` | +| `BartForConditionalGeneration` | BART | `facebook/bart-large-cnn` | | `BertForSequenceClassification` | BERT-based | `textattack/bert-base-uncased-yelp-polarity` | | `Cohere2ForCausalLM` | Command A | `CohereLabs/c4ai-command-a-03-2025` | | `DeciLMForCausalLM` | Nemotron | `nvidia/Llama-3_1-Nemotron-51B-Instruct` | @@ -33,6 +34,7 @@ The following is a table of supported models for the PyTorch backend: | `LagunaForCausalLM` | Laguna-XS | `poolside/laguna-XS.2` | | `LlamaForCausalLM` | Llama 3.1, Llama 3, Llama 2, LLaMA | `meta-llama/Meta-Llama-3.1-70B` | | `Llama4ForConditionalGeneration` | Llama 4 | `meta-llama/Llama-4-Scout-17B-16E-Instruct` | +| `MBartForConditionalGeneration` | mBART | `facebook/mbart-large-50-many-to-one-mmt` | | `MiniCPMV4_6ForConditionalGeneration` [^14]| MiniCPM-V 4.6 | `openbmb/MiniCPM-V-4.6` | | `MiniMaxM2ForCausalLM` [^5] | MiniMax M2/M2.1/M2.7 | `MiniMaxAI/MiniMax-M2.7` | | `MiniMaxM3SparseForConditionalGeneration` [^12]| MiniMax-M3 | `MiniMaxAI/MiniMax-M3` | @@ -57,6 +59,8 @@ The following is a table of supported models for the PyTorch backend: | `SkyworkR1V2ForConditionalGeneration` [^5] | Skywork R1V2, Skywork SWE | `Skywork/Skywork-R1V2-38B` | | `SmolLM3ForCausalLM` [^5] | SmolLM3 | `HuggingFaceTB/SmolLM3-3B` | | `Step3p7ForConditionalGeneration` [^8]| Step-3.7-Flash | `stepfun-ai/Step-3.7-Flash` | +| `T5ForConditionalGeneration` | T5, Flan-T5, ByT5 | `google/flan-t5-small` | +| `WhisperForConditionalGeneration` | Whisper | `openai/whisper-large-v3` | ## Model-Feature Support Matrix (Key Models) @@ -95,6 +99,24 @@ Note: Support for other models may vary. Features marked "N/A" are not applicabl [^13]: The Cosmos 3 family also supports visual generation through the VisualGen API. See [Visual Generation Models](#visual-generation-models). [^14]: Requires `transformers>=5.7.0`: MiniCPM-V 4.6 was upstreamed into transformers as a native model type (`minicpmv4_6`) and the checkpoint ships no remote code (`auto_map`) to fall back on. The Qwen3.5-hybrid text tower runs in BF16. Image, video, and text inputs are supported in this release (video reuses the same NaViT-packed vision path as image via `MiniCPMV4_6InputProcessor`). +# Encoder-Decoder Feature Support Matrix (PyTorch Backend) + +The following capabilities apply to the supported encoder-decoder architectures. For configuration guidance and +limitations, see [Use encoder-decoder models with the PyTorch backend](./encoder-decoder.md). + +| Model Architecture/Feature | Overlap Scheduler | Decoder CUDA Graph | Encoder CUDA Graph | KV Cache Manager V1 | KV Cache Manager V2 | Beam Search | Tensor Parallelism | Pipeline Parallelism | Chunked Prefill | +| --- | --- | --- | --- | --- | --- | --- | --- | --- | --- | +| `BartForConditionalGeneration` | Yes | Yes (except FP32) | Yes | Yes | Yes (single beam) | Yes (V1 only) | Yes | No | No (encoder phase) | +| `MBartForConditionalGeneration` | Yes | Yes (except FP32) | Yes | Yes | Yes (single beam) | Yes (V1 only) | Yes | No | No (encoder phase) | +| `T5ForConditionalGeneration` | Yes | Yes (except FP32) | Yes | Yes | Yes (single beam) | Yes (V1 only) | Yes | No | No (encoder phase) | +| `WhisperForConditionalGeneration` | Yes | Yes (except FP32) | No (feature inputs) | Yes | Yes (single beam) | Yes (V1 only) | Yes | No | No (encoder phase) | + +Decoder CUDA graphs support greedy and beam-search decoding with KV cache manager V1 and single-beam decoding with +V2. Encoder CUDA graphs support the token-input BART, mBART, and T5 families; Whisper's feature-driven audio encoder +runs eagerly. Use the `TRTLLM` attention backend for encoder-decoder models; tensor parallelism also requires attention +head counts divisible by the tensor parallel size. Chunked prefill is not supported for the encoder phase, so the +complete encoder input must fit in the iteration token budget. + # Multimodal Feature Support Matrix (PyTorch Backend) | Model Architecture/Feature | Overlap Scheduler | CUDA Graph | Chunked Prefill | Torch Sampler | TLLM C++ Sampler | KV Cache Reuse | Logits Post Processor | EPD Disaggregated Serving | Modality | diff --git a/tensorrt_llm/_torch/attention_backend/trtllm.py b/tensorrt_llm/_torch/attention_backend/trtllm.py index cb2d5e308d61..56af4c1e56a5 100644 --- a/tensorrt_llm/_torch/attention_backend/trtllm.py +++ b/tensorrt_llm/_torch/attention_backend/trtllm.py @@ -642,6 +642,56 @@ def prepare(self) -> None: host_request_types=self.host_request_types[:self.num_seqs], ) + def prepare_encoder_decoder(self, prompt_lens: torch.Tensor, + kv_lens: torch.Tensor, context_kv_tokens: int, + generation_kv_tokens: int, + max_kv_len: int) -> None: + """Prepare simple encoder-decoder attention from native host buffers.""" + super().prepare() + extra_attrs = get_model_extra_attrs() + if extra_attrs is None: + get_global_attrs().attention_metadata = weakref.ref(self) + + assert self.kv_cache_manager is not None + assert self.draft_kv_cache_manager is None + assert not self.is_spec_decoding_enabled + assert self.kv_cache_params.num_extra_kv_tokens == 0 + assert not self.enable_flash_mla + assert not self.enable_helix + assert not self.enable_context_mla_with_cached_kv + assert self.request_ids is not None + assert max_kv_len <= self.kv_cache_manager.max_seq_len, ( + f"The max KV cache length of input sequences ({max_kv_len}) " + "exceeds the KV cache manager's maximum supported length " + f"({self.kv_cache_manager.max_seq_len}).") + + num_seqs = self.num_seqs + self.prompt_lens_cuda[:num_seqs].copy_(prompt_lens, non_blocking=True) + self.kv_lens_cuda[:num_seqs].copy_(kv_lens, non_blocking=True) + self.host_total_kv_lens[0] = context_kv_tokens + self.host_total_kv_lens[1] = generation_kv_tokens + self.host_request_types[:self.num_contexts].fill_(0) + self.host_request_types[self.num_contexts:num_seqs].fill_(1) + + max_blocks = None + if self.kv_cache_manager.tokens_per_block: + max_blocks = ceil_div(max_kv_len, + self.kv_cache_manager.tokens_per_block) + self.kv_cache_manager.copy_batch_block_offsets( + self.kv_cache_block_offsets, + self.request_ids, + self.beam_width, + self.num_contexts, + num_seqs, + max_blocks=max_blocks) + self._bind_runtime_views( + kv_lens_cuda=self.kv_lens_cuda[:num_seqs], + kv_lens=kv_lens, + prompt_lens_cuda=self.prompt_lens_cuda[:num_seqs], + prompt_lens_cpu=prompt_lens, + host_request_types=self.host_request_types[:num_seqs], + ) + def prepare_encoder_only(self) -> None: """Fast path for encoder-only forward (eager + CUDA graph capture).""" extra_attrs = get_model_extra_attrs() diff --git a/tensorrt_llm/_torch/pyexecutor/cuda_graph_runner.py b/tensorrt_llm/_torch/pyexecutor/cuda_graph_runner.py index 05343d0deddd..62f0c334a8fb 100644 --- a/tensorrt_llm/_torch/pyexecutor/cuda_graph_runner.py +++ b/tensorrt_llm/_torch/pyexecutor/cuda_graph_runner.py @@ -16,7 +16,7 @@ from ..attention_backend.trtllm import TrtllmAttentionMetadata from ..distributed import Distributed from ..expert_statistic import ExpertStatistic -from ..memory_buffer_utils import get_memory_buffers +from ..memory_buffer_utils import Buffers, get_memory_buffers from ..modules.multi_stream_utils import with_multi_stream from ..speculative.eagle3 import Eagle3ResourceManager from ..speculative.interface import SpecMetadata @@ -36,7 +36,8 @@ # as a one-token context chunk to write its cross-KV cache, so enc-dec # dummies need one prompt token plus one generated token. ENC_DEC_CUDA_GRAPH_DUMMY_TOKEN_NUM = 2 -KeyType: TypeAlias = Tuple[int, int, bool, bool, bool] +KeyType: TypeAlias = Tuple[int, int, bool, bool, bool, Tuple[int, ...], + Tuple[int, ...]] def _save_spec_decode_capture_state( @@ -111,6 +112,7 @@ class CUDAGraphRunnerConfig: kv_cache_manager_key: Any dynamic_draft_len_mapping: Optional[Dict[int, int]] = None sparse_attention_config: Optional[BaseSparseAttentionConfig] = None + enable_encoder_decoder_mixed_cuda_graph: bool = False class CUDAGraphRunner: @@ -135,6 +137,8 @@ def __init__(self, config: CUDAGraphRunnerConfig): self.spec_config = config.spec_config self.sparse_config = config.sparse_attention_config self.is_encoder_decoder = config.is_encoder_decoder + self.enable_encoder_decoder_mixed_cuda_graph = ( + config.enable_encoder_decoder_mixed_cuda_graph) self.graphs: Dict[KeyType, torch.cuda.CUDAGraph] = {} self.graph_outputs: Dict[KeyType, @@ -163,6 +167,10 @@ def _create_shared_static_tensors(self): token_per_request = runtime_draft_token_buffer_width + 1 max_total_tokens = (self.max_supported_batch_size * self.max_beam_width * token_per_request) + if self.enable_encoder_decoder_mixed_cuda_graph: + # A mixed encoder-decoder batch can contain multiple decoder + # context tokens per request, unlike a pure generation batch. + max_total_tokens = self.config.max_num_tokens max_total_tokens = min(max_total_tokens, self.config.max_num_tokens) self.shared_static_tensors = { @@ -179,6 +187,51 @@ def _create_shared_static_tensors(self): "mrope_delta_read_seq_slots"] = torch.zeros( (max_total_tokens, ), device="cuda", dtype=torch.long) + def _get_static_encoder_hidden_states( + self, + encoder_hidden_states: torch.Tensor, + num_encoder_tokens: int, + *, + allow_allocate: bool, + ) -> torch.Tensor: + """Return the stable mixed-graph encoder input, allocating at warmup.""" + if encoder_hidden_states.ndim != 2: + raise RuntimeError( + "Mixed encoder-decoder CUDA graphs require rank-2 packed " + "encoder hidden states.") + + static_encoder_hidden_states = self.shared_static_tensors.get( + "encoder_hidden_states") + if static_encoder_hidden_states is None: + if not allow_allocate: + raise RuntimeError( + "Mixed encoder-decoder CUDA graph replay requires the " + "encoder hidden-state buffer initialized during warmup.") + if self.graphs: + raise RuntimeError( + "Mixed encoder-decoder CUDA graph encoder hidden-state " + "buffer cannot be allocated after graph capture.") + static_encoder_hidden_states = encoder_hidden_states.new_empty( + (num_encoder_tokens, encoder_hidden_states.shape[1])) + self.shared_static_tensors[ + "encoder_hidden_states"] = static_encoder_hidden_states + + if static_encoder_hidden_states.shape[0] < num_encoder_tokens: + raise RuntimeError( + "Mixed encoder-decoder CUDA graph encoder hidden-state buffer " + f"has capacity {static_encoder_hidden_states.shape[0]}, but " + f"{num_encoder_tokens} tokens were requested.") + return static_encoder_hidden_states[:num_encoder_tokens] + + def _is_mixed_encoder_decoder_batch(self, batch: ScheduledRequests) -> bool: + return (self.enable_encoder_decoder_mixed_cuda_graph + and batch.num_context_requests > 0 + and batch.num_generation_requests > 0) + + def _can_run_cuda_graph_batch(self, batch: ScheduledRequests) -> bool: + return batch.can_run_cuda_graph or self._is_mixed_encoder_decoder_batch( + batch) + def _get_seq_len_mode( self, batch: ScheduledRequests, @@ -274,7 +327,7 @@ def get_graph_key( # Because we will pad the input to 'max_draft_len' length for the first draft layer. draft_len = self.config.original_max_draft_len if spec_resource_manager.is_first_draft else 0 key = (batch_size, draft_len, spec_resource_manager.is_first_draft, - short_seq_len_mode, is_all_greedy_sample) + short_seq_len_mode, is_all_greedy_sample, (), ()) else: # With dynamic spec decode, the draft length may be zero even when enable_spec_decode is True, # so we need to get the draft length from the batch instead of using enable_spec_decode. @@ -284,10 +337,34 @@ def get_graph_key( draft_len = max(draft_len_list) assert len( set(draft_len_list)) == 1, "All draft lengths must be the same" + context_query_lens = tuple( + int(request.context_chunk_size) + for request in batch.context_requests) + encoder_input_lens = (sum( + int(request.encoder_output_len) + for request in batch.context_requests + if not request.py_skip_cross_kv_projection), ) key = (batch_size, draft_len, False, short_seq_len_mode, - is_all_greedy_sample) + is_all_greedy_sample, context_query_lens, encoder_input_lens) return key + def _get_compatible_mixed_encoder_decoder_key(self, + key: KeyType) -> KeyType: + """Round the packed encoder extent up to a captured graph key.""" + if (not self.padding_enabled or self._capture_allowed + or key in self.graph_metadata or len(key[6]) != 1): + return key + + num_encoder_tokens = key[6][0] + compatible_keys = [ + captured_key for captured_key in self.graph_outputs + if captured_key[:6] == key[:6] and len(captured_key[6]) == 1 + and captured_key[6][0] >= num_encoder_tokens + ] + if not compatible_keys: + return key + return min(compatible_keys, key=lambda captured_key: captured_key[6][0]) + @staticmethod def _get_mrope_position_delta(request: Any) -> Optional[Any]: mrope_position_delta = getattr(request, "py_mrope_position_delta", None) @@ -326,6 +403,7 @@ def maybe_get_cuda_graph( draft_tokens_cuda: Optional[torch.Tensor] = None, new_tensors_device: Optional[SampleStateTensors] = None, spec_resource_manager: Optional[BaseResourceManager] = None, + allow_mixed_encoder_decoder: bool = False, promoted_context_request_ids: frozenset[int] = frozenset(), ) -> Tuple[Optional[Any], Optional[Any], Optional[KeyType]]: """ @@ -344,7 +422,10 @@ def maybe_get_cuda_graph( if ExpertStatistic.should_record(): return None, None, None - can_run_cuda_graph = batch.can_run_cuda_graph + is_mixed_encoder_decoder = self._is_mixed_encoder_decoder_batch(batch) + can_run_cuda_graph = (batch.can_run_cuda_graph + or (is_mixed_encoder_decoder + and allow_mixed_encoder_decoder)) batch_size = batch.batch_size if self.enabled and self.config.enable_attention_dp and self.config.mapping.tp_size > 1: all_can_graph_batch = self.config.dist.tp_allgather( @@ -373,13 +454,17 @@ def maybe_get_cuda_graph( key = self.get_graph_key(batch, new_tensors_device, spec_resource_manager, spec_metadata, promoted_context_request_ids) + if is_mixed_encoder_decoder: + key = self._get_compatible_mixed_encoder_decoder_key(key) if key in self.graph_metadata: return self.graph_metadata[key][ "attn_metadata"], self.graph_metadata[key]["spec_metadata"], key - # Graph doesn't exist yet. If on-the-fly capture is not allowed, - # fall back to eager so the caller doesn't need a separate check. + # Capturing a mixed graph on a live batch would execute its KV-cache + # writes during graph warmup/capture and could resize shared attention + # workspace after older graph pointers have been fixed. Only shapes + # captured by the two-pass startup warmup may replay. if not self._capture_allowed: return None, None, None @@ -389,6 +474,15 @@ def maybe_get_cuda_graph( num_sequences_in_batch = batch_size * self.max_beam_width graph_attn_metadata = attn_metadata.create_cuda_graph_metadata( num_sequences_in_batch, False, key[1], self.cuda_graph_meta_buffers) + if is_mixed_encoder_decoder: + context_query_lens = key[5] + generation_query_len = key[1] + 1 + graph_attn_metadata.seq_lens = torch.tensor( + context_query_lens + (generation_query_len, ) * + (num_sequences_in_batch - len(context_query_lens)), + dtype=torch.int, + ) + graph_attn_metadata.num_contexts = len(context_query_lens) assert graph_attn_metadata.is_cuda_graph if enable_spec_decode: @@ -457,6 +551,15 @@ def get_graph_pool(self): """ return self.memory_pool + def _get_num_tokens_for_key(self, key: KeyType) -> int: + batch_size = key[0] + token_per_generation = key[1] + 1 + context_query_lens = key[5] + num_contexts = len(context_query_lens) + return (sum(context_query_lens) + + (batch_size * self.max_beam_width - num_contexts) * + token_per_generation) + def capture(self, key: KeyType, forward_fn: Callable, @@ -468,10 +571,7 @@ def capture(self, # [CUDA graph spec decode padding] # We pad input IDs/position IDs to the maximum draft length (token per request). # We're forced to do this because we cannot reallocate inputs over many graph runs. - max_draft_len = key[1] - token_per_request = max_draft_len + 1 - num_tokens_for_capture = (batch_size * self.max_beam_width * - token_per_request) + num_tokens_for_capture = self._get_num_tokens_for_key(key) sliced_static_tensors = { "input_ids": @@ -491,6 +591,30 @@ def capture(self, capture_inputs = initial_inputs.copy() capture_inputs.update(sliced_static_tensors) + encoder_input_lens = key[6] + num_encoder_tokens = sum(encoder_input_lens) + if num_encoder_tokens: + encoder_hidden_states = initial_inputs.get("encoder_hidden_states") + if encoder_hidden_states is None: + raise RuntimeError("Mixed encoder-decoder CUDA graph capture " + "requires encoder hidden states.") + static_encoder_hidden_states = ( + self._get_static_encoder_hidden_states( + encoder_hidden_states, + num_encoder_tokens, + allow_allocate=True, + )) + actual_num_encoder_tokens = encoder_hidden_states.shape[0] + if actual_num_encoder_tokens > num_encoder_tokens: + raise RuntimeError( + "Mixed encoder-decoder CUDA graph capture received " + f"{actual_num_encoder_tokens} encoder tokens for a " + f"{num_encoder_tokens}-token graph.") + static_encoder_hidden_states[:actual_num_encoder_tokens].copy_( + encoder_hidden_states) + static_encoder_hidden_states[actual_num_encoder_tokens:].zero_() + capture_inputs[ + "encoder_hidden_states"] = static_encoder_hidden_states attn_metadata = capture_inputs["attn_metadata"] saved_kv_lens_cuda = _save_spec_decode_capture_state( attn_metadata, enable_spec_decode) @@ -571,6 +695,28 @@ def replay(self, key: KeyType, else: static_tensors["position_ids"][:, :seqlen].copy_(position_ids) + num_encoder_tokens = sum(key[6]) + if num_encoder_tokens: + encoder_hidden_states = current_inputs.get("encoder_hidden_states") + if encoder_hidden_states is None: + raise RuntimeError("Mixed encoder-decoder CUDA graph replay " + "requires encoder hidden states.") + actual_num_encoder_tokens = encoder_hidden_states.shape[0] + if actual_num_encoder_tokens > num_encoder_tokens: + raise RuntimeError( + "Mixed encoder-decoder CUDA graph replay received " + f"{actual_num_encoder_tokens} encoder tokens for a " + f"{num_encoder_tokens}-token graph.") + static_encoder_hidden_states = ( + self._get_static_encoder_hidden_states( + encoder_hidden_states, + num_encoder_tokens, + allow_allocate=False, + )) + static_encoder_hidden_states[:actual_num_encoder_tokens].copy_( + encoder_hidden_states) + static_encoder_hidden_states[actual_num_encoder_tokens:].zero_() + self.graphs[key].replay() output_ref = self.graph_outputs[key] @@ -581,7 +727,7 @@ def _get_padded_batch(self, batch: ScheduledRequests, runtime_draft_len: int) -> int: kv_cache_manager = resource_manager.get_resource_manager( self.config.kv_cache_manager_key) - can_run_cuda_graph = batch.can_run_cuda_graph + can_run_cuda_graph = self._can_run_cuda_graph_batch(batch) batch_size = batch.batch_size new_batch_size = batch_size @@ -775,6 +921,8 @@ def clear(self): EncoderKeyType: TypeAlias = Tuple[int, int, int] +_ENCODER_SOURCE_SEQ_LENS = "_encoder_source_seq_lens" +_ENCODER_SOURCE_TO_SLOT = "_encoder_source_to_slot" @dataclass @@ -790,6 +938,8 @@ class EncoderCUDAGraphRunnerConfig: max_num_tokens: int max_seq_len: int cuda_graph_mem_pool: Any + is_encoder_decoder: bool = False + use_fixed_sequence_slots: bool = False class EncoderCUDAGraphRunner: @@ -797,9 +947,11 @@ class EncoderCUDAGraphRunner: Designed for encoder inputs with `input_ids` (flat [total_tokens]) and `seq_lens` ([batch_size]). Encoder CUDA graphs are keyed on the 3-tuple - (padded_batch_size, padded_num_tokens, padded_max_seq_len). + (padded_batch_size, padded_total_tokens, max_seq_len_bucket) for dynamic + encoder-decoder batches when padding is enabled. - Restricted to `TrtllmAttentionMetadata` — FlashInfer's per-batch planner state is not compatible with CUDA graph capture/replay. + Restricted to `TrtllmAttentionMetadata`: FlashInfer's per-batch planner + state is not compatible with CUDA graph capture/replay. """ WARMUP_STEPS = 1 @@ -814,6 +966,17 @@ def __init__(self, config: EncoderCUDAGraphRunnerConfig): self.supported_num_tokens = sorted(config.cuda_graph_num_tokens) self.max_supported_num_tokens = config.max_cuda_graph_num_tokens self.supported_seq_lens = sorted(config.cuda_graph_seq_lens) + self.is_encoder_decoder = config.is_encoder_decoder + self.use_fixed_sequence_slots = config.use_fixed_sequence_slots + self.capture_keys: frozenset[EncoderKeyType] = frozenset() + self._capture_sequence_lengths: Dict[EncoderKeyType, List[int]] = {} + if self.is_encoder_decoder: + self._capture_sequence_lengths = ( + self._build_encoder_decoder_capture_layouts()) + self.capture_keys = frozenset(self._capture_sequence_lengths) + self._capture_keys_by_batch_size: Dict[int, List[EncoderKeyType]] = {} + for key in sorted(self.capture_keys): + self._capture_keys_by_batch_size.setdefault(key[0], []).append(key) self.graphs: Dict[EncoderKeyType, torch.cuda.CUDAGraph] = {} self.graph_outputs: Dict[EncoderKeyType, Callable[[], @@ -825,10 +988,12 @@ def __init__(self, config: EncoderCUDAGraphRunnerConfig): self.shared_static_tensors_cpu: Dict[str, torch.Tensor] = {} if self.enabled: self._create_shared_static_tensors() - self.cuda_graph_meta_buffers = get_memory_buffers() + self.cuda_graph_meta_buffers = (Buffers() if self.is_encoder_decoder + else get_memory_buffers()) self._capture_allowed = False self.is_warmup_only = False + self._staging_retirement_event: Optional[torch.cuda.Event] = None # CUDA graph H2D memcpy nodes require pinned host sources. In CC mode # prefer_pinned() is false: pageable host buffers are preferred, so the @@ -837,8 +1002,9 @@ def __init__(self, config: EncoderCUDAGraphRunnerConfig): def _create_shared_static_tensors(self): """Allocates static tensors sized for the largest supported num_tokens.""" - max_total_tokens = min(self.max_supported_num_tokens, - self.config.max_num_tokens) + max_total_tokens = ( + self.config.max_num_tokens if self.is_encoder_decoder else min( + self.max_supported_num_tokens, self.config.max_num_tokens)) max_batch_size = self.max_supported_batch_size self.shared_static_tensors = { @@ -882,6 +1048,153 @@ def _round_up(value: int, supported: List[int]) -> int: return 0 return supported[idx] + @staticmethod + def build_capture_sequence_lengths(batch_size: int, num_tokens: int, + max_seq_len: int) -> Optional[List[int]]: + """Build a real sequence layout for a configured encoder bucket.""" + if (batch_size <= 0 or num_tokens < batch_size + or num_tokens > batch_size * max_seq_len): + return None + + if batch_size == 1: + return [num_tokens] + + if num_tokens >= max_seq_len + batch_size - 1: + remaining_tokens = num_tokens - max_seq_len + base, extra = divmod(remaining_tokens, batch_size - 1) + return ([max_seq_len] + [base + 1] * extra + [base] * + (batch_size - 1 - extra)) + + return [num_tokens - batch_size + 1] + [1] * (batch_size - 1) + + def _build_encoder_decoder_capture_layouts( + self) -> Dict[EncoderKeyType, List[int]]: + """Map each reachable graph key to one physical sequence-slot layout. + + Sequence layout is deliberately not part of ``EncoderKeyType``. Multiple + configured shapes may normalize to the same three-dimensional key, so + the first deterministic layout becomes that graph's fixed slot + capacities. Runtime sequences are assigned to those slots during replay. + """ + capture_layouts: Dict[EncoderKeyType, List[int]] = {} + for batch_size in self.supported_batch_sizes: + for num_tokens in self.supported_num_tokens: + for max_seq_len in self.supported_seq_lens: + sequence_lengths = self.build_capture_sequence_lengths( + batch_size, num_tokens, max_seq_len) + if sequence_lengths is None: + continue + + key, _, is_valid = self.get_graph_key( + {"seq_lens": sequence_lengths}) + if is_valid: + # Capture at most one graph/layout for a normalized key; + # alternative runtime layouts do not create more keys. + capture_layouts.setdefault(key, sequence_lengths) + + return capture_layouts + + def _get_dynamic_capture_key( + self, + sequence_lengths: List[int], + allow_batch_padding: bool, + ) -> Optional[EncoderKeyType]: + """Return the smallest key whose token bucket and slots fit the batch. + + Keys are ordered from smaller to larger buckets. For fixed-slot replay, + aggregate token and maximum-length checks are insufficient: every + runtime sequence (including batch-padding dummies) must also fit in a + distinct capture-time slot. An incompatible key is skipped in favor of + a larger existing key; no new layout-specific key is created. + """ + batch_size = len(sequence_lengths) + sum(sequence_lengths) + max_seq_len = max(sequence_lengths) if sequence_lengths else 0 + candidate_batch_sizes = (self.supported_batch_sizes + if allow_batch_padding else [batch_size]) + for padded_batch_size in candidate_batch_sizes: + if padded_batch_size < batch_size: + continue + + padded_sequence_lengths = (sequence_lengths + [1] * + (padded_batch_size - batch_size)) + required_num_tokens = sum(padded_sequence_lengths) + for key in self._capture_keys_by_batch_size.get( + padded_batch_size, []): + _, padded_num_tokens, padded_max_seq_len = key + if (padded_num_tokens < required_num_tokens + or padded_num_tokens > self.max_supported_num_tokens + or padded_max_seq_len < max_seq_len + or padded_max_seq_len not in self.supported_seq_lens + or padded_num_tokens + > padded_batch_size * padded_max_seq_len): + continue + if (self.use_fixed_sequence_slots + and self._get_sequence_slot_mapping( + key, padded_sequence_lengths) is None): + # The batch fits the aggregate bucket but not this key's + # individual slot capacities. Try the next captured key. + continue + return key + + return None + + def _get_sequence_slot_mapping( + self, + key: EncoderKeyType, + sequence_lengths: List[int], + ) -> Optional[List[int]]: + """Assign each runtime sequence to one compatible physical graph slot. + + The returned list maps source request index to capture slot index. It + changes only physical placement: source request order is retained + separately and restored after replay. + """ + capture_lengths = self._capture_sequence_lengths.get(key) + if (capture_lengths is None + or len(capture_lengths) != len(sequence_lengths)): + return None + + # Preserve physical order when every request already fits its + # corresponding slot, avoiding unnecessary scatter/gather permutation. + if all(sequence_length <= capture_length + for sequence_length, capture_length in zip( + sequence_lengths, capture_lengths)): + return list(range(len(sequence_lengths))) + + # Largest-to-largest matching is sufficient for one-to-one scalar + # capacities: if any sorted request exceeds its paired slot, no + # permutation can make the layout fit. + sequence_order = sorted(range(len(sequence_lengths)), + key=lambda index: + (-sequence_lengths[index], index)) + capture_order = sorted(range(len(capture_lengths)), + key=lambda index: + (-capture_lengths[index], index)) + source_to_slot = [0] * len(sequence_lengths) + for source_index, slot_index in zip(sequence_order, capture_order): + if sequence_lengths[source_index] > capture_lengths[slot_index]: + return None + source_to_slot[source_index] = slot_index + return source_to_slot + + def get_capture_warmup_sequence_lengths( + self, key: EncoderKeyType) -> Optional[List[int]]: + """Return the representative sequence layout for a capture key.""" + sequence_lengths = self._capture_sequence_lengths.get(key) + return list(sequence_lengths) if sequence_lengths is not None else None + + def _get_capture_sequence_offsets(self, key: EncoderKeyType) -> List[int]: + """Return cumulative fixed-slot offsets for a capture layout.""" + offsets = [0] + for sequence_length in self._capture_sequence_lengths[key]: + offsets.append(offsets[-1] + sequence_length) + if offsets[-1] != key[1]: + raise ValueError( + f"Encoder CUDA graph layout for key {key} contains " + f"{offsets[-1]} tokens.") + return offsets + def _get_valid_graph_key(self, batch_size: int, num_tokens: int, max_seq_len: int) -> EncoderKeyType: num_tokens_idx = bisect.bisect_left(self.supported_num_tokens, @@ -918,8 +1231,28 @@ def get_graph_key( batch_size = len(seq_lens) max_seq_len = max(seq_lens) if batch_size > 0 else 0 + if self.is_encoder_decoder: + if self.padding_enabled and self.capture_keys: + padded_key = self._get_dynamic_capture_key( + seq_lens, + allow_batch_padding=False, + ) + if padded_key is None: + return (batch_size, 0, 0), False, False + is_padding_performed = (padded_key[1] != num_tokens + or padded_key[2] != max_seq_len) + return padded_key, is_padding_performed, True + + max_seq_len_bucket = self._round_up(max_seq_len, + self.supported_seq_lens) + key: EncoderKeyType = (batch_size, num_tokens, max_seq_len_bucket) + is_valid = (num_tokens <= self.max_supported_num_tokens + and max_seq_len_bucket > 0) + return key, False, is_valid + key = self._get_valid_graph_key(batch_size, num_tokens, max_seq_len) - _, padded_num_tokens, padded_max_seq_len = key + padded_num_tokens = key[1] + padded_max_seq_len = key[2] is_padding_performed = (padded_num_tokens != num_tokens or padded_max_seq_len != max_seq_len) @@ -932,9 +1265,8 @@ def get_graph_key( def allow_capture(self): """Context manager that enables CUDA graph capture. - Capture is disabled by default. On-the-fly captures outside this - context are prevented — unseen keys fall back to eager instead of - incurring a multi-millisecond capture latency spike at runtime. + All encoder graphs are captured during explicit startup warmup through + this context. Unseen runtime keys fall back to eager execution. """ self._capture_allowed = True try: @@ -949,8 +1281,16 @@ def pad_batch(self, inputs: Dict[str, Any], yield inputs return - padded_batch_size = self._round_up(batch_size, - self.supported_batch_sizes) + if self.is_encoder_decoder and self.capture_keys: + seq_lens = inputs['seq_lens'] + padded_key = self._get_dynamic_capture_key( + seq_lens, + allow_batch_padding=True, + ) + padded_batch_size = padded_key[0] if padded_key is not None else 0 + else: + padded_batch_size = self._round_up(batch_size, + self.supported_batch_sizes) if padded_batch_size == 0 or padded_batch_size == batch_size: yield inputs return @@ -972,6 +1312,48 @@ def pad_batch(self, inputs: Dict[str, Any], yield padded_inputs + def prepare_encoder_decoder_inputs( + self, + inputs: Dict[str, Any], + key: EncoderKeyType, + source_sequence_lengths: List[int], + ) -> Dict[str, Any]: + """Arrange runtime sequence metadata in capture-time slot order.""" + if not self.is_encoder_decoder: + return inputs + + if not self.use_fixed_sequence_slots: + prepared_inputs = dict(inputs) + prepared_inputs[_ENCODER_SOURCE_SEQ_LENS] = list( + source_sequence_lengths) + return prepared_inputs + + sequence_lengths = inputs["seq_lens"] + if (sequence_lengths[:len(source_sequence_lengths)] + != source_sequence_lengths): + raise ValueError("Encoder source sequence lengths must be the " + "unpadded prefix of graph sequence lengths.") + + source_to_slot = self._get_sequence_slot_mapping(key, sequence_lengths) + if source_to_slot is None: + raise ValueError( + f"Encoder sequence lengths {sequence_lengths} are not " + f"compatible with CUDA graph key {key}.") + + # Attention metadata follows physical slot order, while the packed + # source tensors and the final returned output remain in request order. + slot_sequence_lengths = [0] * len(sequence_lengths) + for source_index, slot_index in enumerate(source_to_slot): + slot_sequence_lengths[slot_index] = sequence_lengths[source_index] + + prepared_inputs = dict(inputs) + prepared_inputs["seq_lens"] = slot_sequence_lengths + prepared_inputs[_ENCODER_SOURCE_SEQ_LENS] = list( + source_sequence_lengths) + prepared_inputs[_ENCODER_SOURCE_TO_SLOT] = source_to_slot[:len( + source_sequence_lengths)] + return prepared_inputs + def maybe_get_cuda_graph( self, inputs: Dict[str, Any], @@ -1010,16 +1392,21 @@ def maybe_get_cuda_graph( key, is_padding_performed, is_padding_successful = self.get_graph_key( inputs) - _, _, padded_max_seq_len = key + if self.is_encoder_decoder and key not in self.capture_keys: + return None, None + padded_max_seq_len = key[2] if (not self.padding_enabled and is_padding_performed) \ or not is_padding_successful: return None, None if key in self.graph_metadata: + # Every graph key aliases the same host staging buffers. Retire a + # prior graph's captured reads before the caller updates them. + self.retire_staging() return self.graph_metadata[key]["attn_metadata"], key - # New key not yet captured. Only create metadata if capture is - # allowed (warmup time); otherwise fall back to eager. + # New key not yet captured. Only create graph metadata during explicit + # startup warmup; unseen runtime keys fall back to eager execution. if not self._capture_allowed: return None, None @@ -1057,9 +1444,20 @@ def maybe_get_cuda_graph( # be pinned or pageable; only captured H2D copies require pinned memory. graph_attn_metadata.bind_encoder_cuda_graph_seq_lens( self.shared_static_tensors_cpu["seq_lens"], padded_batch_size) + if self.use_fixed_sequence_slots: + # CUDA graph replay keeps each request in its capture-time token + # slot. Explicit boundaries let attention combine those fixed + # offsets with the per-replay logical sequence lengths above. + capture_offsets = self._get_capture_sequence_offsets(key) + capture_offsets_cuda = torch.tensor(capture_offsets, + dtype=torch.int32, + device="cuda") + graph_attn_metadata.cu_q_seqlens = capture_offsets_cuda + graph_attn_metadata.cu_kv_seqlens = capture_offsets_cuda graph_attn_metadata.max_seq_len = self.config.max_seq_len graph_attn_metadata.request_ids = list(range(padded_batch_size)) + self.retire_staging() return graph_attn_metadata, key def _contains_nested_tensor(self, x: Any) -> bool: @@ -1074,6 +1472,137 @@ def _contains_nested_tensor(self, x: Any) -> bool: def needs_capture(self, key: EncoderKeyType) -> bool: return self._capture_allowed and key not in self.graphs + def _stage_inputs(self, key: EncoderKeyType, inputs: Dict[str, + Any]) -> None: + """Stage input and position IDs for capture or replay.""" + padded_num_tokens = key[1] + + # Captured H2D nodes read pinned host buffers. In CC mode, where H2D + # is not captured, stage directly into the graph-resident CUDA buffers. + static_tensors = self.shared_static_tensors_cpu if self._capture_h2d_copy else self.shared_static_tensors + + if self.is_encoder_decoder and _ENCODER_SOURCE_TO_SLOT in inputs: + self._stage_encoder_decoder_inputs(key, inputs, static_tensors) + return + + input_ids = inputs["input_ids"] + if isinstance(input_ids, list): + actual_tokens = len(input_ids) + static_tensors["input_ids"][:actual_tokens].copy_( + torch.tensor(input_ids, dtype=torch.int32)) + elif isinstance(input_ids, torch.Tensor): + actual_tokens = int(input_ids.shape[0]) + static_tensors["input_ids"][:actual_tokens].copy_(input_ids) + else: + raise TypeError(f"Unsupported input_ids type: {type(input_ids)}") + static_tensors["input_ids"][actual_tokens:padded_num_tokens].fill_(0) + + # Auto-generate packed position IDs without allocating one concatenated + # tensor, or copy caller-provided values into the stable staging buffer. + staged_position_ids = static_tensors["position_ids"][0] + position_ids = inputs.get("position_ids") + if position_ids is None: + offset = 0 + for seq_len in inputs["seq_lens"]: + staged_position_ids[offset:offset + seq_len].copy_( + self._arange_max[:seq_len]) + offset += seq_len + else: + if isinstance(position_ids, list): + staged_position_ids[:actual_tokens].copy_( + torch.tensor(position_ids, dtype=torch.int32)) + elif isinstance(position_ids, torch.Tensor): + staged_position_ids[:actual_tokens].copy_( + position_ids.flatten()) + else: + raise TypeError( + f"Unsupported position_ids type: {type(position_ids)}") + offset = actual_tokens + + staged_position_ids[offset:padded_num_tokens].fill_(0) + + def _stage_encoder_decoder_inputs( + self, + key: EncoderKeyType, + inputs: Dict[str, Any], + static_tensors: Dict[str, torch.Tensor], + ) -> None: + """Scatter packed request inputs into fixed capture-time slots.""" + source_sequence_lengths = inputs[_ENCODER_SOURCE_SEQ_LENS] + source_to_slot = inputs[_ENCODER_SOURCE_TO_SLOT] + + input_ids = inputs["input_ids"] + if isinstance(input_ids, list): + source_input_ids = torch.tensor(input_ids, dtype=torch.int32) + elif isinstance(input_ids, torch.Tensor): + source_input_ids = input_ids + else: + raise TypeError(f"Unsupported input_ids type: {type(input_ids)}") + + actual_num_tokens = sum(source_sequence_lengths) + if int(source_input_ids.shape[0]) != actual_num_tokens: + raise ValueError( + "Packed encoder input IDs must match source sequence lengths.") + + position_ids = inputs.get("position_ids") + if isinstance(position_ids, list): + source_position_ids = torch.tensor(position_ids, dtype=torch.int32) + elif isinstance(position_ids, torch.Tensor): + source_position_ids = position_ids.flatten() + elif position_ids is None: + source_position_ids = None + else: + raise TypeError( + f"Unsupported position_ids type: {type(position_ids)}") + if (source_position_ids is not None + and int(source_position_ids.shape[0]) != actual_num_tokens): + raise ValueError("Packed encoder position IDs must match source " + "sequence lengths.") + + static_tensors["input_ids"][:key[1]].zero_() + static_tensors["position_ids"][:, :key[1]].zero_() + + capture_offsets = self._get_capture_sequence_offsets(key) + + source_offset = 0 + for source_index, sequence_length in enumerate(source_sequence_lengths): + slot_index = source_to_slot[source_index] + destination_offset = capture_offsets[slot_index] + source_slice = slice(source_offset, source_offset + sequence_length) + destination_slice = slice(destination_offset, + destination_offset + sequence_length) + static_tensors["input_ids"][destination_slice].copy_( + source_input_ids[source_slice]) + if source_position_ids is None: + static_tensors["position_ids"][0, destination_slice].copy_( + self._arange_max[:sequence_length]) + else: + static_tensors["position_ids"][0, destination_slice].copy_( + source_position_ids[source_slice]) + source_offset += sequence_length + + def restore_encoder_decoder_output( + self, + key: EncoderKeyType, + output: torch.Tensor, + inputs: Dict[str, Any], + ) -> torch.Tensor: + """Compact fixed-slot graph output back into request order.""" + source_sequence_lengths = inputs[_ENCODER_SOURCE_SEQ_LENS] + if _ENCODER_SOURCE_TO_SLOT not in inputs: + return output[:sum(source_sequence_lengths)].clone() + + source_to_slot = inputs[_ENCODER_SOURCE_TO_SLOT] + + capture_offsets = self._get_capture_sequence_offsets(key) + + output_slices = [] + for source_index, sequence_length in enumerate(source_sequence_lengths): + source_offset = capture_offsets[source_to_slot[source_index]] + output_slices.append(output[source_offset:source_offset + + sequence_length]) + return torch.cat(output_slices, dim=0) + def capture( self, key: EncoderKeyType, @@ -1081,7 +1610,7 @@ def capture( inputs: Dict[str, Any], ) -> Any: """Warm up and/or capture the forward pass for a graph key.""" - _, padded_num_tokens, _ = key + padded_num_tokens = key[1] sliced_static_tensors = { "input_ids": @@ -1104,6 +1633,20 @@ def capture( self.graph_metadata[key] = {"attn_metadata": attn_md} + # Warmup must see the same runtime data as capture. In particular, + # graph metadata initializes _seq_lens_cuda to ones, while + # prepare_encoder_cuda_graph_replay updates its stable host buffer. + # Populate every device input before warmup so packed-token counts and + # sequence boundaries are consistent. + self._stage_inputs(key, inputs) + if self._capture_h2d_copy: + capture_inputs["input_ids"].copy_( + sliced_static_tensors_cpu["input_ids"], non_blocking=True) + capture_inputs["position_ids"].copy_( + sliced_static_tensors_cpu["position_ids"], non_blocking=True) + attn_md._seq_lens_cuda.copy_(attn_md._seq_lens, non_blocking=True) + torch.cuda.current_stream().synchronize() + output = None with with_multi_stream(True), piecewise_cuda_graph(False): # Warmup runs required by CUDA graph semantics. See @@ -1117,7 +1660,9 @@ def capture( return output graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph, pool=self.memory_pool): + with torch.cuda.graph(graph, + pool=self.memory_pool, + capture_error_mode="thread_local"): if self._capture_h2d_copy: # H2D copies for captured inside the graph: at replay # time it re-issues from the pinned static buffer without @@ -1142,65 +1687,32 @@ def capture( self.memory_pool = graph.pool() return graph_output + def retire_staging(self) -> None: + """Wait until a prior replay no longer reads shared staging buffers.""" + if self._staging_retirement_event is not None: + self._staging_retirement_event.synchronize() + self._staging_retirement_event = None + def replay( self, key: EncoderKeyType, inputs: Dict[str, Any], ) -> Any: """Replay a captured graph with current inputs.""" + self.retire_staging() + stored_meta = self.graph_metadata[key] assert inputs["attn_metadata"] is stored_meta["attn_metadata"] - _, padded_num_tokens, _ = key - - # According to prefer_pinned(), CC forces most transfers to be synchronous. - # So we don't put non_blocking=True here. - static_tensors = self.shared_static_tensors_cpu if self._capture_h2d_copy else self.shared_static_tensors - - # input_ids: convert (if list) and write into pinned active region in - # one allocation + one memcpy. Padding region is zero-filled below. - input_ids = inputs["input_ids"] - if isinstance(input_ids, list): - actual_tokens = len(input_ids) - static_tensors["input_ids"][:actual_tokens].copy_( - torch.tensor(input_ids, dtype=torch.int32)) - elif isinstance(input_ids, torch.Tensor): - actual_tokens = int(input_ids.shape[0]) - static_tensors["input_ids"][:actual_tokens].copy_(input_ids) - else: - raise TypeError(f"Unsupported input_ids type: {type(input_ids)}") - static_tensors["input_ids"][actual_tokens:padded_num_tokens].fill_(0) - - # position_ids: pinned buffer is shape [1, max_total_tokens]; use the - # 1-D row view. Auto-generate via the cached arange (zero allocations, - # N small memcpys) or copy user-provided values. - pinned_pos = static_tensors["position_ids"][0] - position_ids = inputs.get("position_ids") - if position_ids is None: - # Pad entries (seq_len=1) get arange[:1] = [0], the correct - # position for a 1-token dummy request. - offset = 0 - for s in inputs["seq_lens"]: - pinned_pos[offset:offset + s].copy_(self._arange_max[:s]) - offset += s - else: - if isinstance(position_ids, list): - pinned_pos[:actual_tokens].copy_( - torch.tensor(position_ids, dtype=torch.int32)) - elif isinstance(position_ids, torch.Tensor): - pinned_pos[:actual_tokens].copy_(position_ids.flatten()) - else: - raise TypeError( - f"Unsupported position_ids type: {type(position_ids)}") - offset = actual_tokens - - pinned_pos[offset:padded_num_tokens].fill_(0) + self._stage_inputs(key, inputs) if not self._capture_h2d_copy: stored_meta["attn_metadata"]._seq_lens_cuda.copy_( stored_meta["attn_metadata"]._seq_lens, non_blocking=True) self.graphs[key].replay() + self._staging_retirement_event = torch.cuda.Event() + self._staging_retirement_event.record(torch.cuda.current_stream()) return self.graph_outputs[key] diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index e71c09e2ca6a..c0591b35516c 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -11,7 +11,7 @@ import weakref from abc import ABC, abstractmethod from contextlib import contextmanager -from typing import Any, Callable, Dict, List, Optional, Tuple, Union +from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple, Union import torch import torch._dynamo.config @@ -21,6 +21,8 @@ from tensorrt_llm._utils import (is_trace_enabled, maybe_pin_memory, nvtx_range, prefer_pinned, release_gc, torch_dtype_to_str, trace_func) +from tensorrt_llm.bindings.internal import \ + batch_manager as batch_manager_bindings from tensorrt_llm.bindings.internal.runtime import TaskLayerModuleConfig from tensorrt_llm.inputs.multimodal import (MultimodalParams, MultimodalRuntimeData, @@ -487,13 +489,21 @@ def __init__( self.cuda_graph_config = self.llm_args.cuda_graph_config self._is_encode_only = (self.llm_args.encode_only and not self.llm_args.mm_encoder_only) + if (self._is_encode_only + and isinstance(self.cuda_graph_config, EncodeCudaGraphConfig)): + self.encoder_cuda_graph_config = self.cuda_graph_config + else: + self.encoder_cuda_graph_config = ( + self.llm_args.encoder_cuda_graph_config) if (isinstance(self.cuda_graph_config, EncodeCudaGraphConfig) and self._is_encoder_decoder_model()): logger.warning( "EncodeCudaGraphConfig is not supported for encoder-decoder " - "models. Use DecodeCudaGraphConfig or CudaGraphConfig for " - "decoder CUDA graphs. CUDA graphs will be disabled.") + "models through cuda_graph_config. Use DecodeCudaGraphConfig " + "for cuda_graph_config and configure encoder graphs through " + "encoder_cuda_graph_config. Decoder CUDA graphs will be " + "disabled.") self.cuda_graph_config = None if (self.cuda_graph_config is not None and self.dtype == torch.float32 @@ -505,10 +515,10 @@ def __init__( # overruns the allocation (surfaces as cublas EXECUTION_FAILED). # Keep eager until the upstream size query is fixed. logger.warning( - "CUDA graphs are not supported for float32 encoder-decoder " - "models. CUDA graphs will be disabled; use a half-precision " - "checkpoint or model_kwargs={'torch_dtype': ...} to enable " - "them.") + "Decoder CUDA graphs are not supported for float32 " + "encoder-decoder models. Decoder CUDA graphs will be disabled; " + "use a half-precision checkpoint or " + "model_kwargs={'torch_dtype': ...} to enable them.") self.cuda_graph_config = None cuda_graph_batch_sizes = self.cuda_graph_config.batch_sizes if self.cuda_graph_config else CudaGraphConfig.model_fields[ @@ -516,29 +526,78 @@ def __init__( cuda_graph_padding_enabled = self.cuda_graph_config.enable_padding if self.cuda_graph_config else CudaGraphConfig.model_fields[ 'enable_padding'].default - # Encode-only CUDA graph detection. Decode configs do not define these - # encoder-specific bucket fields. - cuda_graph_num_tokens = [] - cuda_graph_seq_lens = [] - if isinstance(self.cuda_graph_config, EncodeCudaGraphConfig): - cuda_graph_num_tokens = self.cuda_graph_config.num_tokens or [] - cuda_graph_seq_lens = self.cuda_graph_config.seq_lens or [] - - if (self._is_encode_only and self.cuda_graph_config is not None - and (not cuda_graph_num_tokens or not cuda_graph_seq_lens)): + # CUDA graph detection for encoder-decoder models and encoder-only models. + # Decode configs do not define these encoder-specific bucket fields. + encoder_cuda_graph_batch_sizes = ( + self.encoder_cuda_graph_config.batch_sizes + if self.encoder_cuda_graph_config is not None else []) + encoder_cuda_graph_num_tokens = ( + self.encoder_cuda_graph_config.num_tokens + if self.encoder_cuda_graph_config is not None else []) + encoder_cuda_graph_seq_lens = (self.encoder_cuda_graph_config.seq_lens + if self.encoder_cuda_graph_config + is not None else []) + encoder_cuda_graph_padding_enabled = ( + self.encoder_cuda_graph_config.enable_padding + if self.encoder_cuda_graph_config is not None else False) + + if (self.encoder_cuda_graph_config is not None + and (not encoder_cuda_graph_num_tokens + or not encoder_cuda_graph_seq_lens)): missing = [] - if not cuda_graph_num_tokens: + if not encoder_cuda_graph_num_tokens: missing.append("num_tokens/max_num_token") - if not cuda_graph_seq_lens: + if not encoder_cuda_graph_seq_lens: missing.append("seq_lens/max_seq_len") logger.warning( - f"encode_only=True with a CudaGraphConfig, but " - f"{' and '.join(missing)} not set. Encoder CUDA graphs " - f"require both. Encoder CUDA graphs will be disabled. " + f"Encoder CUDA graph configuration has " + f"{' and '.join(missing)} unset. Encoder CUDA graphs require " + f"both dimensions and will be disabled. " f"To enable them, specify e.g. " f"EncodeCudaGraphConfig(max_batch_size=64, num_tokens=[128, 256, " f"512], max_seq_len=128, enable_padding=True).") + self._cuda_graph_padding_enabled = cuda_graph_padding_enabled + + self._cuda_graph_batch_sizes = _filter_cuda_graph_batch_sizes( + cuda_graph_batch_sizes, self.batch_size, self.max_num_tokens, + self.original_max_total_draft_tokens, + self._cuda_graph_padding_enabled) if cuda_graph_batch_sizes else [] + + self._max_cuda_graph_batch_size = (self._cuda_graph_batch_sizes[-1] if + self._cuda_graph_batch_sizes else 0) + + self._encoder_cuda_graph_padding_enabled = ( + encoder_cuda_graph_padding_enabled) + self._encoder_cuda_graph_batch_sizes = (_filter_cuda_graph_batch_sizes( + encoder_cuda_graph_batch_sizes, self.encoder_batch_size, + self.encoder_max_num_tokens, 0, + self._encoder_cuda_graph_padding_enabled) if + encoder_cuda_graph_batch_sizes + else []) + + # Encoder CUDA graph bucket lists + self._cuda_graph_num_tokens = (_filter_cuda_graph_num_tokens( + encoder_cuda_graph_num_tokens, self.encoder_max_num_tokens, + self._encoder_cuda_graph_padding_enabled) + if encoder_cuda_graph_num_tokens else []) + + self._max_cuda_graph_num_tokens = (self._cuda_graph_num_tokens[-1] if + self._cuda_graph_num_tokens else 0) + self._cuda_graph_seq_lens = (_filter_cuda_graph_seq_lens( + encoder_cuda_graph_seq_lens, self.max_seq_len, + self._encoder_cuda_graph_padding_enabled) + if encoder_cuda_graph_seq_lens else []) + + self._max_cuda_graph_seq_len = (self._cuda_graph_seq_lens[-1] + if self._cuda_graph_seq_lens else 0) + + use_encoder_cuda_graph = ((self._is_encoder_decoder_model() + or self._is_encode_only) + and self.encoder_cuda_graph_config is not None + and bool(self._cuda_graph_num_tokens) + and bool(self._cuda_graph_seq_lens)) + self.torch_compile_config = self.llm_args.torch_compile_config torch_compile_enabled = bool(self.torch_compile_config is not None) torch_compile_fullgraph = self.torch_compile_config.enable_fullgraph if self.torch_compile_config is not None else TorchCompileConfig.model_fields[ @@ -697,36 +756,13 @@ def __init__( self.iter_states = {} self._cuda_graph_mem_pool = self._torch_compile_backend._graph_pool_handle if self._torch_compile_enabled else None - self._cuda_graph_padding_enabled = cuda_graph_padding_enabled - - self._cuda_graph_batch_sizes = _filter_cuda_graph_batch_sizes( - cuda_graph_batch_sizes, self.batch_size, self.max_num_tokens, - self.original_max_total_draft_tokens, - self._cuda_graph_padding_enabled) if cuda_graph_batch_sizes else [] - - self._max_cuda_graph_batch_size = (self._cuda_graph_batch_sizes[-1] if - self._cuda_graph_batch_sizes else 0) - - # Encoder CUDA graph bucket lists - self._cuda_graph_num_tokens = _filter_cuda_graph_num_tokens( - cuda_graph_num_tokens, self.max_num_tokens, - self._cuda_graph_padding_enabled) if cuda_graph_num_tokens else [] - - self._max_cuda_graph_num_tokens = (self._cuda_graph_num_tokens[-1] if - self._cuda_graph_num_tokens else 0) - self._cuda_graph_seq_lens = _filter_cuda_graph_seq_lens( - cuda_graph_seq_lens, self.max_seq_len, - self._cuda_graph_padding_enabled) if cuda_graph_seq_lens else [] - - self._max_cuda_graph_seq_len = (self._cuda_graph_seq_lens[-1] - if self._cuda_graph_seq_lens else 0) - self._dynamic_draft_len_mapping = self._compute_dynamic_draft_len_mapping( ) self.previous_batch_indices_cuda = torch.empty((self.max_num_tokens, ), dtype=torch.int, device='cuda') + self._encoder_decoder_staged_request_ids: Optional[List[int]] = None self.input_ids_cuda = torch.empty((self.max_num_tokens, ), dtype=torch.int, device='cuda') @@ -795,7 +831,45 @@ def __init__( self.lora_model_config: Optional[LoraModelConfig] = None self._trtllm_gen_jit_warmup = False - # Create config and runner + # Create the encoder runner first. For encoder-decoder models it derives + # every reachable startup capture key through get_graph_key(). + encoder_graph_batch_sizes = self._encoder_cuda_graph_batch_sizes + encoder_graph_max_batch_size = (encoder_graph_batch_sizes[-1] + if encoder_graph_batch_sizes else 0) + encoder_graph_max_num_tokens = self._max_cuda_graph_num_tokens + encoder_cuda_graph_runner_config = EncoderCUDAGraphRunnerConfig( + use_cuda_graph=use_encoder_cuda_graph, + cuda_graph_padding_enabled=( + self._encoder_cuda_graph_padding_enabled), + cuda_graph_batch_sizes=encoder_graph_batch_sizes, + cuda_graph_num_tokens=self._cuda_graph_num_tokens, + cuda_graph_seq_lens=self._cuda_graph_seq_lens, + max_cuda_graph_batch_size=encoder_graph_max_batch_size, + max_cuda_graph_num_tokens=encoder_graph_max_num_tokens, + max_num_tokens=self.encoder_max_num_tokens, + max_seq_len=self.max_seq_len, + cuda_graph_mem_pool=self._cuda_graph_mem_pool, + is_encoder_decoder=self._is_encoder_decoder_model(), + use_fixed_sequence_slots=(self._is_encoder_decoder_model() + and hasattr( + pretrained_config, + "relative_attention_num_buckets")), + ) + self.encoder_cuda_graph_runner = EncoderCUDAGraphRunner( + encoder_cuda_graph_runner_config) + + # Once encoder CUDA graphs are usable, enable mixed decoder graphs by + # default unless the user explicitly opts out. + encoder_decoder_cuda_graph_enabled = ( + self.encoder_cuda_graph_runner.enabled + and self.encoder_cuda_graph_runner.is_encoder_decoder + and bool(self.encoder_cuda_graph_runner.capture_keys)) + enable_encoder_decoder_mixed_cuda_graph = ( + encoder_decoder_cuda_graph_enabled + and self.cuda_graph_config is not None + and self.llm_args.enable_encoder_decoder_mixed_cuda_graph) + + # Create decoder CUDA graph config and runner. cuda_graph_runner_config = CUDAGraphRunnerConfig( use_cuda_graph=(not self._is_encode_only and self.cuda_graph_config is not None), @@ -819,28 +893,11 @@ def __init__( dist=self.dist, kv_cache_manager_key=self.kv_cache_manager_key, sparse_attention_config=self.sparse_attention_config, + enable_encoder_decoder_mixed_cuda_graph=( + enable_encoder_decoder_mixed_cuda_graph), ) self.cuda_graph_runner = CUDAGraphRunner(cuda_graph_runner_config) - # Create Encoder CUDA graph config and runner. - encoder_cuda_graph_runner_config = EncoderCUDAGraphRunnerConfig( - use_cuda_graph=(self._is_encode_only - and self.cuda_graph_config is not None - and bool(self._cuda_graph_num_tokens) - and bool(self._cuda_graph_seq_lens)), - cuda_graph_padding_enabled=self._cuda_graph_padding_enabled, - cuda_graph_batch_sizes=self._cuda_graph_batch_sizes, - cuda_graph_num_tokens=self._cuda_graph_num_tokens, - cuda_graph_seq_lens=self._cuda_graph_seq_lens, - max_cuda_graph_batch_size=self._max_cuda_graph_batch_size, - max_cuda_graph_num_tokens=self._max_cuda_graph_num_tokens, - max_num_tokens=self.max_num_tokens, - max_seq_len=self.max_seq_len, - cuda_graph_mem_pool=self._cuda_graph_mem_pool, - ) - self.encoder_cuda_graph_runner = EncoderCUDAGraphRunner( - encoder_cuda_graph_runner_config) - # Initialize CUDA Graph LoRA manager if LoRA is enabled self.cuda_graph_lora_manager: Optional[CudaGraphLoraManager] = None @@ -1861,11 +1918,81 @@ def _run_cuda_graph_warmup(self, resource_manager: ResourceManager): return self._capture_generation_cuda_graphs(resource_manager) + self._capture_mixed_encoder_decoder_cuda_graphs(resource_manager) # Piecewise graphs have separate capture machinery and do not use the # whole-model attention workspace. Capture them only on the second pass. if not self.cuda_graph_runner.is_warmup_only: self._capture_piecewise_cuda_graphs(resource_manager) + @torch.inference_mode() + @with_warmup_flag + def _warmup_encoder_cuda_graphs_enc_dec( + self, resource_manager: ResourceManager) -> None: + """Capture encoder-decoder encoder graphs on their runtime host thread.""" + runner = self.encoder_cuda_graph_runner + if not runner.is_encoder_decoder: + return + + capture = functools.partial( + self._capture_encoder_cuda_graphs_enc_dec, + resource_manager, + ) + self._warmup_and_capture_encoder_cuda_graphs(capture) + + def _warmup_and_capture_encoder_cuda_graphs( + self, capture: Callable[[], None]) -> None: + """Warm up every encoder graph shape, then capture those shapes.""" + runner = self.encoder_cuda_graph_runner + if not runner.enabled: + return + + with runner.allow_capture(): + runner.is_warmup_only = True + try: + capture() + finally: + runner.is_warmup_only = False + capture() + + def _capture_encoder_cuda_graphs_enc_dec( + self, resource_manager: ResourceManager) -> None: + """Warm up or capture encoder graphs used by encoder-decoder models.""" + runner = self.encoder_cuda_graph_runner + if not runner.enabled or not runner.is_encoder_decoder: + return + + operation = "warmup" if runner.is_warmup_only else "capture" + num_processed = 0 + logger.info( + f"Running encoder-decoder encoder CUDA graph {operation} ...") + for key in sorted(runner.capture_keys, reverse=True): + sequence_lengths = runner.get_capture_warmup_sequence_lengths(key) + if sequence_lengths is None: + continue + + encoder_input_ids = [0] * sum(sequence_lengths) + encoder_position_ids = [] + for sequence_length in sequence_lengths: + encoder_position_ids.extend( + self._apply_position_id_offset(list( + range(sequence_length)))) + inputs = self._prepare_encoder_decoder_encoder_inputs( + encoder_input_ids=encoder_input_ids, + encoder_position_ids=encoder_position_ids, + sequence_lengths=sequence_lengths, + request_ids=list(range(len(sequence_lengths))), + resource_manager=resource_manager, + ) + + logger.info("Encoder-decoder encoder CUDA graph " + f"{operation}: key={key}") + self._encoder_forward_enc_dec(inputs) + torch.cuda.synchronize() + num_processed += 1 + + logger.info("Completed encoder-decoder encoder CUDA graph " + f"{operation} for {num_processed} graph shape(s).") + def _capture_generation_cuda_graphs(self, resource_manager: ResourceManager): """Warm up or capture pure-generation CUDA graph shapes.""" @@ -2068,6 +2195,123 @@ def _run_capture_pass(force_non_greedy: bool, label: str) -> None: if self.spec_metadata is not None: self.spec_metadata.is_all_greedy_sample = True + def _capture_mixed_encoder_decoder_cuda_graphs( + self, resource_manager: ResourceManager) -> None: + """Warm and capture reachable mixed encoder-decoder graph shapes. + + The first global CUDA-graph pass warms every shape so shared attention + workspace reaches its final size. The second pass captures the same + shapes. Runtime capture is deliberately disabled because graph capture + executes KV-cache writes and must never run against live requests. + """ + runner = self.cuda_graph_runner + if not runner.enable_encoder_decoder_mixed_cuda_graph: + return + + max_encoder_output_len = self._get_max_encoder_output_len( + resource_manager) + context_shapes = {(batch_size, total_tokens) + for batch_size, total_tokens, _ in + self.encoder_cuda_graph_runner.capture_keys} + if not context_shapes: + logger.warning("Skipping mixed encoder-decoder CUDA graph capture: " + "no encoder CUDA graph shapes were captured.") + return + + max_encoder_batch_size = max(batch_size + for batch_size, _ in context_shapes) + max_batch_token_counts = { + total_tokens + for batch_size, total_tokens in context_shapes + if batch_size == max_encoder_batch_size + } + paired_context_count = 2 * max_encoder_batch_size + if runner.max_supported_batch_size > paired_context_count: + paired_token_counts = { + first + second + for first in max_batch_token_counts + for second in max_batch_token_counts + } + context_shapes.update((paired_context_count, token_count) + for token_count in paired_token_counts) + + operation = ("warmup" if runner.is_warmup_only else "capture") + hidden_size = self._get_enc_dec_hidden_size() + max_num_encoder_tokens = max( + (total_encoder_tokens + for num_contexts, total_encoder_tokens in context_shapes + if total_encoder_tokens <= num_contexts * max_encoder_output_len + and any(batch_size > num_contexts + for batch_size in runner.supported_batch_sizes)), + default=0) + if max_num_encoder_tokens == 0: + return + model_config = self.model.model_config.pretrained_config + # BART/mBART prepend a forced BOS token after decoder_start; T5 uses + # decoder_start alone. Match the LLM API's decoder-prefix construction. + mixed_context_query_len = (2 if getattr( + model_config, "model_type", None) in ("bart", "mbart") else 1) + for num_contexts, total_encoder_tokens in sorted( + context_shapes, key=lambda shape: shape[1], reverse=True): + if total_encoder_tokens > num_contexts * max_encoder_output_len: + continue + base_encoder_len, remainder = divmod(total_encoder_tokens, + num_contexts) + encoder_output_lens = ([base_encoder_len + 1] * remainder + + [base_encoder_len] * + (num_contexts - remainder)) + if not encoder_output_lens or encoder_output_lens[-1] <= 0: + continue + + for batch_size in runner.supported_batch_sizes: + if batch_size <= num_contexts: + continue + warmup_request = self._create_cuda_graph_warmup_request( + resource_manager, + batch_size, + draft_len=0, + mixed_context_encoder_output_lens=encoder_output_lens, + mixed_context_query_len=mixed_context_query_len) + with self._release_batch_context(warmup_request, + resource_manager) as batch: + if batch is None: + logger.warning( + "Skipping mixed encoder-decoder CUDA graph " + f"{operation}: not enough KV cache space for " + f"batch size={batch_size}.") + continue + + context_requests = batch.context_requests + for request, encoder_output_len in zip( + context_requests, encoder_output_lens): + request.state = LlmRequestState.CONTEXT_INIT + request.context_current_position = 0 + request.context_chunk_size = mixed_context_query_len + request.cached_tokens = 0 + request.py_batch_idx = None + request.py_encoder_output = torch.ones( + (encoder_output_len, hidden_size), + device="cuda", + dtype=self.dtype, + ) + request.py_skip_cross_kv_projection = False + + runner._get_static_encoder_hidden_states( + context_requests[0].py_encoder_output, + max_num_encoder_tokens, + allow_allocate=True, + ) + logger.info("Run mixed encoder-decoder CUDA graph " + f"{operation} for batch size={batch_size}, " + f"context requests={num_contexts}, " + f"packed encoder tokens={total_encoder_tokens}") + self.enable_spec_decode = False + self.runtime_draft_len = 0 + self.forward(batch, + new_tensors_device=None, + resource_manager=resource_manager) + torch.cuda.synchronize() + def _capture_piecewise_cuda_graphs(self, resource_manager: ResourceManager): """Captures piecewise CUDA graphs for context/prefill steps via torch.compile.""" if not (self._torch_compile_piecewise_cuda_graph @@ -2289,11 +2533,14 @@ def _create_warmup_request( return result def _create_cuda_graph_warmup_request( - self, - resource_manager: ResourceManager, - batch_size: int, - draft_len: int, - max_seq_len: int = None) -> Optional[ScheduledRequests]: + self, + resource_manager: ResourceManager, + batch_size: int, + draft_len: int, + max_seq_len: int = None, + mixed_context_encoder_output_lens: Optional[Sequence[int]] = None, + mixed_context_query_len: int = ENC_DEC_CUDA_GRAPH_DUMMY_TOKEN_NUM, + ) -> Optional[ScheduledRequests]: """Creates a dummy ScheduledRequests tailored for CUDA graph capture.""" kv_cache_manager = resource_manager.get_resource_manager( self.kv_cache_manager_key) @@ -2316,26 +2563,73 @@ def _create_cuda_graph_warmup_request( max_encoder_output_len = ( self._get_max_encoder_output_len(resource_manager) if is_enc_dec else None) + num_mixed_contexts = len(mixed_context_encoder_output_lens or + ()) if is_enc_dec else 0 + if num_mixed_contexts >= batch_size: + return None - # Add (batch_size - 1) dummy requests with the minimal seq_len. - token_nums = ([ENC_DEC_CUDA_GRAPH_DUMMY_TOKEN_NUM] * - (batch_size - 1)) if is_enc_dec else None - encoder_output_lens = ([max_encoder_output_len] * - (batch_size - 1)) if is_enc_dec else None - requests = kv_cache_manager.add_dummy_requests( - list(range(batch_size - 1)), - token_nums=token_nums, - is_gen=True, - max_num_draft_tokens=runtime_draft_token_buffer_width, - kv_reserve_draft_tokens=self.max_draft_loop_tokens, - use_mrope=self.use_mrope, - max_beam_width=self.max_beam_width, - encoder_output_lens=encoder_output_lens, - num_extra_decoding_steps=num_extra_decoding_steps, - draft_kv_cache_manager=draft_kv_cache_manager) + # Add (batch_size - 1) dummy requests with the minimal sequence + # length. Mixed capture must create its context rows as real context + # requests; converting generation dummies afterward leaves their + # native prompt/context bookkeeping at one token. + if mixed_context_encoder_output_lens: + context_request_ids = list(range(num_mixed_contexts)) + context_requests = kv_cache_manager.add_dummy_requests( + context_request_ids, + token_nums=[mixed_context_query_len] * num_mixed_contexts, + is_gen=False, + max_num_draft_tokens=runtime_draft_token_buffer_width, + kv_reserve_draft_tokens=self.max_draft_loop_tokens, + use_mrope=self.use_mrope, + max_beam_width=self.max_beam_width, + encoder_output_lens=list(mixed_context_encoder_output_lens), + num_extra_decoding_steps=num_extra_decoding_steps, + draft_kv_cache_manager=draft_kv_cache_manager) + if context_requests is None: + return None - if requests is None: - return None + generation_request_ids = list( + range(num_mixed_contexts, batch_size - 1)) + generation_requests = [] + if generation_request_ids: + generation_requests = kv_cache_manager.add_dummy_requests( + generation_request_ids, + token_nums=[ENC_DEC_CUDA_GRAPH_DUMMY_TOKEN_NUM] * + len(generation_request_ids), + is_gen=True, + max_num_draft_tokens=runtime_draft_token_buffer_width, + kv_reserve_draft_tokens=self.max_draft_loop_tokens, + use_mrope=self.use_mrope, + max_beam_width=self.max_beam_width, + encoder_output_lens=[max_encoder_output_len] * + len(generation_request_ids), + num_extra_decoding_steps=num_extra_decoding_steps, + draft_kv_cache_manager=draft_kv_cache_manager) + if generation_requests is None: + for request in context_requests: + kv_cache_manager.free_resources(request) + if draft_kv_cache_manager is not None: + draft_kv_cache_manager.free_resources(request) + return None + requests = context_requests + generation_requests + else: + token_nums = ([ENC_DEC_CUDA_GRAPH_DUMMY_TOKEN_NUM] * + (batch_size - 1)) if is_enc_dec else None + encoder_output_lens = ([max_encoder_output_len] * + (batch_size - 1)) if is_enc_dec else None + requests = kv_cache_manager.add_dummy_requests( + list(range(batch_size - 1)), + token_nums=token_nums, + is_gen=True, + max_num_draft_tokens=runtime_draft_token_buffer_width, + kv_reserve_draft_tokens=self.max_draft_loop_tokens, + use_mrope=self.use_mrope, + max_beam_width=self.max_beam_width, + encoder_output_lens=encoder_output_lens, + num_extra_decoding_steps=num_extra_decoding_steps, + draft_kv_cache_manager=draft_kv_cache_manager) + if requests is None: + return None def free_warmup_requests() -> None: for r in requests: @@ -2411,14 +2705,26 @@ def free_warmup_requests() -> None: else: max_seq_len_request = max_seq_len_request[0] - # Insert the longest request first to simulate padding for the CUDA graph. - requests.insert(0, max_seq_len_request) - result.generation_requests = requests + if mixed_context_encoder_output_lens: + requests.append(max_seq_len_request) + for request in requests[:num_mixed_contexts]: + request.state = LlmRequestState.CONTEXT_INIT + request.context_current_position = 0 + request.context_chunk_size = mixed_context_query_len + request.cached_tokens = 0 + request.py_batch_idx = None + result.context_requests_last_chunk = requests[:num_mixed_contexts] + result.generation_requests = requests[num_mixed_contexts:] + else: + # Insert the longest request first to simulate padding for the CUDA + # graph. + requests.insert(0, max_seq_len_request) + result.generation_requests = requests if spec_resource_manager is not None: spec_resource_manager.add_dummy_requests( request_ids=list(range(batch_size))) if self._is_encoder_decoder_model(): - if not self._add_cross_dummy_requests(result.generation_requests, + if not self._add_cross_dummy_requests(result.all_requests(), resource_manager): return None return result @@ -3174,6 +3480,10 @@ def _prepare_enc_dec_cross_attn_inputs( encoder_num_cached_tokens_per_seq: List[int], attn_metadata: AttentionMetadata, resource_manager: Optional[ResourceManager], + encoder_kv_lens: Optional[torch.Tensor] = None, + context_encoder_kv_tokens: int = 0, + generation_encoder_kv_tokens: int = 0, + max_encoder_kv_len: int = 0, ) -> Dict[str, Any]: if not encoder_seq_lens: return {} @@ -3214,6 +3524,19 @@ def _prepare_enc_dec_cross_attn_inputs( packed_encoder_hidden_states = None skip_cross_kv_projection = True + def prepare_cross_metadata( + cross_attn_metadata: AttentionMetadata) -> None: + if encoder_kv_lens is None: + cross_attn_metadata.prepare() + return + assert isinstance(cross_attn_metadata, TrtllmAttentionMetadata) + cross_attn_metadata.prepare_encoder_decoder( + prompt_lens=attn_metadata.prompt_lens, + kv_lens=encoder_kv_lens, + context_kv_tokens=context_encoder_kv_tokens, + generation_kv_tokens=generation_encoder_kv_tokens, + max_kv_len=max_encoder_kv_len) + if attn_metadata.is_cuda_graph and attn_metadata.has_cross_sub_metadata: # Fast path for stable CUDA-graph generation steps: the encoder # KV lengths (kv_lens_cuda) and the frozen prompt lengths @@ -3243,7 +3566,7 @@ def _prepare_enc_dec_cross_attn_inputs( encoder_num_cached_tokens_per_seq= encoder_num_cached_tokens_per_seq, ) - cross_attn_metadata.prepare() + prepare_cross_metadata(cross_attn_metadata) if new_encoder_tokens == 0: # Record this stable state for future fast-path use. self._cross_attn_stable_cached_tokens = list( @@ -3274,7 +3597,7 @@ def _prepare_enc_dec_cross_attn_inputs( else: self._cross_attn_stable_cached_tokens = None self._cross_attn_stable_request_ids = None - cross_attn_metadata.prepare() + prepare_cross_metadata(cross_attn_metadata) return { "encoder_hidden_states": packed_encoder_hidden_states, @@ -3310,6 +3633,259 @@ def _ship_multimodal_indices( inputs['text_token_indices'] = text_token_indices_cpu.to( "cuda", non_blocking=True) + def _can_use_encoder_decoder_input_fast_path( + self, scheduled_requests: ScheduledRequests, + new_tokens_device: Optional[torch.Tensor], + next_draft_tokens_device: Optional[torch.Tensor]) -> bool: + """Return whether the TRT-like persistent input path is sufficient.""" + static_eligible = getattr( + self, '_encoder_decoder_input_fast_path_static_eligible', None) + if static_eligible is None: + static_eligible = ( + hasattr(batch_manager_bindings, + "prepare_encoder_decoder_inputs") + and self._is_encoder_decoder_model() and not self.is_draft_model + and self.max_beam_width == 1 + and self.sparse_attention_config is None and not self.use_mrope + and not self.enable_attention_dp + and not self.mapping.has_cp_helix() and not self.is_multimodal + and not self.attn_runtime_features.chunked_prefill + and not self.attn_runtime_features.cache_reuse + and not self.attn_runtime_features.has_speculative_draft_tokens) + self._encoder_decoder_input_fast_path_static_eligible = \ + static_eligible + if (not static_eligible or self.enable_spec_decode + or self.lora_model_config is not None + or new_tokens_device is None + or next_draft_tokens_device is not None + or self.guided_decoder is not None): + return False + + if scheduled_requests.batch_size == 0: + return False + for request in scheduled_requests.generation_requests: + if request.py_batch_idx is None and not request.is_dummy: + return False + return True + + def _acquire_encoder_decoder_host_buffers(self) -> Dict[str, Any]: + """Acquire pinned staging whose preceding asynchronous copies finished.""" + pool = getattr(self, '_encoder_decoder_host_buffer_pool', None) + if pool is None: + pool = [] + self._encoder_decoder_host_buffer_pool = pool + for buffers in pool: + event = buffers['event'] + if event is None or event.query(): + return buffers + + buffers = { + 'input_ids': + torch.empty(self.max_num_tokens, + dtype=torch.int, + pin_memory=prefer_pinned()), + 'position_ids': + torch.empty(self.max_num_tokens, + dtype=torch.int, + pin_memory=prefer_pinned()), + 'sequence_lengths': + torch.empty(self.batch_size, + dtype=torch.int, + pin_memory=prefer_pinned()), + 'prompt_lengths': + torch.empty(self.batch_size, + dtype=torch.int, + pin_memory=prefer_pinned()), + 'cached_token_lengths': + torch.empty(self.batch_size, + dtype=torch.int, + pin_memory=prefer_pinned()), + 'kv_lengths': + torch.empty(self.batch_size, + dtype=torch.int, + pin_memory=prefer_pinned()), + 'encoder_kv_lengths': + torch.empty(self.batch_size, + dtype=torch.int, + pin_memory=prefer_pinned()), + 'previous_batch_indices': + torch.empty(self.batch_size, + dtype=torch.int, + pin_memory=prefer_pinned()), + 'event': + None, + } + pool.append(buffers) + return buffers + + @nvtx_range("_prepare_encoder_decoder_inputs_fast") + def _prepare_encoder_decoder_inputs_fast( + self, scheduled_requests: ScheduledRequests, + kv_cache_manager: Union[KVCacheManager, KVCacheManagerV2], + attn_metadata: AttentionMetadata, new_tokens_device: torch.Tensor, + resource_manager: Optional[ResourceManager]): + """Prepare a simple BART batch with native collation and reused buffers.""" + buffers = self._acquire_encoder_decoder_host_buffers() + position_id_offset = getattr(self, + '_encoder_decoder_position_id_offset', + None) + if position_id_offset is None: + position_id_offset = self._get_position_id_offset() + self._encoder_decoder_position_id_offset = position_id_offset + (request_ids, encoder_seq_lens, encoder_cached_token_lengths, + total_num_tokens, num_context_tokens, num_previous_batch_requests, + cached_kv_tokens, context_kv_tokens, generation_kv_tokens, max_kv_len, + context_encoder_kv_tokens, generation_encoder_kv_tokens, + max_encoder_kv_len + ) = batch_manager_bindings.prepare_encoder_decoder_inputs( + scheduled_requests.context_requests, + scheduled_requests.generation_requests, + buffers['input_ids'], + buffers['position_ids'], + buffers['sequence_lengths'], + buffers['prompt_lengths'], + buffers['cached_token_lengths'], + buffers['kv_lengths'], + buffers['encoder_kv_lengths'], + buffers['previous_batch_indices'], + position_id_offset, + ) + + num_sequences = scheduled_requests.batch_size + num_context_requests = scheduled_requests.num_context_requests + num_generation_requests = scheduled_requests.num_generation_requests + generation_request_ids = request_ids[num_context_requests:] + if num_context_tokens: + self.input_ids_cuda[:num_context_tokens].copy_( + buffers['input_ids'][:num_context_tokens], non_blocking=True) + if num_previous_batch_requests: + previous_slots = self.previous_batch_indices_cuda[: + num_previous_batch_requests] + staged_request_ids = generation_request_ids[: + num_previous_batch_requests] + # Sequence slots are stable for a request's lifetime, so the + # device indices remain valid while this ordered batch does. + if self._encoder_decoder_staged_request_ids != staged_request_ids: + previous_slots.copy_(buffers['previous_batch_indices'] + [:num_previous_batch_requests], + non_blocking=True) + self._encoder_decoder_staged_request_ids = staged_request_ids + generation_begin = num_context_tokens + generation_end = generation_begin + num_previous_batch_requests + torch.index_select( + new_tokens_device[0, :, 0], + dim=0, + index=previous_slots, + out=self.input_ids_cuda[generation_begin:generation_end]) + else: + self._encoder_decoder_staged_request_ids = None + dummy_begin = num_context_tokens + num_previous_batch_requests + if dummy_begin < total_num_tokens: + self.input_ids_cuda[dummy_begin:total_num_tokens].fill_(0) + + self.position_ids_cuda[:total_num_tokens].copy_( + buffers['position_ids'][:total_num_tokens], non_blocking=True) + final_position_ids = self.position_ids_cuda[: + total_num_tokens].unsqueeze( + 0) + + sequence_lengths = buffers['sequence_lengths'][:num_sequences] + attn_metadata._seq_lens = sequence_lengths + if (attn_metadata.is_cuda_graph + and attn_metadata._seq_lens_cuda is not None): + attn_metadata._seq_lens_cuda.copy_(sequence_lengths, + non_blocking=True) + else: + attn_metadata._seq_lens_cuda = sequence_lengths.cuda( + non_blocking=True) + + attn_metadata._num_contexts = scheduled_requests.num_context_requests + attn_metadata._num_ctx_tokens = num_context_tokens + attn_metadata._num_generations = num_generation_requests + attn_metadata._num_tokens = total_num_tokens + attn_metadata.beam_width = 1 + attn_metadata.request_ids = request_ids + attn_metadata.prompt_lens = buffers['prompt_lengths'][:num_sequences] + attn_metadata.num_chunked_ctx_requests = 0 + attn_metadata.kv_cache_params = KVCacheParams( + use_cache=True, + num_cached_tokens_per_seq=buffers['cached_token_lengths'] + [:num_sequences], + num_extra_kv_tokens=0) + attn_metadata.kv_cache_manager = kv_cache_manager + assert isinstance(attn_metadata, TrtllmAttentionMetadata) + attn_metadata.prepare_encoder_decoder( + prompt_lens=buffers['prompt_lengths'][:num_sequences], + kv_lens=buffers['kv_lengths'][:num_sequences], + context_kv_tokens=context_kv_tokens, + generation_kv_tokens=generation_kv_tokens, + max_kv_len=max_kv_len) + + encoder_hidden_states = [] + for request in scheduled_requests.context_requests: + encoder_output = request.py_encoder_output + if encoder_output is None: + raise RuntimeError( + f"Decoder context request {request.py_request_id} has no " + "encoder output.") + encoder_hidden_states.append(encoder_output) + request.py_batch_idx = request.py_seq_slot + + cross_attention_inputs = self._prepare_enc_dec_cross_attn_inputs( + encoder_hidden_states, + encoder_seq_lens, + encoder_cached_token_lengths, + attn_metadata, + resource_manager, + encoder_kv_lens=buffers['encoder_kv_lengths'][:num_sequences], + context_encoder_kv_tokens=context_encoder_kv_tokens, + generation_encoder_kv_tokens=generation_encoder_kv_tokens, + max_encoder_kv_len=max_encoder_kv_len, + ) + + attn_all_rank_num_tokens = self._get_all_rank_num_tokens(attn_metadata) + (padded_num_tokens, can_run_piecewise_cuda_graph, + attn_all_rank_num_tokens) = self._get_padding_params( + total_num_tokens, scheduled_requests.num_context_requests, + attn_all_rank_num_tokens) + set_per_request_piecewise_cuda_graph_flag(can_run_piecewise_cuda_graph) + attn_metadata.padded_num_tokens = (padded_num_tokens + if padded_num_tokens + != total_num_tokens else None) + + virtual_num_tokens = total_num_tokens + if attn_metadata.padded_num_tokens is not None: + self.input_ids_cuda[total_num_tokens:padded_num_tokens].fill_(0) + self.position_ids_cuda[total_num_tokens:padded_num_tokens].fill_(0) + virtual_num_tokens = padded_num_tokens + final_position_ids = self.position_ids_cuda[: + virtual_num_tokens].unsqueeze( + 0) + + inputs = { + 'attn_metadata': attn_metadata, + 'input_ids': self.input_ids_cuda[:virtual_num_tokens], + 'position_ids': final_position_ids, + 'inputs_embeds': None, + 'multimodal_params': [], + 'resource_manager': resource_manager, + } + inputs.update(cross_attention_inputs) + + self.iter_states[ + 'num_ctx_requests'] = scheduled_requests.num_context_requests + self.iter_states['num_ctx_tokens'] = num_context_tokens + self.iter_states['num_generation_tokens'] = num_generation_requests + self.iter_states['cached_kv_tokens'] = cached_kv_tokens + if not self.is_warmup: + self.previous_request_ids = generation_request_ids + self.has_previous_device_draft = False + + event = torch.cuda.Event() + event.record(torch.cuda.current_stream()) + buffers['event'] = event + return inputs, None + def _can_use_incremental_update( self, scheduled_requests: ScheduledRequests, new_tokens_device: Optional[torch.Tensor], @@ -3985,12 +4561,23 @@ def _prepare_tp_inputs( # defensively so the two fast paths can never interleave if the # gates ever evolve. self._steady_gen_cache = None + self._encoder_decoder_staged_request_ids = None return self._apply_incremental_update( scheduled_requests, kv_cache_manager, attn_metadata, spec_metadata, new_tensors_device, cache_indirection_buffer, num_accepted_tokens_device, req_id_to_old_request, resource_manager) + if (not promoted_context_request_ids + and type(attn_metadata) is TrtllmAttentionMetadata + and self._can_use_encoder_decoder_input_fast_path( + scheduled_requests, new_tokens_device, + next_draft_tokens_device)): + return self._prepare_encoder_decoder_inputs_fast( + scheduled_requests, kv_cache_manager, attn_metadata, + new_tokens_device, resource_manager) + + self._encoder_decoder_staged_request_ids = None if (not promoted_context_request_ids and self._can_use_steady_gen_fast_prepare( scheduled_requests, new_tokens_device, @@ -5965,43 +6552,16 @@ def _create_encoder_warmup_inputs( """Synthesize an inputs dict that will bucket exactly at (batch_size, num_tokens, max_seq_len). - Uses two distribution strategies: - - Case A: `total >= max_seq_len + (batch_size - 1)` — one request at - `max_seq_len` tokens, remaining tokens distributed evenly across - the other `batch_size - 1` requests. - - Case B: `total < max_seq_len + (batch_size - 1)` — one request of - `total - (batch_size - 1)` tokens, the rest at 1 token each. - Returns None for infeasible combinations (e.g., batch_size <= 0). """ - if batch_size <= 0 or num_tokens <= 0 or max_seq_len <= 0: - return None - - total = min(num_tokens, batch_size * max_seq_len) - - if batch_size == 1: - lengths = [total] - elif total >= max_seq_len + batch_size - 1: - # Case A - remaining = total - max_seq_len - base = remaining // (batch_size - 1) - extra = remaining % (batch_size - 1) - lengths = [max_seq_len] - lengths += [base + 1] * extra + [base] * (batch_size - 1 - extra) - else: - # Case B - first_len = total - (batch_size - 1) - lengths = [first_len] + [1] * (batch_size - 1) - - # Sanity: every length must be >= 1. - if any(length <= 0 for length in lengths): + lengths = ( + self.encoder_cuda_graph_runner.build_capture_sequence_lengths( + batch_size, num_tokens, max_seq_len)) + if lengths is None: return None - actual_num_tokens = sum(lengths) - input_ids = [0] * actual_num_tokens - inputs: Dict[str, Any] = { - 'input_ids': input_ids, + 'input_ids': [0] * sum(lengths), 'seq_lens': lengths, } return inputs @@ -6050,13 +6610,8 @@ def warmup_encoder(self) -> None: # a larger workspace, so the first pass grows the workspace to its # maximum size. The second pass runs the final per-shape warmup and # captures without resizing the workspace. - with self.encoder_cuda_graph_runner.allow_capture(): - self.encoder_cuda_graph_runner.is_warmup_only = True - try: - self._run_cuda_graph_warmup_encoder() - finally: - self.encoder_cuda_graph_runner.is_warmup_only = False - self._run_cuda_graph_warmup_encoder() + self._warmup_and_capture_encoder_cuda_graphs( + self._capture_encoder_cuda_graphs) # Pre-populate the memory pool with max-shape allocations to reduce # fragmentation at runtime. @@ -6104,13 +6659,6 @@ def _run_autotuner_warmup_encoder(self) -> None: f"{len(AutoTuner.get().profiling_cache)}") AutoTuner.get().print_profiling_cache() - def _run_cuda_graph_warmup_encoder(self) -> None: - """Warm up or capture whole-model encode-only CUDA graphs.""" - if not self.encoder_cuda_graph_runner.enabled: - return - - self._capture_encoder_cuda_graphs() - def _capture_encoder_cuda_graphs(self) -> None: """Warm up or capture encoder CUDA graphs for all feasible keys. @@ -6124,7 +6672,7 @@ def _capture_encoder_cuda_graphs(self) -> None: if not runner.enabled: return - batch_sizes = sorted(self._cuda_graph_batch_sizes, reverse=True) + batch_sizes = sorted(self._encoder_cuda_graph_batch_sizes, reverse=True) num_tokens_list = sorted(self._cuda_graph_num_tokens) seq_lens_list = sorted(self._cuda_graph_seq_lens) @@ -6132,7 +6680,7 @@ def _capture_encoder_cuda_graphs(self) -> None: num_processed = 0 logger.info(f"Running encoder CUDA graph {operation} ...") for bs in batch_sizes: - if bs > self.batch_size: + if bs > self.encoder_batch_size: continue for sl_idx, sl in reversed(list(enumerate(seq_lens_list))): prev_sl = seq_lens_list[sl_idx - 1] if sl_idx > 0 else 0 @@ -6389,6 +6937,8 @@ def forward(self, padded_graph_requests.all_requests()) self._sync_group_all_greedy_sample(spec_metadata) + allow_mixed_encoder_decoder_graph = ( + self.cuda_graph_runner.enable_encoder_decoder_mixed_cuda_graph) maybe_attn_metadata, maybe_spec_metadata, key = self.cuda_graph_runner.maybe_get_cuda_graph( padded_graph_requests, enable_spec_decode=self.enable_spec_decode, @@ -6398,6 +6948,7 @@ def forward(self, if self.is_spec_decode else None, new_tensors_device=new_tensors_device, spec_resource_manager=spec_resource_manager, + allow_mixed_encoder_decoder=(allow_mixed_encoder_decoder_graph), promoted_context_request_ids=promoted_context_request_ids, ) @@ -6665,14 +7216,17 @@ def _make_encoder_attn_metadata( """Build fresh, no-cache attention metadata for one packed encoder batch. ``self.attn_metadata`` is not reused because that object is bound to the decoder's KV-cache manager.""" + if len(sequence_lengths) != len(request_ids): + raise ValueError("Encoder sequence lengths and request IDs must " + "have the same length.") sparse_metadata_params = ( self.sparse_attention_config.to_sparse_metadata_params( pretrained_config=self.model.model_config.pretrained_config) if self.sparse_attention_config is not None else None) 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, + max_num_requests=self.encoder_batch_size, + max_num_tokens=self.encoder_max_num_tokens, + max_num_sequences=self.encoder_batch_size * self.max_beam_width, kv_cache_manager=None, mapping=self.mapping, runtime_features=self.attn_runtime_features, @@ -6698,6 +7252,58 @@ def _make_encoder_attn_metadata( encoder_attn_metadata.prepare_encoder_only() return encoder_attn_metadata + def _prepare_encoder_decoder_encoder_inputs( + self, + encoder_input_ids: List[int], + encoder_position_ids: List[int], + sequence_lengths: List[int], + request_ids: List[int], + resource_manager: Optional[ResourceManager] = None, + ) -> Dict[str, Any]: + num_tokens = len(encoder_input_ids) + if num_tokens != len(encoder_position_ids): + raise ValueError("Encoder input IDs and position IDs must have " + "the same length.") + assert num_tokens <= self.encoder_max_num_tokens, ( + f"encoder packed length ({num_tokens}) exceeds " + f"encoder_max_num_tokens ({self.encoder_max_num_tokens})") + + encoder_attn_metadata = self._make_encoder_attn_metadata( + sequence_lengths, request_ids) + 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()) + encoder_graph_runner = self.encoder_cuda_graph_runner + encoder_batch_size = len(sequence_lengths) + use_graph_staging = ( + encoder_graph_runner.enabled and + (encoder_batch_size in encoder_graph_runner.supported_batch_sizes or + (encoder_graph_runner.padding_enabled and encoder_batch_size + <= encoder_graph_runner.max_supported_batch_size))) + + return { + 'encoder_input_ids': + (encoder_input_ids_t if use_graph_staging else + encoder_input_ids_t.to('cuda', non_blocking=True)), + 'encoder_position_ids': + ((encoder_position_ids_t + if use_graph_staging else encoder_position_ids_t.to( + 'cuda', non_blocking=True)).unsqueeze(0)), + 'encoder_attn_metadata': + encoder_attn_metadata, + 'encoder_seq_lens': + sequence_lengths, + 'encoder_input_ids_host': + encoder_input_ids_t, + 'encoder_position_ids_host': + encoder_position_ids_t, + 'resource_manager': + resource_manager, + } + @nvtx_range("_prepare_tp_inputs_encoder_features") def _prepare_tp_inputs_encoder_features( self, @@ -6727,9 +7333,9 @@ def _prepare_tp_inputs_encoder_features( request_ids.append(request.py_request_id) num_tokens = sum(sequence_lengths) - assert num_tokens <= self.max_num_tokens, ( - f"encoder packed length ({num_tokens}) exceeds max_num_tokens " - f"({self.max_num_tokens})") + assert num_tokens <= self.encoder_max_num_tokens, ( + f"encoder packed length ({num_tokens}) exceeds " + f"encoder_max_num_tokens ({self.encoder_max_num_tokens})") encoder_attn_metadata = self._make_encoder_attn_metadata( sequence_lengths, request_ids) @@ -6798,34 +7404,13 @@ def _prepare_tp_inputs_encoder( 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})") - - encoder_attn_metadata = self._make_encoder_attn_metadata( - sequence_lengths, request_ids) - - 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 + return self._prepare_encoder_decoder_encoder_inputs( + encoder_input_ids=encoder_input_ids, + encoder_position_ids=encoder_position_ids, + sequence_lengths=sequence_lengths, + request_ids=request_ids, + resource_manager=resource_manager, + ) @nvtx_range("_forward_step_encoder") def _forward_step_encoder( @@ -6891,6 +7476,90 @@ def _forward_step_encoder( ) return encoder_hidden_states + def _forward_step_encoder_cuda_graph( + self, + inputs: Dict[str, Any], + ) -> torch.Tensor: + return self._forward_step_encoder({ + 'encoder_input_ids': + inputs['input_ids'], + 'encoder_position_ids': + inputs.get('position_ids'), + 'encoder_attn_metadata': + inputs['attn_metadata'], + 'resource_manager': + inputs.get('resource_manager'), + }) + + def _encoder_forward_enc_dec( + self, + inputs: Dict[str, Any], + ) -> torch.Tensor: + """Run the encoder-decoder encoder, using a CUDA graph when eligible.""" + input_ids = inputs.get('encoder_input_ids_host') + position_ids = inputs.get('encoder_position_ids_host') + seq_lens = inputs['encoder_seq_lens'] + runner = self.encoder_cuda_graph_runner + + if input_ids is None or position_ids is None: + return self._forward_step_encoder(inputs) + + runner_inputs = { + 'input_ids': input_ids, + 'position_ids': position_ids, + 'seq_lens': seq_lens, + 'resource_manager': inputs.get('resource_manager'), + } + with runner.pad_batch(runner_inputs, + len(seq_lens)) as padded_runner_inputs: + graph_attn_metadata, key = runner.maybe_get_cuda_graph( + padded_runner_inputs, inputs['encoder_attn_metadata']) + if key is None: + if inputs['encoder_input_ids'].device.type == 'cpu': + inputs = dict(inputs) + inputs['encoder_input_ids'] = inputs[ + 'encoder_input_ids'].to('cuda', non_blocking=True) + inputs['encoder_position_ids'] = inputs[ + 'encoder_position_ids'].to('cuda', non_blocking=True) + return self._forward_step_encoder(inputs) + + # Every graph key aliases the same pinned staging allocation. Retire + # the previous captured H2D before updating seq_lens or any other + # shared host input for this replay. + runner.retire_staging() + model_inputs = runner.prepare_encoder_decoder_inputs( + padded_runner_inputs, key, seq_lens) + graph_attn_metadata.prepare_encoder_cuda_graph_replay( + model_inputs['seq_lens'], key[1]) + model_inputs['attn_metadata'] = graph_attn_metadata + + moe_load_balancer: MoeLoadBalancer = getattr( + self, 'moe_load_balancer', None) + with with_shared_pool(runner.get_graph_pool()): + capture_outputs = None + if runner.needs_capture(key): + + def capture_forward_fn( + capture_inputs: Dict[str, Any]) -> torch.Tensor: + with MoeLoadBalancerIterContext(moe_load_balancer): + return self._forward_step_encoder_cuda_graph( + capture_inputs) + + capture_outputs = runner.capture(key, capture_forward_fn, + model_inputs) + + if runner.is_warmup_only: + graph_outputs = capture_outputs + else: + with MoeLoadBalancerIterContext(moe_load_balancer): + graph_outputs = runner.replay(key, model_inputs) + + if not isinstance(graph_outputs, torch.Tensor): + raise TypeError("Encoder-decoder CUDA graph replay must return " + "a tensor of encoder hidden states.") + return runner.restore_encoder_decoder_output(key, graph_outputs, + model_inputs) + @nvtx_range("forward_encoder") def forward_encoder( self, @@ -6917,7 +7586,7 @@ def forward_encoder( with torch.inference_mode(): inputs = self._prepare_tp_inputs_encoder( encoder_requests, resource_manager=resource_manager) - encoder_hidden_states = self._forward_step_encoder(inputs) + encoder_hidden_states = self._encoder_forward_enc_dec(inputs) return encoder_hidden_states, inputs['encoder_seq_lens'] diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 6e3252b48688..e47d45c9ade5 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -8,6 +8,7 @@ import threading import time import traceback +from concurrent.futures import Future, ThreadPoolExecutor from contextlib import contextmanager from enum import IntEnum from queue import Queue @@ -382,6 +383,20 @@ class BatchStatePP(BatchState): microbatch_id: int = -1 +@dataclasses.dataclass +class EncoderStepResult: + hidden_states: torch.Tensor + sequence_lengths: List[int] + ready_event: torch.cuda.Event + + +@dataclasses.dataclass +class PendingEncoderStep: + requests: List[LlmRequest] + future: Future[EncoderStepResult] + result: Optional[EncoderStepResult] = None + + class AsyncTransferManager: """ Handle asynchronous transfer of KV cache after a request has completed. @@ -771,10 +786,10 @@ def __init__( # models, so encoder PP send/recv support is not implemented in the # PyTorch path for now. Reject pp_size > 1. # TODO: Add support for pp + encoder models - is_encoder_decoder = bool( + self.is_encoder_decoder = bool( getattr(getattr(self.model_engine.model, "model_config", None), "is_encoder_decoder", False)) - if is_encoder_decoder: + if self.is_encoder_decoder: if self.dist.pp_size > 1: raise NotImplementedError( "pp_size > 1 is not supported for encoder-decoder models " @@ -814,6 +829,17 @@ def __init__( # Ensure the default stream waits for execution_stream to complete # before subsequent operations. torch.cuda.current_stream().wait_stream(self.execution_stream) + self.encoder_launch_executor = (ThreadPoolExecutor( + max_workers=1, thread_name_prefix="encoder-launch") + if self.is_encoder_decoder else None) + self.pending_encoder_steps: List[PendingEncoderStep] = [] + if self.encoder_launch_executor is not None: + # CUDA graph capture inherits per-thread CUDA library state. Capture + # on the same single worker that owns every runtime encoder replay. + self.encoder_stream.wait_stream(self.execution_stream) + self.encoder_launch_executor.submit( + self._warmup_encoder_cuda_graphs_enc_dec).result() + self.is_warmup = False # Snapshot some cumulative KV cache counters so that stats reported to @@ -831,6 +857,7 @@ def __init__( self.adp_ctx_waiting_iters_count = 0 self.adp_ctx_batching_wait_iters_count = 0 self.batch_wait_iters_count = 0 + self.encoder_batch_wait_iters_count = 0 def on_detected(): # The graceful shutdown path can itself deadlock on collectives @@ -1472,6 +1499,11 @@ def shutdown(self): # executor loops have already processed the shutdown broadcast and are # no longer driving NCCL, so the send cannot deadlock. self._shutdown_sleep_wakeup_listeners() + encoder_launch_executor = getattr(self, "encoder_launch_executor", None) + if encoder_launch_executor is not None: + encoder_launch_executor.shutdown(wait=True) + self.encoder_launch_executor = None + self.worker_started = False # Release CUDA graphs before resource managers free their GPU memory. # Resource managers (e.g. SuffixAutomatonManager) allocate GPU workspace @@ -3596,6 +3628,7 @@ def _sync_gen_only_benchmark_has_insufficient_kv( def _prepare_and_schedule_batch(self): self._sync_disagg_transfer_made_progress = False + self._poll_encoder_steps() new_requests = self._fetch_and_activate_new_requests() if self.should_stop_processing: return None, None @@ -4035,15 +4068,6 @@ 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: @@ -4073,6 +4097,9 @@ def _executor_loop(self): self._revert_gen_alloc(scheduled_batch) self._finalize_adp_dummy_allocation(can_queue) + if not can_queue and scheduled_batch.encoder_requests: + self._run_encoder_step(scheduled_batch.encoder_requests) + if can_queue: # init_disagg_gen_requests must be before drafter loop, otherwise draft requests do not have initialized matchers. # init_disagg_gen_requests must be before engine forward, where the prev_seq_slot is updated. @@ -4117,6 +4144,10 @@ def _executor_loop(self): ) gpu_forward_events_from_perf_pool = True + if scheduled_batch.encoder_requests: + self._submit_encoder_step( + scheduled_batch.encoder_requests) + with self.perf_manager.record_perf_events( gpu_forward_start, gpu_forward_end) as fwd_timing: if self.dwdp_manager is not None: @@ -4503,9 +4534,6 @@ def _executor_loop_overlap(self): if not self._is_kv_manager_v2: self._terminate_requests(scheduled_batch.paused_requests) - if scheduled_batch.encoder_requests: - self._run_encoder_step(scheduled_batch.encoder_requests) - gpu_forward_events_from_perf_pool = False can_queue, can_queue_this_rank = self._can_queue( scheduled_batch) @@ -4551,6 +4579,9 @@ def _executor_loop_overlap(self): self._revert_gen_alloc(scheduled_batch) self._finalize_adp_dummy_allocation(can_queue) + if not can_queue and scheduled_batch.encoder_requests: + self._run_encoder_step(scheduled_batch.encoder_requests) + # If the batch is not empty on this rank, but empty on other ranks, # we need to delay the update of the previous batch's sample state, # and let the later iteration to update it. @@ -4616,6 +4647,9 @@ def _executor_loop_overlap(self): gpu_forward_start, gpu_forward_end = self.perf_manager.borrow_forward_timing_events( ) gpu_forward_events_from_perf_pool = True + if scheduled_batch.encoder_requests: + self._submit_encoder_step( + scheduled_batch.encoder_requests) with self.perf_manager.record_perf_events( gpu_forward_start, gpu_forward_end) as fwd_timing: @@ -5295,6 +5329,95 @@ def _waiting_requests(self, context_requests: list[LlmRequest], self.batch_wait_iters_count = 0 return context_requests + def _waiting_encoder_requests( + self, encoder_requests: list[LlmRequest], + context_requests: list[LlmRequest], + generation_requests: list[LlmRequest]) -> list[LlmRequest]: + """Accumulate encoder work while an admitted decode batch progresses. + + Encoder-decoder serving can otherwise launch one eager encoder forward + for each replacement request. Use the existing iteration deadline and + token threshold to form a larger encoder microbatch without blocking + the executor thread. Decoder generation continues while the encoder + requests wait. The encoder has its own counter because the resulting + decoder-context requests are already coalesced and must not wait for a + second window. CUDA-graph microbatch admission returns at most one + supported graph batch so scheduler overfill cannot turn a target-eight + batch into an eager batch of nine or more requests. + """ + if not encoder_requests: + self.encoder_batch_wait_iters_count = 0 + return encoder_requests + + encoder_max_batch_size = self.llm_args.encoder_max_batch_size + encoder_cuda_graph_config = self.llm_args.encoder_cuda_graph_config + if (encoder_max_batch_size is not None + and encoder_cuda_graph_config is not None + and bool(encoder_cuda_graph_config.num_tokens) + and bool(encoder_cuda_graph_config.seq_lens)): + encoder_batch_size_limit = min(encoder_max_batch_size, + self.max_batch_size) + configured_batch_sizes = (encoder_cuda_graph_config.batch_sizes + or []) + supported_batch_sizes = [ + batch_size for batch_size in configured_batch_sizes + if batch_size <= encoder_batch_size_limit + ] + if (encoder_cuda_graph_config.enable_padding + and any(batch_size > encoder_batch_size_limit + for batch_size in configured_batch_sizes) and + (not supported_batch_sizes + or supported_batch_sizes[-1] != encoder_batch_size_limit)): + supported_batch_sizes.append(encoder_batch_size_limit) + + if supported_batch_sizes: + microbatch_target = supported_batch_sizes[-1] + decoder_occupancy = (len(context_requests) + + len(generation_requests)) + decoder_low_watermark = (self.max_batch_size - + microbatch_target) + deadline_reached = (self.encoder_batch_wait_iters_count + >= self.batch_wait_timeout_iters) + + if decoder_occupancy <= decoder_low_watermark: + if len(encoder_requests) >= microbatch_target: + self.encoder_batch_wait_iters_count = 0 + return encoder_requests[:microbatch_target] + + if deadline_reached: + releasable_batch_sizes = [ + batch_size for batch_size in supported_batch_sizes + if batch_size <= len(encoder_requests) + ] + fallback_batch_size = ( + releasable_batch_sizes[-1] if releasable_batch_sizes + else min(len(encoder_requests), microbatch_target)) + self.encoder_batch_wait_iters_count = 0 + return encoder_requests[:fallback_batch_size] + + self.encoder_batch_wait_iters_count += 1 + return [] + + num_scheduled_tokens = sum(request.encoder_output_len + for request in encoder_requests) + num_scheduled_tokens += sum(1 + request.num_draft_tokens + for request in generation_requests) + has_decoder_work = bool(context_requests or generation_requests) + if not has_decoder_work: + has_decoder_work = any( + request.request_id in self.inflight_req_ids + and request.state != LlmRequestState.ENCODER_INIT + for request in self.active_requests) + should_wait = (has_decoder_work and self.encoder_batch_wait_iters_count + < self.batch_wait_timeout_iters and num_scheduled_tokens + < self.batch_wait_max_tokens_ratio * self.max_num_tokens) + if should_wait: + self.encoder_batch_wait_iters_count += 1 + return [] + + self.encoder_batch_wait_iters_count = 0 + return encoder_requests + def _get_ctx_mla_kv_len_cap(self): """Cap on the summed context attended-KV length (total_kv_len) per forward step, cached. @@ -5361,6 +5484,16 @@ def _schedule(self): scheduler_output = self.scheduler.schedule_request( self.active_requests, self.inflight_req_ids) + scheduled_encoder_requests = scheduler_output.encoder_requests + should_batch_encoder_requests = (self.is_encoder_decoder + and not self.enable_attention_dp + and self.enable_batch_waiting) + if should_batch_encoder_requests: + scheduled_encoder_requests = self._waiting_encoder_requests( + scheduler_output.encoder_requests, + scheduler_output.context_requests, + scheduler_output.generation_requests) + scheduled_context_requests = scheduler_output.context_requests if self.enable_attention_dp and self.attention_dp_enable_balance: scheduled_context_requests = self._balance_adp_requests( @@ -5368,9 +5501,12 @@ def _schedule(self): scheduler_output.generation_requests) # If no generation requests, no need to wait, to avoid dead waiting - should_check_waiting = not self.enable_attention_dp and self.enable_batch_waiting and len( - scheduler_output.context_requests) > 0 and len( - scheduler_output.generation_requests) > 0 + should_check_waiting = (not self.is_encoder_decoder + and not self.enable_attention_dp + and self.enable_batch_waiting + and len(scheduler_output.context_requests) > 0 + and len( + scheduler_output.generation_requests) > 0) if should_check_waiting: # With KV cache manager V2, scheduling has already grown context request KV cache capacity. Requests dropped # for batch waiting still occupy KV cache and may reduce the batch size available for generation requests. @@ -5394,7 +5530,7 @@ def _schedule(self): scheduled_context_requests) scheduled_requests = ScheduledRequests() - scheduled_requests.encoder_requests = scheduler_output.encoder_requests + scheduled_requests.encoder_requests = scheduled_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 @@ -5421,44 +5557,170 @@ def _schedule(self): # 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``. + def _warmup_encoder_cuda_graphs_enc_dec(self) -> None: + """Capture encoder graphs on the worker used for runtime replay.""" + warmup = getattr( + self.model_engine, + "_warmup_encoder_cuda_graphs_enc_dec", + None, + ) + if not callable(warmup): + return + + torch.cuda.set_device(self.device_id) + with torch.cuda.stream(self.encoder_stream): + warmup(self.resource_manager) + + def _submit_encoder_step(self, encoder_requests: List[LlmRequest]) -> None: + """Queue encoder work, serializing it with decoder work under TP.""" + executor = self.encoder_launch_executor + if executor is None: + raise RuntimeError("Encoder launch executor is unavailable.") + + requests = list(encoder_requests) + for request in requests: + self.inflight_req_ids.insert(request.request_id) + + serialize_tp = self.dist.tp_size > 1 + if serialize_tp: + self.encoder_stream.wait_stream(self.execution_stream) + + try: + future = executor.submit(self._run_encoder_step_unchecked, requests) + except Exception: + for request in requests: + self.inflight_req_ids.erase(request.request_id) + raise + + if serialize_tp: + # Encoder and decoder forwards share TP communicators and + # workspaces. Complete encoder GPU work before the caller launches + # decoder work, while retaining the worker thread affinity needed + # by encoder CUDA graph replay. + try: + result = future.result() + result.ready_event.synchronize() + self._publish_encoder_step(requests, result) + except Exception as e: + self._finish_failed_encoder_step(requests, e) + return + for request in requests: + self.inflight_req_ids.erase(request.request_id) + return + + self.pending_encoder_steps.append( + PendingEncoderStep(requests=requests, future=future)) + + @nvtx_range("_poll_encoder_steps") + def _poll_encoder_steps(self) -> None: + """Publish ready encoder batches without waiting for their futures. - 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. + Encoder request IDs stay in ``inflight_req_ids`` from submission + until both the launch worker and its CUDA completion event finish. + Consequently the scheduler cannot submit an encoder request twice or + admit its decoder-context step before the encoder output is ready. """ - if not encoder_requests: + pending_steps = getattr(self, "pending_encoder_steps", None) + if not pending_steps: return + while pending_steps: + pending = pending_steps[0] + if pending.result is None: + if not pending.future.done(): + break + try: + pending.result = pending.future.result() + except Exception as e: + pending_steps.pop(0) + self._finish_failed_encoder_step(pending.requests, e) + continue + + if not pending.result.ready_event.query(): + break + + pending_steps.pop(0) + try: + self._publish_encoder_step(pending.requests, pending.result) + except Exception as e: + self._finish_failed_encoder_step(pending.requests, e) + continue + + for request in pending.requests: + self.inflight_req_ids.erase(request.request_id) + + def _finish_failed_encoder_step(self, encoder_requests: List[LlmRequest], + error: Exception) -> None: + for request in encoder_requests: + self.inflight_req_ids.erase(request.request_id) + traceback.print_exception(error) + error_msg = str(error) + logger.error(f"Encountered an error in encoder forward: {error_msg}") + failed_requests = [ + request for request in encoder_requests + if request.state != LlmRequestState.GENERATION_COMPLETE + ] + if failed_requests: + self._handle_errors(error_msg, requests=failed_requests) + + def _run_encoder_step(self, encoder_requests: List[LlmRequest]) -> None: 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, - )) + executor = self.encoder_launch_executor + if executor is None: + raise RuntimeError("Encoder launch executor is unavailable.") + future = executor.submit(self._run_encoder_step_unchecked, + encoder_requests) + result = future.result() + self._publish_encoder_step(encoder_requests, result) 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._finish_failed_encoder_step(encoder_requests, e) - 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("_run_encoder_step") + def _run_encoder_step_unchecked( + self, encoder_requests: List[LlmRequest]) -> EncoderStepResult: + """Drive one encoder iteration for ``encoder_requests``. + + Runs the encoder stack on its independent stream and returns packed + output plus a completion event. Request state remains owned by the + main executor thread and is updated by ``_poll_encoder_steps`` only + after this event reports ready. + + The caller submits encoder work immediately before decoder forward so + their independent CPU launch and GPU execution can overlap. Encoder + inputs are staged inside this stream, and request-level completion + events guard the only downstream dependency. Successive encoder + batches remain ordered by the stream itself. + """ + if not encoder_requests: + raise ValueError("Encoder step requires at least one request.") + + torch.cuda.set_device(self.device_id) + 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, + )) + + ready_event = torch.cuda.Event() + ready_event.record(self.encoder_stream) + return EncoderStepResult( + hidden_states=encoder_hidden_states, + sequence_lengths=encoder_seq_lens, + ready_event=ready_event, + ) + + def _publish_encoder_step(self, encoder_requests: List[LlmRequest], + result: EncoderStepResult) -> None: + """Make a completed encoder batch visible to decoder scheduling.""" + self._scatter_encoder_output( + encoder_requests, + result.hidden_states, + result.sequence_lengths, + result.ready_event, + ) + # 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( @@ -5466,6 +5728,7 @@ def _scatter_encoder_output( encoder_requests: List[LlmRequest], encoder_hidden_states: torch.Tensor, encoder_seq_lens: List[int], + ready_event: torch.cuda.Event, ) -> None: """Slice packed encoder hidden states into per-request tensors. @@ -5494,10 +5757,12 @@ def _scatter_encoder_output( 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 + if req.state == LlmRequestState.ENCODER_INIT: + req.py_encoder_output = encoder_hidden_states[offset:offset + + seq_len] + req.py_skip_cross_kv_projection = False + req.py_encoder_output_ready_event = ready_event + req.state = LlmRequestState.CONTEXT_INIT offset += seq_len @nvtx_range("_attach_encoder_output_to_execution_stream") diff --git a/tensorrt_llm/_torch/pyexecutor/sampler/sampler.py b/tensorrt_llm/_torch/pyexecutor/sampler/sampler.py index a71558740f67..ae55460fd511 100644 --- a/tensorrt_llm/_torch/pyexecutor/sampler/sampler.py +++ b/tensorrt_llm/_torch/pyexecutor/sampler/sampler.py @@ -1024,6 +1024,7 @@ def finish_reasons_list(self) -> FinishReasonsList: @dataclass(kw_only=True) class SampleStateTorch(SampleState[SampleStateTensorsHostTorch, SampleStateTensors]): beam_history_builders: list[BeamHistoryBuilder | None] | None = None + single_step_greedy: bool = False @dataclass(kw_only=True, frozen=True) @@ -1474,6 +1475,9 @@ def __init__(self, args: Args): self._prev_first_finish_reasons_host: list[torch.Tensor | None] = [ None ] * self.max_num_sequences + self._stable_greedy_request_ids: list[int] = [] + self._stable_greedy_seq_slots_host: Optional[torch.Tensor] = None + self._stable_greedy_seq_slots_cuda: Optional[torch.Tensor] = None @staticmethod def _is_draft_batch(requests: list[LlmRequest]) -> bool: @@ -2557,6 +2561,12 @@ def update_requests( self._pending_steps[slot] -= 1 assert state.host is not None + # Reuse sample_async's qualification instead of rechecking every + # request after the asynchronous sample completes. + if state.single_step_greedy: + self._update_requests_single_beam_single_step(state) + return + new_tokens = state.host.new_tokens finish_reasons = state.host.finish_reasons_list() first_finish_reasons = ( @@ -2707,6 +2717,48 @@ def _maybe_build_beam_history(req_idx: int) -> BeamHistory | None: self._penalty_handler.update_token_counts(finalized_token_updates) + @nvtx_range("_update_requests_single_beam_single_step") + def _update_requests_single_beam_single_step(self, state: SampleStateTorch) -> None: + """Update the common greedy, single-token case without draft machinery.""" + assert state.host is not None + requests = [ + request + for request in state.requests + if request.state != LlmRequestState.GENERATION_COMPLETE + ] + if not requests: + return + + all_new_tokens = state.host.new_tokens.tolist() + if len(requests) == len(state.requests): + new_tokens = all_new_tokens + else: + new_tokens = [ + new_token + for request, new_token in zip(state.requests, all_new_tokens) + if request.state != LlmRequestState.GENERATION_COMPLETE + ] + add_new_tokens_to_requests(requests, new_tokens, DEFAULT_BEAM_IDX) + + # sample_async deliberately omits the device finish-reason tensor for + # this qualified path; completion is derived from compact host tokens. + assert state.host.finish_reasons is None + for request, new_token in zip(requests, new_tokens): + # The stable greedy path excludes stop words. Keep EOS ahead of the + # length check so a terminal EOS at the token limit is reported as + # END_ID, matching _handle_stop_criteria. + if new_token == request.py_end_id: + request.finish_by(FinishReason.END_ID, DEFAULT_BEAM_IDX) + elif ( + request.max_beam_num_tokens - request.py_orig_prompt_len + >= request.py_max_new_tokens + or request.max_beam_num_tokens >= self.max_seq_len + ): + request.finish_by(FinishReason.LENGTH, DEFAULT_BEAM_IDX) + request.py_num_accepted_draft_tokens = 0 + request.py_rewind_len = 0 + request.py_decoding_iter += 1 + def _return_log_probs(self, requests: list[LlmRequest]) -> bool: return any(req.py_return_log_probs for req in requests) @@ -2764,6 +2816,7 @@ def sample_async( seq_slots_cuda, seq_lens_cuda, new_tokens_host, + single_step_greedy, ) = self._process_requests( scheduled_requests, model_outputs, @@ -2790,22 +2843,26 @@ def sample_async( # their buffers in the store. # Assume that either all requests are drafts or none are drafts is_draft_batch = requests[0].py_is_draft - finish_reasons_device = self._finish_reasons_handler.write_finish_reasons( - seq_slots_host=seq_slots_host, - is_draft_batch=is_draft_batch, - seq_slots_cuda=seq_slots_cuda, - seq_lens_cuda=seq_lens_cuda, - new_tokens_cuda=new_tokens, - first_finish_reasons_cuda=( - beam_search_store.first_finish_reasons - if beam_search_store is not None - else None - ), - ) - finish_reasons_host = self._copy_to_host(finish_reasons_device) + if not single_step_greedy: + assert seq_lens_host is not None + assert seq_lens_cuda is not None + finish_reasons_device = self._finish_reasons_handler.write_finish_reasons( + seq_slots_host=seq_slots_host, + is_draft_batch=is_draft_batch, + seq_slots_cuda=seq_slots_cuda, + seq_lens_cuda=seq_lens_cuda, + new_tokens_cuda=new_tokens, + first_finish_reasons_cuda=( + beam_search_store.first_finish_reasons + if beam_search_store is not None + else None + ), + ) + finish_reasons_host = self._copy_to_host(finish_reasons_device) if self._use_beam_search: assert beam_search_store is not None + assert seq_lens_cuda is not None first_finish_reasons = beam_search_store.first_finish_reasons first_finish_reasons_host = self._copy_to_host(first_finish_reasons) self._update_original_tokens( @@ -2848,6 +2905,7 @@ def sample_async( ), sampler_event=sampler_event, beam_history_builders=beam_history_builders, + single_step_greedy=single_step_greedy, ) @staticmethod @@ -2868,7 +2926,7 @@ def _fast_greedy_sample_kernel( batch_dest_indices: torch.Tensor, max_beam_width: int, d2t: torch.Tensor | None, - ) -> None: + ) -> torch.Tensor: """Applies fast greedy sampling to the logits. Performs argmax, applies d2t translation if present, and scatters @@ -2887,6 +2945,7 @@ def _fast_greedy_sample_kernel( new_tokens_cuda.view(-1, *new_tokens_cuda.shape[2:]).scatter_( 0, batch_dest_indices_expanded, next_tokens_expanded ) + return next_tokens @staticmethod def _apply_embedding_bias( @@ -3781,10 +3840,83 @@ def _process_requests( new_tokens_cuda: torch.Tensor, num_context_logits_prefix_sum: list[int], ) -> tuple[ - list[LlmRequest], torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor + list[LlmRequest], + torch.Tensor, + Optional[torch.Tensor], + torch.Tensor, + Optional[torch.Tensor], + torch.Tensor, + bool, ]: raw_logits_cuda = model_outputs["logits"] + generation_requests = scheduled_requests.generation_requests + request_ids = [request.py_request_id for request in generation_requests] + has_stable_request_ids = self._stable_greedy_request_ids == request_ids + can_use_stable_greedy_path = ( + bool(generation_requests) + and self.max_beam_width == 1 + and scheduled_requests.num_context_requests == 0 + and len(generation_requests) <= raw_logits_cuda.shape[0] + and model_outputs.get("d2t") is None + and all( + not request.is_dummy and get_draft_token_length(request) == 0 + for request in generation_requests + ) + and ( + has_stable_request_ids + or all( + request._py_embedding_bias_1d is None + and not getattr(request, "py_bad_words", None) + and not getattr(request, "py_no_repeat_ngram_size", None) + and not request.py_min_length + and not request.py_return_log_probs + and not request.py_stop_words_list + and _request_strategy(request, vocab_size=2**31) == GREEDY + for request in generation_requests + ) + ) + ) + if can_use_stable_greedy_path: + if has_stable_request_ids: + assert self._stable_greedy_seq_slots_host is not None + assert self._stable_greedy_seq_slots_cuda is not None + seq_slots_host = self._stable_greedy_seq_slots_host + seq_slots_cuda = self._stable_greedy_seq_slots_cuda + else: + maybe_seq_slots = [request.py_seq_slot for request in generation_requests] + assert all(seq_slot is not None for seq_slot in maybe_seq_slots) + seq_slots = [cast(int, seq_slot) for seq_slot in maybe_seq_slots] + seq_slots_host = torch.tensor( + seq_slots, dtype=torch.int32, pin_memory=prefer_pinned() + ) + seq_slots_cuda = seq_slots_host.to( + device="cuda", dtype=torch.int64, non_blocking=True + ) + self._stable_greedy_request_ids = request_ids + self._stable_greedy_seq_slots_host = seq_slots_host + self._stable_greedy_seq_slots_cuda = seq_slots_cuda + + next_tokens = self._fast_greedy_sample_kernel( + raw_logits_cuda[: len(generation_requests)], + new_tokens_cuda, + seq_slots_cuda, + self.max_beam_width, + None, + ) + new_tokens_host = self._copy_to_host(next_tokens) + return ( + generation_requests, + seq_slots_host, + None, + seq_slots_cuda, + None, + new_tokens_host, + True, + ) + + self._stable_greedy_request_ids = [] + sampling_requests, sampling_requests_metadata, logits_cuda = self._select_generated_logits( scheduled_requests, raw_logits_cuda, @@ -3906,6 +4038,7 @@ def _process_requests( seq_slots_cuda, seq_lens_cuda, new_tokens_host, + False, ) # Indexer for accessing tokens in 'logits_cuda', corresponding to the @@ -3961,6 +4094,7 @@ def _process_requests( seq_slots_cuda, seq_lens_cuda, new_tokens_host, + False, ) @override diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index d325b69f65d0..a7b6fa6abcfa 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -4981,6 +4981,25 @@ class TorchLlmArgs(BaseLlmArgs): Note that each CUDA graph can use up to 200 MB of extra memory.", status="beta") + encoder_cuda_graph_config: Optional[EncodeCudaGraphConfig] = Field( + default=None, + description=( + "CUDA graph configuration for the encoder forward pass of an " + "encoder-decoder model. Use `cuda_graph_config` for the decoder " + "and this field for the encoder. Encoder CUDA graphs require " + "`encoder_max_batch_size` to be set."), + status="prototype") + + enable_encoder_decoder_mixed_cuda_graph: bool = Field( + default=True, + description=( + "Enable the mixed-batch CUDA graph performance optimization for " + "encoder-decoder models. The graph handles decoder iterations " + "containing both context and generation requests. It is enabled " + "by default when both `cuda_graph_config` and " + "`encoder_cuda_graph_config` produce usable graph shapes."), + status="prototype") + @field_validator('cuda_graph_config', mode='before') @classmethod def infer_cuda_graph_config_mode(cls, v): @@ -5034,21 +5053,21 @@ def init_multimodal_config(cls, v): encoder_max_batch_size: Optional[int] = Field( default=None, - description=( - "Maximum batch size for the multimodal encoder's AttentionMetadata. " - "Falls back to `max_batch_size` when unset. This budget is shared " - "proportionately across all modalities the model encodes, not set " - "per modality; per-modality knobs may be added later."), + description= + ("Maximum encoder batch size. For encoder-decoder models, this also " + "controls encoder microbatch admission and limits encoder CUDA graph " + "batch sizes. For multimodal models, this is the shared " + "AttentionMetadata budget across all encoded modalities. Falls back " + "to `max_batch_size` when unset."), status="prototype") encoder_max_num_tokens: Optional[int] = Field( default=None, description=( - "Maximum number of tokens for the multimodal encoder's " - "AttentionMetadata. Falls back to `max_num_tokens` when unset. This " - "budget is shared proportionately across all modalities the model " - "encodes, not set per modality; per-modality knobs may be added " - "later."), + "Maximum number of encoder tokens. For encoder-decoder models, this " + "limits encoder CUDA graph total-token buckets. For multimodal " + "models, this is the shared AttentionMetadata budget across all " + "encoded modalities. Falls back to `max_num_tokens` when unset."), status="prototype") @field_validator("encoder_max_batch_size", "encoder_max_num_tokens") @@ -5058,6 +5077,28 @@ def validate_encoder_runtime_sizes(cls, v: Optional[int]) -> Optional[int]: raise ValueError("must be a positive integer when set") return v + @model_validator(mode="after") + def validate_encoder_cuda_graph_config(self) -> 'TorchLlmArgs': + if self.encoder_cuda_graph_config is None: + return self + if self.encode_only: + raise ValueError( + "Use cuda_graph_config=EncodeCudaGraphConfig(...) when " + "encode_only=True; encoder_cuda_graph_config is for " + "encoder-decoder models.") + if self.encoder_max_batch_size is None: + raise ValueError( + "encoder_cuda_graph_config requires encoder_max_batch_size.") + missing = [] + if not self.encoder_cuda_graph_config.num_tokens: + missing.append("num_tokens/max_num_token") + if not self.encoder_cuda_graph_config.seq_lens: + missing.append("seq_lens/max_seq_len") + if missing: + raise ValueError("encoder_cuda_graph_config requires " + f"{' and '.join(missing)}.") + return self + attn_backend: str = Field( default='TRTLLM', description="Attention backend to use.", diff --git a/tensorrt_llm/usage/llm_args_golden_manifest.json b/tensorrt_llm/usage/llm_args_golden_manifest.json index 60db88c0e67b..82592b65e8b0 100644 --- a/tensorrt_llm/usage/llm_args_golden_manifest.json +++ b/tensorrt_llm/usage/llm_args_golden_manifest.json @@ -407,6 +407,13 @@ "kind": "value", "path": "enable_early_first_token_response" }, + { + "allowed_values": [], + "annotation": "", + "converter": "", + "kind": "value", + "path": "enable_encoder_decoder_mixed_cuda_graph" + }, { "allowed_values": [], "annotation": "", @@ -484,6 +491,64 @@ "kind": "value", "path": "encode_only" }, + { + "allowed_values": [], + "annotation": "Optional[List[int]]", + "converter": "", + "kind": "value", + "path": "encoder_cuda_graph_config.batch_sizes" + }, + { + "allowed_values": [], + "annotation": "", + "converter": "", + "kind": "value", + "path": "encoder_cuda_graph_config.enable_padding" + }, + { + "allowed_values": [], + "annotation": "", + "converter": "", + "kind": "value", + "path": "encoder_cuda_graph_config.max_batch_size" + }, + { + "allowed_values": [], + "annotation": "", + "converter": "", + "kind": "value", + "path": "encoder_cuda_graph_config.max_num_token" + }, + { + "allowed_values": [], + "annotation": "", + "converter": "", + "kind": "value", + "path": "encoder_cuda_graph_config.max_seq_len" + }, + { + "allowed_values": [ + "encode" + ], + "annotation": "Literal['encode']", + "converter": "", + "kind": "categorical", + "path": "encoder_cuda_graph_config.mode" + }, + { + "allowed_values": [], + "annotation": "Optional[List[Annotated[int, Gt(gt=0)]]]", + "converter": "", + "kind": "value", + "path": "encoder_cuda_graph_config.num_tokens" + }, + { + "allowed_values": [], + "annotation": "Optional[List[Annotated[int, Gt(gt=0)]]]", + "converter": "", + "kind": "value", + "path": "encoder_cuda_graph_config.seq_lens" + }, { "allowed_values": [], "annotation": "Optional[int]", diff --git a/tests/integration/defs/kv_cache/test_final_single_token_context_cuda_graph.py b/tests/integration/defs/kv_cache/test_final_single_token_context_cuda_graph.py index 6b0643788b43..db1b6314aa9f 100644 --- a/tests/integration/defs/kv_cache/test_final_single_token_context_cuda_graph.py +++ b/tests/integration/defs/kv_cache/test_final_single_token_context_cuda_graph.py @@ -91,6 +91,7 @@ def maybe_get_cuda_graph( draft_tokens_cuda: torch.Tensor | None = None, new_tensors_device: SampleStateTensors | None = None, spec_resource_manager: BaseResourceManager | None = None, + allow_mixed_encoder_decoder: bool = False, promoted_context_request_ids: frozenset[int] = frozenset(), ) -> tuple[Any | None, Any | None, KeyType | None]: # A new decision means the preceding one reached eager execution if it @@ -104,7 +105,8 @@ def maybe_get_cuda_graph( draft_tokens_cuda, new_tensors_device, spec_resource_manager, - promoted_context_request_ids, + allow_mixed_encoder_decoder=allow_mixed_encoder_decoder, + promoted_context_request_ids=promoted_context_request_ids, ) if promoted_context_request_ids: execution = _PromotedContextGraphExecution( diff --git a/tests/integration/defs/llmapi/test_llm_api_pytorch_bart.py b/tests/integration/defs/llmapi/test_llm_api_pytorch_bart.py index 6f94d06b6efa..8185497178ed 100644 --- a/tests/integration/defs/llmapi/test_llm_api_pytorch_bart.py +++ b/tests/integration/defs/llmapi/test_llm_api_pytorch_bart.py @@ -13,6 +13,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import time from pathlib import Path import pytest @@ -21,6 +22,7 @@ from tensorrt_llm.llmapi import ( LLM, CudaGraphConfig, + EncodeCudaGraphConfig, KvCacheConfig, RequestOutput, SamplingParams, @@ -46,6 +48,7 @@ _MBART_TARGET_LANG = "en_XX" _MBART_SOURCE_TEXT = "Şeful ONU spune că nu există o soluţie militară în Siria." _MAX_NEW_TOKENS = 10 +_CONTINUOUS_ADMISSION_MAX_NEW_TOKENS = 16 _MAX_SEQUENCE_LENGTH = 128 _MAX_KV_TOKENS = 384 _MIN_GPU_MEMORY_MB = 16_000 @@ -303,6 +306,14 @@ def _decoder_cuda_graph_config( ) +class _SleepLogitsProcessor: + def __init__(self, delay_seconds: float) -> None: + self.delay_seconds = delay_seconds + + def __call__(self, req_id, logits, token_ids, stream_ptr, client_id) -> None: + time.sleep(self.delay_seconds) + + def _assert_bart_response( response: RequestOutput, num_return_sequences: int, @@ -593,3 +604,144 @@ def test_bart_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch( request_idx ], ) + + +@pytest.mark.parametrize( + "tensor_parallel_size", + [ + pytest.param(1, id="tp1"), + pytest.param(2, id="tp2", marks=pytest.mark.skip_less_device(2)), + ], +) +def test_bart_pytorch_continuous_admission_replays_encoder_and_mixed_cuda_graphs( + monkeypatch: pytest.MonkeyPatch, + tensor_parallel_size: int, +) -> None: + """Preserve mixed-length encoder outputs during continuous admission.""" + monkeypatch.setenv("TRTLLM_SKIP_KV_CACHE_ESTIMATION", "1") + if tensor_parallel_size == 1: + monkeypatch.setenv("TLLM_WORKER_USE_SINGLE_PROCESS", "1") + else: + monkeypatch.delenv("TLLM_WORKER_USE_SINGLE_PROCESS", raising=False) + + model_path = _get_model_path(_MODEL_NAME) + tokenizer = AutoTokenizer.from_pretrained(model_path) + first_sampling_params = SamplingParams( + max_tokens=_CONTINUOUS_ADMISSION_MAX_NEW_TOKENS, + temperature=0.0, + ignore_eos=True, + logits_processor=_SleepLogitsProcessor(delay_seconds=0.02), + ) + second_sampling_params = _sampling_params( + num_beams=1, + num_return_sequences=1, + ) + encoder_graph_config = EncodeCudaGraphConfig( + batch_sizes=[1, 2], + num_tokens=[64, 128], + seq_lens=[64], + enable_padding=True, + ) + + with LLM( + model_path, + backend="pytorch", + attn_backend="TRTLLM", + cuda_graph_config=_decoder_cuda_graph_config([3]), + encoder_cuda_graph_config=encoder_graph_config, + enable_encoder_decoder_mixed_cuda_graph=True, + disable_overlap_scheduler=False, + dtype="bfloat16", + enable_chunked_prefill=False, + encoder_max_batch_size=2, + encoder_max_num_tokens=128, + 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=False, + ), + max_batch_size=3, + max_beam_width=1, + max_input_len=_MAX_SEQUENCE_LENGTH, + max_num_tokens=_MAX_SEQUENCE_LENGTH, + max_seq_len=_MAX_SEQUENCE_LENGTH, + model_kwargs={"torch_dtype": "bfloat16"}, + scheduler_config=SchedulerConfig(use_python_scheduler=True), + batch_wait_timeout_iters=2, + tensor_parallel_size=tensor_parallel_size, + ) as llm: + encoder_replay_keys = [] + decoder_replay_keys = [] + if tensor_parallel_size == 1: + model_engine = llm._executor.engine.model_engine + encoder_runner = model_engine.encoder_cuda_graph_runner + decoder_runner = model_engine.cuda_graph_runner + + assert encoder_runner.enabled + assert encoder_runner.graphs + captured_mixed_keys = {key for key in decoder_runner.graphs if key[5] and key[6]} + assert captured_mixed_keys + + original_encoder_replay = encoder_runner.replay + original_decoder_replay = decoder_runner.replay + + def record_encoder_replay(key, inputs): + encoder_replay_keys.append(key) + return original_encoder_replay(key, inputs) + + def record_decoder_replay(key, inputs): + decoder_replay_keys.append(key) + return original_decoder_replay(key, inputs) + + monkeypatch.setattr(encoder_runner, "replay", record_encoder_replay) + monkeypatch.setattr(decoder_runner, "replay", record_decoder_replay) + + first_response = llm.generate_async( + _SOURCE_TEXT, + sampling_params=first_sampling_params, + streaming=True, + ) + first_stream_step = next(first_response) + assert not first_stream_step.finished + + encoder_replay_count_before_admission = len(encoder_replay_keys) + admitted_responses = [ + llm.generate_async( + source_text, + sampling_params=second_sampling_params, + streaming=False, + ) + for source_text in _MIXED_ENCODER_SOURCE_TEXTS + ] + + first_response.result() + for response in admitted_responses: + response.result() + + first_token_ids = _assert_bart_response( + first_response, + num_return_sequences=1, + max_tokens=_CONTINUOUS_ADMISSION_MAX_NEW_TOKENS, + ) + assert first_token_ids[0][:_MAX_NEW_TOKENS] == _EXPECTED_GREEDY_OUTPUT_TOKEN_IDS + + for request_idx, response in enumerate(admitted_responses): + token_ids = _assert_bart_response(response, num_return_sequences=1) + _assert_expected_generation( + tokenizer, + token_ids, + exact_match=True, + expected_token_ids_by_output=( + _MIXED_ENCODER_EXPECTED_TOKEN_IDS_BY_REQUEST[request_idx] + ), + ) + + if tensor_parallel_size == 1: + admitted_encoder_keys = encoder_replay_keys[encoder_replay_count_before_admission:] + assert any(key[0] == 2 for key in admitted_encoder_keys) + assert set(encoder_replay_keys) <= set(encoder_runner.graphs) + replayed_mixed_keys = {key for key in decoder_replay_keys if key[5] and key[6]} + assert replayed_mixed_keys + assert replayed_mixed_keys <= captured_mixed_keys diff --git a/tests/integration/defs/llmapi/test_llm_api_pytorch_t5.py b/tests/integration/defs/llmapi/test_llm_api_pytorch_t5.py index 492b88ff3d26..46f525565ec8 100644 --- a/tests/integration/defs/llmapi/test_llm_api_pytorch_t5.py +++ b/tests/integration/defs/llmapi/test_llm_api_pytorch_t5.py @@ -22,6 +22,7 @@ from tensorrt_llm.llmapi import ( LLM, CudaGraphConfig, + EncodeCudaGraphConfig, KvCacheConfig, RequestOutput, SamplingParams, @@ -819,13 +820,16 @@ def test_t5_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch( ) -def test_t5_pytorch_generate_encoder_decoder_mixed_context_generation_batch( +def test_t5_pytorch_continuous_admission_replays_encoder_and_mixed_cuda_graphs( monkeypatch: pytest.MonkeyPatch, ) -> None: + """Preserve mixed-length encoder outputs during continuous admission.""" monkeypatch.setenv("TRTLLM_SKIP_KV_CACHE_ESTIMATION", "1") + monkeypatch.setenv("TLLM_WORKER_USE_SINGLE_PROCESS", "1") model_name = "t5-small" model_path = _get_t5_model_path(model_name) + tokenizer = AutoTokenizer.from_pretrained(model_path) first_sampling_params = SamplingParams( max_tokens=_MIXED_CONTEXT_GENERATION_MAX_NEW_TOKENS, temperature=0.0, @@ -836,15 +840,25 @@ def test_t5_pytorch_generate_encoder_decoder_mixed_context_generation_batch( max_tokens=_MAX_NEW_TOKENS, temperature=0.0, ) + encoder_graph_config = EncodeCudaGraphConfig( + batch_sizes=[1, 2], + num_tokens=[64, 128], + seq_lens=[64], + enable_padding=True, + ) with LLM( model_path, backend="pytorch", attn_backend="TRTLLM", - cuda_graph_config=_decoder_cuda_graph_config([2]), - disable_overlap_scheduler=True, + cuda_graph_config=_decoder_cuda_graph_config([3]), + encoder_cuda_graph_config=encoder_graph_config, + enable_encoder_decoder_mixed_cuda_graph=True, + disable_overlap_scheduler=False, dtype="bfloat16", enable_chunked_prefill=False, + encoder_max_batch_size=2, + encoder_max_num_tokens=128, kv_cache_config=KvCacheConfig( enable_block_reuse=False, max_tokens=_MAX_KV_TOKENS, @@ -852,14 +866,44 @@ def test_t5_pytorch_generate_encoder_decoder_mixed_context_generation_batch( cross_kv_cache_fraction=_CROSS_KV_CACHE_FRACTION, use_kv_cache_manager_v2=False, ), - max_batch_size=2, + max_batch_size=3, max_beam_width=1, max_input_len=_MAX_SEQUENCE_LENGTH, max_num_tokens=_MAX_SEQUENCE_LENGTH, max_seq_len=_MAX_SEQUENCE_LENGTH, model_kwargs={"torch_dtype": "bfloat16"}, scheduler_config=SchedulerConfig(use_python_scheduler=True), + batch_wait_timeout_iters=2, ) as llm: + model_engine = llm._executor.engine.model_engine + encoder_runner = model_engine.encoder_cuda_graph_runner + decoder_runner = model_engine.cuda_graph_runner + + assert encoder_runner.enabled + assert encoder_runner.graphs + assert encoder_runner.use_fixed_sequence_slots + captured_mixed_keys = {key for key in decoder_runner.graphs if key[5] and key[6]} + assert captured_mixed_keys + + encoder_replay_keys = [] + decoder_replay_keys = [] + original_encoder_replay = encoder_runner.replay + original_decoder_replay = decoder_runner.replay + + def record_encoder_replay(key, inputs): + encoder_replay_keys.append(key) + return original_encoder_replay(key, inputs) + + def record_decoder_replay(key, inputs): + decoder_replay_keys.append(key) + return original_decoder_replay(key, inputs) + + monkeypatch.setattr(encoder_runner, "replay", record_encoder_replay) + monkeypatch.setattr(decoder_runner, "replay", record_decoder_replay) + + expected_token_ids_by_request = _MIXED_ENCODER_OUTPUT_TOKEN_IDS_BY_MODEL_AND_BEAMS[ + (model_name, 1) + ] first_response = llm.generate_async( _SOURCE_TEXT, sampling_params=first_sampling_params, @@ -868,18 +912,42 @@ def test_t5_pytorch_generate_encoder_decoder_mixed_context_generation_batch( first_stream_step = next(first_response) assert not first_stream_step.finished - second_response = llm.generate_async( - _MIXED_ENCODER_SOURCE_TEXTS[1], - sampling_params=second_sampling_params, - streaming=False, - ) + encoder_replay_count_before_admission = len(encoder_replay_keys) + admitted_responses = [ + llm.generate_async( + source_text, + sampling_params=second_sampling_params, + streaming=False, + ) + for source_text in _MIXED_ENCODER_SOURCE_TEXTS + ] first_response.result() - second_response.result() + for response in admitted_responses: + response.result() - _assert_t5_response( + first_token_ids = _assert_t5_response( first_response, num_return_sequences=1, max_tokens=_MIXED_CONTEXT_GENERATION_MAX_NEW_TOKENS, ) - _assert_t5_response(second_response, num_return_sequences=1) + assert first_token_ids[0][:_MAX_NEW_TOKENS] == expected_token_ids_by_request[0][0] + + for request_idx, response in enumerate(admitted_responses): + token_ids = _assert_t5_response(response, num_return_sequences=1) + _assert_expected_generation( + tokenizer, + token_ids, + exact_match=True, + expected_token_ids_by_output=expected_token_ids_by_request[request_idx], + expected_text_fragment=( + _MIXED_ENCODER_EXPECTED_TEXT_FRAGMENTS_BY_MODEL[model_name][request_idx] + ), + ) + + admitted_encoder_keys = encoder_replay_keys[encoder_replay_count_before_admission:] + assert any(key[0] == 2 for key in admitted_encoder_keys) + assert set(encoder_replay_keys) <= set(encoder_runner.graphs) + replayed_mixed_keys = {key for key in decoder_replay_keys if key[5] and key[6]} + assert replayed_mixed_keys + assert replayed_mixed_keys <= captured_mixed_keys diff --git a/tests/integration/test_lists/test-db/l0_dgx_h100.yml b/tests/integration/test_lists/test-db/l0_dgx_h100.yml index 3cda7d22ca40..33894c49c25b 100644 --- a/tests/integration/test_lists/test-db/l0_dgx_h100.yml +++ b/tests/integration/test_lists/test-db/l0_dgx_h100.yml @@ -25,6 +25,7 @@ l0_dgx_h100: # ------------- Encoder-decoder TP tests --------------- - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-on-greedy-tp2-t5-small] - llmapi/test_llm_api_pytorch_bart.py::test_bart_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-on-greedy-tp2-bart-large-cnn] + - llmapi/test_llm_api_pytorch_bart.py::test_bart_pytorch_continuous_admission_replays_encoder_and_mixed_cuda_graphs[tp2] - llmapi/test_llm_api_pytorch_whisper.py::test_whisper_pytorch_feature_combinations[fp32-kv-v1-graphs-off-greedy-tp2] # ------------- Disaggregated serving tests --------------- - accuracy/test_disaggregated_serving.py::TestLlama3_1_8BInstruct::test_eagle3[eagle3_one_model=True-overlap_scheduler=True] diff --git a/tests/integration/test_lists/test-db/l0_h100.yml b/tests/integration/test_lists/test-db/l0_h100.yml index ac3903686689..187179d76220 100644 --- a/tests/integration/test_lists/test-db/l0_h100.yml +++ b/tests/integration/test_lists/test-db/l0_h100.yml @@ -334,6 +334,7 @@ l0_h100: - llmapi/test_llm_api_pytorch_bart.py::test_bart_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-on-beam2-overlap-bart-large-cnn] - llmapi/test_llm_api_pytorch_bart.py::test_bart_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v2-cuda-graph-on-greedy-overlap-bart-large-cnn] - llmapi/test_llm_api_pytorch_bart.py::test_bart_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch[bf16-kv-v2-decoder-cuda-graph-on-greedy-batch2-bart-large-cnn] + - llmapi/test_llm_api_pytorch_bart.py::test_bart_pytorch_continuous_admission_replays_encoder_and_mixed_cuda_graphs[tp1] - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-on-beam2-t5-base] - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-on-beam2-t5-large] - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-on-beam2-flan-t5-base] @@ -355,6 +356,7 @@ l0_h100: - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch[bf16-kv-v1-decoder-cuda-graph-on-beam2-batch2-t5-small] - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch[bf16-kv-v1-decoder-cuda-graph-on-beam2-batch2-flan-t5-small] - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch[bf16-kv-v2-decoder-cuda-graph-on-greedy-batch2-t5-small] + - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_continuous_admission_replays_encoder_and_mixed_cuda_graphs - llmapi/test_llm_api_pytorch_whisper.py::test_whisper_pytorch_beam_search[fp32-kv-v1-graphs-off-beam2] - llmapi/test_llm_api_pytorch_whisper.py::test_whisper_pytorch_feature_combinations[fp32-kv-v2-graphs-off-greedy] - llmapi/test_llm_api_pytorch_whisper.py::test_whisper_pytorch_feature_combinations[fp32-kv-v1-graphs-requested-greedy] diff --git a/tests/integration/test_lists/test-db/l0_l40s.yml b/tests/integration/test_lists/test-db/l0_l40s.yml index 18fa6b86be18..1bba99fa161f 100644 --- a/tests/integration/test_lists/test-db/l0_l40s.yml +++ b/tests/integration/test_lists/test-db/l0_l40s.yml @@ -53,7 +53,6 @@ l0_l40s: - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v2-cuda-graph-on-greedy-t5-small] - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_end_to_end[bf16-kv-v1-cuda-graph-on-greedy-overlap-t5-small] - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_mixed_encoder_lengths_batch[bf16-kv-v1-decoder-cuda-graph-on-greedy-batch2-t5-small] - - llmapi/test_llm_api_pytorch_t5.py::test_t5_pytorch_generate_encoder_decoder_mixed_context_generation_batch # Whisper (encoder-decoder) — customer-side deployment targets L40S/H200 - llmapi/test_llm_api_pytorch_whisper.py::test_whisper_pytorch_transcribe_end_to_end - llmapi/test_llm_api_pytorch_whisper.py::test_whisper_pytorch_feature_combinations[bf16-kv-v2-decoder-graphs-on-greedy] diff --git a/tests/unittest/_torch/executor/test_py_executor.py b/tests/unittest/_torch/executor/test_py_executor.py index 50f33d321c2f..ea13fdc9935e 100644 --- a/tests/unittest/_torch/executor/test_py_executor.py +++ b/tests/unittest/_torch/executor/test_py_executor.py @@ -16,9 +16,10 @@ import threading import time import types -from unittest.mock import MagicMock, Mock +from unittest.mock import MagicMock, Mock, patch import pytest +import torch from tensorrt_llm._torch.distributed.communicator import ReduceOp from tensorrt_llm._torch.pyexecutor.executor_request_queue import ( @@ -26,7 +27,11 @@ RequestQueueItem, ) from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest, LlmRequestState, SamplingConfig -from tensorrt_llm._torch.pyexecutor.py_executor import DisaggTransferAdmissionController, PyExecutor +from tensorrt_llm._torch.pyexecutor.py_executor import ( + DisaggTransferAdmissionController, + EncoderStepResult, + PyExecutor, +) from tensorrt_llm._torch.pyexecutor.resource_manager import NoFreeSlotsError, ResourceManagerType from tensorrt_llm._torch.pyexecutor.scheduler import ( FCFSWaitingQueue, @@ -35,6 +40,20 @@ ) +class _InflightRequestIds: + def __init__(self): + self.ids = set() + + def insert(self, request_id): + self.ids.add(request_id) + + def erase(self, request_id): + self.ids.discard(request_id) + + def __contains__(self, request_id): + return request_id in self.ids + + class MockPyExecutor: """A mock PyExecutor class for testing request handling logic. @@ -117,6 +136,263 @@ def mock_dist(): return mock_dist +def _make_async_encoder_executor(future): + executor = object.__new__(PyExecutor) + executor.dist = types.SimpleNamespace(tp_size=1) + executor.encoder_launch_executor = Mock() + executor.encoder_launch_executor.submit.return_value = future + executor.pending_encoder_steps = [] + executor.inflight_req_ids = _InflightRequestIds() + executor._run_encoder_step_unchecked = Mock() + executor._publish_encoder_step = Mock() + executor._handle_errors = Mock() + return executor + + +def _make_encoder_batch_wait_executor(batch_sizes=None, encoder_max_batch_size=8): + executor = object.__new__(PyExecutor) + executor.max_batch_size = 32 + batch_sizes = batch_sizes or [1, 2, 4, 8] + executor.llm_args = types.SimpleNamespace( + encoder_cuda_graph_config=types.SimpleNamespace( + batch_sizes=batch_sizes, + enable_padding=True, + num_tokens=[96], + seq_lens=[512], + ), + encoder_max_batch_size=encoder_max_batch_size, + ) + executor.batch_wait_timeout_iters = 48 + executor.encoder_batch_wait_iters_count = 0 + return executor + + +def _make_encoder_fallback_batch_wait_executor(): + executor = object.__new__(PyExecutor) + executor.llm_args = types.SimpleNamespace( + encoder_cuda_graph_config=None, + encoder_max_batch_size=None, + ) + executor.batch_wait_timeout_iters = 48 + executor.encoder_batch_wait_iters_count = 0 + executor.batch_wait_max_tokens_ratio = 0.5 + executor.max_num_tokens = 32 + executor.active_requests = [] + executor.inflight_req_ids = _InflightRequestIds() + return executor + + +def _make_encoder_request(request_id): + return types.SimpleNamespace( + request_id=request_id, + state=LlmRequestState.ENCODER_INIT, + encoder_output_len=4, + ) + + +def test_encoder_graph_warmup_uses_runtime_encoder_stream(): + executor = object.__new__(PyExecutor) + executor.device_id = 3 + executor.encoder_stream = Mock() + executor.resource_manager = object() + executor.model_engine = Mock() + stream_context = MagicMock() + + with ( + patch("torch.cuda.set_device") as set_device, + patch("torch.cuda.stream", return_value=stream_context) as cuda_stream, + ): + executor._warmup_encoder_cuda_graphs_enc_dec() + + set_device.assert_called_once_with(3) + cuda_stream.assert_called_once_with(executor.encoder_stream) + executor.model_engine._warmup_encoder_cuda_graphs_enc_dec.assert_called_once_with( + executor.resource_manager + ) + + +def test_encoder_microbatch_graph_admission_boundaries(): + executor = _make_encoder_batch_wait_executor() + encoder_requests = [object()] * 7 + scheduled = executor._waiting_encoder_requests( + encoder_requests, + [], + [object()] * 24, + ) + assert scheduled == [] + assert executor.encoder_batch_wait_iters_count == 1 + + executor = _make_encoder_batch_wait_executor() + encoder_requests = [object() for _ in range(12)] + scheduled = executor._waiting_encoder_requests( + encoder_requests, + [], + [object()] * 20, + ) + assert scheduled == encoder_requests[:8] + assert executor.encoder_batch_wait_iters_count == 0 + + executor = _make_encoder_batch_wait_executor( + batch_sizes=[1, 3, 6], + encoder_max_batch_size=8, + ) + executor.encoder_batch_wait_iters_count = executor.batch_wait_timeout_iters + encoder_requests = [object() for _ in range(5)] + scheduled = executor._waiting_encoder_requests( + encoder_requests, + [], + [], + ) + assert scheduled == encoder_requests[:3] + assert executor.encoder_batch_wait_iters_count == 0 + + executor = _make_encoder_batch_wait_executor() + executor.encoder_batch_wait_iters_count = executor.batch_wait_timeout_iters + encoder_requests = [object() for _ in range(8)] + scheduled = executor._waiting_encoder_requests( + encoder_requests, + [], + [object() for _ in range(25)], + ) + assert scheduled == [] + assert executor.encoder_batch_wait_iters_count == executor.batch_wait_timeout_iters + 1 + + scheduled = executor._waiting_encoder_requests( + encoder_requests, + [], + [object() for _ in range(24)], + ) + + assert scheduled == encoder_requests + assert executor.encoder_batch_wait_iters_count == 0 + + +def test_encoder_fallback_distinguishes_inflight_encoder_and_decoder_work(): + executor = _make_encoder_fallback_batch_wait_executor() + encoder_requests = [_make_encoder_request(1)] + inflight_encoder_request = _make_encoder_request(2) + executor.active_requests.append(inflight_encoder_request) + executor.inflight_req_ids.insert(inflight_encoder_request.request_id) + + scheduled = executor._waiting_encoder_requests(encoder_requests, [], []) + assert scheduled == encoder_requests + assert executor.encoder_batch_wait_iters_count == 0 + + decoder_request = types.SimpleNamespace( + request_id=3, + state=LlmRequestState.GENERATION_IN_PROGRESS, + ) + executor.active_requests = [decoder_request] + executor.inflight_req_ids.erase(inflight_encoder_request.request_id) + executor.inflight_req_ids.insert(decoder_request.request_id) + + scheduled = executor._waiting_encoder_requests(encoder_requests, [], []) + assert scheduled == [] + assert executor.encoder_batch_wait_iters_count == 1 + + +def test_async_encoder_step_lifecycle(): + ready_event = Mock() + ready_event.query.side_effect = [False, True] + result = EncoderStepResult( + hidden_states=torch.arange(12).reshape(6, 2), + sequence_lengths=[2, 4], + ready_event=ready_event, + ) + future = Mock() + future.done.side_effect = [False, True] + future.result.return_value = result + executor = _make_async_encoder_executor(future) + active_request = types.SimpleNamespace( + request_id=11, + state=LlmRequestState.ENCODER_INIT, + ) + completed_request = types.SimpleNamespace( + request_id=12, + state=LlmRequestState.GENERATION_COMPLETE, + ) + requests = [active_request, completed_request] + executor._publish_encoder_step.side_effect = ( + lambda encoder_requests, encoder_result: PyExecutor._publish_encoder_step( + executor, + encoder_requests, + encoder_result, + ) + ) + + executor._submit_encoder_step(requests) + executor._poll_encoder_steps() + + future.result.assert_not_called() + executor._publish_encoder_step.assert_not_called() + assert executor.inflight_req_ids.ids == {11, 12} + assert len(executor.pending_encoder_steps) == 1 + executor.encoder_launch_executor.submit.assert_called_once_with( + executor._run_encoder_step_unchecked, + requests, + ) + + executor._poll_encoder_steps() + future.result.assert_called_once_with() + ready_event.query.assert_called_once_with() + executor._publish_encoder_step.assert_not_called() + assert executor.inflight_req_ids.ids == {11, 12} + assert len(executor.pending_encoder_steps) == 1 + + executor._poll_encoder_steps() + future.result.assert_called_once_with() + assert ready_event.query.call_count == 2 + executor._publish_encoder_step.assert_called_once_with(requests, result) + assert executor.inflight_req_ids.ids == set() + assert executor.pending_encoder_steps == [] + assert active_request.state == LlmRequestState.CONTEXT_INIT + assert active_request.py_encoder_output_ready_event is ready_event + assert torch.equal(active_request.py_encoder_output, result.hidden_states[:2]) + assert completed_request.state == LlmRequestState.GENERATION_COMPLETE + assert not hasattr(completed_request, "py_encoder_output") + + executor.execution_stream = Mock() + encoder_output = Mock() + active_request.py_encoder_output = encoder_output + scheduled_requests = types.SimpleNamespace(context_requests=[active_request]) + executor._attach_encoder_output_to_execution_stream(scheduled_requests) + + executor.execution_stream.wait_event.assert_not_called() + encoder_output.record_stream.assert_called_once_with(executor.execution_stream) + assert active_request.py_encoder_output_ready_event is None + + +def test_tp_encoder_step_synchronizes_and_publishes_inline(): + call_order = [] + execution_stream = Mock() + encoder_stream = Mock() + encoder_stream.wait_stream.side_effect = lambda stream: call_order.append("wait_stream") + ready_event = Mock() + ready_event.synchronize.side_effect = lambda: call_order.append("synchronize") + result = EncoderStepResult( + hidden_states=torch.empty((1, 2)), + sequence_lengths=[1], + ready_event=ready_event, + ) + future = Mock() + future.result.side_effect = lambda: (call_order.append("result"), result)[1] + executor = _make_async_encoder_executor(future) + executor.dist.tp_size = 2 + executor.execution_stream = execution_stream + executor.encoder_stream = encoder_stream + executor._publish_encoder_step.side_effect = lambda requests, encoder_result: call_order.append( + "publish" + ) + request = types.SimpleNamespace(request_id=13, state=LlmRequestState.ENCODER_INIT) + + executor._submit_encoder_step([request]) + + assert call_order == ["wait_stream", "result", "synchronize", "publish"] + encoder_stream.wait_stream.assert_called_once_with(execution_stream) + assert executor.inflight_req_ids.ids == set() + assert executor.pending_encoder_steps == [] + + @pytest.fixture def mock_executor(mock_dist): """Create a MockPyExecutor instance for testing.""" diff --git a/tests/unittest/_torch/executor/test_pytorch_model_engine.py b/tests/unittest/_torch/executor/test_pytorch_model_engine.py index 884fb8e2eebb..efab11083058 100644 --- a/tests/unittest/_torch/executor/test_pytorch_model_engine.py +++ b/tests/unittest/_torch/executor/test_pytorch_model_engine.py @@ -16,7 +16,7 @@ from tensorrt_llm._torch.pyexecutor.connectors.kv_cache_connector import \ KvCacheConnectorWorker from tensorrt_llm._torch.pyexecutor.cuda_graph_runner import ( - CUDAGraphRunner, _restore_spec_decode_capture_state, + CUDAGraphRunner, EncoderCUDAGraphRunner, _restore_spec_decode_capture_state, _save_spec_decode_capture_state) from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest from tensorrt_llm._torch.pyexecutor.model_engine import ( @@ -188,7 +188,8 @@ def _make_request_stub(req_id: int, prompt_len: int = 4) -> SimpleNamespace: def _make_forward_only_engine( - graph_key: tuple[int, int, bool, bool, bool] | None, + graph_key: tuple[int, int, bool, bool, bool, tuple[int, ...], + tuple[int, ...]] | None, runner_enabled: bool = True, ) -> tuple[PyTorchModelEngine, Mock, Mock, Mock, dict[str, object]]: engine = object.__new__(PyTorchModelEngine) @@ -610,7 +611,7 @@ def test_graph_key_forwards_promoted_context_ids(self) -> None: runner._get_seq_len_mode.assert_called_once_with( batch, None, promoted_ids) - self.assertEqual(key, (1, 0, False, True, True)) + self.assertEqual(key, (1, 0, False, True, True, (), (0, ))) def test_graph_lookup_forwards_promoted_context_ids(self) -> None: runner = Mock() @@ -619,7 +620,7 @@ def test_graph_lookup_forwards_promoted_context_ids(self) -> None: enable_attention_dp=False, use_mrope=False, ) - key = (1, 0, False, True, True) + key = (1, 0, False, True, True, (), (0, )) graph_attn_metadata = object() graph_spec_metadata = object() runner.get_graph_key.return_value = key @@ -630,6 +631,7 @@ def test_graph_lookup_forwards_promoted_context_ids(self) -> None: "spec_metadata": graph_spec_metadata, } } + runner._is_mixed_encoder_decoder_batch.return_value = False request = _make_request_stub(7) batch = ScheduledRequests() batch.generation_requests = [request] @@ -652,7 +654,7 @@ def test_graph_lookup_forwards_promoted_context_ids(self) -> None: (graph_attn_metadata, graph_spec_metadata, key)) def test_forward_commits_candidate_only_on_graph_hit(self) -> None: - key = (2, 0, False, False, True) + key = (2, 0, False, False, True, (), (0, )) engine, runner, resource_manager, _, outputs = \ _make_forward_only_engine(key) context = _make_request_stub(1) @@ -711,7 +713,7 @@ def test_forward_graph_miss_uses_semantic_eager_batch(self) -> None: def test_zero_runtime_draft_speculation_commits_graph_candidate( self) -> None: - key = (2, 0, False, False, True) + key = (2, 0, False, False, True, (), (0, )) engine, runner, resource_manager, semantic_attn_metadata, outputs = \ _make_forward_only_engine(key) engine.enable_spec_decode = True @@ -807,7 +809,7 @@ def test_zero_runtime_non_linear_tree_speculation_uses_semantic_eager_batch( runner.replay.assert_not_called() def test_forward_allows_guided_context_logits_on_graph_hit(self) -> None: - key = (1, 0, False, False, True) + key = (1, 0, False, False, True, (), (0, )) engine, runner, resource_manager, _, outputs = \ _make_forward_only_engine(key) engine.guided_decoder = Mock() @@ -860,7 +862,7 @@ def test_multimodal_graph_miss_preserves_semantic_payload(self) -> None: self.assertIn("multimodal_embedding", multimodal_data) def test_generation_only_forward_does_not_call_new_selector(self) -> None: - key = (1, 0, False, False, True) + key = (1, 0, False, False, True, (), (0, )) engine, runner, resource_manager, _, _ = _make_forward_only_engine(key) generation = _make_request_stub(2) batch = ScheduledRequests() @@ -932,6 +934,70 @@ def test_global_incompatibilities_bypass_candidate_selection(self) -> None: class PyTorchModelEngineTestCase(unittest.TestCase): + def test_encoder_cuda_graph_stages_and_restores_fixed_sequence_slots( + self) -> None: + runner = EncoderCUDAGraphRunner.__new__(EncoderCUDAGraphRunner) + runner.is_encoder_decoder = True + runner.use_fixed_sequence_slots = True + runner.supported_batch_sizes = [2] + runner.supported_seq_lens = [512] + runner.max_supported_num_tokens = 1024 + small_key = (2, 512, 512) + compatible_key = (2, 1024, 512) + runner._capture_sequence_lengths = { + small_key: [511, 1], + compatible_key: [512, 512], + } + runner._capture_keys_by_batch_size = { + 2: [small_key, compatible_key], + } + runner._arange_max = torch.arange(1024, dtype=torch.int32) + + self.assertEqual( + runner._get_dynamic_capture_key([200, 300], + allow_batch_padding=False), + compatible_key) + + source_sequence_lengths = [1, 400] + key = runner._get_dynamic_capture_key(source_sequence_lengths, + allow_batch_padding=False) + self.assertEqual(key, small_key) + self.assertEqual(runner._get_capture_sequence_offsets(key), + [0, 511, 512]) + + input_ids = torch.arange(401, dtype=torch.int32) + inputs = runner.prepare_encoder_decoder_inputs( + { + "input_ids": input_ids, + "position_ids": input_ids, + "seq_lens": source_sequence_lengths, + }, + key, + source_sequence_lengths, + ) + self.assertEqual(inputs["seq_lens"], [400, 1]) + self.assertEqual(inputs["_encoder_source_to_slot"], [1, 0]) + + static_tensors = { + "input_ids": torch.empty(512, dtype=torch.int32), + "position_ids": torch.empty((1, 512), dtype=torch.int32), + } + runner._stage_encoder_decoder_inputs(key, inputs, static_tensors) + expected_staged_ids = torch.zeros(512, dtype=torch.int32) + expected_staged_ids[:400] = input_ids[1:] + expected_staged_ids[511] = input_ids[0] + torch.testing.assert_close(static_tensors["input_ids"], + expected_staged_ids) + torch.testing.assert_close(static_tensors["position_ids"][0], + expected_staged_ids) + + fixed_slot_output = torch.arange(512).unsqueeze(1) + restored_output = runner.restore_encoder_decoder_output( + key, fixed_slot_output, inputs) + expected_output = torch.cat( + (fixed_slot_output[511:512], fixed_slot_output[:400])) + torch.testing.assert_close(restored_output, expected_output) + def test_prepare_multimodal_indices_uses_mixin_token_ids(self) -> None: engine = object.__new__(PyTorchModelEngine) engine.model = DummyMultimodalIndexModel() diff --git a/tests/unittest/_torch/executor/test_pytorch_model_engine_warmup.py b/tests/unittest/_torch/executor/test_pytorch_model_engine_warmup.py index 9724ad3d0a5e..8d95fb9f1a53 100644 --- a/tests/unittest/_torch/executor/test_pytorch_model_engine_warmup.py +++ b/tests/unittest/_torch/executor/test_pytorch_model_engine_warmup.py @@ -1,3 +1,6 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + """Unit tests for warmup-cleanup behavior in PyTorchModelEngine.warmup(). Locks in that gc.collect() + torch.cuda.empty_cache() fire immediately after @@ -10,6 +13,7 @@ import contextlib import unittest from dataclasses import dataclass +from types import SimpleNamespace from unittest.mock import patch import torch @@ -174,6 +178,48 @@ def _record(*msg): class TestWarmupCleanup(unittest.TestCase): """Lock in warmup-cleanup behavior introduced by PR #14609 (Plan B).""" + def test_encoder_decoder_encoder_warmup_is_deferred_and_uses_two_passes(self): + model_engine = object.__new__(PyTorchModelEngine) + model_engine.cuda_graph_runner = SimpleNamespace( + enabled=True, + is_warmup_only=True, + ) + model_engine._torch_compile_piecewise_cuda_graph = False + model_engine.is_warmup = False + + @contextlib.contextmanager + def allow_capture(): + yield + + runner = SimpleNamespace( + enabled=True, + is_encoder_decoder=True, + is_warmup_only=False, + allow_capture=allow_capture, + ) + model_engine.encoder_cuda_graph_runner = runner + resource_manager = object() + warmup_states = [] + + with ( + patch.object(model_engine, "_capture_generation_cuda_graphs") as generation, + patch.object(model_engine, "_capture_mixed_encoder_decoder_cuda_graphs") as mixed, + patch.object( + model_engine, + "_capture_encoder_cuda_graphs_enc_dec", + side_effect=lambda _: warmup_states.append(runner.is_warmup_only), + ) as encoder, + ): + model_engine._run_cuda_graph_warmup(resource_manager) + generation.assert_called_once_with(resource_manager) + mixed.assert_called_once_with(resource_manager) + encoder.assert_not_called() + model_engine._warmup_encoder_cuda_graphs_enc_dec(resource_manager) + + assert encoder.call_count == 2 + assert warmup_states == [True, False] + assert not runner.is_warmup_only + def test_empty_cache_fires_immediately_after_autotuner(self): """Change 1 placement: empty_cache must be the call right after _run_autotuner_warmup.""" diff --git a/tests/unittest/_torch/sampler/test_torch_sampler.py b/tests/unittest/_torch/sampler/test_torch_sampler.py index f8e520061f79..e9e720335500 100644 --- a/tests/unittest/_torch/sampler/test_torch_sampler.py +++ b/tests/unittest/_torch/sampler/test_torch_sampler.py @@ -42,6 +42,8 @@ get_draft_token_length, ) from tensorrt_llm._torch.pyexecutor.sampler import ( + SampleStateTensorsHostTorch, + SampleStateTorch, TorchSampler, _BatchedSamplingResult, _request_get_sampling_params, @@ -747,12 +749,105 @@ def _uut(res=res): run_test_with_warmup(_test_runner, max_sync_s=0.3) +@force_ampere +def test_greedy_no_repeat_ngram_uses_token_ban_path(): + sampler = TorchSampler( + TorchSampler.Args( + max_seq_len=16, + max_draft_len=0, + max_num_sequences=1, + max_beam_width=1, + max_total_draft_tokens=0, + disable_overlap_scheduler=True, + ) + ) + request = LlmRequest( + request_id=0, + max_new_tokens=4, + input_tokens=[1, 2, 1], + sampling_config=SamplingConfig(), + seq_slot=0, + is_streaming=False, + ) + request.py_no_repeat_ngram_size = 2 + scheduled_requests = ScheduledRequests() + scheduled_requests.generation_requests = [request] + logits = torch.tensor([[0.0, 0.0, 10.0, 9.0]], device="cuda") + + *_, new_tokens_host, single_step_greedy = sampler._process_requests( + scheduled_requests, + {"logits": logits}, + sampler.store.new_tokens, + [0], + ) + torch.cuda.synchronize() + + assert not single_step_greedy + assert new_tokens_host.reshape(-1)[0].item() == 3 + + class TestFinishReasons: NOT_FINISHED = FinishReason.NOT_FINISHED STOP_WORDS = FinishReason.STOP_WORDS END_ID = FinishReason.END_ID LENGTH = FinishReason.LENGTH + def test_single_step_greedy_updates_finish_reasons_and_filters_completed_requests(self): + sampler = object.__new__(TorchSampler) + sampler.max_seq_len = 20 + sampler._track_pending_steps = False + requests = [ + LlmRequest( + request_id=0, + seq_slot=0, + input_tokens=[2, 0], + max_new_tokens=1, + end_id=2, + sampling_config=SamplingConfig(), + is_streaming=False, + ), + LlmRequest( + request_id=1, + seq_slot=1, + input_tokens=[2, 0], + max_new_tokens=1, + end_id=2, + sampling_config=SamplingConfig(), + is_streaming=False, + ), + LlmRequest( + request_id=2, + seq_slot=2, + input_tokens=[2, 0], + max_new_tokens=10, + end_id=2, + sampling_config=SamplingConfig(), + is_streaming=False, + ), + ] + requests[2].finish_by(FinishReason.LENGTH, 0) + new_tokens = torch.tensor([2, 7, 99], dtype=torch.int32) + state = SampleStateTorch( + requests=requests, + device=None, + host=SampleStateTensorsHostTorch( + new_tokens=new_tokens, + finish_reasons=None, + first_finish_reasons=None, + ), + single_step_greedy=True, + ) + + sampler.update_requests(state) + + assert all(request.is_finished for request in requests) + # The first request reaches EOS and length together; EOS takes precedence. + assert not requests[0].is_finished_due_to_length + assert requests[1].is_finished_due_to_length + assert requests[0].get_tokens(0)[-1] == 2 + assert requests[1].get_tokens(0)[-1] == 7 + assert requests[2].get_tokens(0) == [2, 0] + class RequestCase: MAX_NEW_TOKENS = 10 MAX_NUM_SEQUENCES = 128 diff --git a/tests/unittest/api_stability/references/llm.yaml b/tests/unittest/api_stability/references/llm.yaml index e30864532cb3..9351334eae99 100644 --- a/tests/unittest/api_stability/references/llm.yaml +++ b/tests/unittest/api_stability/references/llm.yaml @@ -95,6 +95,14 @@ methods: annotation: Union[tensorrt_llm.llmapi.llm_args.DecodeCudaGraphConfig, tensorrt_llm.llmapi.llm_args.EncodeCudaGraphConfig, NoneType] default: null status: beta + encoder_cuda_graph_config: + annotation: Optional[tensorrt_llm.llmapi.llm_args.EncodeCudaGraphConfig] + default: null + status: beta + enable_encoder_decoder_mixed_cuda_graph: + annotation: bool + default: True + status: beta multimodal_config: annotation: tensorrt_llm.llmapi.llm_args.MultimodalConfig default: null diff --git a/tests/unittest/llmapi/test_llm_args.py b/tests/unittest/llmapi/test_llm_args.py index 5b417fab2b6f..2f3982acf422 100644 --- a/tests/unittest/llmapi/test_llm_args.py +++ b/tests/unittest/llmapi/test_llm_args.py @@ -1746,6 +1746,73 @@ def test_cuda_graph_config_accepts_encoder_config(self): assert args.cuda_graph_config.seq_lens == [8, 32] assert args.cuda_graph_config.max_seq_len == 32 + def test_encoder_decoder_cuda_graph_user_interface(self): + encoder_config = EncodeCudaGraphConfig( + batch_sizes=[1, 4], + num_tokens=[16, 64], + seq_lens=[8, 32], + enable_padding=True, + ) + args = TorchLlmArgs( + model=llama_model_path, + encoder_max_batch_size=4, + cuda_graph_config=DecodeCudaGraphConfig( + batch_sizes=[1, 4], + enable_padding=True, + ), + encoder_cuda_graph_config=encoder_config, + ) + + assert isinstance(args.cuda_graph_config, DecodeCudaGraphConfig) + assert isinstance(args.encoder_cuda_graph_config, EncodeCudaGraphConfig) + assert args.encoder_cuda_graph_config.batch_sizes == [1, 4] + assert args.enable_encoder_decoder_mixed_cuda_graph + + disabled_args = TorchLlmArgs( + model=llama_model_path, + encoder_max_batch_size=4, + encoder_cuda_graph_config=encoder_config, + enable_encoder_decoder_mixed_cuda_graph=False, + ) + + assert not disabled_args.enable_encoder_decoder_mixed_cuda_graph + + def test_encoder_cuda_graph_config_validation(self): + invalid_cases = [ + ( + { + "encoder_cuda_graph_config": + EncodeCudaGraphConfig( + batch_sizes=[1, 4], + num_tokens=[16, 64], + seq_lens=[8, 32], + enable_padding=True, + ), + }, + "encoder_cuda_graph_config requires encoder_max_batch_size", + ), + ( + { + "encoder_max_batch_size": + 4, + "encoder_cuda_graph_config": + EncodeCudaGraphConfig( + batch_sizes=[1, 4], + enable_padding=True, + ), + }, + ("encoder_cuda_graph_config requires " + "num_tokens/max_num_token and seq_lens/max_seq_len"), + ), + ] + + for kwargs, error_match in invalid_cases: + with pytest.raises(ValidationError, match=error_match): + TorchLlmArgs( + model=llama_model_path, + **kwargs, + ) + def test_cuda_graph_config_infers_encode_mode_from_raw_dict(self): args = TorchLlmArgs( model=llama_model_path,