From f86ae1f59d94721b45d9d82d4454c90f7b97edbe Mon Sep 17 00:00:00 2001 From: ZhaoyangWang Date: Mon, 11 May 2026 02:27:38 -0700 Subject: [PATCH 01/11] [TRTLLM-11508][refactor] merge Eagle3 and MTP-eagle Unify the Eagle3 one-model and MTP-eagle one-model speculative-decoding workers into a single Eagle3OneModelWorker in eagle3.py, branching on self.is_mtp_eagle. MTPEagleWorker becomes a thin backward-compatible subclass; mtp.py keeps a module-level __getattr__ shim so the historical import path continues to resolve. Key changes: - Eagle3OneModelSpecMetadata gains slot_ids and subseq_all_rank_num_tokens; prepare() skips num_tokens adjustment for MTP-eagle and populates slot_ids from the resource manager. - Eagle3ResourceManager owns the relaxed-acceptance delta pool for both modes. - New helpers _get_step_all_rank_num_tokens, _run_draft_forward, and _prepare_flash_mla_generation_layout encapsulate the per-step branching. - sample_and_accept_draft_tokens takes input_ids and supports the relaxed-thinking path previously exclusive to MTPEagleWorker. - EagleDecodingConfig grows the relaxed-acceptance fields mirrored from MTPDecodingConfig. - SpeculativeDecodingMode.is_mtp_one_model() now means vanilla MTP only; predicates and per-model checks are extended to recognize MTP_EAGLE_ONE_MODEL as a first-class one-model mode. - The Eagle3 _saved_kv_lens_cuda save/restore is dropped (relying on attn_metadata.update_for_spec_dec() instead) - needs verification under Eagle3 regression tests. - Factory routing in utils.py routes MTP_EAGLE_ONE_MODEL to the unified Eagle3 metadata, sampler, resource manager, and worker. Signed-off-by: ZhaoyangWang --- tensorrt_llm/_torch/model_config.py | 5 +- .../_torch/models/modeling_deepseekv3.py | 4 +- .../_torch/models/modeling_exaone_moe.py | 12 +- tensorrt_llm/_torch/models/modeling_glm.py | 6 +- .../_torch/models/modeling_nemotron_h.py | 10 +- .../_torch/models/modeling_qwen3_next.py | 10 +- .../_torch/models/modeling_speculative.py | 6 +- tensorrt_llm/_torch/speculative/__init__.py | 4 +- tensorrt_llm/_torch/speculative/eagle3.py | 589 ++++++++++++++---- .../_torch/speculative/eagle3_dynamic_tree.py | 15 +- tensorrt_llm/_torch/speculative/interface.py | 31 +- tensorrt_llm/_torch/speculative/mtp.py | 283 +-------- tensorrt_llm/_torch/speculative/utils.py | 47 +- tensorrt_llm/llmapi/llm_args.py | 28 + 14 files changed, 620 insertions(+), 430 deletions(-) diff --git a/tensorrt_llm/_torch/model_config.py b/tensorrt_llm/_torch/model_config.py index fa448d8876c2..cadd21543647 100644 --- a/tensorrt_llm/_torch/model_config.py +++ b/tensorrt_llm/_torch/model_config.py @@ -779,8 +779,9 @@ def ceil_div(a, b): hidden_size = ceil_div(self.pretrained_config.hidden_size, attn_tp_size) num_layers = self.pretrained_config.num_hidden_layers num_attention_layers = self.get_num_attention_layers() - if (self.spec_config is not None - and self.spec_config.spec_dec_mode.is_mtp_one_model()): + if (self.spec_config is not None and + (self.spec_config.spec_dec_mode.is_mtp_one_model() + or self.spec_config.spec_dec_mode.is_mtp_eagle_one_model())): assert self.spec_config.num_nextn_predict_layers is not None, ( "num_nextn_predict_layers must be set from model config before building ModelConfig. " "Ensure update_spec_config_from_model_config() has been called." diff --git a/tensorrt_llm/_torch/models/modeling_deepseekv3.py b/tensorrt_llm/_torch/models/modeling_deepseekv3.py index e981dc233010..fc01588ae2e0 100755 --- a/tensorrt_llm/_torch/models/modeling_deepseekv3.py +++ b/tensorrt_llm/_torch/models/modeling_deepseekv3.py @@ -1854,7 +1854,9 @@ def __init__(self, model_config: ModelConfig[PretrainedConfig]): model_config=model_config) self.model_nextn = 0 - if model_config.spec_config is not None and model_config.spec_config.spec_dec_mode.is_mtp_one_model( + if model_config.spec_config is not None and ( + model_config.spec_config.spec_dec_mode.is_mtp_one_model() or + model_config.spec_config.spec_dec_mode.is_mtp_eagle_one_model() ): model_nextn = self.config.num_nextn_predict_layers ckpt_nextn = self.config.num_nextn_predict_layers diff --git a/tensorrt_llm/_torch/models/modeling_exaone_moe.py b/tensorrt_llm/_torch/models/modeling_exaone_moe.py index 9df138259b57..25b578caf5e7 100644 --- a/tensorrt_llm/_torch/models/modeling_exaone_moe.py +++ b/tensorrt_llm/_torch/models/modeling_exaone_moe.py @@ -640,9 +640,9 @@ def __init__( self, model_config: ModelConfig[ExaoneMoeConfig], ): - if ( - model_config.spec_config is not None - and model_config.spec_config.spec_dec_mode.is_mtp_one_model() + if model_config.spec_config is not None and ( + model_config.spec_config.spec_dec_mode.is_mtp_one_model() + or model_config.spec_config.spec_dec_mode.is_mtp_eagle_one_model() ): # NOTE: K-EXAONE does not contain the 'num_nextn_predict_layers' field, # which should be equal to 1. Manually set the value here if not present. @@ -654,9 +654,9 @@ def __init__( model_config=model_config, ) - if ( - model_config.spec_config is not None - and model_config.spec_config.spec_dec_mode.is_mtp_one_model() + if model_config.spec_config is not None and ( + model_config.spec_config.spec_dec_mode.is_mtp_one_model() + or model_config.spec_config.spec_dec_mode.is_mtp_eagle_one_model() ): model_nextn = self.config.num_nextn_predict_layers ckpt_nextn = self.config.num_nextn_predict_layers diff --git a/tensorrt_llm/_torch/models/modeling_glm.py b/tensorrt_llm/_torch/models/modeling_glm.py index 293510b65099..018001a8b4a9 100644 --- a/tensorrt_llm/_torch/models/modeling_glm.py +++ b/tensorrt_llm/_torch/models/modeling_glm.py @@ -1020,9 +1020,9 @@ def __init__(self, model_config: ModelConfig[PretrainedConfig]): super().__init__(model=Glm4Model(model_config), model_config=model_config) self.model_nextn = 0 - if ( - model_config.spec_config is not None - and model_config.spec_config.spec_dec_mode.is_mtp_one_model() + if model_config.spec_config is not None and ( + model_config.spec_config.spec_dec_mode.is_mtp_one_model() + or model_config.spec_config.spec_dec_mode.is_mtp_eagle_one_model() ): model_nextn = self.config.num_nextn_predict_layers ckpt_nextn = self.config.num_nextn_predict_layers diff --git a/tensorrt_llm/_torch/models/modeling_nemotron_h.py b/tensorrt_llm/_torch/models/modeling_nemotron_h.py index f643eff0ef1e..fedaefc1debe 100644 --- a/tensorrt_llm/_torch/models/modeling_nemotron_h.py +++ b/tensorrt_llm/_torch/models/modeling_nemotron_h.py @@ -907,8 +907,9 @@ def __init__( model_config=model_config, ) self.model_nextn = 0 - if (model_config.spec_config is not None - and model_config.spec_config.spec_dec_mode.is_mtp_one_model()): + if (model_config.spec_config is not None and + (model_config.spec_config.spec_dec_mode.is_mtp_one_model() or + model_config.spec_config.spec_dec_mode.is_mtp_eagle_one_model())): model_nextn = self.config.num_nextn_predict_layers ckpt_nextn = self.config.num_nextn_predict_layers self.num_hidden_layers = self.config.num_hidden_layers @@ -1041,8 +1042,9 @@ def __init__( use_custom_cublas_mm=use_custom_cublas_mm, ) self.model_nextn = 0 - if (model_config.spec_config is not None - and model_config.spec_config.spec_dec_mode.is_mtp_one_model()): + if (model_config.spec_config is not None and + (model_config.spec_config.spec_dec_mode.is_mtp_one_model() or + model_config.spec_config.spec_dec_mode.is_mtp_eagle_one_model())): model_nextn = self.config.num_nextn_predict_layers ckpt_nextn = self.config.num_nextn_predict_layers self.num_hidden_layers = self.config.num_hidden_layers diff --git a/tensorrt_llm/_torch/models/modeling_qwen3_next.py b/tensorrt_llm/_torch/models/modeling_qwen3_next.py index d6f4fd57794f..4bbb6854239d 100644 --- a/tensorrt_llm/_torch/models/modeling_qwen3_next.py +++ b/tensorrt_llm/_torch/models/modeling_qwen3_next.py @@ -944,8 +944,9 @@ def __init__( self, model_config: ModelConfig[Qwen3NextConfig], ): - if (model_config.spec_config is not None - and model_config.spec_config.spec_dec_mode.is_mtp_one_model()): + if (model_config.spec_config is not None and + (model_config.spec_config.spec_dec_mode.is_mtp_one_model() or + model_config.spec_config.spec_dec_mode.is_mtp_eagle_one_model())): ckpt_num_nextn = getattr(model_config.pretrained_config, "num_nextn_predict_layers", None) if ckpt_num_nextn not in (None, 1): @@ -962,8 +963,9 @@ def __init__( ) self.preload_weight_modules = self.model.preload_weight_modules - if (model_config.spec_config is not None - and model_config.spec_config.spec_dec_mode.is_mtp_one_model()): + if (model_config.spec_config is not None and + (model_config.spec_config.spec_dec_mode.is_mtp_one_model() or + model_config.spec_config.spec_dec_mode.is_mtp_eagle_one_model())): self.model.layers.extend(self.draft_model.mtp_layers) diff --git a/tensorrt_llm/_torch/models/modeling_speculative.py b/tensorrt_llm/_torch/models/modeling_speculative.py index 62b24956d077..efaa2f294e4b 100755 --- a/tensorrt_llm/_torch/models/modeling_speculative.py +++ b/tensorrt_llm/_torch/models/modeling_speculative.py @@ -1428,7 +1428,8 @@ def __init__( f"Model type {model_type} not supported for MTP") spec_dec_mode = model_config.spec_config.spec_dec_mode - assert spec_dec_mode.is_mtp_one_model() + assert (spec_dec_mode.is_mtp_one_model() + or spec_dec_mode.is_mtp_eagle_one_model()) checkpoint_mtp_num_layers = model_config.pretrained_config.num_nextn_predict_layers if spec_dec_mode.is_mtp_eagle_one_model(): mtp_num_layers = 1 @@ -1614,7 +1615,8 @@ def get_draft_model(model_config, draft_config, lm_head, model): f"Unsupported eagle3 model architecture: {spec_dec_mode.eagle3_model_arch}" ) - elif spec_dec_mode.is_mtp_one_model(): + elif (spec_dec_mode.is_mtp_one_model() + or spec_dec_mode.is_mtp_eagle_one_model()): return MTPForCausalLM(model_config, model_config.pretrained_config.num_hidden_layers, lm_head, model) diff --git a/tensorrt_llm/_torch/speculative/__init__.py b/tensorrt_llm/_torch/speculative/__init__.py index 0f16df6baffd..d1c1b2605283 100644 --- a/tensorrt_llm/_torch/speculative/__init__.py +++ b/tensorrt_llm/_torch/speculative/__init__.py @@ -2,12 +2,12 @@ from .dflash import DFlashSpecMetadata, DFlashWorker from .draft_target import (DraftTargetOneModelSpecMetadata, DraftTargetOneModelWorker) -from .eagle3 import Eagle3SpecMetadata +from .eagle3 import Eagle3SpecMetadata, MTPEagleWorker from .interface import (SpecMetadata, SpecWorkerBase, prepare_attn_metadata_for_draft_replay, restore_attn_metadata_after_draft_replay, should_use_separate_draft_kv_cache) -from .mtp import MTPEagleWorker, MTPSampler, MTPSpecMetadata, MTPWorker +from .mtp import MTPSampler, MTPSpecMetadata, MTPWorker from .ngram import NGramDrafter, NGramPoolManager from .pard import PARDSpecMetadata, PARDWorker from .sa_enhancer import SADraftEnhancer diff --git a/tensorrt_llm/_torch/speculative/eagle3.py b/tensorrt_llm/_torch/speculative/eagle3.py index b5b0b9e24877..46dd3faec7d5 100644 --- a/tensorrt_llm/_torch/speculative/eagle3.py +++ b/tensorrt_llm/_torch/speculative/eagle3.py @@ -2,6 +2,7 @@ from typing import TYPE_CHECKING, Dict, List, Optional, Set import torch +import torch.nn.functional as F from torch import nn from tensorrt_llm._torch.custom_ops import inplace_slice_copy @@ -9,12 +10,15 @@ from tensorrt_llm.mapping import Mapping from ..attention_backend import AttentionMetadata +from ..distributed.ops import allgather +from ..model_config import ModelConfig from ..pyexecutor.llm_request import LlmRequest +from ..pyexecutor.mamba_cache_manager import MambaHybridCacheManager from ..pyexecutor.resource_manager import BaseResourceManager, SlotManager from ..pyexecutor.sampler import TorchSampler from ..pyexecutor.scheduler import ScheduledRequests from .interface import SpecMetadata, SpecWorkerBase -from .mtp import MTPSampler +from .mtp import MTPSampler, _select_mtp_position_ids from .sa_enhancer import SADraftEnhancer from .spec_tree_manager import SpecTreeManager @@ -72,6 +76,15 @@ def __init__(self, ) # sequence length, only used for metadata preparation self.seq_lens = {i: 0 for i in range(slot_size)} + + # Per-request delta pool tracking whether the request is in the + # thinking phase; mirrors MTPHiddenStatesManager.mtp_relaxed_delta_pool. + self.use_relaxed_acceptance_for_thinking = getattr( + config, 'use_relaxed_acceptance_for_thinking', False) + if self.use_relaxed_acceptance_for_thinking: + self.relaxed_delta_pool = torch.zeros((slot_size, ), + dtype=torch.float, + device='cuda') # start indices of each slot self.start_indices = {i: 0 for i in range(slot_size)} # whether the next draft forward is the first @@ -98,6 +111,8 @@ def prepare_resources(self, scheduled_batch: ScheduledRequests): if req.is_first_context_chunk: slot_id = self.slot_manager.add_slot(req.request_id) self.slot_ids.append(slot_id) + if self.use_relaxed_acceptance_for_thinking: + self.relaxed_delta_pool[slot_id].fill_(0) # reset the flag before model forward self.is_first_draft = True @@ -108,6 +123,8 @@ def free_resources(self, request: LlmRequest): slot_id = self.slot_manager.get_slot(request.request_id) self.seq_lens[slot_id] = 0 self.start_indices[slot_id] = 0 + if self.use_relaxed_acceptance_for_thinking: + self.relaxed_delta_pool[slot_id].fill_(0) self.slot_manager.remove_slot(request.request_id) if self.sa_manager is not None: self.sa_manager.remove_request(request.request_id) @@ -365,15 +382,24 @@ class Eagle3OneModelSpecMetadata(SpecMetadata): dtype: torch.dtype = torch.bfloat16 # The index of the batch inputs batch_indices_cuda: Optional[torch.Tensor] = None - # Optional resource manager (used to access SA manager for EAGLE3+SA) + # Optional resource manager (used to access SA manager and relaxed-acceptance + # delta pool for Eagle3+SA / Eagle3+relaxed-thinking / MTP Eagle modes) spec_resource_manager: Optional[Eagle3ResourceManager] = None # Dynamic tree flags use_dynamic_tree: bool = False eagle_choices: Optional[List[List[int]]] = None + # Slot IDs for each request; populated in prepare() when spec_resource_manager + # is present (required for relaxed acceptance, mirrors MTPSpecMetadata.slot_ids). + slot_ids: Optional[torch.Tensor] = None + # One-model speculative decoding uses the first draft forward token counts + # for the first loop iteration and per-sequence token counts for + # subsequent iterations. + subseq_all_rank_num_tokens: Optional[List[int]] = None def __post_init__(self): if self.layers_to_capture is None: - if self.num_layers == 1: + if self.spec_dec_mode.is_mtp_eagle_one_model( + ) or self.num_layers == 1: self.layers_to_capture = (self.num_layers - 1, ) else: if self.num_layers <= 5: @@ -415,6 +441,13 @@ def __post_init__(self): dtype=torch.int, device='cuda', ) + # Pre-allocate slot_ids; filled in prepare() when spec_resource_manager + # is present. Mirrors MTPSpecMetadata.slot_ids allocation pattern. + self.slot_ids = torch.empty( + [self.max_num_requests], + dtype=torch.long, + device='cuda', + ) # Set tree flags based on config if self.use_dynamic_tree: @@ -441,11 +474,28 @@ def prepare(self): pin_memory=prefer_pinned()) self.batch_indices_cuda[:num_seqs].copy_(batch_indices, non_blocking=True) - if self.is_spec_dec_tree: - self.num_tokens -= ( - self.num_generations) * self.max_total_draft_tokens - else: - self.num_tokens -= (self.num_generations) * self.max_draft_len + # MTP Eagle uses max_draft_len + 1 tokens in the first draft forward so + # it must not subtract here; Eagle3 follows the standard tree/linear path. + if not self.spec_dec_mode.is_mtp_eagle_one_model(): + if self.is_spec_dec_tree: + self.num_tokens -= ( + self.num_generations) * self.max_total_draft_tokens + else: + self.num_tokens -= (self.num_generations) * self.max_draft_len + + if self.spec_resource_manager is not None: + # Populate slot_ids for all requests in this batch. Used by relaxed + # acceptance (relaxed_delta_pool indexing), mirroring the pattern + # in MTPSpecMetadata.prepare(). + eagle_slot_ids = [ + self.spec_resource_manager.slot_manager.get_slot(rid) + for rid in self.request_ids + ] + eagle_slot_ids_tensor = torch.tensor(eagle_slot_ids, + dtype=torch.int, + pin_memory=prefer_pinned()) + self.slot_ids[:num_seqs].copy_(eagle_slot_ids_tensor, + non_blocking=True) sa_manager = getattr(self.spec_resource_manager, 'sa_manager', None) if sa_manager is not None: @@ -485,41 +535,56 @@ def _get_max_new_tokens(self, args: TorchSampler.Args, class Eagle3OneModelWorker(SpecWorkerBase): - """Eagle3 one-model worker for linear tree speculative decoding. + """Unified one-model worker for Eagle3 and MTP Eagle speculative decoding. - For dynamic tree mode, use Eagle3OneModelDynamicTreeWorker from - eagle3_dynamic_tree.py instead. + The operating mode is determined by ``spec_config.spec_dec_mode``: + - EAGLE3_ONE_MODEL: multi-layer hidden states from Eagle3, apply_eagle3_fc + projection, independent EAGLE draft model network. + - MTP_EAGLE_ONE_MODEL: single last-layer hidden states, MTP layer called + repeatedly, supports TP-aware sampling and Mamba hybrid cache. + + Where the two modes differ, ``self.is_mtp_eagle`` is used to branch. + For dynamic tree Eagle3, use ``Eagle3OneModelDynamicTreeWorker`` from + ``eagle3_dynamic_tree.py``. """ def __init__(self, spec_config: "EagleDecodingConfig", - mapping: Mapping, + mapping: Optional[Mapping] = None, + model_config: Optional[ModelConfig] = None, use_separate_draft_kv_cache: bool = False): super().__init__(use_separate_draft_kv_cache) self.spec_config = spec_config self.mapping = mapping + # model_config is required for MTP Eagle TP / ADP / Mamba support; the + # Eagle3 path can leave it as None. + self.model_config = model_config + + # Mode flag: True = MTP Eagle one-model, False = Eagle3 one-model. + self.is_mtp_eagle = spec_config.spec_dec_mode.is_mtp_eagle_one_model() + + # SA enhancer (common to both modes) self.sa_enhancer: Optional[SADraftEnhancer] = None if getattr(spec_config, 'sa_config', None) is not None: self.sa_enhancer = SADraftEnhancer(spec_config.sa_config.threshold) self.use_dynamic_tree = getattr(spec_config, 'use_dynamic_tree', False) self.spec_tree_manager = None + # MTP Eagle: lazily-resolved flag for Mamba hybrid cache support + self._is_mamba_hybrid_cache = None + @property def max_draft_len(self) -> int: return self.spec_config.max_draft_len def _prepare_attn_metadata_for_spec_dec(self, attn_metadata): attn_metadata.prepare_for_spec_dec("_seq_lens", "_seq_lens_cuda") - # Save kv_lens_cuda values separately instead of routing through - # prepare_for_spec_dec, which would clone the tensor and break the - # kv_lens_cuda_runtime view that TRTLLM attention reads from. + # NOTE(TRTLLM-11508): the previous kv_lens_cuda save/restore was removed + # during the Eagle3/MTP-eagle merge. The drafting loop now updates + # kv_lens_cuda incrementally and calls attn_metadata.update_for_spec_dec() + # to keep the runtime view consistent. Verify under Eagle3 regressions + # if any kv-lens drift is observed. batch_size = attn_metadata.num_seqs - if hasattr(attn_metadata, 'kv_lens_cuda'): - self._saved_kv_lens_cuda = attn_metadata.kv_lens_cuda[: - batch_size].clone( - ) - else: - self._saved_kv_lens_cuda = None # Save spec-dec params that the drafting loop will overwrite. # Without this, CUDA graph warmup's second iteration would run @@ -547,12 +612,6 @@ def _prepare_attn_metadata_for_spec_dec(self, attn_metadata): def _restore_attn_metadata_from_spec_dec(self, attn_metadata): super()._restore_attn_metadata_from_spec_dec(attn_metadata) - if self._saved_kv_lens_cuda is not None: - batch_size = self._saved_kv_lens_cuda.shape[0] - attn_metadata.kv_lens_cuda[:batch_size].copy_( - self._saved_kv_lens_cuda) - self._saved_kv_lens_cuda = None - if self._saved_packed_mask is not None: batch_size = self._saved_packed_mask.shape[0] attn_metadata.spec_decoding_packed_mask[:batch_size].copy_( @@ -599,9 +658,23 @@ def forward(self, self._execute_guided_decoder_if_present(logits) - # Sample and accept tokens + # Sample and accept tokens. ``input_ids`` is required by the relaxed- + # acceptance path (scans for thinking-phase tokens); ignored otherwise. accepted_tokens, num_accepted_tokens = self.sample_and_accept_draft_tokens( - logits, attn_metadata, spec_metadata) + input_ids, logits, attn_metadata, spec_metadata) + + # MTP Eagle only: Mamba hybrid models need state updates after token + # acceptance because accepted token count affects which Mamba states + # are valid; Eagle3 does not use Mamba layers. + if self.is_mtp_eagle: + if self._is_mamba_hybrid_cache is None: + self._is_mamba_hybrid_cache = isinstance( + attn_metadata.kv_cache_manager, MambaHybridCacheManager) + if num_gens > 0 and self._is_mamba_hybrid_cache: + attn_metadata.kv_cache_manager.update_mamba_states( + attn_metadata=attn_metadata, + num_accepted_tokens=num_accepted_tokens, + state_indices=attn_metadata.mamba_metadata.state_indices) sa_manager = getattr(spec_metadata.spec_resource_manager, 'sa_manager', None) @@ -630,7 +703,9 @@ def forward(self, spec_metadata=spec_metadata, draft_model=draft_model) - # Predict draft tokens + # Predict draft tokens. ``original_all_rank_num_tokens`` is saved here + # so the post-loop restore (below) can put attn_metadata back into a + # state the target model expects. original_all_rank_num_tokens = attn_metadata.all_rank_num_tokens # Get the draft KV cache manager if using separate layouts @@ -679,28 +754,40 @@ def _forward_linear_draft_loop(self, inputs, attn_metadata, spec_metadata, num_contexts, batch_size, num_accepted_tokens, original_all_rank_num_tokens): - """Original linear draft loop (1 token per layer).""" + """Linear draft loop, unified for Eagle3 and MTP Eagle.""" runtime_draft_len = spec_metadata.runtime_draft_len + num_gens = batch_size - num_contexts next_draft_tokens = [] draft_logits_list = [] - position_ids = inputs["position_ids"] + last_tokens_idx = torch.cumsum(attn_metadata.seq_lens_cuda, + dim=0, + dtype=torch.long) - 1 with self.draft_kv_cache_context(attn_metadata, draft_kv_cache_manager): for i in range(runtime_draft_len): + # Run draft model (mode-specific via helper). For Eagle3 the + # helper mutates ``attn_metadata.all_rank_num_tokens`` to the + # right per-step value; for MTP Eagle the value is passed as + # a kwarg to ``mtp_layers[0]`` directly. + hidden_states, hidden_states_to_save = self._run_draft_forward( + draft_model, inputs, spec_metadata, i, + original_all_rank_num_tokens) + + # Compute gather_ids: on the first draft step each generation + # request may have accepted multiple tokens, so we index into + # the flattened token sequence to find the last accepted one. + # From step 1 onwards every sequence has length 1, so + # ``batch_indices_cuda`` is sufficient. if i == 0: - num_gens = batch_size - num_contexts start_ids_gen = ( spec_metadata.batch_indices_cuda[:num_gens] * (runtime_draft_len + 1)).long() gather_ids_gen = (start_ids_gen + num_accepted_tokens[num_contexts:] - 1 + attn_metadata.num_ctx_tokens) - gather_ids = torch.concat([ - spec_metadata.gather_ids[:num_contexts], gather_ids_gen - ], - dim=0) + gather_ids = torch.concat( + [last_tokens_idx[:num_contexts], gather_ids_gen], dim=0) else: - # All of the seq_len are 1, use batch_indices_cuda as gather_ids gather_ids = spec_metadata.batch_indices_cuda[:batch_size] if self.guided_decoder is not None: @@ -709,59 +796,133 @@ def _forward_linear_draft_loop(self, inputs, attn_metadata, spec_metadata, num_accepted_tokens, draft_step=i) - # Update attn_metadata.all_rank_num_tokens for attention DP - if original_all_rank_num_tokens is not None: - if i == 0: - attn_metadata.all_rank_num_tokens = original_all_rank_num_tokens - elif spec_metadata.all_rank_num_seqs is not None: - attn_metadata.all_rank_num_tokens = spec_metadata.all_rank_num_seqs - - hidden_states, hidden_states_to_save = draft_model.model( - **inputs) - - # FIXME (jhaotingc): Currently we disable use_spec_decoding mode for Eagle engine nth steps except 1st step. - # Eagle engine takes in draft_len tokens from the previous step, run spec-dec mode with those tokens, - # then the following step can use regular decoding mode to generate 1 tokens per step. - # Currently the spec-dec mask for chained tree is not implemented yet. - # When token tree is supported, this can be removed and all steps may use spec-dec mode as well. - attn_metadata.use_spec_decoding = False - - logits = draft_model.logits_processor(hidden_states[gather_ids], - draft_model.lm_head, - attn_metadata, True) + # Compute logits. + # MTP Eagle: shared_head of the MTP layer, with optional + # ADP+LM-head-TP padding to ``max_num_requests`` so every TP + # rank produces logits of the same shape. + # Eagle3: logits_processor of the EAGLE draft model. + use_lm_head_tp_in_adp = ( + self.is_mtp_eagle and self.model_config is not None + and self.model_config.mapping.enable_attention_dp + and getattr(self.model_config.mapping, + 'enable_lm_head_tp_in_adp', False)) + if self.is_mtp_eagle: + if use_lm_head_tp_in_adp: + hidden_states_gathered = hidden_states[gather_ids] + token_count = hidden_states_gathered.view( + -1, hidden_states_gathered.shape[-1]).shape[0] + max_num_requests = spec_metadata.max_num_requests + pad_len = max_num_requests - token_count + if pad_len > 0: + padded_hidden_states = F.pad( + hidden_states_gathered.view( + -1, hidden_states_gathered.shape[-1]), + (0, 0, 0, pad_len), + mode="constant", + value=0) + elif pad_len == 0: + padded_hidden_states = hidden_states_gathered.view( + -1, hidden_states_gathered.shape[-1]) + else: + raise ValueError( + "Eagle3OneModelWorker (MTP Eagle mode): " + "token_count > max_num_requests, which is not supported" + ) + logits = draft_model.mtp_layers[0].shared_head( + padded_hidden_states, draft_model.lm_head, + attn_metadata, True) + else: + logits = draft_model.mtp_layers[0].shared_head( + hidden_states[gather_ids], draft_model.lm_head, + attn_metadata, True) + else: + logits = draft_model.logits_processor( + hidden_states[gather_ids], draft_model.lm_head, + attn_metadata, True) + if self.guided_decoder is not None: + if self.is_mtp_eagle: + self.guided_decoder.execute_draft_batch(logits, + draft_step=i) + else: + d2t = getattr(draft_model.model, "d2t", None) + self.guided_decoder.execute_draft_batch(logits, + d2t, + draft_step=i) + + # Sample the next draft token. + # MTP Eagle: TP-aware sampler; when ADP+LM-head-TP is active + # logits are padded to max_num_requests across TP ranks, so + # the result must be trimmed back to token_count. + # Eagle3: simple greedy sampling; d2t remaps vocab indices when + # the draft model uses a compressed vocabulary. + if self.is_mtp_eagle: + if use_lm_head_tp_in_adp: + mapping_lm_head_tp = draft_model.mtp_layers[ + 0].shared_head.mapping_lm_head_tp + new_draft_token = self.draft_sampler( + logits, mapping_lm_head_tp) + new_draft_token = new_draft_token[:token_count] + else: + new_draft_token = self.draft_sampler(logits) + else: d2t = getattr(draft_model.model, "d2t", None) - self.guided_decoder.execute_draft_batch(logits, - d2t, - draft_step=i) + new_draft_token = self._draft_sampler_greedy(logits, d2t) - if spec_metadata.use_rejection_sampling: + # Stash unpadded Eagle3 draft logits for rejection sampling on + # the next iteration. MTP Eagle's logits may be ADP-padded to + # max_num_requests, so we skip them here. + if not self.is_mtp_eagle and spec_metadata.use_rejection_sampling: draft_logits_list.append(logits.clone()) - new_draft_token = self.draft_decoder(logits, draft_model) next_draft_tokens.append(new_draft_token) - # update inputs - hidden_states = hidden_states_to_save[gather_ids] - position_ids = inputs["position_ids"][gather_ids] + 1 - # update attn_metadata + + # Update hidden states for the next iteration. + # MTP Eagle: the MTP layer returns a single tensor; slice by + # gather_ids to get one hidden state per request. + # Eagle3: the EAGLE draft model returns a secondary + # ``hidden_states_to_save`` specifically for this purpose. + if self.is_mtp_eagle: + hidden_states = hidden_states[gather_ids] + else: + hidden_states = hidden_states_to_save[gather_ids] + position_ids = ( + _select_mtp_position_ids(inputs["position_ids"], gather_ids) + + 1) + + # Update attn_metadata for the next iteration. if i == 0: attn_metadata._seq_lens[:batch_size].fill_(1) attn_metadata._seq_lens_cuda[:batch_size].fill_(1) attn_metadata.on_update() - # cannot run generation if there is no kv cache - if inputs["attn_metadata"].kv_cache_manager is not None: + has_kv_cache = inputs[ + "attn_metadata"].kv_cache_manager is not None + if has_kv_cache: attn_metadata.host_request_types[:attn_metadata. num_contexts].fill_(1) attn_metadata.num_contexts = 0 - # update kv_lens_cuda if hasattr(attn_metadata, 'kv_lens_cuda'): attn_metadata.kv_lens_cuda[num_contexts:batch_size] -= ( runtime_draft_len - num_accepted_tokens[num_contexts:]) attn_metadata.kv_lens_cuda[:num_contexts] += 1 - elif hasattr(attn_metadata, 'kv_lens_cuda'): - attn_metadata.kv_lens_cuda[:batch_size] += 1 - # support attention dp + + if has_kv_cache: + self._prepare_flash_mla_generation_layout( + attn_metadata, num_contexts, batch_size) + if hasattr(attn_metadata, 'kv_lens_cuda'): + attn_metadata.update_for_spec_dec() + + # Eagle engine takes ``draft_len`` tokens from the previous + # step, runs spec-dec mode with those tokens, then later + # steps use regular decoding mode. Disable spec_decoding so + # the masks/positions stay correct on subsequent iters. + attn_metadata.use_spec_decoding = False + else: + if hasattr(attn_metadata, 'kv_lens_cuda'): + attn_metadata.kv_lens_cuda[:batch_size] += 1 + attn_metadata.update_for_spec_dec() + inputs = { "input_ids": new_draft_token, "position_ids": position_ids, @@ -787,53 +948,234 @@ def _forward_linear_draft_loop(self, inputs, attn_metadata, spec_metadata, return next_draft_tokens + def _get_step_all_rank_num_tokens(self, spec_metadata, step_idx: int, + fallback): + """Pick the right ``all_rank_num_tokens`` for this draft iteration. + + Step 0 uses ``spec_metadata.all_rank_num_tokens`` (or the fallback + captured from attn_metadata before the loop); subsequent steps use + ``spec_metadata.subseq_all_rank_num_tokens`` since every sequence + contributes a single token per iteration. + """ + if step_idx == 0: + return (spec_metadata.all_rank_num_tokens + if spec_metadata.all_rank_num_tokens is not None else + fallback) + return spec_metadata.subseq_all_rank_num_tokens + + def _run_draft_forward(self, draft_model, inputs, spec_metadata, + step_idx: int, original_all_rank_num_tokens): + """Invoke the draft model for one iteration, branching on mode. + + For MTP Eagle, ``all_rank_num_tokens`` is passed directly as a kwarg + to ``mtp_layers[0]``. For Eagle3, it is set on ``attn_metadata`` so + the EAGLE draft model picks it up; the outer ``forward()`` is + responsible for restoring the original value after the loop. + """ + all_rank_num_tokens = self._get_step_all_rank_num_tokens( + spec_metadata, step_idx, original_all_rank_num_tokens) + + if self.is_mtp_eagle: + hidden_states = draft_model.mtp_layers[0]( + embed_tokens=draft_model.embed_tokens, + all_rank_num_tokens=all_rank_num_tokens, + **inputs) + return hidden_states, None + + # Eagle3: route per-step token counts through attn_metadata. Commit 2 + # of this refactor migrates this to a kwarg + try/finally inside + # ``Eagle3DraftModel.forward``. + attn_metadata = inputs["attn_metadata"] + if all_rank_num_tokens is not None: + attn_metadata.all_rank_num_tokens = all_rank_num_tokens + hidden_states, hidden_states_to_save = draft_model.model(**inputs) + return hidden_states, hidden_states_to_save + + def _prepare_flash_mla_generation_layout(self, attn_metadata, num_contexts, + batch_size): + """Reorder ``kv_block_ids_per_seq`` so gen requests precede context. + + Flash MLA on first-step expects the layout used during normal + generation; both Eagle3 and MTP Eagle hit this when context requests + share the batch with gen requests. + """ + if num_contexts <= 0 or not attn_metadata.enable_flash_mla: + return + reorder_block_ids_per_seq = torch.cat([ + attn_metadata.kv_block_ids_per_seq[num_contexts:batch_size], + attn_metadata.kv_block_ids_per_seq[:num_contexts] + ]) + attn_metadata.block_ids_per_seq[:batch_size, :].copy_( + reorder_block_ids_per_seq, non_blocking=True) + + @torch.compile(options={"max-autotune": True}) + def _get_local_max_and_combined(self, logits, mapping_lm_tp=None): + local_max_values, local_argmax = torch.max(logits, dim=-1, keepdim=True) + vocab_per_rank = logits.shape[-1] + mapping_lm_tp = mapping_lm_tp if mapping_lm_tp is not None else \ + self.model_config.mapping + max_index_per_rank = local_argmax.type( + torch.int32) + (mapping_lm_tp.tp_rank * vocab_per_rank) + max_index_per_rank_float = max_index_per_rank.float() + local_max_values_float32 = local_max_values.float() + combined = torch.stack( + [max_index_per_rank_float, local_max_values_float32], + dim=-1).flatten(-2) + return combined + + @torch.compile(options={"max-autotune": True}) + def _get_draft_tokens_from_gathered(self, gathered): + gathered_indices_float = gathered[..., 0::2] + gathered_values_float = gathered[..., 1::2] + max_indices = torch.argmax(gathered_values_float, dim=-1, keepdim=True) + draft_tokens = torch.gather(gathered_indices_float, -1, + max_indices).squeeze(-1).type(torch.int32) + return draft_tokens + + def draft_sampler( + self, + logits: torch.Tensor, + mapping_lm_head_tp=None, + ): + """TP-aware greedy draft token sampler (MTP Eagle path). + + Falls back to simple argmax when no tensor parallelism is active or + when only attention DP is enabled without LM-head TP. + """ + if (self.model_config is not None + and hasattr(self.model_config, 'mapping') + and self.model_config.mapping.tp_size > 1 + and not self.model_config.mapping.enable_attention_dp): + combined = self._get_local_max_and_combined(logits) + gathered = allgather(combined, self.model_config.mapping, dim=-1) + return self._get_draft_tokens_from_gathered(gathered) + elif (self.model_config is not None + and hasattr(self.model_config, 'mapping') + and self.model_config.mapping.tp_size > 1 + and self.model_config.mapping.enable_lm_head_tp_in_adp): + combined = self._get_local_max_and_combined(logits, + mapping_lm_head_tp) + gathered = allgather(combined, mapping_lm_head_tp, dim=-1) + batch_size = logits.shape[0] + local_batch_size = batch_size // mapping_lm_head_tp.tp_size + gathered = gathered.view(mapping_lm_head_tp.tp_size, + local_batch_size, -1) + sliced_gathered = gathered[mapping_lm_head_tp.tp_rank] + return self._get_draft_tokens_from_gathered(sliced_gathered) + else: + return self._draft_sampler_greedy(logits) + + @torch.compile(options={"max-autotune": True}) + def _topk_kernel(self, gen_logprobs, num_gens, mtp_num_modules, + spec_metadata): + topk_value, topk_indices = torch.topk(gen_logprobs, + k=self.spec_config.relaxed_topk, + dim=-1) + topk_indices = topk_indices.reshape(num_gens, mtp_num_modules + 1, + self.spec_config.relaxed_topk) + topk_value = topk_value.reshape(num_gens, mtp_num_modules + 1, + self.spec_config.relaxed_topk) + draft_tokens = spec_metadata.draft_tokens.reshape( + num_gens, mtp_num_modules) + return topk_value, topk_indices, draft_tokens + + @torch.compile(options={"max-autotune": True}) + def _process_generation_logits(self, logits, num_contexts): + gen_logits = logits[num_contexts:] + gen_logprobs = torch.softmax(gen_logits, dim=-1) + return gen_logprobs + def sample_and_accept_draft_tokens( self, + input_ids: torch.IntTensor, logits: torch.Tensor, attn_metadata: AttentionMetadata, spec_metadata: Eagle3OneModelSpecMetadata, ): + """Sample the golden token and verify previously proposed draft tokens. + + ``input_ids`` is scanned for thinking-phase tokens when relaxed + acceptance is enabled (both Eagle3 and MTP Eagle); ignored otherwise. + """ batch_size = attn_metadata.num_seqs num_contexts = attn_metadata.num_contexts num_gens = batch_size - num_contexts - # Linear mode: reshape draft tokens for base implementation + runtime_draft_len = spec_metadata.runtime_draft_len + + if getattr(self.spec_config, 'use_relaxed_acceptance_for_thinking', + False): + # Relaxed acceptance — common path for Eagle3 and MTP Eagle. + # Accepts draft tokens that fall within the top-K candidates of the + # target distribution during the thinking phase. + if logits.dim() == 1: + logits = logits.unsqueeze(0) + + accepted_tokens = torch.ones((batch_size, runtime_draft_len + 1), + dtype=torch.int, + device=logits.device) + num_accepted_tokens = torch.ones(batch_size, + dtype=torch.int, + device=logits.device) + + resource_manager = spec_metadata.spec_resource_manager + relaxed_delta_pool = resource_manager.relaxed_delta_pool + + # Context phase: detect thinking tokens and update the delta pool + con_logits = logits[:num_contexts] + con_target_tokens = torch.argmax(con_logits, dim=-1) + accepted_tokens[:num_contexts, 0] = con_target_tokens[:num_contexts] + last_tokens_idx_for_thinking = torch.cumsum( + attn_metadata.seq_lens_cuda, dim=0, dtype=torch.long) - 1 + ctx_input_ids = input_ids[:attn_metadata.num_ctx_tokens] + ctx_is_think = (ctx_input_ids == + self.spec_config.begin_thinking_phase_token).int() + ctx_is_think_cumsum = torch.cumsum(ctx_is_think, dim=0) + ctx_last_cumsum = ctx_is_think_cumsum[ + last_tokens_idx_for_thinking[:num_contexts]] + ctx_think_tokens_num = torch.diff( + ctx_last_cumsum, + dim=0, + prepend=torch.zeros(1, + dtype=torch.int, + device=ctx_last_cumsum.device)) + ctx_delta = (ctx_think_tokens_num + >= 1).int() * self.spec_config.relaxed_delta + ctx_slot_ids = spec_metadata.slot_ids[:num_contexts] + relaxed_delta_pool.index_copy_(0, ctx_slot_ids, ctx_delta) + + # Generation phase: top-k logprobs + relaxed acceptance op + gen_logprobs = self._process_generation_logits(logits, num_contexts) + topk_value, topk_indices, draft_tokens = self._topk_kernel( + gen_logprobs, num_gens, runtime_draft_len, spec_metadata) + + accepted_tokens, num_accepted_tokens = torch.ops.trtllm.mtp_relaxed_acceptance_op( + spec_metadata.slot_ids, topk_value, topk_indices, draft_tokens, + relaxed_delta_pool, num_accepted_tokens, accepted_tokens, + runtime_draft_len, batch_size, num_contexts, + self.spec_config.relaxed_topk, self.spec_config.relaxed_delta, + self.spec_config.begin_thinking_phase_token, + self.spec_config.end_thinking_phase_token) + + num_accepted_tokens = self._apply_force_accepted_tokens( + num_accepted_tokens, num_contexts, runtime_draft_len) + + return accepted_tokens, num_accepted_tokens + + # Strict acceptance — common path for Eagle3 and MTP Eagle. Both modes + # use runtime_draft_len for dynamic draft length support. + if logits.dim() == 1: + logits = logits.unsqueeze(0) draft_tokens = spec_metadata.draft_tokens.reshape( num_gens, - spec_metadata.runtime_draft_len) if num_gens > 0 else torch.empty( + runtime_draft_len) if num_gens > 0 else torch.empty( 0, - spec_metadata.runtime_draft_len, + runtime_draft_len, dtype=torch.int, device=logits.device) return self._accept_draft_tokens(logits, draft_tokens, num_contexts, batch_size, spec_metadata) - def draft_decoder( - self, - logits: torch.Tensor, - draft_model: nn.Module, - ): - ''' - Sampling draft tokens with support for non-greedy sampling. - - Args: - logits: torch.Tensor - [num_tokens, vocab_size] - Logits produced by the draft model. - draft_model: nn.Module - The draft model. - - Returns: - draft_tokens: torch.Tensor - [batch_size * max_draft_len] - Draft token ids. Flattened. - ''' - - d2t = getattr(draft_model.model, "d2t", None) - draft_tokens = self._draft_sampler_greedy(logits, d2t) - - return draft_tokens - def prepare_1st_drafter_inputs( self, input_ids: torch.LongTensor, @@ -844,15 +1186,24 @@ def prepare_1st_drafter_inputs( spec_metadata: Eagle3OneModelSpecMetadata, draft_model: nn.Module, ): + """Prepare inputs for the first draft model forward. + + Branching: + - Eagle3: applies ``apply_eagle3_fc`` on multi-layer concatenated + hidden states. + - MTP Eagle: uses ``hidden_states`` directly (single last layer); + no FC projection. + """ num_contexts = attn_metadata.num_contexts num_tokens = input_ids.shape[0] - # prepare hidden states - hidden_size_up = spec_metadata.hidden_size * len( - spec_metadata.layers_to_capture) - hidden_states = spec_metadata.hidden_states[:num_tokens, : - hidden_size_up] - hidden_states = draft_model.apply_eagle3_fc(hidden_states) + if not self.is_mtp_eagle: + # Eagle3: project the multi-layer concatenated hidden states. + hidden_size_up = spec_metadata.hidden_size * len( + spec_metadata.layers_to_capture) + hidden_states = spec_metadata.hidden_states[:num_tokens, : + hidden_size_up] + hidden_states = draft_model.apply_eagle3_fc(hidden_states) # context input_ids_ctx = self._prepare_context_input_ids( @@ -873,3 +1224,25 @@ def prepare_1st_drafter_inputs( "attn_metadata": attn_metadata, "spec_metadata": spec_metadata, } + + +class MTPEagleWorker(Eagle3OneModelWorker): + """Backward-compatible alias for ``Eagle3OneModelWorker`` in MTP Eagle mode. + + The constructor matches the historical positional signature + ``(spec_config, model_config, use_separate_draft_kv_cache)`` so callers + that import ``MTPEagleWorker`` from ``mtp.py`` or instantiate it directly + keep working. All logic is inherited from :class:`Eagle3OneModelWorker`. + """ + + def __init__(self, + spec_config, + model_config: Optional[ModelConfig] = None, + use_separate_draft_kv_cache: bool = False): + super().__init__( + spec_config, + mapping=None, + model_config=model_config, + use_separate_draft_kv_cache=use_separate_draft_kv_cache) + # Preserved for callers/tests that still expect this attribute. + self.is_thop = False diff --git a/tensorrt_llm/_torch/speculative/eagle3_dynamic_tree.py b/tensorrt_llm/_torch/speculative/eagle3_dynamic_tree.py index 90afb921d0a8..510abf64a509 100644 --- a/tensorrt_llm/_torch/speculative/eagle3_dynamic_tree.py +++ b/tensorrt_llm/_torch/speculative/eagle3_dynamic_tree.py @@ -180,7 +180,11 @@ def __init__( self, spec_config: "EagleDecodingConfig", mapping, use_separate_draft_kv_cache: bool = False ): """Initialize dynamic-tree specific buffers and helper ops.""" - super().__init__(spec_config, mapping, use_separate_draft_kv_cache) + super().__init__( + spec_config, + mapping=mapping, + use_separate_draft_kv_cache=use_separate_draft_kv_cache, + ) assert self.use_dynamic_tree, ( "Eagle3OneModelDynamicTreeWorker requires use_dynamic_tree=True" ) @@ -453,8 +457,13 @@ def _relocate_kv_eagerly(self, attn_metadata, batch_size): ) @nvtx_range("eagle3_dyn.sample_and_accept_draft_tokens") - def sample_and_accept_draft_tokens(self, logits, attn_metadata, spec_metadata): - """Override to handle dynamic tree verification.""" + def sample_and_accept_draft_tokens(self, input_ids, logits, attn_metadata, spec_metadata): + """Override to handle dynamic tree verification. + + ``input_ids`` is unused here (relaxed acceptance is not supported in + dynamic-tree mode); accepted to match the base class signature. + """ + del input_ids batch_size = attn_metadata.num_seqs num_contexts = attn_metadata.num_contexts num_gens = batch_size - num_contexts diff --git a/tensorrt_llm/_torch/speculative/interface.py b/tensorrt_llm/_torch/speculative/interface.py index c7a8358124cf..eefd06e8179d 100644 --- a/tensorrt_llm/_torch/speculative/interface.py +++ b/tensorrt_llm/_torch/speculative/interface.py @@ -220,7 +220,7 @@ class SpeculativeDecodingMode(IntEnum): AUTO = auto() def is_mtp_one_model(self): - return self == SpeculativeDecodingMode.MTP or self == SpeculativeDecodingMode.MTP_EAGLE_ONE_MODEL + return self == SpeculativeDecodingMode.MTP def is_mtp_eagle_one_model(self): return self == SpeculativeDecodingMode.MTP_EAGLE_ONE_MODEL @@ -236,7 +236,8 @@ def is_eagle3(self): def use_one_engine(self): return self.is_eagle3_one_model() or self.is_mtp_one_model( - ) or self.is_external_drafter() or self.is_sa() + ) or self.is_mtp_eagle_one_model() or self.is_external_drafter( + ) or self.is_sa() def is_eagle3_one_model(self): return self == SpeculativeDecodingMode.EAGLE3_ONE_MODEL @@ -275,28 +276,31 @@ def is_external_drafter(self): return self.is_parallel_draft() or self.is_draft_target_one_model() def without_logits(self): - return self.is_mtp_one_model() or self.is_eagle3_one_model( - ) or self.is_external_drafter() or self.is_sa() + return self.is_mtp_one_model() or self.is_mtp_eagle_one_model( + ) or self.is_eagle3_one_model() or self.is_external_drafter( + ) or self.is_sa() def needs_kv_cache_rewind(self): - return self.is_mtp_one_model() or self.is_eagle3_one_model( - ) or self.is_ngram() or self.is_sa() or self.is_external_drafter() + return self.is_mtp_one_model() or self.is_mtp_eagle_one_model( + ) or self.is_eagle3_one_model() or self.is_ngram() or self.is_sa( + ) or self.is_external_drafter() def support_overlap_scheduler(self): - return self.is_mtp_one_model() or self.is_eagle3_one_model( - ) or self.is_sa() or self.has_draft_model() or self.is_external_drafter( - ) + return self.is_mtp_one_model() or self.is_mtp_eagle_one_model( + ) or self.is_eagle3_one_model() or self.is_sa() or self.has_draft_model( + ) or self.is_external_drafter() def support_guided_decoder(self): return self.is_none() or self.has_spec_drafter() def support_capturable_guided_decoder(self): - return self.is_mtp_one_model() or self.is_eagle3_one_model( - ) or self.is_external_drafter() or self.is_sa() + return self.is_mtp_one_model() or self.is_mtp_eagle_one_model( + ) or self.is_eagle3_one_model() or self.is_external_drafter( + ) or self.is_sa() def support_dynamic_draft_len(self): # TODO: expand to all one-model algorithms - return self.is_eagle3_one_model() + return self.is_eagle3_one_model() or self.is_mtp_eagle_one_model() def has_draft_model(self): return self.is_eagle3() or self.is_draft_target() or self.is_mtp_eagle() @@ -317,7 +321,8 @@ def need_load_draft_weights(self): return self.is_eagle3_one_model() or self.is_external_drafter() def has_spec_decoder(self): - return self.is_mtp_one_model() or self.is_mtp_eagle() or self.is_eagle3( + return self.is_mtp_one_model() or self.is_mtp_eagle_one_model( + ) or self.is_mtp_eagle() or self.is_eagle3( ) or self.is_eagle3_one_model() or self.is_external_drafter( ) or self.is_sa() diff --git a/tensorrt_llm/_torch/speculative/mtp.py b/tensorrt_llm/_torch/speculative/mtp.py index 20181bf29179..93bb5e36d981 100644 --- a/tensorrt_llm/_torch/speculative/mtp.py +++ b/tensorrt_llm/_torch/speculative/mtp.py @@ -3,10 +3,7 @@ from typing import TYPE_CHECKING, List, Optional import torch -import torch.nn.functional as F -from tensorrt_llm._torch.pyexecutor.mamba_cache_manager import \ - MambaHybridCacheManager from tensorrt_llm._utils import prefer_pinned from tensorrt_llm.mapping import Mapping @@ -1141,275 +1138,13 @@ def draft_sampler( return draft_tokens -class MTPEagleWorker(MTPWorker): +# ``MTPEagleWorker`` moved to ``eagle3.py`` as part of the Eagle3/MTP-Eagle +# merge (TRTLLM-11508). Preserve the historical import path so external +# callers like ``from tensorrt_llm._torch.speculative.mtp import MTPEagleWorker`` +# keep working without a hard dependency cycle. +def __getattr__(name): + if name == "MTPEagleWorker": + from .eagle3 import MTPEagleWorker + return MTPEagleWorker + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") - def __init__(self, - spec_config: "MTPDecodingConfig", - model_config: Optional[ModelConfig] = None, - use_separate_draft_kv_cache: bool = False): - super().__init__(spec_config, model_config, use_separate_draft_kv_cache) - self.model_config = model_config - self.mtp_num_modules = spec_config.max_draft_len - self._is_mamba_hybrid_cache = None - - @torch.compile(options={"max-autotune": True}) - def update_draft_tokens(self, next_draft_tokens, new_draft_token, - hidden_states, gather_ids, inputs): - next_draft_tokens.append(new_draft_token) - # update inputs - hidden_states = hidden_states[gather_ids] - position_ids = ( - _select_mtp_position_ids(inputs["position_ids"], gather_ids) + 1) - return hidden_states, position_ids - - @torch.compile(options={"max-autotune": True}) - def prepare_position_ids_and_last_tokens(self, position_ids, seq_lens_cuda): - position_ids = position_ids.squeeze(0) - last_tokens_idx = torch.cumsum(seq_lens_cuda, dim=0, - dtype=torch.long) - 1 - return position_ids, last_tokens_idx - - def forward( - self, - input_ids, - position_ids, - hidden_states, - logits, - attn_metadata, - spec_metadata, - draft_model, - resource_manager=None, - ): - - batch_size = attn_metadata.num_seqs - num_contexts = attn_metadata.num_contexts - num_gens = batch_size - num_contexts - - raw_logits = logits - - self._execute_guided_decoder_if_present(logits) - - # Sample and verify draft tokens - accepted_tokens, num_accepted_tokens = self.sample_and_accept_draft_tokens( - input_ids, logits, spec_metadata, attn_metadata) - - if self._is_mamba_hybrid_cache is None: - self._is_mamba_hybrid_cache = isinstance( - attn_metadata.kv_cache_manager, MambaHybridCacheManager) - if num_gens > 0 and self._is_mamba_hybrid_cache: - attn_metadata.kv_cache_manager.update_mamba_states( - attn_metadata=attn_metadata, - num_accepted_tokens=num_accepted_tokens, - state_indices=attn_metadata.mamba_metadata.state_indices) - - # Save the old attn_metadata and spec_metadata - self._prepare_attn_metadata_for_spec_dec(attn_metadata) - - position_ids, last_tokens_idx = self.prepare_position_ids_and_last_tokens( - position_ids, attn_metadata.seq_lens_cuda) - inputs = self.prepare_drafter_inputs(input_ids=input_ids, - position_ids=position_ids, - last_tokens_idx=last_tokens_idx, - hidden_states=hidden_states, - accepted_tokens=accepted_tokens, - attn_metadata=attn_metadata, - spec_metadata=spec_metadata) - - # Get the draft KV cache manager if using separate layouts - draft_kv_cache_manager = self.get_draft_kv_cache_manager( - resource_manager) - - # Predict draft tokens - next_draft_tokens = [] - with self.draft_kv_cache_context(attn_metadata, draft_kv_cache_manager): - for i in range(self.mtp_num_modules): - if i == 0: - hidden_states = draft_model.mtp_layers[0]( - embed_tokens=draft_model.embed_tokens, - all_rank_num_tokens=spec_metadata.all_rank_num_tokens, - **inputs) - - start_ids_gen = ( - spec_metadata.batch_indices_cuda[:num_gens] * - (self.mtp_num_modules + 1)).long() - gather_ids_gen = (start_ids_gen + - num_accepted_tokens[num_contexts:] - 1 + - attn_metadata.num_ctx_tokens) - gather_ids = torch.concat( - [last_tokens_idx[:num_contexts], gather_ids_gen], dim=0) - else: - hidden_states = draft_model.mtp_layers[0]( - embed_tokens=draft_model.embed_tokens, - all_rank_num_tokens=spec_metadata. - subseq_all_rank_num_tokens, - **inputs) - - # All of the seq_len are 1, use batch_indices_cuda as gather_ids - gather_ids = spec_metadata.batch_indices_cuda[:batch_size] - - if self.guided_decoder is not None: - new_tokens = inputs["input_ids"][gather_ids] - self.guided_decoder.add_draft_batch(new_tokens, - num_accepted_tokens, - draft_step=i) - if self.model_config.mapping.enable_attention_dp and \ - getattr(self.model_config.mapping, 'enable_lm_head_tp_in_adp', False): - hidden_states_gathered = hidden_states[gather_ids] - token_count = hidden_states_gathered.view( - -1, hidden_states_gathered.shape[-1]).shape[0] - max_num_requests = spec_metadata.max_num_requests - pad_len = max_num_requests - token_count - if pad_len > 0: - padded_hidden_states = F.pad( - hidden_states_gathered.view( - -1, hidden_states_gathered.shape[-1]), - (0, 0, 0, pad_len), - mode="constant", - value=0) - elif pad_len == 0: - padded_hidden_states = hidden_states_gathered.view( - -1, hidden_states_gathered.shape[-1]) - else: - raise ValueError( - "In MTPEagleWorker.forward(), token_count > max_num_requests, which is not supported" - ) - logits = draft_model.mtp_layers[0].shared_head( - padded_hidden_states, draft_model.lm_head, - attn_metadata, True) - else: - logits = draft_model.mtp_layers[0].shared_head( - hidden_states[gather_ids], draft_model.lm_head, - attn_metadata, True) - if self.guided_decoder is not None: - self.guided_decoder.execute_draft_batch(logits, - draft_step=i) - - if self.model_config.mapping.enable_attention_dp and \ - getattr(self.model_config.mapping, 'enable_lm_head_tp_in_adp', False): - mapping_lm_head_tp = draft_model.mtp_layers[ - 0].shared_head.mapping_lm_head_tp - new_draft_token = self.draft_sampler( - logits, mapping_lm_head_tp) - new_draft_token = new_draft_token[:token_count] - else: - new_draft_token = self.draft_sampler(logits) - - hidden_states, position_ids = self.update_draft_tokens( - next_draft_tokens, new_draft_token, hidden_states, - gather_ids, inputs) - # update attn_metadata - if i == 0: - attn_metadata._seq_lens[:batch_size].fill_(1) - attn_metadata._seq_lens_cuda[:batch_size].fill_(1) - attn_metadata.on_update() - # cannot run generation if there is no kv cache - has_kv_cache = inputs[ - "attn_metadata"].kv_cache_manager is not None - if has_kv_cache: - attn_metadata.host_request_types[:attn_metadata. - num_contexts].fill_(1) - attn_metadata.num_contexts = 0 - # update kv_lens_cuda - if hasattr(attn_metadata, 'kv_lens_cuda'): - attn_metadata.kv_lens_cuda[num_contexts:batch_size] -= ( - self.mtp_num_modules - - num_accepted_tokens[num_contexts:]) - attn_metadata.kv_lens_cuda[:num_contexts] += 1 - # update metadata for flash mla - if has_kv_cache and num_contexts > 0 and attn_metadata.enable_flash_mla: - reorder_block_ids_per_seq = torch.cat([ - attn_metadata. - kv_block_ids_per_seq[num_contexts:batch_size], - attn_metadata.kv_block_ids_per_seq[:num_contexts] - ]) - attn_metadata.block_ids_per_seq[:batch_size, :].copy_( - reorder_block_ids_per_seq, non_blocking=True) - # update metadata - # some attention metadata needs to be updated when changing seq_lens/kv_lens - attn_metadata.update_for_spec_dec() - # Disable spec-dec mode for subsequent iterations (i>0) - # as draft model only infer 1 token for the subsequent inference. - attn_metadata.use_spec_decoding = False - elif hasattr(attn_metadata, 'kv_lens_cuda'): - # update kv_lens_cuda - attn_metadata.kv_lens_cuda[:batch_size] += 1 - - # update metadata - # some attention metadata needs to be updated when changing kv_lens - attn_metadata.update_for_spec_dec() - inputs = { - "input_ids": new_draft_token, - "position_ids": position_ids, - "hidden_states": hidden_states, - "attn_metadata": attn_metadata, - } - - # restore attn_metadata to support cuda graph - self._restore_attn_metadata_from_spec_dec(attn_metadata) - attn_metadata.use_spec_decoding = True - - # Override with SA draft tokens after all MTP layers have run, - # so that MTP layers never see SA tokens in their inputs. - # Must happen before stacking since next_draft_tokens is still a list. - if self.sa_enhancer is not None: - stacked = torch.stack(next_draft_tokens, dim=1) - gen_draft_tokens = stacked[num_contexts:] - gen_draft_tokens = self.sa_enhancer.maybe_override_all_draft_tokens( - gen_draft_tokens) - stacked[num_contexts:] = gen_draft_tokens - next_draft_tokens = [stacked[:, i] for i in range(stacked.shape[1])] - - next_draft_tokens, next_new_tokens = self._prepare_next_tokens( - next_draft_tokens, accepted_tokens, spec_metadata, batch_size, - num_accepted_tokens) - - return { - 'logits': raw_logits, - 'new_tokens': accepted_tokens, - 'new_tokens_lens': num_accepted_tokens, - 'next_draft_tokens': next_draft_tokens, - 'next_new_tokens': next_new_tokens - } - - @torch.compile(options={"max-autotune": True}) - def _prepare_next_tokens(self, next_draft_tokens, accepted_tokens, - spec_metadata, batch_size, num_accepted_tokens): - """ - Stack draft tokens and prepare next_new_tokens for overlap scheduler. - """ - next_draft_tokens = torch.stack(next_draft_tokens, dim=1) - next_new_tokens = self._prepare_next_new_tokens( - accepted_tokens, next_draft_tokens, - spec_metadata.batch_indices_cuda, batch_size, num_accepted_tokens) - return next_draft_tokens, next_new_tokens - - @torch.compile(options={"max-autotune": True}) - def prepare_drafter_inputs( - self, - input_ids: torch.IntTensor, - position_ids: torch.IntTensor, - last_tokens_idx: torch.LongTensor, - hidden_states: torch.Tensor, - accepted_tokens: torch.Tensor, - attn_metadata: AttentionMetadata, - spec_metadata: MTPSpecMetadata, - ): - num_contexts = attn_metadata.num_contexts - - # context - input_ids_ctx = self._prepare_context_input_ids( - input_ids, attn_metadata.num_ctx_tokens, last_tokens_idx, - accepted_tokens, num_contexts) - - # generation - input_ids_gen = accepted_tokens[num_contexts:, :].flatten() - - # get draft inputs - input_ids = torch.concat([input_ids_ctx, input_ids_gen], dim=0) - - return { - "input_ids": input_ids, - "position_ids": position_ids, - "hidden_states": hidden_states, - "attn_metadata": attn_metadata, - } diff --git a/tensorrt_llm/_torch/speculative/utils.py b/tensorrt_llm/_torch/speculative/utils.py index 68ef347b0533..bbce5c33e747 100644 --- a/tensorrt_llm/_torch/speculative/utils.py +++ b/tensorrt_llm/_torch/speculative/utils.py @@ -20,11 +20,10 @@ from .eagle3 import (Eagle3OneModelDynamicTreeResourceManager, Eagle3OneModelSampler, Eagle3OneModelSpecMetadata, Eagle3OneModelWorker, Eagle3ResourceManager, - Eagle3SpecMetadata) + Eagle3SpecMetadata, MTPEagleWorker) from .eagle3_dynamic_tree import Eagle3OneModelDynamicTreeWorker from .model_drafter import ModelDrafter -from .mtp import (MTPEagleWorker, MTPHiddenStatesManager, MTPSampler, - MTPSpecMetadata, MTPWorker) +from .mtp import MTPHiddenStatesManager, MTPSampler, MTPSpecMetadata, MTPWorker from .ngram import NGramDrafter, NGramPoolManager from .pard import PARDSpecMetadata, PARDWorker from .sa_worker import SASampler, SASpecMetadata, SAWorker @@ -43,6 +42,24 @@ def get_spec_metadata(spec_config, use_rejection_sampling = getattr(spec_config, "use_rejection_sampling", False) vocab_size = getattr(model_config, "vocab_size", 0) + if spec_config.spec_dec_mode.is_mtp_eagle_one_model(): + # MTP Eagle one-model now reuses Eagle3 one-model metadata so it + # picks up the unified worker, sampler, and slot_ids/subseq plumbing. + # Capture only the final layer's hidden state (MTP Eagle behavior). + return Eagle3OneModelSpecMetadata( + max_draft_len=spec_config.max_draft_len, + max_total_draft_tokens=spec_config.tokens_per_gen_step - 1, + spec_dec_mode=spec_config.spec_dec_mode, + max_num_requests=max_num_requests, + num_layers=model_config.num_hidden_layers, + hidden_size=model_config.hidden_size, + max_num_tokens=max_num_tokens, + layers_to_capture={model_config.num_hidden_layers - 1}, + allow_advanced_sampling=spec_config.allow_advanced_sampling, + use_rejection_sampling=use_rejection_sampling, + vocab_size=vocab_size, + spec_resource_manager=spec_resource_manager, + ) if spec_config.spec_dec_mode.is_mtp_one_model(): return MTPSpecMetadata( max_draft_len=spec_config.max_draft_len, @@ -185,11 +202,16 @@ def get_spec_resource_manager(model_engine, draft_model_engine=None): sa_manager = SuffixAutomatonManager(sa_cfg, max_num_requests, max_seq_len) if spec_config.use_relaxed_acceptance_for_thinking or sa_manager is not None: - return MTPHiddenStatesManager( + # Unified resource manager: the unified worker reads + # ``relaxed_delta_pool`` from ``Eagle3ResourceManager`` (mirrors the + # pool ``MTPHiddenStatesManager`` used to provide). + return Eagle3ResourceManager( spec_config, model_config.torch_dtype, model_config.hidden_size, max_num_requests, + max_seq_len, + max_num_tokens, sa_manager=sa_manager, ) else: @@ -263,6 +285,9 @@ def get_spec_decoder( sampler_args: TorchSampler.Args, spec_config: "DecodingBaseConfig", ): + if spec_config.spec_dec_mode.is_mtp_eagle_one_model(): + # MTP Eagle one-model now uses the same sampler as Eagle3 one-model. + return Eagle3OneModelSampler(sampler_args, spec_config=spec_config) if spec_config.spec_dec_mode.is_mtp_one_model(): return MTPSampler(sampler_args, nextn=spec_config.max_draft_len) if spec_config.spec_dec_mode.is_eagle3( @@ -314,6 +339,8 @@ def get_spec_drafter(model_engine, def get_num_spec_layers(spec_config): + if spec_config.spec_dec_mode.is_mtp_eagle_one_model(): + return 1 if spec_config.spec_dec_mode.is_mtp_one_model(): return spec_config.num_nextn_predict_layers if spec_config.spec_dec_mode.is_eagle3_one_model(): @@ -330,14 +357,18 @@ def get_spec_worker(spec_config, if spec_dec_mode.is_mtp_vanilla(): return MTPWorker(spec_config, model_config, use_separate_draft_kv_cache) if spec_dec_mode.is_mtp_eagle_one_model(): - return MTPEagleWorker(spec_config, model_config, - use_separate_draft_kv_cache) + return MTPEagleWorker( + spec_config, + model_config=model_config, + use_separate_draft_kv_cache=use_separate_draft_kv_cache) if spec_dec_mode.is_eagle3_one_model(): if getattr(spec_config, 'use_dynamic_tree', False): return Eagle3OneModelDynamicTreeWorker(spec_config, mapping, use_separate_draft_kv_cache) - return Eagle3OneModelWorker(spec_config, mapping, - use_separate_draft_kv_cache) + return Eagle3OneModelWorker( + spec_config, + mapping=mapping, + use_separate_draft_kv_cache=use_separate_draft_kv_cache) if spec_dec_mode.is_pard(): return PARDWorker(spec_config, mapping, use_separate_draft_kv_cache) if spec_dec_mode.is_dflash(): diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 2c3aad405890..cc8a4ca44ce3 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -1179,6 +1179,34 @@ class EagleDecodingConfig(DecodingBaseConfig): default="llama3", description="The model architecture of the eagle3 model.") + # Relaxed acceptance settings (mirrors MTPDecodingConfig for thinking models) + use_relaxed_acceptance_for_thinking: bool = Field( + default=False, + description= + "Enable relaxed acceptance during thinking phase for reasoning models. " + "Accepts draft tokens matching any top-K candidate instead of exact top-1." + ) + relaxed_topk: int = Field( + default=1, + description= + "Number of top candidate tokens to consider for relaxed acceptance. " + "Draft token is accepted if it matches any of these.") + relaxed_delta: float = Field( + default=0., + description= + "Probability threshold for relaxed acceptance. Only candidates with " + "prob >= (top-1 prob - delta) are kept.") + begin_thinking_phase_token: int = Field( + default=128798, + description= + "Token ID marking start of thinking phase. Relaxed acceptance only applies within this phase." + ) + end_thinking_phase_token: int = Field( + default=128799, + description= + "Token ID marking end of thinking phase. Strict acceptance resumes after this." + ) + @field_validator('eagle_choices', mode='before') @classmethod def validate_eagle_choices(cls, v): From 1be4656939288d25793a81d3b21118258fa958ba Mon Sep 17 00:00:00 2001 From: ZhaoyangWang Date: Mon, 11 May 2026 02:30:57 -0700 Subject: [PATCH 02/11] [TRTLLM-11508][fix] unify and restore all_rank_num_tokens Move the per-step ``all_rank_num_tokens`` plumbing into the draft model itself so the unified Eagle3OneModelWorker no longer mutates ``attn_metadata`` on the way into the draft loop. - Eagle3DraftModel.forward takes an optional ``all_rank_num_tokens`` kwarg and wraps its body in try/finally that restores ``attn_metadata.all_rank_num_tokens`` on exit. - _run_draft_forward in eagle3.py passes ``all_rank_num_tokens`` via ``inputs`` for Eagle3 (kwarg to Eagle3DraftModel) and as a direct kwarg to ``mtp_layers[0]`` for MTP Eagle; the worker no longer needs the old fallback parameter. - _get_step_all_rank_num_tokens reads only from spec_metadata (all_rank_num_tokens at step 0, subseq_all_rank_num_tokens otherwise). - model_engine.py populates ``spec_metadata.subseq_all_rank_num_tokens`` for both Eagle3 one-model and MTP-eagle one-model at all three sites that allgather per-rank token counts. Signed-off-by: ZhaoyangWang --- .../_torch/models/modeling_speculative.py | 85 +++++++++++-------- .../_torch/pyexecutor/model_engine.py | 14 +++ tensorrt_llm/_torch/speculative/eagle3.py | 45 ++++------ 3 files changed, 80 insertions(+), 64 deletions(-) diff --git a/tensorrt_llm/_torch/models/modeling_speculative.py b/tensorrt_llm/_torch/models/modeling_speculative.py index efaa2f294e4b..a8c5ecce59ed 100755 --- a/tensorrt_llm/_torch/models/modeling_speculative.py +++ b/tensorrt_llm/_torch/models/modeling_speculative.py @@ -402,53 +402,66 @@ def forward( inputs_embeds: Optional[torch.FloatTensor] = None, spec_metadata: Optional[SpecMetadata] = None, hidden_states: Optional[torch.Tensor] = None, + all_rank_num_tokens: Optional[List[int]] = None, ) -> torch.Tensor: - assert self.embed_tokens is not None + # When ``all_rank_num_tokens`` is supplied the caller wants this draft + # forward to run with a different attention-DP token distribution + # (e.g. the worker's per-step value); restore the original on exit so + # the next call sees the same attn_metadata it had on entry. + previous_all_rank_num_tokens = attn_metadata.all_rank_num_tokens + if all_rank_num_tokens is not None: + attn_metadata.all_rank_num_tokens = all_rank_num_tokens - if (input_ids is None) ^ (inputs_embeds is not None): - raise ValueError( - "You cannot specify both input_ids and inputs_embeds at the same time, and must specify either one" - ) - - if inputs_embeds is None: - inputs_embeds = self.embed_tokens(input_ids).to(self.dtype) + try: + assert self.embed_tokens is not None - assert hidden_states is not None - # NOTE: If hidden states from the target model have to be concatenated, - # ideally, we expect that to happen outside the model definition. This - # helps us avoid data-dependent control flow and gives us better CUDA - # graph coverage. - if self._eh_proj_before_attn: - input_embeds = self.enorm(inputs_embeds) - hidden_states = torch.cat([input_embeds, hidden_states], dim=-1) - hidden_states = self.eh_proj(hidden_states) + if (input_ids is None) ^ (inputs_embeds is not None): + raise ValueError( + "You cannot specify both input_ids and inputs_embeds at the same time, and must specify either one" + ) - residual = None - if self.num_layers > 1: - for layer in self.midlayer: - if residual is not None: - hidden_states = hidden_states + residual - hidden_states, residual = layer( + if inputs_embeds is None: + inputs_embeds = self.embed_tokens(input_ids).to(self.dtype) + + assert hidden_states is not None + # NOTE: If hidden states from the target model have to be concatenated, + # ideally, we expect that to happen outside the model definition. This + # helps us avoid data-dependent control flow and gives us better CUDA + # graph coverage. + if self._eh_proj_before_attn: + input_embeds = self.enorm(inputs_embeds) + hidden_states = torch.cat([input_embeds, hidden_states], dim=-1) + hidden_states = self.eh_proj(hidden_states) + + residual = None + if self.num_layers > 1: + for layer in self.midlayer: + if residual is not None: + hidden_states = hidden_states + residual + hidden_states, residual = layer( + position_ids=position_ids, + embeds=inputs_embeds, + hidden_states=hidden_states, + attn_metadata=attn_metadata, + spec_metadata=spec_metadata, + ) + else: + hidden_states, residual = self.midlayer( position_ids=position_ids, embeds=inputs_embeds, hidden_states=hidden_states, attn_metadata=attn_metadata, spec_metadata=spec_metadata, ) - else: - hidden_states, residual = self.midlayer( - position_ids=position_ids, - embeds=inputs_embeds, - hidden_states=hidden_states, - attn_metadata=attn_metadata, - spec_metadata=spec_metadata, - ) - hidden_states, hidden_states_to_save = self.norm( - hidden_states, residual) - if self._return_hidden_post_norm: - return hidden_states, hidden_states - return hidden_states, hidden_states_to_save + hidden_states, hidden_states_to_save = self.norm( + hidden_states, residual) + if self._return_hidden_post_norm: + return hidden_states, hidden_states + return hidden_states, hidden_states_to_save + finally: + if all_rank_num_tokens is not None: + attn_metadata.all_rank_num_tokens = previous_all_rank_num_tokens # We use Llama3 as the base architecture for EAGLE3 draft layers diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index ae88b6d62e80..0f9e9f637c77 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -2081,6 +2081,14 @@ def _prepare_incremental_update_metadata( spec_metadata.all_rank_num_seqs = [ item[1] for item in all_rank_num_tokens ] + # Both Eagle3 one-model and MTP-eagle one-model use + # subseq_all_rank_num_tokens for draft loop iterations i>0 + # (per-sequence counts since each sequence contributes one + # token per iteration). + if (spec_metadata.spec_dec_mode.is_mtp_eagle_one_model() + or spec_metadata.spec_dec_mode.is_eagle3_one_model()): + spec_metadata.subseq_all_rank_num_tokens = ( + spec_metadata.all_rank_num_seqs) # Set iteration states - batch dictionary updates self.iter_states.update({ @@ -3309,6 +3317,9 @@ def previous_seq_slots_device(): all_rank_num_seqs = [item[1] for item in all_rank_num_tokens] spec_metadata.all_rank_num_tokens = spec_all_rank_num_tokens spec_metadata.all_rank_num_seqs = all_rank_num_seqs + if (spec_metadata.spec_dec_mode.is_mtp_eagle_one_model() + or spec_metadata.spec_dec_mode.is_eagle3_one_model()): + spec_metadata.subseq_all_rank_num_tokens = all_rank_num_seqs if mm_token_indices is not None: mask = torch.ones(total_num_tokens, dtype=torch.bool) @@ -3480,6 +3491,9 @@ def _prepare_tp_inputs_no_cache( attn_metadata.all_rank_num_tokens = attn_all_rank_num_tokens spec_metadata.all_rank_num_tokens = spec_all_rank_num_tokens spec_metadata.all_rank_num_seqs = all_rank_num_seqs + if (spec_metadata.spec_dec_mode.is_mtp_eagle_one_model() + or spec_metadata.spec_dec_mode.is_eagle3_one_model()): + spec_metadata.subseq_all_rank_num_tokens = all_rank_num_seqs else: all_rank_num_tokens = self.dist.tp_cp_allgather( attn_metadata.num_tokens) diff --git a/tensorrt_llm/_torch/speculative/eagle3.py b/tensorrt_llm/_torch/speculative/eagle3.py index 46dd3faec7d5..582494da4f73 100644 --- a/tensorrt_llm/_torch/speculative/eagle3.py +++ b/tensorrt_llm/_torch/speculative/eagle3.py @@ -765,13 +765,12 @@ def _forward_linear_draft_loop(self, inputs, attn_metadata, spec_metadata, with self.draft_kv_cache_context(attn_metadata, draft_kv_cache_manager): for i in range(runtime_draft_len): - # Run draft model (mode-specific via helper). For Eagle3 the - # helper mutates ``attn_metadata.all_rank_num_tokens`` to the - # right per-step value; for MTP Eagle the value is passed as - # a kwarg to ``mtp_layers[0]`` directly. + # Run draft model (mode-specific via helper). The helper + # passes ``all_rank_num_tokens`` as a kwarg so the draft model + # handles save/restore internally (Eagle3DraftModel.forward + # uses try/finally); attn_metadata is left untouched here. hidden_states, hidden_states_to_save = self._run_draft_forward( - draft_model, inputs, spec_metadata, i, - original_all_rank_num_tokens) + draft_model, inputs, spec_metadata, i) # Compute gather_ids: on the first draft step each generation # request may have accepted multiple tokens, so we index into @@ -948,32 +947,27 @@ def _forward_linear_draft_loop(self, inputs, attn_metadata, spec_metadata, return next_draft_tokens - def _get_step_all_rank_num_tokens(self, spec_metadata, step_idx: int, - fallback): + def _get_step_all_rank_num_tokens(self, spec_metadata, step_idx: int): """Pick the right ``all_rank_num_tokens`` for this draft iteration. - Step 0 uses ``spec_metadata.all_rank_num_tokens`` (or the fallback - captured from attn_metadata before the loop); subsequent steps use - ``spec_metadata.subseq_all_rank_num_tokens`` since every sequence + Step 0 uses ``spec_metadata.all_rank_num_tokens``; subsequent steps + use ``spec_metadata.subseq_all_rank_num_tokens`` since every sequence contributes a single token per iteration. """ - if step_idx == 0: - return (spec_metadata.all_rank_num_tokens - if spec_metadata.all_rank_num_tokens is not None else - fallback) - return spec_metadata.subseq_all_rank_num_tokens + return (spec_metadata.all_rank_num_tokens if step_idx == 0 else + spec_metadata.subseq_all_rank_num_tokens) def _run_draft_forward(self, draft_model, inputs, spec_metadata, - step_idx: int, original_all_rank_num_tokens): + step_idx: int): """Invoke the draft model for one iteration, branching on mode. - For MTP Eagle, ``all_rank_num_tokens`` is passed directly as a kwarg - to ``mtp_layers[0]``. For Eagle3, it is set on ``attn_metadata`` so - the EAGLE draft model picks it up; the outer ``forward()`` is - responsible for restoring the original value after the loop. + ``all_rank_num_tokens`` is passed as a kwarg in both modes. For MTP + Eagle it goes directly to ``mtp_layers[0]``; for Eagle3 it goes to + ``Eagle3DraftModel.forward`` which guards it with a try/finally so + attn_metadata sees the original value on return. """ all_rank_num_tokens = self._get_step_all_rank_num_tokens( - spec_metadata, step_idx, original_all_rank_num_tokens) + spec_metadata, step_idx) if self.is_mtp_eagle: hidden_states = draft_model.mtp_layers[0]( @@ -982,12 +976,7 @@ def _run_draft_forward(self, draft_model, inputs, spec_metadata, **inputs) return hidden_states, None - # Eagle3: route per-step token counts through attn_metadata. Commit 2 - # of this refactor migrates this to a kwarg + try/finally inside - # ``Eagle3DraftModel.forward``. - attn_metadata = inputs["attn_metadata"] - if all_rank_num_tokens is not None: - attn_metadata.all_rank_num_tokens = all_rank_num_tokens + inputs["all_rank_num_tokens"] = all_rank_num_tokens hidden_states, hidden_states_to_save = draft_model.model(**inputs) return hidden_states, hidden_states_to_save From 842e0141d4c758dac29a75a47a03330b7ae69d82 Mon Sep 17 00:00:00 2001 From: ZhaoyangWang Date: Tue, 12 May 2026 03:05:20 -0700 Subject: [PATCH 03/11] clean code Signed-off-by: ZhaoyangWang --- tensorrt_llm/_torch/speculative/eagle3.py | 17 ++++++----------- tensorrt_llm/_torch/speculative/mtp.py | 2 -- tensorrt_llm/_torch/speculative/utils.py | 6 ++---- 3 files changed, 8 insertions(+), 17 deletions(-) diff --git a/tensorrt_llm/_torch/speculative/eagle3.py b/tensorrt_llm/_torch/speculative/eagle3.py index 582494da4f73..7e8ebcd1071d 100644 --- a/tensorrt_llm/_torch/speculative/eagle3.py +++ b/tensorrt_llm/_torch/speculative/eagle3.py @@ -759,9 +759,8 @@ def _forward_linear_draft_loop(self, inputs, attn_metadata, spec_metadata, num_gens = batch_size - num_contexts next_draft_tokens = [] draft_logits_list = [] - last_tokens_idx = torch.cumsum(attn_metadata.seq_lens_cuda, - dim=0, - dtype=torch.long) - 1 + last_tokens_idx = torch.cumsum( + attn_metadata.seq_lens_cuda, dim=0, dtype=torch.long) - 1 with self.draft_kv_cache_context(attn_metadata, draft_kv_cache_manager): for i in range(runtime_draft_len): @@ -954,8 +953,8 @@ def _get_step_all_rank_num_tokens(self, spec_metadata, step_idx: int): use ``spec_metadata.subseq_all_rank_num_tokens`` since every sequence contributes a single token per iteration. """ - return (spec_metadata.all_rank_num_tokens if step_idx == 0 else - spec_metadata.subseq_all_rank_num_tokens) + return (spec_metadata.all_rank_num_tokens + if step_idx == 0 else spec_metadata.subseq_all_rank_num_tokens) def _run_draft_forward(self, draft_model, inputs, spec_metadata, step_idx: int): @@ -1156,12 +1155,8 @@ def sample_and_accept_draft_tokens( if logits.dim() == 1: logits = logits.unsqueeze(0) draft_tokens = spec_metadata.draft_tokens.reshape( - num_gens, - runtime_draft_len) if num_gens > 0 else torch.empty( - 0, - runtime_draft_len, - dtype=torch.int, - device=logits.device) + num_gens, runtime_draft_len) if num_gens > 0 else torch.empty( + 0, runtime_draft_len, dtype=torch.int, device=logits.device) return self._accept_draft_tokens(logits, draft_tokens, num_contexts, batch_size, spec_metadata) diff --git a/tensorrt_llm/_torch/speculative/mtp.py b/tensorrt_llm/_torch/speculative/mtp.py index 93bb5e36d981..3f182f7b3442 100644 --- a/tensorrt_llm/_torch/speculative/mtp.py +++ b/tensorrt_llm/_torch/speculative/mtp.py @@ -9,7 +9,6 @@ from ..attention_backend import AttentionMetadata from ..distributed.ops import allgather -from ..model_config import ModelConfig from ..pyexecutor.llm_request import LlmRequest from ..pyexecutor.resource_manager import BaseResourceManager, SlotManager from ..pyexecutor.sampler import TorchSampler @@ -1147,4 +1146,3 @@ def __getattr__(name): from .eagle3 import MTPEagleWorker return MTPEagleWorker raise AttributeError(f"module {__name__!r} has no attribute {name!r}") - diff --git a/tensorrt_llm/_torch/speculative/utils.py b/tensorrt_llm/_torch/speculative/utils.py index bbce5c33e747..7d28b5a5634a 100644 --- a/tensorrt_llm/_torch/speculative/utils.py +++ b/tensorrt_llm/_torch/speculative/utils.py @@ -357,10 +357,8 @@ def get_spec_worker(spec_config, if spec_dec_mode.is_mtp_vanilla(): return MTPWorker(spec_config, model_config, use_separate_draft_kv_cache) if spec_dec_mode.is_mtp_eagle_one_model(): - return MTPEagleWorker( - spec_config, - model_config=model_config, - use_separate_draft_kv_cache=use_separate_draft_kv_cache) + return MTPEagleWorker(spec_config, model_config, + use_separate_draft_kv_cache) if spec_dec_mode.is_eagle3_one_model(): if getattr(spec_config, 'use_dynamic_tree', False): return Eagle3OneModelDynamicTreeWorker(spec_config, mapping, From d42764b7c3e42b1442009ccd21b9e8f8466325a7 Mon Sep 17 00:00:00 2001 From: ZhaoyangWang Date: Tue, 12 May 2026 04:13:02 -0700 Subject: [PATCH 04/11] fix some issue Signed-off-by: ZhaoyangWang --- .../_torch/models/modeling_exaone_moe.py | 15 ++++++++ .../_torch/models/modeling_speculative.py | 6 ++-- tensorrt_llm/_torch/speculative/eagle3.py | 3 +- .../_torch/speculative/eagle3_dynamic_tree.py | 8 +++++ tensorrt_llm/_torch/speculative/interface.py | 15 ++++++++ tensorrt_llm/llmapi/llm_args.py | 35 +++++++++++++------ 6 files changed, 67 insertions(+), 15 deletions(-) diff --git a/tensorrt_llm/_torch/models/modeling_exaone_moe.py b/tensorrt_llm/_torch/models/modeling_exaone_moe.py index 25b578caf5e7..533a542cc731 100644 --- a/tensorrt_llm/_torch/models/modeling_exaone_moe.py +++ b/tensorrt_llm/_torch/models/modeling_exaone_moe.py @@ -1,3 +1,18 @@ +# SPDX-FileCopyrightText: Copyright (c) 2022-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + import math import os from typing import Dict, List, Optional, Tuple diff --git a/tensorrt_llm/_torch/models/modeling_speculative.py b/tensorrt_llm/_torch/models/modeling_speculative.py index a8c5ecce59ed..8c2ab0c8e42e 100755 --- a/tensorrt_llm/_torch/models/modeling_speculative.py +++ b/tensorrt_llm/_torch/models/modeling_speculative.py @@ -413,14 +413,13 @@ def forward( attn_metadata.all_rank_num_tokens = all_rank_num_tokens try: - assert self.embed_tokens is not None - if (input_ids is None) ^ (inputs_embeds is not None): raise ValueError( "You cannot specify both input_ids and inputs_embeds at the same time, and must specify either one" ) if inputs_embeds is None: + assert self.embed_tokens is not None inputs_embeds = self.embed_tokens(input_ids).to(self.dtype) assert hidden_states is not None @@ -645,14 +644,13 @@ def forward( spec_metadata: SpecMetadata | None = None, hidden_states: torch.Tensor | None = None, ) -> torch.Tensor: - assert self.embed_tokens is not None - if (input_ids is None) ^ (inputs_embeds is not None): raise ValueError( "You cannot specify both input_ids and inputs_embeds at the same time, and must specify either one" ) if inputs_embeds is None: + assert self.embed_tokens is not None inputs_embeds = self.embed_tokens(input_ids).to(self.dtype) assert hidden_states is not None diff --git a/tensorrt_llm/_torch/speculative/eagle3.py b/tensorrt_llm/_torch/speculative/eagle3.py index 7e8ebcd1071d..1d27a0756471 100644 --- a/tensorrt_llm/_torch/speculative/eagle3.py +++ b/tensorrt_llm/_torch/speculative/eagle3.py @@ -483,7 +483,8 @@ def prepare(self): else: self.num_tokens -= (self.num_generations) * self.max_draft_len - if self.spec_resource_manager is not None: + if getattr(self.spec_resource_manager, "slot_manager", + None) is not None: # Populate slot_ids for all requests in this batch. Used by relaxed # acceptance (relaxed_delta_pool indexing), mirroring the pattern # in MTPSpecMetadata.prepare(). diff --git a/tensorrt_llm/_torch/speculative/eagle3_dynamic_tree.py b/tensorrt_llm/_torch/speculative/eagle3_dynamic_tree.py index 510abf64a509..47376001166d 100644 --- a/tensorrt_llm/_torch/speculative/eagle3_dynamic_tree.py +++ b/tensorrt_llm/_torch/speculative/eagle3_dynamic_tree.py @@ -185,6 +185,14 @@ def __init__( mapping=mapping, use_separate_draft_kv_cache=use_separate_draft_kv_cache, ) + if ( + getattr(spec_config, "use_relaxed_acceptance_for_thinking", False) + or getattr(spec_config, "sa_config", None) is not None + ): + raise ValueError( + "Dynamic tree mode does not support relaxed acceptance or " + "suffix-automaton enhancement." + ) assert self.use_dynamic_tree, ( "Eagle3OneModelDynamicTreeWorker requires use_dynamic_tree=True" ) diff --git a/tensorrt_llm/_torch/speculative/interface.py b/tensorrt_llm/_torch/speculative/interface.py index eefd06e8179d..19aa2062812e 100644 --- a/tensorrt_llm/_torch/speculative/interface.py +++ b/tensorrt_llm/_torch/speculative/interface.py @@ -1,3 +1,18 @@ +# SPDX-FileCopyrightText: Copyright (c) 2022-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + import copy import os from abc import ABC, abstractmethod diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index cc8a4ca44ce3..b2841ae3a44e 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -1,3 +1,18 @@ +# SPDX-FileCopyrightText: Copyright (c) 2022-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + import ast import functools import json @@ -1186,22 +1201,22 @@ class EagleDecodingConfig(DecodingBaseConfig): "Enable relaxed acceptance during thinking phase for reasoning models. " "Accepts draft tokens matching any top-K candidate instead of exact top-1." ) - relaxed_topk: int = Field( + relaxed_topk: PositiveInt = Field( default=1, description= "Number of top candidate tokens to consider for relaxed acceptance. " "Draft token is accepted if it matches any of these.") - relaxed_delta: float = Field( - default=0., + relaxed_delta: NonNegativeFloat = Field( + default=0.0, description= "Probability threshold for relaxed acceptance. Only candidates with " "prob >= (top-1 prob - delta) are kept.") - begin_thinking_phase_token: int = Field( + begin_thinking_phase_token: NonNegativeInt = Field( default=128798, description= "Token ID marking start of thinking phase. Relaxed acceptance only applies within this phase." ) - end_thinking_phase_token: int = Field( + end_thinking_phase_token: NonNegativeInt = Field( default=128799, description= "Token ID marking end of thinking phase. Strict acceptance resumes after this." @@ -1616,13 +1631,13 @@ class MTPDecodingConfig(DecodingBaseConfig): description= "Enable relaxed acceptance during thinking phase for reasoning models. Accepts draft tokens matching any top-K candidate instead of exact top-1." ) - relaxed_topk: int = Field( + relaxed_topk: PositiveInt = Field( default=1, description= "Number of top candidate tokens to consider for relaxed acceptance. Draft token is accepted if it matches any of these." ) - relaxed_delta: float = Field( - default=0., + relaxed_delta: NonNegativeFloat = Field( + default=0.0, description= "Probability threshold for relaxed acceptance. Only candidates with prob >= (top-1 prob - delta) are kept." ) @@ -1653,12 +1668,12 @@ class MTPDecodingConfig(DecodingBaseConfig): "Auto-populated from the model's pretrained config. Do not set manually." ) - begin_thinking_phase_token: int = Field( + begin_thinking_phase_token: NonNegativeInt = Field( default=128798, description= "Token ID marking start of thinking phase. Relaxed acceptance only applies within this phase." ) - end_thinking_phase_token: int = Field( + end_thinking_phase_token: NonNegativeInt = Field( default=128799, description= "Token ID marking end of thinking phase. Strict acceptance resumes after this." From 19e5da9327177378413d7c3c56c9a1ed5e166cc8 Mon Sep 17 00:00:00 2001 From: ZhaoyangWang Date: Sun, 24 May 2026 23:15:11 -0700 Subject: [PATCH 05/11] [TRTLLM-11508][refactor] restore is_mtp_one_model union semantics Address review feedback: keep is_mtp_one_model() covering both MTP and MTP_EAGLE_ONE_MODEL (matches main), and use is_mtp_vanilla() only where the call should match vanilla MTP exclusively. Drop the JIRA tag from the NOTE comment in eagle3.py and simplify the now-redundant "is_mtp_one_model() or is_mtp_eagle_one_model()" patterns introduced by the merge. Signed-off-by: ZhaoyangWang --- tensorrt_llm/_torch/model_config.py | 5 ++- .../_torch/models/modeling_deepseekv3.py | 4 +-- .../_torch/models/modeling_exaone_moe.py | 12 +++---- tensorrt_llm/_torch/models/modeling_glm.py | 6 ++-- .../_torch/models/modeling_nemotron_h.py | 10 +++--- .../_torch/models/modeling_qwen3_next.py | 10 +++--- .../_torch/models/modeling_speculative.py | 6 ++-- tensorrt_llm/_torch/speculative/eagle3.py | 10 +++--- tensorrt_llm/_torch/speculative/interface.py | 32 +++++++++---------- tensorrt_llm/_torch/speculative/mtp.py | 2 +- tensorrt_llm/_torch/speculative/utils.py | 8 ++--- 11 files changed, 47 insertions(+), 58 deletions(-) diff --git a/tensorrt_llm/_torch/model_config.py b/tensorrt_llm/_torch/model_config.py index cadd21543647..fa448d8876c2 100644 --- a/tensorrt_llm/_torch/model_config.py +++ b/tensorrt_llm/_torch/model_config.py @@ -779,9 +779,8 @@ def ceil_div(a, b): hidden_size = ceil_div(self.pretrained_config.hidden_size, attn_tp_size) num_layers = self.pretrained_config.num_hidden_layers num_attention_layers = self.get_num_attention_layers() - if (self.spec_config is not None and - (self.spec_config.spec_dec_mode.is_mtp_one_model() - or self.spec_config.spec_dec_mode.is_mtp_eagle_one_model())): + if (self.spec_config is not None + and self.spec_config.spec_dec_mode.is_mtp_one_model()): assert self.spec_config.num_nextn_predict_layers is not None, ( "num_nextn_predict_layers must be set from model config before building ModelConfig. " "Ensure update_spec_config_from_model_config() has been called." diff --git a/tensorrt_llm/_torch/models/modeling_deepseekv3.py b/tensorrt_llm/_torch/models/modeling_deepseekv3.py index fc01588ae2e0..e981dc233010 100755 --- a/tensorrt_llm/_torch/models/modeling_deepseekv3.py +++ b/tensorrt_llm/_torch/models/modeling_deepseekv3.py @@ -1854,9 +1854,7 @@ def __init__(self, model_config: ModelConfig[PretrainedConfig]): model_config=model_config) self.model_nextn = 0 - if model_config.spec_config is not None and ( - model_config.spec_config.spec_dec_mode.is_mtp_one_model() or - model_config.spec_config.spec_dec_mode.is_mtp_eagle_one_model() + if model_config.spec_config is not None and model_config.spec_config.spec_dec_mode.is_mtp_one_model( ): model_nextn = self.config.num_nextn_predict_layers ckpt_nextn = self.config.num_nextn_predict_layers diff --git a/tensorrt_llm/_torch/models/modeling_exaone_moe.py b/tensorrt_llm/_torch/models/modeling_exaone_moe.py index 533a542cc731..ba8577da9613 100644 --- a/tensorrt_llm/_torch/models/modeling_exaone_moe.py +++ b/tensorrt_llm/_torch/models/modeling_exaone_moe.py @@ -655,9 +655,9 @@ def __init__( self, model_config: ModelConfig[ExaoneMoeConfig], ): - if model_config.spec_config is not None and ( - model_config.spec_config.spec_dec_mode.is_mtp_one_model() - or model_config.spec_config.spec_dec_mode.is_mtp_eagle_one_model() + if ( + model_config.spec_config is not None + and model_config.spec_config.spec_dec_mode.is_mtp_one_model() ): # NOTE: K-EXAONE does not contain the 'num_nextn_predict_layers' field, # which should be equal to 1. Manually set the value here if not present. @@ -669,9 +669,9 @@ def __init__( model_config=model_config, ) - if model_config.spec_config is not None and ( - model_config.spec_config.spec_dec_mode.is_mtp_one_model() - or model_config.spec_config.spec_dec_mode.is_mtp_eagle_one_model() + if ( + model_config.spec_config is not None + and model_config.spec_config.spec_dec_mode.is_mtp_one_model() ): model_nextn = self.config.num_nextn_predict_layers ckpt_nextn = self.config.num_nextn_predict_layers diff --git a/tensorrt_llm/_torch/models/modeling_glm.py b/tensorrt_llm/_torch/models/modeling_glm.py index 018001a8b4a9..293510b65099 100644 --- a/tensorrt_llm/_torch/models/modeling_glm.py +++ b/tensorrt_llm/_torch/models/modeling_glm.py @@ -1020,9 +1020,9 @@ def __init__(self, model_config: ModelConfig[PretrainedConfig]): super().__init__(model=Glm4Model(model_config), model_config=model_config) self.model_nextn = 0 - if model_config.spec_config is not None and ( - model_config.spec_config.spec_dec_mode.is_mtp_one_model() - or model_config.spec_config.spec_dec_mode.is_mtp_eagle_one_model() + if ( + model_config.spec_config is not None + and model_config.spec_config.spec_dec_mode.is_mtp_one_model() ): model_nextn = self.config.num_nextn_predict_layers ckpt_nextn = self.config.num_nextn_predict_layers diff --git a/tensorrt_llm/_torch/models/modeling_nemotron_h.py b/tensorrt_llm/_torch/models/modeling_nemotron_h.py index fedaefc1debe..f643eff0ef1e 100644 --- a/tensorrt_llm/_torch/models/modeling_nemotron_h.py +++ b/tensorrt_llm/_torch/models/modeling_nemotron_h.py @@ -907,9 +907,8 @@ def __init__( model_config=model_config, ) self.model_nextn = 0 - if (model_config.spec_config is not None and - (model_config.spec_config.spec_dec_mode.is_mtp_one_model() or - model_config.spec_config.spec_dec_mode.is_mtp_eagle_one_model())): + if (model_config.spec_config is not None + and model_config.spec_config.spec_dec_mode.is_mtp_one_model()): model_nextn = self.config.num_nextn_predict_layers ckpt_nextn = self.config.num_nextn_predict_layers self.num_hidden_layers = self.config.num_hidden_layers @@ -1042,9 +1041,8 @@ def __init__( use_custom_cublas_mm=use_custom_cublas_mm, ) self.model_nextn = 0 - if (model_config.spec_config is not None and - (model_config.spec_config.spec_dec_mode.is_mtp_one_model() or - model_config.spec_config.spec_dec_mode.is_mtp_eagle_one_model())): + if (model_config.spec_config is not None + and model_config.spec_config.spec_dec_mode.is_mtp_one_model()): model_nextn = self.config.num_nextn_predict_layers ckpt_nextn = self.config.num_nextn_predict_layers self.num_hidden_layers = self.config.num_hidden_layers diff --git a/tensorrt_llm/_torch/models/modeling_qwen3_next.py b/tensorrt_llm/_torch/models/modeling_qwen3_next.py index 4bbb6854239d..d6f4fd57794f 100644 --- a/tensorrt_llm/_torch/models/modeling_qwen3_next.py +++ b/tensorrt_llm/_torch/models/modeling_qwen3_next.py @@ -944,9 +944,8 @@ def __init__( self, model_config: ModelConfig[Qwen3NextConfig], ): - if (model_config.spec_config is not None and - (model_config.spec_config.spec_dec_mode.is_mtp_one_model() or - model_config.spec_config.spec_dec_mode.is_mtp_eagle_one_model())): + if (model_config.spec_config is not None + and model_config.spec_config.spec_dec_mode.is_mtp_one_model()): ckpt_num_nextn = getattr(model_config.pretrained_config, "num_nextn_predict_layers", None) if ckpt_num_nextn not in (None, 1): @@ -963,9 +962,8 @@ def __init__( ) self.preload_weight_modules = self.model.preload_weight_modules - if (model_config.spec_config is not None and - (model_config.spec_config.spec_dec_mode.is_mtp_one_model() or - model_config.spec_config.spec_dec_mode.is_mtp_eagle_one_model())): + if (model_config.spec_config is not None + and model_config.spec_config.spec_dec_mode.is_mtp_one_model()): self.model.layers.extend(self.draft_model.mtp_layers) diff --git a/tensorrt_llm/_torch/models/modeling_speculative.py b/tensorrt_llm/_torch/models/modeling_speculative.py index 8c2ab0c8e42e..ddc2c5687879 100755 --- a/tensorrt_llm/_torch/models/modeling_speculative.py +++ b/tensorrt_llm/_torch/models/modeling_speculative.py @@ -1439,8 +1439,7 @@ def __init__( f"Model type {model_type} not supported for MTP") spec_dec_mode = model_config.spec_config.spec_dec_mode - assert (spec_dec_mode.is_mtp_one_model() - or spec_dec_mode.is_mtp_eagle_one_model()) + assert spec_dec_mode.is_mtp_one_model() checkpoint_mtp_num_layers = model_config.pretrained_config.num_nextn_predict_layers if spec_dec_mode.is_mtp_eagle_one_model(): mtp_num_layers = 1 @@ -1626,8 +1625,7 @@ def get_draft_model(model_config, draft_config, lm_head, model): f"Unsupported eagle3 model architecture: {spec_dec_mode.eagle3_model_arch}" ) - elif (spec_dec_mode.is_mtp_one_model() - or spec_dec_mode.is_mtp_eagle_one_model()): + elif spec_dec_mode.is_mtp_one_model(): return MTPForCausalLM(model_config, model_config.pretrained_config.num_hidden_layers, lm_head, model) diff --git a/tensorrt_llm/_torch/speculative/eagle3.py b/tensorrt_llm/_torch/speculative/eagle3.py index 1d27a0756471..4390a81fa392 100644 --- a/tensorrt_llm/_torch/speculative/eagle3.py +++ b/tensorrt_llm/_torch/speculative/eagle3.py @@ -580,11 +580,11 @@ def max_draft_len(self) -> int: def _prepare_attn_metadata_for_spec_dec(self, attn_metadata): attn_metadata.prepare_for_spec_dec("_seq_lens", "_seq_lens_cuda") - # NOTE(TRTLLM-11508): the previous kv_lens_cuda save/restore was removed - # during the Eagle3/MTP-eagle merge. The drafting loop now updates - # kv_lens_cuda incrementally and calls attn_metadata.update_for_spec_dec() - # to keep the runtime view consistent. Verify under Eagle3 regressions - # if any kv-lens drift is observed. + # NOTE: the previous kv_lens_cuda save/restore was removed during the + # Eagle3/MTP-eagle merge. The drafting loop now updates kv_lens_cuda + # incrementally and calls attn_metadata.update_for_spec_dec() to keep + # the runtime view consistent. Verify under Eagle3 regressions if any + # kv-lens drift is observed. batch_size = attn_metadata.num_seqs # Save spec-dec params that the drafting loop will overwrite. diff --git a/tensorrt_llm/_torch/speculative/interface.py b/tensorrt_llm/_torch/speculative/interface.py index 19aa2062812e..0404373514b4 100644 --- a/tensorrt_llm/_torch/speculative/interface.py +++ b/tensorrt_llm/_torch/speculative/interface.py @@ -235,7 +235,10 @@ class SpeculativeDecodingMode(IntEnum): AUTO = auto() def is_mtp_one_model(self): - return self == SpeculativeDecodingMode.MTP + # Union: covers vanilla MTP and MTP_EAGLE_ONE_MODEL. Use is_mtp_vanilla() + # when only the vanilla MTP variant should match. + return (self == SpeculativeDecodingMode.MTP + or self == SpeculativeDecodingMode.MTP_EAGLE_ONE_MODEL) def is_mtp_eagle_one_model(self): return self == SpeculativeDecodingMode.MTP_EAGLE_ONE_MODEL @@ -251,8 +254,7 @@ def is_eagle3(self): def use_one_engine(self): return self.is_eagle3_one_model() or self.is_mtp_one_model( - ) or self.is_mtp_eagle_one_model() or self.is_external_drafter( - ) or self.is_sa() + ) or self.is_external_drafter() or self.is_sa() def is_eagle3_one_model(self): return self == SpeculativeDecodingMode.EAGLE3_ONE_MODEL @@ -291,27 +293,24 @@ def is_external_drafter(self): return self.is_parallel_draft() or self.is_draft_target_one_model() def without_logits(self): - return self.is_mtp_one_model() or self.is_mtp_eagle_one_model( - ) or self.is_eagle3_one_model() or self.is_external_drafter( - ) or self.is_sa() + return self.is_mtp_one_model() or self.is_eagle3_one_model( + ) or self.is_external_drafter() or self.is_sa() def needs_kv_cache_rewind(self): - return self.is_mtp_one_model() or self.is_mtp_eagle_one_model( - ) or self.is_eagle3_one_model() or self.is_ngram() or self.is_sa( - ) or self.is_external_drafter() + return self.is_mtp_one_model() or self.is_eagle3_one_model( + ) or self.is_ngram() or self.is_sa() or self.is_external_drafter() def support_overlap_scheduler(self): - return self.is_mtp_one_model() or self.is_mtp_eagle_one_model( - ) or self.is_eagle3_one_model() or self.is_sa() or self.has_draft_model( - ) or self.is_external_drafter() + return self.is_mtp_one_model() or self.is_eagle3_one_model( + ) or self.is_sa() or self.has_draft_model() or self.is_external_drafter( + ) def support_guided_decoder(self): return self.is_none() or self.has_spec_drafter() def support_capturable_guided_decoder(self): - return self.is_mtp_one_model() or self.is_mtp_eagle_one_model( - ) or self.is_eagle3_one_model() or self.is_external_drafter( - ) or self.is_sa() + return self.is_mtp_one_model() or self.is_eagle3_one_model( + ) or self.is_external_drafter() or self.is_sa() def support_dynamic_draft_len(self): # TODO: expand to all one-model algorithms @@ -336,8 +335,7 @@ def need_load_draft_weights(self): return self.is_eagle3_one_model() or self.is_external_drafter() def has_spec_decoder(self): - return self.is_mtp_one_model() or self.is_mtp_eagle_one_model( - ) or self.is_mtp_eagle() or self.is_eagle3( + return self.is_mtp_one_model() or self.is_mtp_eagle() or self.is_eagle3( ) or self.is_eagle3_one_model() or self.is_external_drafter( ) or self.is_sa() diff --git a/tensorrt_llm/_torch/speculative/mtp.py b/tensorrt_llm/_torch/speculative/mtp.py index 3f182f7b3442..542c6be865ae 100644 --- a/tensorrt_llm/_torch/speculative/mtp.py +++ b/tensorrt_llm/_torch/speculative/mtp.py @@ -207,7 +207,7 @@ def prepare(self): mtp_slot_ids.append(slot_id) # MTP Vanilla: Update mtp hidden states and past tokens - if self.spec_dec_mode.is_mtp_one_model(): + if self.spec_dec_mode.is_mtp_vanilla(): mtp_hidden_states_ptrs = [] mtp_past_tokens_ptrs = [] for slot_id in mtp_slot_ids: diff --git a/tensorrt_llm/_torch/speculative/utils.py b/tensorrt_llm/_torch/speculative/utils.py index 7d28b5a5634a..766a67ff0af4 100644 --- a/tensorrt_llm/_torch/speculative/utils.py +++ b/tensorrt_llm/_torch/speculative/utils.py @@ -60,7 +60,7 @@ def get_spec_metadata(spec_config, vocab_size=vocab_size, spec_resource_manager=spec_resource_manager, ) - if spec_config.spec_dec_mode.is_mtp_one_model(): + if spec_config.spec_dec_mode.is_mtp_vanilla(): return MTPSpecMetadata( max_draft_len=spec_config.max_draft_len, max_total_draft_tokens=spec_config.tokens_per_gen_step - 1, @@ -216,7 +216,7 @@ def get_spec_resource_manager(model_engine, draft_model_engine=None): ) else: return None - if spec_dec_mode.is_mtp_one_model(): + if spec_dec_mode.is_mtp_vanilla(): sa_manager = None sa_cfg = getattr(spec_config, 'sa_config', None) if sa_cfg is not None: @@ -288,7 +288,7 @@ def get_spec_decoder( if spec_config.spec_dec_mode.is_mtp_eagle_one_model(): # MTP Eagle one-model now uses the same sampler as Eagle3 one-model. return Eagle3OneModelSampler(sampler_args, spec_config=spec_config) - if spec_config.spec_dec_mode.is_mtp_one_model(): + if spec_config.spec_dec_mode.is_mtp_vanilla(): return MTPSampler(sampler_args, nextn=spec_config.max_draft_len) if spec_config.spec_dec_mode.is_eagle3( ) or spec_config.spec_dec_mode.is_mtp_eagle(): @@ -341,7 +341,7 @@ def get_spec_drafter(model_engine, def get_num_spec_layers(spec_config): if spec_config.spec_dec_mode.is_mtp_eagle_one_model(): return 1 - if spec_config.spec_dec_mode.is_mtp_one_model(): + if spec_config.spec_dec_mode.is_mtp_vanilla(): return spec_config.num_nextn_predict_layers if spec_config.spec_dec_mode.is_eagle3_one_model(): num_eagle_layers = spec_config.num_eagle_layers From 131ca73c492e0a4efdb4c4897fd67b44f42b53dc Mon Sep 17 00:00:00 2001 From: ZhaoyangWang Date: Wed, 13 May 2026 02:48:56 -0700 Subject: [PATCH 06/11] [TRTLLM-11508][fix] accept resource_manager in SpecWorkerBase.skip_forward The one-model worker forward signature (Eagle3 / MTP-Eagle) takes ``resource_manager``, and modeling_speculative.py forwards it unconditionally to ``self.spec_worker(...)``. On non-last PP ranks, ``forward`` is replaced by ``skip_forward`` via modeling_utils.skip_forward(), which raised ``TypeError: SpecWorkerBase.skip_forward() got an unexpected keyword argument 'resource_manager'`` and silently terminated the executor worker. Before the merge, MTP-Eagle used MTPWorker.skip_forward which already accepted ``resource_manager``; the merged path now inherits SpecWorkerBase.skip_forward, which did not. Add the parameter (unused) to restore PP compatibility. Validated on H200 x4 with TestDeepSeekV3Lite::test_bfloat16_4gpus[tp2pp2-mtp_nextn=2-attention_dp=True-cuda_graph=True-overlap_scheduler=True-torch_compile=False] (was: silent crash during warmup; now: PASSED, GSM8K 63.72). Signed-off-by: ZhaoyangWang --- tensorrt_llm/_torch/speculative/interface.py | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/speculative/interface.py b/tensorrt_llm/_torch/speculative/interface.py index 0404373514b4..c62111f0f511 100644 --- a/tensorrt_llm/_torch/speculative/interface.py +++ b/tensorrt_llm/_torch/speculative/interface.py @@ -734,8 +734,15 @@ def skip_forward( attn_metadata, spec_metadata, draft_model, + resource_manager=None, ): - """Skip spec dec for non-last rank (PP). Returns placeholder outputs.""" + """Skip spec dec for non-last rank (PP). Returns placeholder outputs. + + ``resource_manager`` is accepted but unused; it appears in the + ``forward()`` signature of one-model workers (Eagle3 / MTP-Eagle) and + the caller in ``modeling_speculative.py`` forwards it unconditionally, + so the skip path must accept it as well. + """ batch_size = attn_metadata.num_seqs accepted_tokens = torch.empty((batch_size, (self.max_draft_len + 1)), dtype=torch.int, From a143f6599ae7f37204f0de546deac8eb8445e898 Mon Sep 17 00:00:00 2001 From: ZhaoyangWang Date: Mon, 25 May 2026 03:52:50 -0700 Subject: [PATCH 07/11] [TRTLLM-11508][chore] apply yapf formatting on Eagle3OneModelWorker Pre-commit yapf reformatted the position_ids update in the unified linear draft loop. No functional change. Signed-off-by: ZhaoyangWang --- tensorrt_llm/_torch/speculative/eagle3.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/tensorrt_llm/_torch/speculative/eagle3.py b/tensorrt_llm/_torch/speculative/eagle3.py index 4390a81fa392..afe4ffa64cb0 100644 --- a/tensorrt_llm/_torch/speculative/eagle3.py +++ b/tensorrt_llm/_torch/speculative/eagle3.py @@ -885,9 +885,8 @@ def _forward_linear_draft_loop(self, inputs, attn_metadata, spec_metadata, hidden_states = hidden_states[gather_ids] else: hidden_states = hidden_states_to_save[gather_ids] - position_ids = ( - _select_mtp_position_ids(inputs["position_ids"], gather_ids) - + 1) + position_ids = (_select_mtp_position_ids( + inputs["position_ids"], gather_ids) + 1) # Update attn_metadata for the next iteration. if i == 0: From 1c91e1787c75a6d600c848d46c78c4bdee3edb9f Mon Sep 17 00:00:00 2001 From: ZhaoyangWang Date: Mon, 25 May 2026 23:05:50 -0700 Subject: [PATCH 08/11] [TRTLLM-11508][fix] include MTP_EAGLE_ONE_MODEL in num_capture_layers MTPDecodingConfig.num_capture_layers only matched is_mtp_eagle(), returning 0 for MTP_EAGLE_ONE_MODEL. After unifying MTP-Eagle and Eagle3 onto Eagle3ResourceManager, this caused hidden_states to be allocated with 0 columns and Eagle3OneModelSpecMetadata.__post_init__ to assert at warmup time: AssertionError: hidden_states shape mismatch: Eagle3ResourceManager has torch.Size([N, 0]), but metadata expects (:, hidden_size) from hidden_size=H x capture_layers=[L-1] Match the metadata semantics: both MTP_EAGLE and MTP_EAGLE_ONE_MODEL capture the target's last hidden layer, so num_capture_layers is 1 for both. Signed-off-by: ZhaoyangWang --- tensorrt_llm/llmapi/llm_args.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index b2841ae3a44e..83b05640021a 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -1719,7 +1719,13 @@ def supports_backend(self, backend: str) -> bool: @property def num_capture_layers(self) -> int: - return 1 if self.spec_dec_mode.is_mtp_eagle() else 0 + # MTP_EAGLE and MTP_EAGLE_ONE_MODEL both capture the target model's + # last hidden layer (see Eagle3OneModelSpecMetadata.__post_init__); + # the shared Eagle3ResourceManager must allocate hidden_states with + # matching column count. + mode = self.spec_dec_mode + return 1 if (mode.is_mtp_eagle() + or mode.is_mtp_eagle_one_model()) else 0 @property def spec_dec_mode(self): From 7b411d862efa00ba8bb5ed2f39fbcda83a0dd0a1 Mon Sep 17 00:00:00 2001 From: ZhaoyangWang Date: Tue, 26 May 2026 02:01:02 -0700 Subject: [PATCH 09/11] [TRTLLM-11508][fix] skip hidden-state capture for MTP_EAGLE_ONE_MODEL Eagle3OneModelWorker on the MTP Eagle path consumes the target model's hidden_states argument directly (prepare_1st_drafter_inputs and _run_draft_forward both gate the spec_metadata.hidden_states read on not is_mtp_eagle), so the layer-capture machinery is never read. Configuring layers_to_capture={last} for MTP_EAGLE_ONE_MODEL still had two visible side effects: - modeling_deepseekv3 / modeling_glm / modeling_gpt_oss disable POST_MOE_FUSION / POST_MLP_FUSION whenever spec_metadata.is_layer_capture( layer_idx) is true. MTP Eagle paid that cost despite never using the capture. - Eagle3OneModelSpecMetadata allocated a (max_num_tokens, hidden_size) hidden_states buffer (and Eagle3ResourceManager pre-allocated one too) that the MTP Eagle worker never reads. Opt out of capture for MTP_EAGLE_ONE_MODEL: - MTPDecodingConfig.num_capture_layers returns 0 for one-model (1 stays for two-model MTP_EAGLE, which still feeds captured hidden states into the separate draft engine). - Eagle3OneModelSpecMetadata.__post_init__ defaults layers_to_capture to () for MTP_EAGLE_ONE_MODEL and skips the hidden_states allocation. - get_spec_metadata() stops passing layers_to_capture={last} on the MTP_EAGLE_ONE_MODEL branch and lets the metadata default kick in. With layers_to_capture empty, is_layer_capture() is false for every layer and maybe_capture_hidden_states() is a no-op for-loop, so post-MLP / post-MoE fusion stays enabled and no capture buffer is allocated. Signed-off-by: ZhaoyangWang --- tensorrt_llm/_torch/speculative/eagle3.py | 21 +++++++++++++++++---- tensorrt_llm/_torch/speculative/utils.py | 11 +++++++---- tensorrt_llm/llmapi/llm_args.py | 16 +++++++++------- 3 files changed, 33 insertions(+), 15 deletions(-) diff --git a/tensorrt_llm/_torch/speculative/eagle3.py b/tensorrt_llm/_torch/speculative/eagle3.py index afe4ffa64cb0..3df7b1da0539 100644 --- a/tensorrt_llm/_torch/speculative/eagle3.py +++ b/tensorrt_llm/_torch/speculative/eagle3.py @@ -398,8 +398,16 @@ class Eagle3OneModelSpecMetadata(SpecMetadata): def __post_init__(self): if self.layers_to_capture is None: - if self.spec_dec_mode.is_mtp_eagle_one_model( - ) or self.num_layers == 1: + if self.spec_dec_mode.is_mtp_eagle_one_model(): + # MTP Eagle one-model feeds the target model's hidden_states + # directly to the MTP layer (see Eagle3OneModelWorker + # prepare_1st_drafter_inputs / _run_draft_forward, both gated + # on self.is_mtp_eagle). It never reads spec_metadata.hidden_states, + # so leave layers_to_capture empty: this makes is_layer_capture() + # return False everywhere and avoids the post-MLP/MoE fusion + # disable side effect in modeling_deepseekv3 / glm / etc. + self.layers_to_capture = () + elif self.num_layers == 1: self.layers_to_capture = (self.num_layers - 1, ) else: if self.num_layers <= 5: @@ -411,8 +419,13 @@ def __post_init__(self): else: self.layers_to_capture = sorted(list(self.layers_to_capture)) self.num_capture_layers = len(self.layers_to_capture) - if (self.spec_resource_manager is not None - and self.spec_resource_manager.hidden_states is not None): + if self.num_capture_layers == 0: + # No layers to capture (MTP Eagle one-model). Skip buffer + # allocation entirely; nothing reads self.hidden_states on this + # path. + self.hidden_states = None + elif (self.spec_resource_manager is not None + and self.spec_resource_manager.hidden_states is not None): self.hidden_states = self.spec_resource_manager.hidden_states expected_cols = self.hidden_size * len(self.layers_to_capture) assert self.hidden_states.shape[1] == expected_cols, ( diff --git a/tensorrt_llm/_torch/speculative/utils.py b/tensorrt_llm/_torch/speculative/utils.py index 766a67ff0af4..8bcf20d35b9d 100644 --- a/tensorrt_llm/_torch/speculative/utils.py +++ b/tensorrt_llm/_torch/speculative/utils.py @@ -43,9 +43,13 @@ def get_spec_metadata(spec_config, False) vocab_size = getattr(model_config, "vocab_size", 0) if spec_config.spec_dec_mode.is_mtp_eagle_one_model(): - # MTP Eagle one-model now reuses Eagle3 one-model metadata so it - # picks up the unified worker, sampler, and slot_ids/subseq plumbing. - # Capture only the final layer's hidden state (MTP Eagle behavior). + # MTP Eagle one-model reuses Eagle3 one-model metadata for the + # unified worker/sampler/slot_ids plumbing, but skips per-layer + # hidden-state capture: the worker feeds the target model's + # hidden_states directly into the MTP layer, so we leave + # layers_to_capture unset and let Eagle3OneModelSpecMetadata default + # it to an empty tuple. This also keeps post-MLP/MoE fusion enabled + # on models that gate it on is_layer_capture(). return Eagle3OneModelSpecMetadata( max_draft_len=spec_config.max_draft_len, max_total_draft_tokens=spec_config.tokens_per_gen_step - 1, @@ -54,7 +58,6 @@ def get_spec_metadata(spec_config, num_layers=model_config.num_hidden_layers, hidden_size=model_config.hidden_size, max_num_tokens=max_num_tokens, - layers_to_capture={model_config.num_hidden_layers - 1}, allow_advanced_sampling=spec_config.allow_advanced_sampling, use_rejection_sampling=use_rejection_sampling, vocab_size=vocab_size, diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 83b05640021a..5a72f96f86e2 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -1719,13 +1719,15 @@ def supports_backend(self, backend: str) -> bool: @property def num_capture_layers(self) -> int: - # MTP_EAGLE and MTP_EAGLE_ONE_MODEL both capture the target model's - # last hidden layer (see Eagle3OneModelSpecMetadata.__post_init__); - # the shared Eagle3ResourceManager must allocate hidden_states with - # matching column count. - mode = self.spec_dec_mode - return 1 if (mode.is_mtp_eagle() - or mode.is_mtp_eagle_one_model()) else 0 + # MTP_EAGLE (two-model) feeds captured target hidden states into the + # separate draft engine, so the shared Eagle3ResourceManager must + # allocate a hidden_states buffer for it. MTP_EAGLE_ONE_MODEL passes + # the target model's hidden_states straight to the MTP layer + # (see Eagle3OneModelWorker.prepare_1st_drafter_inputs / _run_draft_forward, + # both gated on self.is_mtp_eagle), so no capture buffer is needed + # and we should skip allocation to avoid disabling post-MLP/MoE + # fusion via the layer-capture hook. + return 1 if self.spec_dec_mode.is_mtp_eagle() else 0 @property def spec_dec_mode(self): From 239c7a3fe464967a4034b4c9eb9e84af8409561b Mon Sep 17 00:00:00 2001 From: ZhaoyangWang Date: Tue, 26 May 2026 22:36:44 -0700 Subject: [PATCH 10/11] [TRTLLM-11508][chore] address PR review feedback Bundle of small fixes from the merge-eagle-mtp PR review: - Drop the historical `kv_lens_cuda` save/restore NOTE in `Eagle3OneModelWorker._prepare_attn_metadata_for_spec_dec()`. The note was internal context for the merge commit and is not useful as inline documentation of the current state. - Remove the `if self.is_mtp_eagle:` gate around the Mamba-hybrid state update in `Eagle3OneModelWorker.forward()`. The `isinstance(kv_cache_manager, MambaHybridCacheManager)` check already no-ops on non-Mamba caches, so spec-mode gating was unnecessary and would prevent Eagle3 from picking up Mamba support if a Mamba-style Eagle draft ever lands. - Delete the `__getattr__` backward-compat shim in `mtp.py` that re-exported `MTPEagleWorker` from `eagle3.py`. No in-tree caller imports `MTPEagleWorker` from `tensorrt_llm._torch.speculative.mtp` anymore (both `speculative/__init__.py` and `speculative/utils.py` import it directly from `.eagle3`). - Drop the relaxed-acceptance fields (`use_relaxed_acceptance_for_thinking`, `relaxed_topk`, `relaxed_delta`, `begin_thinking_phase_token`, `end_thinking_phase_token`) from `EagleDecodingConfig`. The unified worker's relaxed-acceptance branch is already guarded by `getattr(spec_config, 'use_relaxed_acceptance_for_thinking', False)`, so an `EagleDecodingConfig` without these fields naturally short- circuits the path. `MTPDecodingConfig` keeps the same fields unchanged. Limits the user-facing surface of relaxed acceptance to the MTP family, per review preference (relaxed acceptance degrades the spec-dec output distribution and is questionable for Eagle). - Rewrite the "Eagle engine takes ``draft_len`` tokens" comment in `_forward_linear_draft_loop()` to mention both Eagle3 and MTP Eagle, matching the unified worker that path now drives. - Expand the comment around the `num_tokens` subtract in `Eagle3OneModelSpecMetadata.prepare()` to explain that the value is consumed only by the attention-DP allgather in `model_engine`, that both modes share `prepare_1st_drafter_inputs()` and therefore feed the same input_ids tensor to the kernel, and that the subtract asymmetry exists solely to preserve the pre-refactor shape-hint convention of each mode. The previous wording ("MTP Eagle uses max_draft_len + 1 tokens... Eagle3 follows the standard tree/linear path") incorrectly implied an input-shape difference. Signed-off-by: ZhaoyangWang --- tensorrt_llm/_torch/speculative/eagle3.py | 47 ++++++++++++----------- tensorrt_llm/_torch/speculative/mtp.py | 11 ------ tensorrt_llm/llmapi/llm_args.py | 28 -------------- 3 files changed, 24 insertions(+), 62 deletions(-) diff --git a/tensorrt_llm/_torch/speculative/eagle3.py b/tensorrt_llm/_torch/speculative/eagle3.py index 3df7b1da0539..6acef9ed348f 100644 --- a/tensorrt_llm/_torch/speculative/eagle3.py +++ b/tensorrt_llm/_torch/speculative/eagle3.py @@ -487,8 +487,12 @@ def prepare(self): pin_memory=prefer_pinned()) self.batch_indices_cuda[:num_seqs].copy_(batch_indices, non_blocking=True) - # MTP Eagle uses max_draft_len + 1 tokens in the first draft forward so - # it must not subtract here; Eagle3 follows the standard tree/linear path. + # `num_tokens` here only feeds the attention-DP shape hint + # (allgathered in model_engine and overridden into + # `attn_metadata.all_rank_num_tokens` on the step-0 draft forward). + # Each mode uses a different convention: + # - MTP Eagle: keep the 1st-iter shape (matches input_ids). + # - Eagle3: subtract to the subseq shape. if not self.spec_dec_mode.is_mtp_eagle_one_model(): if self.is_spec_dec_tree: self.num_tokens -= ( @@ -593,11 +597,6 @@ def max_draft_len(self) -> int: def _prepare_attn_metadata_for_spec_dec(self, attn_metadata): attn_metadata.prepare_for_spec_dec("_seq_lens", "_seq_lens_cuda") - # NOTE: the previous kv_lens_cuda save/restore was removed during the - # Eagle3/MTP-eagle merge. The drafting loop now updates kv_lens_cuda - # incrementally and calls attn_metadata.update_for_spec_dec() to keep - # the runtime view consistent. Verify under Eagle3 regressions if any - # kv-lens drift is observed. batch_size = attn_metadata.num_seqs # Save spec-dec params that the drafting loop will overwrite. @@ -677,18 +676,19 @@ def forward(self, accepted_tokens, num_accepted_tokens = self.sample_and_accept_draft_tokens( input_ids, logits, attn_metadata, spec_metadata) - # MTP Eagle only: Mamba hybrid models need state updates after token - # acceptance because accepted token count affects which Mamba states - # are valid; Eagle3 does not use Mamba layers. - if self.is_mtp_eagle: - if self._is_mamba_hybrid_cache is None: - self._is_mamba_hybrid_cache = isinstance( - attn_metadata.kv_cache_manager, MambaHybridCacheManager) - if num_gens > 0 and self._is_mamba_hybrid_cache: - attn_metadata.kv_cache_manager.update_mamba_states( - attn_metadata=attn_metadata, - num_accepted_tokens=num_accepted_tokens, - state_indices=attn_metadata.mamba_metadata.state_indices) + # Mamba hybrid models need state updates after token acceptance because + # the accepted token count affects which Mamba states are valid. The + # isinstance check below naturally no-ops on non-Mamba kv_cache_managers, + # so this is safe to run unconditionally regardless of spec mode (Eagle3 + # over a Mamba-style draft is plausible, even if no such draft exists today). + if self._is_mamba_hybrid_cache is None: + self._is_mamba_hybrid_cache = isinstance( + attn_metadata.kv_cache_manager, MambaHybridCacheManager) + if num_gens > 0 and self._is_mamba_hybrid_cache: + attn_metadata.kv_cache_manager.update_mamba_states( + attn_metadata=attn_metadata, + num_accepted_tokens=num_accepted_tokens, + state_indices=attn_metadata.mamba_metadata.state_indices) sa_manager = getattr(spec_metadata.spec_resource_manager, 'sa_manager', None) @@ -924,10 +924,11 @@ def _forward_linear_draft_loop(self, inputs, attn_metadata, spec_metadata, if hasattr(attn_metadata, 'kv_lens_cuda'): attn_metadata.update_for_spec_dec() - # Eagle engine takes ``draft_len`` tokens from the previous - # step, runs spec-dec mode with those tokens, then later - # steps use regular decoding mode. Disable spec_decoding so - # the masks/positions stay correct on subsequent iters. + # Both Eagle3 and MTP Eagle drafters take ``draft_len + 1`` + # tokens in the first draft step (attention runs in spec-dec + # mode), then 1 token per step in subsequent iterations. + # Disable spec_decoding here so the masks/positions stay + # correct on subsequent iters. attn_metadata.use_spec_decoding = False else: if hasattr(attn_metadata, 'kv_lens_cuda'): diff --git a/tensorrt_llm/_torch/speculative/mtp.py b/tensorrt_llm/_torch/speculative/mtp.py index 542c6be865ae..256def91104a 100644 --- a/tensorrt_llm/_torch/speculative/mtp.py +++ b/tensorrt_llm/_torch/speculative/mtp.py @@ -1135,14 +1135,3 @@ def draft_sampler( draft_tokens = self._draft_sampler_greedy(logits) return draft_tokens - - -# ``MTPEagleWorker`` moved to ``eagle3.py`` as part of the Eagle3/MTP-Eagle -# merge (TRTLLM-11508). Preserve the historical import path so external -# callers like ``from tensorrt_llm._torch.speculative.mtp import MTPEagleWorker`` -# keep working without a hard dependency cycle. -def __getattr__(name): - if name == "MTPEagleWorker": - from .eagle3 import MTPEagleWorker - return MTPEagleWorker - raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 5a72f96f86e2..87b5c89ea34e 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -1194,34 +1194,6 @@ class EagleDecodingConfig(DecodingBaseConfig): default="llama3", description="The model architecture of the eagle3 model.") - # Relaxed acceptance settings (mirrors MTPDecodingConfig for thinking models) - use_relaxed_acceptance_for_thinking: bool = Field( - default=False, - description= - "Enable relaxed acceptance during thinking phase for reasoning models. " - "Accepts draft tokens matching any top-K candidate instead of exact top-1." - ) - relaxed_topk: PositiveInt = Field( - default=1, - description= - "Number of top candidate tokens to consider for relaxed acceptance. " - "Draft token is accepted if it matches any of these.") - relaxed_delta: NonNegativeFloat = Field( - default=0.0, - description= - "Probability threshold for relaxed acceptance. Only candidates with " - "prob >= (top-1 prob - delta) are kept.") - begin_thinking_phase_token: NonNegativeInt = Field( - default=128798, - description= - "Token ID marking start of thinking phase. Relaxed acceptance only applies within this phase." - ) - end_thinking_phase_token: NonNegativeInt = Field( - default=128799, - description= - "Token ID marking end of thinking phase. Strict acceptance resumes after this." - ) - @field_validator('eagle_choices', mode='before') @classmethod def validate_eagle_choices(cls, v): From 9051c11b72a513aade3cc1d34c830e788b27f7ab Mon Sep 17 00:00:00 2001 From: ZhaoyangWang Date: Tue, 2 Jun 2026 01:46:01 -0700 Subject: [PATCH 11/11] [TRTLLM-11508][chore] extract spec_metadata all_rank_num_tokens helper Factor the duplicated spec_metadata distributed-token-count assignments (all_rank_num_tokens, all_rank_num_seqs, and the Eagle3/MTP-eagle one-model subseq_all_rank_num_tokens) into a shared helper _set_spec_metadata_all_rank_num_tokens, called from _apply_incremental_update, _prepare_tp_inputs, and _prepare_tp_inputs_no_cache. Signed-off-by: ZhaoyangWang --- .../_torch/pyexecutor/model_engine.py | 58 ++++++++----------- 1 file changed, 23 insertions(+), 35 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 0f9e9f637c77..f5c54c1af25f 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -1867,6 +1867,19 @@ def _get_all_rank_ctx_requests(self, num_ctx_requests: int): return list(self.dist.tp_allgather(num_ctx_requests)) return None + def _set_spec_metadata_all_rank_num_tokens( + self, spec_metadata: SpecMetadata, + spec_all_rank_num_tokens: List[int], + all_rank_num_seqs: List[int]) -> None: + # Eagle3 / MTP-eagle one-model use subseq_all_rank_num_tokens for + # draft loop iterations i>0 (per-sequence counts, since each + # sequence contributes one token per iteration). + spec_metadata.all_rank_num_tokens = spec_all_rank_num_tokens + spec_metadata.all_rank_num_seqs = all_rank_num_seqs + if (spec_metadata.spec_dec_mode.is_mtp_eagle_one_model() + or spec_metadata.spec_dec_mode.is_eagle3_one_model()): + spec_metadata.subseq_all_rank_num_tokens = all_rank_num_seqs + def _get_padding_params( self, total_num_tokens: int, num_ctx_requests: int, attn_all_rank_num_tokens: Optional[List[int]] @@ -2075,20 +2088,9 @@ def _prepare_incremental_update_metadata( all_rank_num_tokens = self.dist.tp_cp_allgather( [spec_metadata.num_tokens, len(sequence_lengths)]) - spec_metadata.all_rank_num_tokens = [ - item[0] for item in all_rank_num_tokens - ] - spec_metadata.all_rank_num_seqs = [ - item[1] for item in all_rank_num_tokens - ] - # Both Eagle3 one-model and MTP-eagle one-model use - # subseq_all_rank_num_tokens for draft loop iterations i>0 - # (per-sequence counts since each sequence contributes one - # token per iteration). - if (spec_metadata.spec_dec_mode.is_mtp_eagle_one_model() - or spec_metadata.spec_dec_mode.is_eagle3_one_model()): - spec_metadata.subseq_all_rank_num_tokens = ( - spec_metadata.all_rank_num_seqs) + self._set_spec_metadata_all_rank_num_tokens( + spec_metadata, [item[0] for item in all_rank_num_tokens], + [item[1] for item in all_rank_num_tokens]) # Set iteration states - batch dictionary updates self.iter_states.update({ @@ -3310,16 +3312,9 @@ def previous_seq_slots_device(): all_rank_num_tokens = self.dist.tp_cp_allgather( [spec_metadata.num_tokens, len(sequence_lengths)]) - - spec_all_rank_num_tokens = [ - item[0] for item in all_rank_num_tokens - ] - all_rank_num_seqs = [item[1] for item in all_rank_num_tokens] - spec_metadata.all_rank_num_tokens = spec_all_rank_num_tokens - spec_metadata.all_rank_num_seqs = all_rank_num_seqs - if (spec_metadata.spec_dec_mode.is_mtp_eagle_one_model() - or spec_metadata.spec_dec_mode.is_eagle3_one_model()): - spec_metadata.subseq_all_rank_num_tokens = all_rank_num_seqs + self._set_spec_metadata_all_rank_num_tokens( + spec_metadata, [item[0] for item in all_rank_num_tokens], + [item[1] for item in all_rank_num_tokens]) if mm_token_indices is not None: mask = torch.ones(total_num_tokens, dtype=torch.bool) @@ -3481,19 +3476,12 @@ def _prepare_tp_inputs_no_cache( attn_metadata.num_tokens, spec_metadata.num_tokens, len(sequence_lengths) ]) - attn_all_rank_num_tokens = [ + attn_metadata.all_rank_num_tokens = [ item[0] for item in all_rank_num_tokens ] - spec_all_rank_num_tokens = [ - item[1] for item in all_rank_num_tokens - ] - all_rank_num_seqs = [item[2] for item in all_rank_num_tokens] - attn_metadata.all_rank_num_tokens = attn_all_rank_num_tokens - spec_metadata.all_rank_num_tokens = spec_all_rank_num_tokens - spec_metadata.all_rank_num_seqs = all_rank_num_seqs - if (spec_metadata.spec_dec_mode.is_mtp_eagle_one_model() - or spec_metadata.spec_dec_mode.is_eagle3_one_model()): - spec_metadata.subseq_all_rank_num_tokens = all_rank_num_seqs + self._set_spec_metadata_all_rank_num_tokens( + spec_metadata, [item[1] for item in all_rank_num_tokens], + [item[2] for item in all_rank_num_tokens]) else: all_rank_num_tokens = self.dist.tp_cp_allgather( attn_metadata.num_tokens)