From d851ffb7f2476311cb99bdbe1651a0f328141992 Mon Sep 17 00:00:00 2001 From: Giancarlo Delfin Date: Thu, 5 Mar 2026 05:22:10 +0000 Subject: [PATCH 1/2] [Model Runner V2] Support multi-modal embeddings for spec decode model Signed-off-by: Giancarlo Delfin --- vllm/v1/worker/gpu/model_runner.py | 25 ++++++++++++++++++ .../gpu/spec_decode/eagle/speculator.py | 26 +++++++++++++++++++ 2 files changed, 51 insertions(+) diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index 7268b8ac191f..5d6b88795730 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -432,6 +432,16 @@ def _dummy_run( # dummy run the eagle speculator's propose to ensure DP/EP sync. if self.speculator is not None: assert self.sampler is not None + mm_inputs: tuple[list[torch.Tensor], torch.Tensor] | None = None + if self.speculator.supports_mm_inputs: + mm_inputs = ( + [], + torch.zeros( + input_batch.num_tokens, + dtype=torch.bool, + device=self.device, + ), + ) self.speculator.propose( input_batch=input_batch, attn_metadata=attn_metadata, @@ -452,6 +462,7 @@ def _dummy_run( num_tokens_across_dp=num_tokens_across_dp, dummy_run=True, skip_attn_for_dummy_run=skip_attn, + mm_inputs=mm_inputs, ) assert hidden_states is not None # Last PP rank always has hidden_states @@ -1126,6 +1137,19 @@ def sample_tokens( copy_event=self.output_copy_event, ) + mm_inputs = None + if self.speculator is not None and self.speculator.supports_mm_inputs: + # Get cached multimodal embeddings for draft forward. + mm_inputs = self.model_state.encoder_runner.gather_mm_embeddings( + input_batch.req_ids, + input_batch.num_tokens, + input_batch.num_scheduled_tokens, + input_batch.query_start_loc_np, + self.req_states.prefill_len.np[input_batch.idx_mapping_np], + self.req_states.num_computed_prefill_tokens[input_batch.idx_mapping_np] + + 1, + ) + # Postprocess results and update request states. # NOTE: This is intentionally done after creating the AsyncOutput, # ensuring that `copy_event` is recorded before calling postprocess. @@ -1150,6 +1174,7 @@ def sample_tokens( self.sampler.sampling_states.seeds.gpu, self.req_states.draft_logits, num_tokens_across_dp=num_tokens_across_dp, + mm_inputs=mm_inputs, ) self.req_states.draft_tokens[input_batch.idx_mapping] = draft_tokens self.draft_tokens_handler.set_draft_tokens(input_batch, draft_tokens) diff --git a/vllm/v1/worker/gpu/spec_decode/eagle/speculator.py b/vllm/v1/worker/gpu/spec_decode/eagle/speculator.py index 922031a52180..b15e8fde7e39 100644 --- a/vllm/v1/worker/gpu/spec_decode/eagle/speculator.py +++ b/vllm/v1/worker/gpu/spec_decode/eagle/speculator.py @@ -9,6 +9,7 @@ from vllm.config.compilation import CUDAGraphMode from vllm.forward_context import BatchDescriptor, set_forward_context from vllm.logger import init_logger +from vllm.multimodal import MULTIMODAL_REGISTRY from vllm.triton_utils import tl, triton from vllm.v1.kv_cache_interface import KVCacheConfig from vllm.v1.worker.gpu.attn_utils import ( @@ -75,6 +76,13 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device): device=device, ) + self.supports_mm_inputs = MULTIMODAL_REGISTRY.supports_multimodal_inputs( + self.draft_model_config + ) + self.inputs_embeds = torch.zeros( + self.max_num_tokens, self.hidden_size, dtype=self.dtype, device=device + ) + # currently we don't support PIECEWISE for Eagle. cudagraph_mode = vllm_config.compilation_config.cudagraph_mode if cudagraph_mode.decode_mode() == CUDAGraphMode.FULL: @@ -109,6 +117,7 @@ def run_model( slot_mappings: dict[str, torch.Tensor] | None, num_tokens_across_dp: torch.Tensor | None, cudagraph_runtime_mode: CUDAGraphMode = CUDAGraphMode.NONE, + mm_inputs: tuple[list[torch.Tensor], torch.Tensor] | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: batch_descriptor = BatchDescriptor(num_tokens=num_tokens) with set_forward_context( @@ -120,10 +129,25 @@ def run_model( slot_mapping=slot_mappings, batch_descriptor=batch_descriptor, ): + inputs_embeds = None + if self.supports_mm_inputs: + # Merge multimodal embeddings with input ids. + mm_embeds, is_mm_embed = mm_inputs or (None, None) + num_input_tokens = ( + is_mm_embed.shape[0] if is_mm_embed is not None else num_tokens + ) + self.inputs_embeds[:num_input_tokens] = self.model.embed_input_ids( + self.input_buffers.input_ids[:num_input_tokens], + multimodal_embeddings=mm_embeds, + is_multimodal=is_mm_embed, + ) + inputs_embeds = self.inputs_embeds[:num_tokens] + ret_hidden_states = self.model( input_ids=self.input_buffers.input_ids[:num_tokens], positions=self.input_buffers.positions[:num_tokens], hidden_states=self.hidden_states[:num_tokens], + inputs_embeds=inputs_embeds, ) if self.method == "mtp": last_hidden_states = ret_hidden_states @@ -228,6 +252,7 @@ def propose( num_tokens_across_dp: torch.Tensor | None = None, dummy_run: bool = False, skip_attn_for_dummy_run: bool = False, + mm_inputs: tuple[list[torch.Tensor], torch.Tensor] | None = None, ) -> torch.Tensor: # NOTE(woosuk): To avoid CPU-GPU synchronization without CPU knowing the # number of rejected tokens, we maintain the size of eagle's input_ids and @@ -262,6 +287,7 @@ def propose( attn_metadata, slot_mappings, num_tokens_across_dp=num_tokens_across_dp, + mm_inputs=mm_inputs, ) sample_hidden_states = last_hidden_states[last_token_indices] logits = self.model.compute_logits(sample_hidden_states) From 662ed5e197e23087640a400b15371de5ade772a4 Mon Sep 17 00:00:00 2001 From: Woosuk Kwon Date: Sun, 22 Mar 2026 08:01:34 +0000 Subject: [PATCH 2/2] minor Signed-off-by: Woosuk Kwon --- vllm/v1/worker/gpu/model_runner.py | 32 ++++++++++++++++++------------ 1 file changed, 19 insertions(+), 13 deletions(-) diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index 23203b4e7642..d10530c95884 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -1145,19 +1145,6 @@ def sample_tokens( copy_event=self.output_copy_event, ) - mm_inputs = None - if self.speculator is not None and self.speculator.supports_mm_inputs: - # Get cached multimodal embeddings for draft forward. - mm_inputs = self.model_state.encoder_runner.gather_mm_embeddings( - input_batch.req_ids, - input_batch.num_tokens, - input_batch.num_scheduled_tokens, - input_batch.query_start_loc_np, - self.req_states.prefill_len.np[input_batch.idx_mapping_np], - self.req_states.num_computed_prefill_tokens[input_batch.idx_mapping_np] - + 1, - ) - # Postprocess results and update request states. # NOTE: This is intentionally done after creating the AsyncOutput, # ensuring that `copy_event` is recorded before calling postprocess. @@ -1166,8 +1153,27 @@ def sample_tokens( self.postprocess( input_batch, sampler_output.sampled_token_ids, num_sampled, num_rejected ) + if self.speculator is not None: assert self.sampler is not None + mm_inputs: tuple[list[torch.Tensor], torch.Tensor] | None = None + if self.speculator.supports_mm_inputs: + # Get cached multimodal embeddings for draft forward. + prefill_lens = self.req_states.prefill_len.np[ + input_batch.idx_mapping_np + ] + computed_prefill_lens = self.req_states.num_computed_prefill_tokens[ + input_batch.idx_mapping_np + ] + mm_inputs = self.model_state.encoder_runner.gather_mm_embeddings( + input_batch.req_ids, + input_batch.num_tokens, + input_batch.num_scheduled_tokens, + input_batch.query_start_loc_np, + prefill_lens, + computed_prefill_lens + 1, # + 1 to consider the skew in eagle + ) + draft_tokens = self.speculator.propose( input_batch, attn_metadata,