diff --git a/tensorrt_llm/_torch/speculative/eagle3.py b/tensorrt_llm/_torch/speculative/eagle3.py index 68dac4c5e9f0..29a88b10f283 100644 --- a/tensorrt_llm/_torch/speculative/eagle3.py +++ b/tensorrt_llm/_torch/speculative/eagle3.py @@ -1159,7 +1159,10 @@ def sample_and_accept_draft_tokens( self.spec_config.end_thinking_phase_token) num_accepted_tokens = self._apply_force_accepted_tokens( - num_accepted_tokens, num_contexts, runtime_draft_len) + num_accepted_tokens, + num_contexts, + runtime_draft_len, + spec_metadata=spec_metadata) return accepted_tokens, num_accepted_tokens diff --git a/tensorrt_llm/_torch/speculative/eagle3_dynamic_tree.py b/tensorrt_llm/_torch/speculative/eagle3_dynamic_tree.py index 446171e6ba69..0e09c0f146d4 100644 --- a/tensorrt_llm/_torch/speculative/eagle3_dynamic_tree.py +++ b/tensorrt_llm/_torch/speculative/eagle3_dynamic_tree.py @@ -842,7 +842,10 @@ def _sample_and_accept_dynamic_tree( ) num_accepted_tokens = self._apply_force_accepted_tokens( - num_accepted_tokens, num_contexts, self.max_draft_len + num_accepted_tokens, + num_contexts, + self.max_draft_len, + spec_metadata=spec_metadata, ) return accepted_tokens, num_accepted_tokens @@ -968,7 +971,10 @@ def _sample_and_accept_dynamic_tree_rejection( ) num_accepted_tokens = self._apply_force_accepted_tokens( - num_accepted_tokens, num_contexts, self.max_draft_len + num_accepted_tokens, + num_contexts, + self.max_draft_len, + spec_metadata=spec_metadata, ) return accepted_tokens, num_accepted_tokens diff --git a/tensorrt_llm/_torch/speculative/interface.py b/tensorrt_llm/_torch/speculative/interface.py index 62e56ffc8446..9aeab730c6ef 100644 --- a/tensorrt_llm/_torch/speculative/interface.py +++ b/tensorrt_llm/_torch/speculative/interface.py @@ -1005,8 +1005,11 @@ def _ensure_force_accept_rng_state(self, device: torch.device) -> None: dtype=torch.int64, device=device) - def _apply_force_accepted_tokens(self, num_accepted_tokens, num_contexts, - runtime_draft_len: int): + def _apply_force_accepted_tokens(self, + num_accepted_tokens, + num_contexts, + runtime_draft_len: int, + spec_metadata=None): """ Apply a forced (synthetic) number of accepted draft tokens if the ``TLLM_SPEC_DECODE_FORCE_NUM_ACCEPTED_TOKENS`` environment variable is @@ -1033,6 +1036,11 @@ def _apply_force_accepted_tokens(self, num_accepted_tokens, num_contexts, accepted counts (target token + accepted draft tokens). num_contexts: Number of context (prefill) requests in the batch. runtime_draft_len: The draft length for the current iteration. + spec_metadata: Optional SpecMetadata. When provided, used to + detect eager CUDA-graph warmup so the override is skipped + there — warmup batches use dummy requests whose KV cache and + draft buffers are not populated for an inflated accepted + count, which would drive downstream MTP ops out-of-bounds. Returns: Modified num_accepted_tokens tensor. @@ -1040,6 +1048,12 @@ def _apply_force_accepted_tokens(self, num_accepted_tokens, num_contexts, if self.force_num_accepted_tokens == 0.0: return num_accepted_tokens + if spec_metadata is not None: + is_warmup = (spec_metadata.is_cuda_graph + and not torch.cuda.is_current_stream_capturing()) + if is_warmup: + return num_accepted_tokens + # Decompose into a deterministic integer part (always accepted) and a # probabilistic fractional part. ``int(...)`` truncates toward zero, # which matches floor for the supported non-negative range. @@ -1151,7 +1165,10 @@ def _sample_and_accept_draft_tokens_base( # Apply force override if set num_accepted_tokens = self._apply_force_accepted_tokens( - num_accepted_tokens, num_contexts, runtime_draft_len) + num_accepted_tokens, + num_contexts, + runtime_draft_len, + spec_metadata=spec_metadata) return accepted_tokens, num_accepted_tokens @@ -1336,7 +1353,10 @@ def _sample_and_accept_draft_tokens_rejection( num_accepted_tokens[num_contexts:] = gen_num_accepted num_accepted_tokens = self._apply_force_accepted_tokens( - num_accepted_tokens, num_contexts, runtime_draft_len) + num_accepted_tokens, + num_contexts, + runtime_draft_len, + spec_metadata=spec_metadata) return accepted_tokens, num_accepted_tokens def _draft_sampler_greedy(self, logits: torch.Tensor, d2t=None): diff --git a/tensorrt_llm/_torch/speculative/mtp.py b/tensorrt_llm/_torch/speculative/mtp.py index dda345844c19..3cc616ce6a6a 100644 --- a/tensorrt_llm/_torch/speculative/mtp.py +++ b/tensorrt_llm/_torch/speculative/mtp.py @@ -826,8 +826,10 @@ def sample_and_accept_draft_tokens( # Apply force override for relaxed acceptance path num_accepted_tokens = self._apply_force_accepted_tokens( - num_accepted_tokens, num_contexts, - spec_metadata.runtime_draft_len) + num_accepted_tokens, + num_contexts, + spec_metadata.runtime_draft_len, + spec_metadata=spec_metadata) # Strict acceptance else: @@ -843,8 +845,10 @@ def sample_and_accept_draft_tokens( # Apply force override for THOP path num_accepted_tokens = self._apply_force_accepted_tokens( - num_accepted_tokens, num_contexts, - spec_metadata.runtime_draft_len) + num_accepted_tokens, + num_contexts, + spec_metadata.runtime_draft_len, + spec_metadata=spec_metadata) else: # Reshape draft tokens for base implementation draft_tokens = spec_metadata.draft_tokens.reshape(