Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
41 changes: 35 additions & 6 deletions tensorrt_llm/_torch/speculative/mtp_dynamic_tree.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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 #
# ------------------------------------------------------------------ #
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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

Expand All @@ -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.
Expand Down
Loading