From 6e90065527230260d5696c3057cf71523596ddca Mon Sep 17 00:00:00 2001 From: Izzy Putterman Date: Wed, 28 Jan 2026 15:48:58 -0800 Subject: [PATCH] MTP Speculative Sampling Signed-off-by: Izzy Putterman --- tensorrt_llm/_torch/pyexecutor/model_engine.py | 5 +++-- tensorrt_llm/_torch/speculative/interface.py | 2 +- tensorrt_llm/_torch/speculative/one_model_sampler.py | 3 ++- tensorrt_llm/_torch/speculative/utils.py | 1 + 4 files changed, 7 insertions(+), 4 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 488079011ecb..7248a8266e64 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -52,7 +52,7 @@ from ..speculative.drafting_loops import BaseDraftingLoopWrapper from ..speculative.eagle3 import (Eagle3OneModelSpecMetadata, Eagle3ResourceManager, Eagle3SpecMetadata) -from ..speculative.mtp import SampleStateTensorsMTP +from ..speculative.mtp import MTPSpecMetadata, SampleStateTensorsMTP from ..speculative.utils import SpecDecodingTensor from ..utils import (get_model_extra_attrs, set_per_request_piecewise_cuda_graph_flag, @@ -2756,7 +2756,8 @@ def previous_seq_slots_device(): num_accepted_draft_tokens)] if isinstance(spec_metadata, Eagle3SpecMetadata): spec_metadata.request_accepted_path = request_accepted_path - if isinstance(spec_metadata, Eagle3OneModelSpecMetadata): + if isinstance(spec_metadata, + (Eagle3OneModelSpecMetadata, MTPSpecMetadata)): spec_metadata.populate_sampling_params_for_one_model( scheduled_requests.all_requests()) spec_metadata.prepare() diff --git a/tensorrt_llm/_torch/speculative/interface.py b/tensorrt_llm/_torch/speculative/interface.py index 6bc084665052..4650e66b9795 100644 --- a/tensorrt_llm/_torch/speculative/interface.py +++ b/tensorrt_llm/_torch/speculative/interface.py @@ -376,7 +376,7 @@ def __init__(self): super().__init__() self.guided_decoder: Optional["CapturableGuidedDecoder"] = None self.force_num_accepted_tokens = get_force_num_accepted_tokens() - self.use_flashinfer = IS_FLASHINFER_AVAILABLE and flashinfer.__version__ >= "0.6.0" + self.use_flashinfer = IS_FLASHINFER_AVAILABLE and flashinfer.__version__ >= "0.6.3" self.seed = 0 self.offset = 0 diff --git a/tensorrt_llm/_torch/speculative/one_model_sampler.py b/tensorrt_llm/_torch/speculative/one_model_sampler.py index 7d49aa85dd17..2afd214dc562 100644 --- a/tensorrt_llm/_torch/speculative/one_model_sampler.py +++ b/tensorrt_llm/_torch/speculative/one_model_sampler.py @@ -73,7 +73,8 @@ def apply_temperature( return logits.div_(temp.unsqueeze(dim=1)) -@torch.compile(options={"max-autotune": True}) +# Broken with current ToT (Jan 27) +# @torch.compile(options={"max-autotune": True}) def sampling_batch_spec_dec_one_model( logits: torch.Tensor, temperatures: torch.Tensor, diff --git a/tensorrt_llm/_torch/speculative/utils.py b/tensorrt_llm/_torch/speculative/utils.py index 6a22ad19bd4b..1eb036de99e7 100644 --- a/tensorrt_llm/_torch/speculative/utils.py +++ b/tensorrt_llm/_torch/speculative/utils.py @@ -31,6 +31,7 @@ def get_spec_metadata(spec_config, mtp_num_modules=spec_config.num_nextn_predict_layers, max_num_requests=max_num_requests, mtp_hidden_states_manager=spec_resource_manager, + allow_advanced_sampling=spec_config.allow_advanced_sampling, ) if spec_config.spec_dec_mode.is_mtp_eagle(): return Eagle3SpecMetadata(