diff --git a/tests/config/test_speculative_draft_hf_overrides.py b/tests/config/test_speculative_draft_hf_overrides.py index 7e425d68eecb..75b4bd2963e0 100644 --- a/tests/config/test_speculative_draft_hf_overrides.py +++ b/tests/config/test_speculative_draft_hf_overrides.py @@ -86,7 +86,7 @@ def record(hf_config: PretrainedConfig) -> PretrainedConfig: @pytest.mark.cpu_test -def test_inkling_override_exposes_only_first_mtp_depth(): +def test_inkling_override_exposes_all_mtp_depths(): text_config = _make_hf_config( architectures=["InklingForCausalLM"], model_type="inkling_model", @@ -107,7 +107,9 @@ def test_inkling_override_exposes_only_first_mtp_depth(): assert out is text_config assert out.model_type == "inkling_mtp" assert out.architectures == ["InklingMTPModel"] - assert out.n_predict == 1 + # Multi-module MTP: every checkpoint depth is exposed (module i drafts + # speculative token i), no longer clamped to the first depth. + assert out.n_predict == 8 assert out.num_nextn_predict_layers == 8 assert out.chain_hidden_post_norm is False assert out.local_layer_ids == [0, 2, 4] diff --git a/tests/v1/core/test_scheduler.py b/tests/v1/core/test_scheduler.py index c8c6fc85480a..c0932d975c18 100644 --- a/tests/v1/core/test_scheduler.py +++ b/tests/v1/core/test_scheduler.py @@ -3701,6 +3701,7 @@ def test_mamba_align_eagle_schedules_encoder_at_boundary(): ) scheduler.need_mamba_block_aligned_split = True scheduler.use_eagle = True + scheduler.num_prefill_lookahead = 1 scheduler.max_num_encoder_input_tokens = 2048 scheduler.encoder_cache_manager = EncoderCacheManager(cache_size=2048) @@ -5341,8 +5342,9 @@ def test_free_encoder_inputs_defers_for_eagle_lookahead(): worker-side token-embedding fallback is only a backstop.""" scheduler = create_scheduler(model="llava-hf/llava-1.5-7b-hf") # create_scheduler only builds ngram spec configs; force the eagle path that - # _free_encoder_inputs keys off (self.use_eagle). + # _free_encoder_inputs keys off (its read-ahead deferral). scheduler.use_eagle = True + scheduler.num_prefill_lookahead = 1 mm_positions = [[PlaceholderRange(offset=50, length=100)]] request = create_requests( num_requests=1, diff --git a/vllm/config/speculative.py b/vllm/config/speculative.py index d8f049755ce4..26c1b3e801c5 100644 --- a/vllm/config/speculative.py +++ b/vllm/config/speculative.py @@ -630,8 +630,7 @@ def hf_config_override(hf_config: PretrainedConfig) -> PretrainedConfig: hf_config.model_type = "inkling_mtp" hf_config.update( { - # Inkling currently exposes only the first checkpoint depth. - "n_predict": 1, + "n_predict": checkpoint_depths, "num_nextn_predict_layers": checkpoint_depths, "chain_hidden_post_norm": mtp_config.get( "chain_hidden_post_norm", False @@ -1083,14 +1082,6 @@ def __post_init__(self): "`num_speculative_tokens` was not provided" ) - if ( - self.draft_model_config.hf_config.model_type == "inkling_mtp" - and self.num_speculative_tokens != 1 - ): - raise ValueError( - "Inkling MTP currently supports exactly one speculative token" - ) - if self.dspark_draft_topk is not None and self.method != "dspark": raise ValueError("dspark_draft_topk is only supported by DSpark") diff --git a/vllm/v1/core/kv_cache_coordinator.py b/vllm/v1/core/kv_cache_coordinator.py index 321bbb0a76ac..f5cd79b285f6 100644 --- a/vllm/v1/core/kv_cache_coordinator.py +++ b/vllm/v1/core/kv_cache_coordinator.py @@ -80,6 +80,7 @@ def __init__( scheduler_block_size: int, hash_block_size: int, metrics_collector: KVCacheMetricsCollector | None = None, + num_prefill_lookahead: int = 0, ): self.kv_cache_config = kv_cache_config self.max_model_len = max_model_len @@ -91,6 +92,7 @@ def __init__( for g in kv_cache_config.kv_cache_groups ) self.scheduler_block_size = scheduler_block_size + self.num_reprefillable_tokens = max(0, num_prefill_lookahead - 1) self.block_pool = BlockPool( num_gpu_blocks=kv_cache_config.num_blocks, @@ -108,6 +110,28 @@ def __init__( if use_eagle and not self.eagle_group_ids: self.eagle_group_ids = set(range(len(kv_cache_config.kv_cache_groups))) + # During chunked prefill with EAGLE, the single next prefill lookahead + # token past the chunk boundary is combined with the final hidden state + # and written to the KV cache. Therefore, the final chunk token must be + # excluded from prefix cache hits to prevent requests from acquiring the + # KV cache slot polluted with the next prefill token, which may or may not + # be present after the matching prefix. The last-block drop handles this + # edge case. During multi-module MTP, the issue generalizes to a prefill + # lookahead of num_speculative_tokens, so the dropped tail must be large + # enough to contain them. Hits land on scheduler-block boundaries (see + # `_cache_hit_alignment_tokens`), so the excluded tail is + # scheduler_block_size, not the group's own block size. + if ( + enable_caching + and self.eagle_group_ids + and scheduler_block_size < num_prefill_lookahead + ): + raise ValueError( + f"Multi-module MTP with prefix caching requires scheduler_block_size" + f" (={scheduler_block_size}) >= num_speculative_tokens" + f" (={num_prefill_lookahead})." + ) + self.single_type_managers = tuple( get_manager_for_kv_cache_spec( kv_cache_spec=kv_cache_group.kv_cache_spec, @@ -286,9 +310,14 @@ def cache_blocks(self, request: Request, num_computed_tokens: int) -> None: (including tokens that are already cached). """ for manager in self.single_type_managers: + # Only cache tokens with finalized KV. The last num_reprefillable_tokens + # tokens can be re-prefilled during multi-module MTP. + num_tokens_to_cache = max( + 0, num_computed_tokens - self.num_reprefillable_tokens + ) manager.cache_blocks( request, - num_computed_tokens, + num_tokens_to_cache, retention_interval=self.retention_interval, ) @@ -407,6 +436,7 @@ def __init__( scheduler_block_size: int, hash_block_size: int, metrics_collector: KVCacheMetricsCollector | None = None, + num_prefill_lookahead: int = 0, ): super().__init__( kv_cache_config, @@ -420,6 +450,7 @@ def __init__( scheduler_block_size=scheduler_block_size, hash_block_size=hash_block_size, metrics_collector=metrics_collector, + num_prefill_lookahead=num_prefill_lookahead, ) self.num_single_type_manager = len(self.single_type_managers) @@ -457,6 +488,7 @@ def __init__( scheduler_block_size: int, hash_block_size: int, metrics_collector: KVCacheMetricsCollector | None = None, + num_prefill_lookahead: int = 0, ): super().__init__( kv_cache_config, @@ -470,6 +502,7 @@ def __init__( scheduler_block_size=scheduler_block_size, hash_block_size=hash_block_size, metrics_collector=metrics_collector, + num_prefill_lookahead=num_prefill_lookahead, ) self.kv_cache_spec = self.kv_cache_config.kv_cache_groups[0].kv_cache_spec self.block_size = self.kv_cache_spec.block_size @@ -542,6 +575,7 @@ def __init__( scheduler_block_size: int, hash_block_size: int, metrics_collector: KVCacheMetricsCollector | None = None, + num_prefill_lookahead: int = 0, ): super().__init__( kv_cache_config, @@ -555,6 +589,7 @@ def __init__( scheduler_block_size=scheduler_block_size, hash_block_size=hash_block_size, metrics_collector=metrics_collector, + num_prefill_lookahead=num_prefill_lookahead, ) # hash_block_size: the block size used to compute block hashes. # The actual block size usually equals hash_block_size, but in cases where @@ -688,9 +723,20 @@ def cache_blocks(self, request: Request, num_computed_tokens: int) -> None: # EAGLE groups match one block past each aligned boundary and drop # it, so make that lookahead block eligible to be cached. if manager.use_eagle and aligned_num_computed_tokens > 0: + # Only cache tokens with finalized KV. The last + # num_reprefillable_tokens tokens can be re-prefilled during + # multi-module MTP. + num_finalized_computed_tokens = max( + 0, num_computed_tokens - self.num_reprefillable_tokens + ) + aligned_num_finalized_computed_tokens = ( + num_finalized_computed_tokens + // self.scheduler_block_size + * self.scheduler_block_size + ) num_tokens_to_cache = min( - num_computed_tokens, - aligned_num_computed_tokens + manager.block_size, + num_finalized_computed_tokens, + aligned_num_finalized_computed_tokens + manager.block_size, ) # The manager already knows the fine hit granularity # (``scheduler_block_size``); retention is passed separately so it @@ -880,6 +926,7 @@ def get_kv_cache_coordinator( scheduler_block_size: int, hash_block_size: int, metrics_collector: KVCacheMetricsCollector | None = None, + num_prefill_lookahead: int = 0, ) -> KVCacheCoordinator: if not enable_caching: return KVCacheCoordinatorNoPrefixCache( @@ -893,6 +940,7 @@ def get_kv_cache_coordinator( scheduler_block_size=scheduler_block_size, hash_block_size=hash_block_size, metrics_collector=metrics_collector, + num_prefill_lookahead=num_prefill_lookahead, ) if len(kv_cache_config.kv_cache_groups) == 1: return UnitaryKVCacheCoordinator( @@ -907,6 +955,7 @@ def get_kv_cache_coordinator( scheduler_block_size=scheduler_block_size, hash_block_size=hash_block_size, metrics_collector=metrics_collector, + num_prefill_lookahead=num_prefill_lookahead, ) return HybridKVCacheCoordinator( kv_cache_config, @@ -920,4 +969,5 @@ def get_kv_cache_coordinator( scheduler_block_size=scheduler_block_size, hash_block_size=hash_block_size, metrics_collector=metrics_collector, + num_prefill_lookahead=num_prefill_lookahead, ) diff --git a/vllm/v1/core/kv_cache_manager.py b/vllm/v1/core/kv_cache_manager.py index ca1fb73420a2..44097c3da276 100644 --- a/vllm/v1/core/kv_cache_manager.py +++ b/vllm/v1/core/kv_cache_manager.py @@ -125,6 +125,7 @@ def __init__( max_in_flight_tokens: int | None = None, enable_caching: bool = True, use_eagle: bool = False, + num_prefill_lookahead: int = 0, log_stats: bool = False, enable_kv_cache_events: bool = False, dcp_world_size: int = 1, @@ -161,6 +162,7 @@ def __init__( scheduler_block_size=scheduler_block_size, hash_block_size=hash_block_size, metrics_collector=self.metrics_collector, + num_prefill_lookahead=num_prefill_lookahead, ) self.num_kv_cache_groups = len(kv_cache_config.kv_cache_groups) self.block_pool = self.coordinator.block_pool diff --git a/vllm/v1/core/kv_cache_utils.py b/vllm/v1/core/kv_cache_utils.py index 57d6600b368e..ad51e4b2b397 100644 --- a/vllm/v1/core/kv_cache_utils.py +++ b/vllm/v1/core/kv_cache_utils.py @@ -1468,6 +1468,9 @@ def promoted_page_size_padded(spec: AttentionSpec, block_size: int) -> int | Non promoted_specs[layer_name] = replace_as( spec, target_cls, + # Promoted specs allocate blocks for all tokens and never free + # below the window, so the trailing-edge extension is moot. + drop=("extra_retained_tokens",), block_size=block_size, page_size_padded=promoted_page_size_padded(spec, block_size), ) @@ -2128,6 +2131,21 @@ def get_kv_cache_configs( # Check if the KV cache specs are registered correctly. # This is to prevent that some layers are initialized with unregistered specs. KVCacheSpecRegistry.check_kv_cache_spec_registry(merged_kv_cache_specs) + + # When speculating with more than 1 speculative module (e.g. multi-layered MTP) + # tag every SlidingWindowSpec with how many extra tokens to retain in the window. + extra_retained_tokens = ( + vllm_config.speculative_config.num_speculative_tokens - 1 + if vllm_config.speculative_config is not None + and vllm_config.speculative_config.use_multi_module_mtp() + else 0 + ) + for layer_name, layer_spec in merged_kv_cache_specs.items(): + if isinstance(layer_spec, SlidingWindowSpec): + merged_kv_cache_specs[layer_name] = replace( + layer_spec, extra_retained_tokens=extra_retained_tokens + ) + # Get global KV cache groups. This also handles spec unification for # hybrid models when disable_hybrid_kv_cache_manager is enabled. # After this call, merged_kv_cache_specs may be modified in-place. diff --git a/vllm/v1/core/sched/scheduler.py b/vllm/v1/core/sched/scheduler.py index 9d14d00edc56..e4a21328660a 100644 --- a/vllm/v1/core/sched/scheduler.py +++ b/vllm/v1/core/sched/scheduler.py @@ -247,6 +247,13 @@ def __init__( self.use_eagle = False self.num_spec_tokens = vllm_config.num_speculative_tokens self.num_lookahead_tokens = vllm_config.num_lookahead_tokens + # Positions past the computed tokens that the drafter reads mid-prefill. + # Eagle-family drafters read 1 ahead, but multi-module MTP reads + # num_spec_tokens ahead at chunked-prefill boundaries. Determines the + # encoder scheduling shift, the deferred encoder free, the KV cache + # manager's re-prefillable window (this minus 1), and how many tokens to + # reserve between a chunk boundary and the prefill end. + self.num_prefill_lookahead = 0 self.dynamic_sd_lookup: list[int] | None = None if speculative_config is not None: if speculative_config.num_speculative_tokens_per_batch_size: @@ -256,6 +263,12 @@ def __init__( vllm_num_speculative_tokens=self.num_spec_tokens, ) self.use_eagle = speculative_config.use_eagle() + if self.use_eagle: + self.num_prefill_lookahead = ( + self.num_spec_tokens + if speculative_config.use_multi_module_mtp() + else 1 + ) # Create the KV cache manager. if hash_block_size is None: @@ -267,6 +280,7 @@ def __init__( max_in_flight_tokens=vllm_config.max_in_flight_tokens, enable_caching=self.cache_config.enable_prefix_caching, use_eagle=self.use_eagle, + num_prefill_lookahead=self.num_prefill_lookahead, log_stats=self.log_stats, enable_kv_cache_events=self.enable_kv_cache_events, dcp_world_size=self.dcp_world_size, @@ -438,6 +452,27 @@ def _get_local_prefix_cache_hit( ) return blocks, num_local, shared_prefix_boundary, False + def _reserve_prefill_lookahead( + self, + request: Request, + num_computed_tokens: int, + num_new_tokens: int, + ) -> int: + """Never end a prefill chunk within num_prefill_lookahead of the + prefill end. + + At a chunked-prefill boundary, the multi-module MTP drafter consumes + the next num_prefill_lookahead known prefill tokens as draft inputs. A + boundary closer to the end than that would make it fall back to + sampled drafts, permanently polluting the trailing modules' KV caches. + Either finish the prefill or leave at least num_prefill_lookahead for + the next chunk. No-op for eagle-family drafters (lookahead 1). + """ + remaining = request.num_tokens - num_computed_tokens - num_new_tokens + if 0 < remaining < self.num_prefill_lookahead: + num_new_tokens -= self.num_prefill_lookahead - remaining + return max(num_new_tokens, 0) + def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: self.current_step += 1 # NOTE(woosuk) on the scheduling algorithm: @@ -561,9 +596,15 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: request.num_computed_tokens, num_new_tokens, encoder_compute_budget, - shift_computed_tokens=1 if self.use_eagle else 0, + shift_computed_tokens=self.num_prefill_lookahead, ) + # Multi-module MTP: avoid ending a prefill chunk within + # num_prefill_lookahead of the prefill end. + num_new_tokens = self._reserve_prefill_lookahead( + request, request.num_computed_tokens, num_new_tokens + ) + if num_new_tokens == 0: # The request cannot be scheduled because one of the following # reasons: @@ -576,6 +617,8 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: # 3. The encoder cache is exhausted. # 4. Insufficient budget for a block-aligned chunk in hybrid # models with mamba cache mode \"align\". + # 5. Insufficient budget to keep a multi-module MTP prefill + # chunk out of the prefill-lookahead window. # NOTE(woosuk): Here, by doing `continue` instead of `break`, # we do not strictly follow the FCFS scheduling policy and # allow the lower-priority requests to be scheduled. @@ -946,11 +989,18 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: num_computed_tokens, num_new_tokens, encoder_compute_budget, - shift_computed_tokens=1 if self.use_eagle else 0, + shift_computed_tokens=self.num_prefill_lookahead, ) - if num_new_tokens == 0: - # The request cannot be scheduled. - break + + # Multi-module MTP: avoid ending a prefill chunk within + # num_prefill_lookahead of the prefill end. + num_new_tokens = self._reserve_prefill_lookahead( + request, num_computed_tokens, num_new_tokens + ) + + if num_new_tokens == 0: + # The request cannot be scheduled. + break # During async KV load, no forward pass is run yet. # Allocate speculative lookahead slots later to avoid @@ -2157,9 +2207,9 @@ def _free_encoder_inputs(self, request: Request) -> None: return # Defer the free by the drafter's look-ahead so an entry stays - # referenced until the drafter's +1 read has also passed it, mirroring - # the shift the encoder scheduling path applies. - spec_lookahead = 1 if self.use_eagle else 0 + # referenced until the drafter's read-ahead has also passed it, + # mirroring the shift the encoder scheduling path applies. + spec_lookahead = self.num_prefill_lookahead # Here, we use list(set) to avoid modifying the set while iterating # over it. diff --git a/vllm/v1/core/single_type_kv_cache_manager.py b/vllm/v1/core/single_type_kv_cache_manager.py index f5f854a5dfd9..ded6bc809a3f 100644 --- a/vllm/v1/core/single_type_kv_cache_manager.py +++ b/vllm/v1/core/single_type_kv_cache_manager.py @@ -561,6 +561,11 @@ def find_longest_cache_hit( return an empty list. If eagle is enabled, drop the last matched block to force recompute the last block to get the required hidden states for eagle drafting head. + For multi-module MTP, this recompute also rewrites the dropped block's + draft-layer KVs, which depend on up to num_speculative_tokens - 1 + tokens past the matched prefix (i.e. on the cache writer's + continuation, which the block hash does not cover); the coordinator + asserts the block size covers that window. Need to be customized for each attention type. Args: @@ -877,6 +882,10 @@ class SlidingWindowManager(SingleTypeKVCacheManager): def __init__(self, kv_cache_spec: SlidingWindowSpec, **kwargs) -> None: super().__init__(kv_cache_spec, **kwargs) self.sliding_window = kv_cache_spec.sliding_window + # Extra trailing tokens to retain below the window (never attended) so a + # multi-module MTP store-side lag can still reconstruct the window from + # cached blocks. + self.extra_retained_tokens = kv_cache_spec.extra_retained_tokens @classmethod def _contiguous_blocks_for_hit( @@ -1072,13 +1081,22 @@ def get_num_skipped_tokens(self, num_computed_tokens: int) -> int: attention computation since they are outside the sliding window. Thus, get_num_skipped_tokens(7) == 4. + The trailing edge of the window is extended by ``extra_retained_tokens`` + so that those extra trailing tokens' blocks are retained (but not + attended). This is needed for multi-module spec decoding which can + re-prefill the last num_spec_prefill_tokens - 1 tokens from the end + of the sequence, and thus needs to delay freeing/caching of blocks. + Args: num_computed_tokens: The number of tokens that have been computed. Returns: The number of tokens that will be skipped for attention computation. """ - return max(0, num_computed_tokens - self.sliding_window + 1) + return max( + 0, + num_computed_tokens - self.sliding_window + 1 - self.extra_retained_tokens, + ) def get_num_common_prefix_blocks(self, running_request_id: str) -> int: """ diff --git a/vllm/v1/kv_cache_interface.py b/vllm/v1/kv_cache_interface.py index dcc8298781b5..1ab3facb4985 100644 --- a/vllm/v1/kv_cache_interface.py +++ b/vllm/v1/kv_cache_interface.py @@ -5,6 +5,7 @@ import copy from collections import Counter +from collections.abc import Collection from dataclasses import dataclass, fields, replace from enum import Enum, IntEnum from math import prod @@ -88,14 +89,24 @@ def is_quantized_kv_cache(kv_cache_dtype: str) -> bool: return get_kv_quant_mode(kv_cache_dtype) != KVQuantMode.NONE -def replace_as(spec: KVCacheSpec, target_cls: type[_SpecT], **changes) -> _SpecT: +def replace_as( + spec: KVCacheSpec, + target_cls: type[_SpecT], + *, + drop: Collection[str] = (), + **changes, +) -> _SpecT: """``dataclasses.replace``, but rebuilding *spec* as *target_cls* e.g. ``SlidingWindowSpec`` -> ``FullAttentionSpec`` - Every field of *spec* must exist on *target_cls*; fields only *target_cls* has keep - their default values. + Every field of *spec* must exist on *target_cls* unless named in *drop*; + fields only *target_cls* has keep their default values. """ - kwargs = {f.name: getattr(spec, f.name) for f in fields(spec) if f.init} + kwargs = { + f.name: getattr(spec, f.name) + for f in fields(spec) + if f.init and f.name not in drop + } kwargs.update(changes) return target_cls(**kwargs) @@ -573,6 +584,12 @@ def is_uniform_with_collection( class SlidingWindowSpec(AttentionSpec): sliding_window: int head_size_v: int = None # type: ignore[assignment] + # The trailing edge of the window is extended by ``extra_retained_tokens`` + # so that those extra trailing tokens' blocks are retained (but not + # attended). This is needed for multi-module spec decoding which can + # re-prefill the last num_spec_prefill_tokens - 1 tokens from the end + # of the sequence, and thus needs to delay freeing/caching of blocks. + extra_retained_tokens: int = 0 def __post_init__(self): if self.head_size_v is None: @@ -614,8 +631,13 @@ def max_admission_blocks_per_request( """ # During chunked prefill, we hold KV for the last `sliding_window-1` # computed tokens plus the in-flight tokens (frees happen on the - # processed-token basis); never more than `max_model_len`. - num_tokens = min(self.sliding_window - 1 + max_in_flight_tokens, max_model_len) + # processed-token basis); never more than `max_model_len`. An additional + # `extra_retained_tokens` trailing tokens are kept alive below the + # window for multi-module spec decoding, and must be accounted here too. + num_tokens = min( + self.sliding_window - 1 + self.extra_retained_tokens + max_in_flight_tokens, + max_model_len, + ) # +1 because the sliding window may not start from the beginning of # the block. E.g. block size 4 and num_token 4 needs two blocks # [XXCD][EF] to store the 6-token window [CDEF]. @@ -686,16 +708,18 @@ def merge(cls, specs: list[Self]) -> Self: model_version_set = set(spec.model_version for spec in specs) sliding_window_set = set(spec.sliding_window for spec in specs) block_stride_set = set(spec.indexes_kv_by_block_stride for spec in specs) + extra_retained_set = set(spec.extra_retained_tokens for spec in specs) assert ( len(cache_dtype_str_set) == 1 and len(compress_ratio_set) == 1 and len(model_version_set) == 1 and len(sliding_window_set) == 1 and len(block_stride_set) == 1 + and len(extra_retained_set) == 1 ), ( "All attention layers in the same KV cache group must use the same " "quantization method, compress ratio, model version, sliding " - "window size, and KV block stride indexing." + "window size, KV block stride indexing, and retained token count." ) return cls( block_size=specs[0].block_size, @@ -705,6 +729,7 @@ def merge(cls, specs: list[Self]) -> Self: page_size_padded=specs[0].page_size_padded, indexes_kv_by_block_stride=block_stride_set.pop(), sliding_window=sliding_window_set.pop(), + extra_retained_tokens=extra_retained_set.pop(), cache_dtype_str=cache_dtype_str_set.pop(), compress_ratio=compress_ratio_set.pop(), model_version=model_version_set.pop(), diff --git a/vllm/v1/worker/gpu/spec_decode/__init__.py b/vllm/v1/worker/gpu/spec_decode/__init__.py index c70f169f7be6..4229696f255c 100644 --- a/vllm/v1/worker/gpu/spec_decode/__init__.py +++ b/vllm/v1/worker/gpu/spec_decode/__init__.py @@ -26,6 +26,12 @@ def init_speculator(vllm_config: VllmConfig, device: torch.device): ) return Gemma4Speculator(vllm_config, device) + elif speculative_config.use_multi_module_mtp(): + from vllm.v1.worker.gpu.spec_decode.multi_module_mtp.speculator import ( + MultiModuleMTPSpeculator, + ) + + return MultiModuleMTPSpeculator(vllm_config, device) elif speculative_config.method == "mtp": from vllm.v1.worker.gpu.spec_decode.mtp.speculator import MTPSpeculator