Skip to content
Open
Show file tree
Hide file tree
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
Original file line number Diff line number Diff line change
Expand Up @@ -775,7 +775,7 @@ def get_output_embeddings(self):

@dataclass
class EagleWrapperOutput(ModelOutput):
"""Output format compatible with Eagle3OneModelSampler/MTPSampler.
"""Output format compatible with SpecSampler.
This output format allows the one-model speculative decoding flow to bypass
logits-based sampling in the sampler. The EagleWrapper performs all sampling
Expand Down
4 changes: 2 additions & 2 deletions tensorrt_llm/_torch/auto_deploy/shim/ad_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,7 +48,7 @@
SimpleScheduler,
)
from tensorrt_llm._torch.pyexecutor.seq_slot_manager import SeqSlotManager
from tensorrt_llm._torch.speculative.eagle3 import Eagle3OneModelSampler
from tensorrt_llm._torch.speculative.spec_sampler_base import SpecSampler
from tensorrt_llm._utils import get_free_port, mpi_rank, mpi_world_size, nvtx_range
from tensorrt_llm.inputs.multimodal import MultimodalRuntimeData, check_mm_embed_cumsum_if_needed
from tensorrt_llm.llmapi.llm_args import ContextChunkingPolicy, MultimodalConfig, SamplerType
Expand Down Expand Up @@ -1113,7 +1113,7 @@ def instantiate_sampler(
max_beam_width=ad_config.max_beam_width,
disable_overlap_scheduler=ad_config.disable_overlap_scheduler,
)
return Eagle3OneModelSampler(sampler_args)
return SpecSampler(sampler_args)

sampler_type = ad_config.sampler_type
if sampler_type == SamplerType.auto:
Expand Down
2 changes: 1 addition & 1 deletion tensorrt_llm/_torch/pyexecutor/sampler/top_p_decay.py
Original file line number Diff line number Diff line change
Expand Up @@ -153,7 +153,7 @@ def validate_request(request: LlmRequest) -> None:
# tokens and produces multiple tokens per step (req_num_steps =
# 1 + draft_token_length). One-model speculation (vanilla MTP, one-model
# Eagle3 / MTP-Eagle, SA, draft-target-one-model) uses its own
# SpecSamplerBase-derived sampler and never reaches TorchSampler; the
# SpecSampler and never reaches TorchSampler; the
# drafter-based modes that DO flow draft tokens through TorchSampler
# (two-model draft-target, NGram, user-provided, two-model Eagle3 /
# MTP-Eagle) are what can make this length non-zero. top-p decay does not
Expand Down
10 changes: 4 additions & 6 deletions tensorrt_llm/_torch/speculative/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,15 +7,15 @@
prepare_attn_metadata_for_draft_replay,
restore_attn_metadata_after_draft_replay,
should_use_separate_draft_kv_cache)
from .mtp import MTPSampler, MTPSpecMetadata, MTPWorker
from .mtp import MTPSpecMetadata, MTPWorker
from .ngram import NGramDrafter, NGramPoolManager
from .pard import PARDSpecMetadata, PARDWorker
from .sa_enhancer import SADraftEnhancer
from .sa_worker import SASampler, SASpecMetadata, SAWorker
from .sa_worker import SASpecMetadata, SAWorker
from .save_hidden_state import (SaveHiddenStatesResourceManager,
SaveHiddenStatesSpecMetadata)
from .spec_sampler_base import (SampleStateSpec, SampleStateTensorsSpec,
SpecSamplerBase)
SpecSampler)
from .spec_tree_manager import SpecTreeManager
from .suffix_automaton import SuffixAutomatonManager
from .utils import (get_draft_kv_cache_manager, get_num_extra_kv_tokens,
Expand All @@ -32,15 +32,13 @@
"DraftTargetOneModelWorker",
"Eagle3SpecMetadata",
"MTPEagleWorker",
"MTPSampler",
"MTPSpecMetadata",
"MTPWorker",
"NGramDrafter",
"NGramPoolManager",
"PARDSpecMetadata",
"PARDWorker",
"SADraftEnhancer",
"SASampler",
"SASpecMetadata",
"SAWorker",
"SuffixAutomatonManager",
Expand All @@ -49,7 +47,7 @@
"SaveHiddenStatesResourceManager",
"SaveHiddenStatesSpecMetadata",
"SpecMetadata",
"SpecSamplerBase",
"SpecSampler",
"SpecWorkerBase",
"get_draft_kv_cache_manager",
"get_num_extra_kv_tokens",
Expand Down
13 changes: 0 additions & 13 deletions tensorrt_llm/_torch/speculative/draft_target.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,9 +30,7 @@
from tensorrt_llm.mapping import Mapping

from ..attention_backend import AttentionMetadata
from ..pyexecutor.sampler import TorchSampler
from .interface import SpecMetadata, SpecWorkerBase
from .mtp import MTPSampler

if TYPE_CHECKING:
from ...llmapi.llm_args import DraftTargetDecodingConfig
Expand Down Expand Up @@ -75,17 +73,6 @@ def prepare(self):
self.is_spec_dec_dynamic_tree = False


class DraftTargetOneModelSampler(MTPSampler):
"""
Sampler for DraftTarget one-model speculative decoding.

Inherits from MTPSampler to reuse the speculative decoding sampling logic.
"""

def __init__(self, args: TorchSampler.Args):
super().__init__(args, nextn=args.max_draft_len)


class DraftTargetOneModelWorker(SpecWorkerBase):
def __init__(
self,
Expand Down
2 changes: 1 addition & 1 deletion tensorrt_llm/_torch/speculative/drafter.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,7 +67,7 @@ def should_use_spec_decode(self, requests: List[LlmRequest],
# Drafters that use TorchSampler (NGram, two-model) compute py_rewind_len
# from len(py_draft_tokens), which includes padding. They must set this
# to True so that extend_capacity_for_tokens is called after padding.
# One-model drafters (MTP / Eagle3 / SA) use SpecSamplerBase which
# One-model drafters (MTP / Eagle3 / SA) use SpecSampler which
# computes rewind from runtime_draft_len, so padding is harmless.
_needs_padding_kv_extension: bool = False

Expand Down
19 changes: 1 addition & 18 deletions tensorrt_llm/_torch/speculative/eagle3.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,10 +14,9 @@
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, _select_mtp_position_ids
from .mtp import _select_mtp_position_ids
from .sa_enhancer import SADraftEnhancer
from .spec_tree_manager import SpecTreeManager

Expand Down Expand Up @@ -591,22 +590,6 @@ def maybe_capture_hidden_states(
break


class Eagle3OneModelSampler(MTPSampler):
"""Sampler for one-model EAGLE3 (linear and dynamic tree modes)."""

def __init__(self, args: TorchSampler.Args, spec_config=None):
self._spec_config = spec_config
super().__init__(args, nextn=args.max_total_draft_tokens)

def _get_max_new_tokens(self, args: TorchSampler.Args,
draft_len: int) -> int:
"""Dynamic tree: accepted path depth <= max_draft_len + 1."""
if (self._spec_config is not None
and getattr(self._spec_config, 'use_dynamic_tree', False)):
return self._spec_config.max_draft_len + 1
return self._get_max_tokens(args, draft_len)


class Eagle3OneModelWorker(SpecWorkerBase):
"""Unified one-model worker for Eagle3 and MTP Eagle speculative decoding.

Expand Down
2 changes: 1 addition & 1 deletion tensorrt_llm/_torch/speculative/interface.py
Original file line number Diff line number Diff line change
Expand Up @@ -750,7 +750,7 @@ def _normalize_request_sampling_params(
# decoding does not support min_p (there is no request_min_p buffer
# nor min_p wiring in the sampling_batch_spec_dec_one_model*
# kernels); a min_p request is rejected at admission by
# SpecSamplerBase.validate_request, so nothing reaching this scan
# SpecSampler.validate_request, so nothing reaching this scan
# carries a min_p that would change its classification. The
# two-model draft/target path honors min_p via _request_strategy.
is_greedy = SamplingParams.params_imply_greedy_decoding(
Expand Down
42 changes: 0 additions & 42 deletions tensorrt_llm/_torch/speculative/mtp.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,3 @@
import sys
from dataclasses import dataclass
from typing import TYPE_CHECKING, List, Optional

Expand All @@ -9,22 +8,13 @@
from ..attention_backend import AttentionMetadata
from ..pyexecutor.llm_request import LlmRequest
from ..pyexecutor.resource_manager import BaseResourceManager, SlotManager
from ..pyexecutor.sampler import TorchSampler
from ..pyexecutor.scheduler import ScheduledRequests
from .interface import SpecMetadata, SpecWorkerBase
from .sa_enhancer import SADraftEnhancer
from .spec_sampler_base import SampleStateSpec, SpecSamplerBase

if TYPE_CHECKING:
from tensorrt_llm.llmapi.llm_args import MTPDecodingConfig

if sys.version_info[:2] >= (3, 12):
from typing import override
else:
from typing_extensions import override

SampleStateMTP = SampleStateSpec


def _normalize_mtp_position_ids(position_ids: torch.Tensor) -> torch.Tensor:
"""Collapse plain [1, N] position IDs while preserving MRoPE axes."""
Expand Down Expand Up @@ -239,38 +229,6 @@ def prepare(self):
sa_manager.prepare(gen_request_ids, self.runtime_draft_len)


class MTPSampler(SpecSamplerBase):
"""
MTP sampler.

Inherits from SpecSamplerBase with overrides for tree-based speculation
using max_total_draft_tokens instead of draft_len.
"""

SampleState = SampleStateMTP

@override
def is_generation_model(self) -> bool:
return True

def setup_sampler_step(self, scheduled_requests: ScheduledRequests):
pass

def __init__(self, args: TorchSampler.Args, *, nextn: int):
super().__init__(args, draft_len=nextn)

@override
def _get_max_tokens(self, args: TorchSampler.Args, draft_len: int) -> int:
"""MTP uses max_total_draft_tokens + 1 for tree-based speculation."""
return args.max_total_draft_tokens + 1

@override
def _get_draft_tokens_storage_size(self, args: TorchSampler.Args,
draft_len: int) -> int:
"""MTP uses max_total_draft_tokens for draft token storage."""
return args.max_total_draft_tokens


class MTPWorker(SpecWorkerBase):

def __init__(self,
Expand Down
20 changes: 0 additions & 20 deletions tensorrt_llm/_torch/speculative/sa_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,6 @@
Key components:
- SASpecMetadata: Metadata for SA speculative decoding
- SAWorker: Spec worker that uses suffix automaton for draft generation
- SASampler: Sampler that handles GPU->CPU result extraction
"""

from dataclasses import dataclass, field
Expand All @@ -32,9 +31,7 @@
from tensorrt_llm._utils import prefer_pinned

from ..pyexecutor.mamba_cache_manager import MambaHybridCacheManager
from ..pyexecutor.sampler import TorchSampler
from .interface import SpecMetadata, SpecWorkerBase
from .spec_sampler_base import SampleStateSpec, SpecSamplerBase
from .suffix_automaton import SuffixAutomatonManager

if TYPE_CHECKING:
Expand Down Expand Up @@ -334,20 +331,3 @@ def _generate_draft_tokens(
draft_tokens = draft_tokens * mask

return draft_tokens # [batch_size, max_draft_len] GPU tensor


class SASampler(SpecSamplerBase):
"""
Sampler for SA that extracts GPU results to CPU after graph replay.

Uses SpecSamplerBase with default behavior (draft_len + 1 storage,
adds dummy draft tokens for context requests).
"""

SampleState = SampleStateSpec

def __init__(self, args: TorchSampler.Args, *, max_draft_len: int):
super().__init__(args, draft_len=max_draft_len)

def is_generation_model(self) -> bool:
return True
Loading
Loading