From 2872dcba318ef05892f8c17e8381f402fde70b34 Mon Sep 17 00:00:00 2001 From: trtllm-agent <296075020+trtllm-agent@users.noreply.github.com> Date: Wed, 22 Jul 2026 04:03:21 -0700 Subject: [PATCH] [nvbugs/6478723][fix] MTPEagleDynamicTreeWorker: rename forward to _forward_impl and route cleanup through base class MTPEagleDynamicTreeWorker (added by #15582) overrode SpecWorkerBase.forward directly, which SpecWorkerBase.__init_subclass__ (nvbugs/6442074, #16382) forbids: subclasses must implement _forward_impl so the base-class forward can guarantee spec-dec attn-metadata cleanup on failure. The class-creation TypeError blocked pytest collection for every test in tests/integration/defs/accuracy/test_llm_api_pytorch.py. Fix: * Rename MTPEagleDynamicTreeWorker.forward -> _forward_impl. Signature already matches parent MTPEagleWorker._forward_impl (eagle3.py:712), SpecWorkerBase.forward delegates via *args/**kwargs, and the @nvtx_range("mtp_dyn.forward") decorator is preserved for profiler- label stability. * Stash the extra dynamic-tree-only attn_metadata state (all_rank_num_tokens, force_prepare_spec_dec_tree_mask) on self and override _ensure_spec_dec_state_restored so a tolerated forward failure between _prepare_attn_metadata_for_spec_dec and _restore_attn_metadata_from_spec_dec doesn't leak the pre-draft-loop state into subsequent iterations. Same pattern PARD/DFlash use. Signed-off-by: trtllm-agent <296075020+trtllm-agent@users.noreply.github.com> --- .../_torch/speculative/mtp_dynamic_tree.py | 41 ++++++++++++++++--- 1 file changed, 35 insertions(+), 6 deletions(-) diff --git a/tensorrt_llm/_torch/speculative/mtp_dynamic_tree.py b/tensorrt_llm/_torch/speculative/mtp_dynamic_tree.py index 5dca84abae8e..a08677c48372 100644 --- a/tensorrt_llm/_torch/speculative/mtp_dynamic_tree.py +++ b/tensorrt_llm/_torch/speculative/mtp_dynamic_tree.py @@ -81,6 +81,14 @@ def __init__( self.spec_tree_manager = None self._d2t = None + # (all_rank_num_tokens, force_prepare_spec_dec_tree_mask) snapshot taken by + # _forward_impl before mutating attn_metadata; None means "not saved / already + # restored" (all_rank_num_tokens is legitimately None outside ADP, so we can't + # use it directly as a sentinel). Restored by _ensure_spec_dec_state_restored + # on the failure path so a tolerated forward exception does not leak state + # into subsequent iterations. + self._saved_extra_state = None + # Pre-allocated draft-loop buffers (CUDA-graph safe). self.draft_tokens_buffer = torch.zeros( max_batch_size, loop_max_tokens, dtype=torch.int32, device="cuda" @@ -227,6 +235,20 @@ def _restore_attn_metadata_from_spec_dec(self, attn_metadata): ) self._saved_generation_lengths = None + def _ensure_spec_dec_state_restored(self, attn_metadata, spec_metadata): + # Base cleanup handles the has_spec_dec_saved_state path. Undo the + # dynamic-tree-only mutations (all_rank_num_tokens, + # force_prepare_spec_dec_tree_mask) so a tolerated forward failure + # does not leak the pre-draft-loop state into later iterations. + super()._ensure_spec_dec_state_restored(attn_metadata, spec_metadata) + if attn_metadata is None or self._saved_extra_state is None: + return + ( + attn_metadata.all_rank_num_tokens, + attn_metadata.force_prepare_spec_dec_tree_mask, + ) = self._saved_extra_state + self._saved_extra_state = None + # ------------------------------------------------------------------ # # Helpers # # ------------------------------------------------------------------ # @@ -636,7 +658,7 @@ def _relocate_kv_eagerly(self, attn_metadata, batch_size): # Top-level forward # # ------------------------------------------------------------------ # @nvtx_range("mtp_dyn.forward") - def forward( + def _forward_impl( self, input_ids, position_ids, @@ -681,9 +703,13 @@ def forward( accepted_leaf_positions=accepted_leaf_positions, ) - # Save attn/spec metadata before the draft loop mutates it. - original_all_rank_num_tokens = attn_metadata.all_rank_num_tokens - original_force_prepare_spec_dec_tree_mask = attn_metadata.force_prepare_spec_dec_tree_mask + # Save attn/spec metadata before the draft loop mutates it. Stashed on + # self so _ensure_spec_dec_state_restored can undo the mutation if the + # draft loop raises (see SpecWorkerBase.forward try/finally). + self._saved_extra_state = ( + attn_metadata.all_rank_num_tokens, + attn_metadata.force_prepare_spec_dec_tree_mask, + ) self._prepare_attn_metadata_for_spec_dec(attn_metadata) attn_metadata.force_prepare_spec_dec_tree_mask = True @@ -706,8 +732,11 @@ def forward( # Restore attn metadata to support cuda graph. self._restore_attn_metadata_from_spec_dec(attn_metadata) - attn_metadata.all_rank_num_tokens = original_all_rank_num_tokens - attn_metadata.force_prepare_spec_dec_tree_mask = original_force_prepare_spec_dec_tree_mask + ( + attn_metadata.all_rank_num_tokens, + attn_metadata.force_prepare_spec_dec_tree_mask, + ) = self._saved_extra_state + self._saved_extra_state = None attn_metadata.use_spec_decoding = True # (d) Prepare next_new_tokens for overlap scheduler.