Skip to content

[bugfix] Mark draft tokens to rebuilt their embeddings. - #57356

Merged
DarkLight1337 merged 1 commit into
vllm-project:mainfrom
HieDean:main
Sep 17, 2026
Merged

DarkLight1337 merged 1 commit into
vllm-project:mainfrom
HieDean:main

Conversation

@HieDean

@HieDean HieDean commented Sep 17, 2026

Copy link
Copy Markdown
Contributor

Purpose

Fix an inconsistency between input_ids and is_token_ids when enable_prompt_embeds is used together with speculative decoding.

Speculative draft token IDs are scattered directly into input_ids.gpu, but the corresponding positions in is_token_ids were not marked as token IDs. As a result, the prompt-embedding path treated these positions as external embeddings and reused stale embeddings during target verification, causing incorrect target predictions and very low acceptance lengths.

This change marks spec_flattened_indices as token-ID positions before synchronizing is_token_ids to the GPU.

Test Plan

import torch

from transformers import AutoModelForCausalLM
from datasets import load_dataset
from vllm import LLM, SamplingParams

TARGET_MODEL = "Qwen3-8B"
DRAFT_MODEL = "Qwen3-8B-DFlash-b16"
DATASET = "gsm8k"

# Create a sampling params object.
sampling_params = SamplingParams(temperature=0.0, max_tokens=128)

def main():
    hf_target_model = AutoModelForCausalLM.from_pretrained(TARGET_MODEL)

    # Create an LLM.
    llm = LLM(
        model=TARGET_MODEL,
        enable_prompt_embeds=True,
        speculative_config={"model": DRAFT_MODEL, "num_speculative_tokens": 16},
        disable_log_stats=False,
    )
    tokenizer = llm.get_tokenizer()

    dataset = load_dataset(DATASET, "main")
    dataset = list(dataset["test"])[:10]
    for data in dataset:
        messages = [
            [{"role": "user", "content": f"{data["question"]}"}]
        ]
        input_ids = tokenizer.apply_chat_template(
            messages,
            tokenize=True,
            add_generation_prompt=False,
            return_tensors="pt"
        )

        embeddings = hf_target_model.get_input_embeddings()(input_ids["input_ids"])
        prompt_len = embeddings.shape[1]

        inputs = {
            "prompt_embeds": embeddings,
            "prompt_token_ids": [tokenizer.vocab_size] * prompt_len,
            "prompt_is_token_ids": [False] * prompt_len,
        }
        outputs = llm.generate(inputs, sampling_params, use_tqdm=False)

        # Print the outputs.
        print("\nGenerated Outputs:\n" + "-" * 60)
        for output in outputs:
            prompt = output.prompt
            generated_text = output.outputs[0].text
            print(f"Prompt:    {prompt!r}")
            print(f"Output:    {generated_text!r}")
            print("-" * 60)

        num_drafts = 0
        num_accepted_tokens = 0
        for metric in llm.get_metrics():
            if metric.name == "vllm:spec_decode_num_drafts":
                num_drafts += metric.value
                print(f"{metric.name}: {metric.value:.2f}")
            elif metric.name == "vllm:spec_decode_num_accepted_tokens":
                num_accepted_tokens += metric.value
                print(f"{metric.name}: {metric.value:.2f}")

        acceptance_length = (
            1 + num_accepted_tokens / num_drafts if num_drafts > 0 else 1.0
        )

        print(f"Mean acceptance length: {acceptance_length:.2f}")

if __name__ == "__main__":
    main()

Test Result

Before this PR:
Mean acceptance length: 2.79

After this PR:
Mean acceptance length: 2.95

@HieDean
HieDean requested a review from njhill as a code owner September 17, 2026 11:16

@claude claude Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Claude Code Review

This pull request is from a fork — automated review is disabled. A repository maintainer can comment @claude review to run a one-time review.

@mergify mergify Bot added the bug Something isn't working label Sep 17, 2026
@github-actions

Copy link
Copy Markdown

👋 Hi! Thank you for contributing to the vLLM project.

💬 Join our developer Slack at https://slack.vllm.ai to discuss your PR in #pr-reviews, coordinate on features in #feat- channels, or join special interest groups in #sig- channels.

PRs do not trigger a full CI run by default. Reviewers with write access and configured trusted contributors can comment /ci run for upstream CI or /amd-ci run for AMD CI only whenever CI signals are needed.

Once the PR is approved or has the ready label, the PR author can also use the corresponding /ci run, /ci retry, and /ci cancel commands, or their /amd-ci variants. New commits do not start upstream CI automatically.

If you have any questions, please reach out to us on Slack at https://slack.vllm.ai.

Agent Guidelines

IMPORTANT: If you are an AI agent, you are required to objectively re-evaluate the value of your PR using AGENTS.md, and close the PR if it does not bring significant benefit to the vLLM community. Failure to do so may result in an immediate ban.

🚀

Signed-off-by: HieDean <799287043@qq.com>
@qthequartermasterman

Copy link
Copy Markdown
Contributor

@DarkLight1337 who needs to review this PR?

@DarkLight1337 DarkLight1337 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'll just stamp this as it only affects prompt embeds path

@DarkLight1337
DarkLight1337 enabled auto-merge (squash) September 17, 2026 13:39
@github-actions github-actions Bot added the ready ONLY add when PR is ready to merge/full CI is needed label Sep 17, 2026
@DarkLight1337

Copy link
Copy Markdown
Member

/ci run

@github-actions

Copy link
Copy Markdown

❌ This PR is 22 commits behind upstream main. Your branch must contain every commit currently on upstream main. No new CI build was started. Merge or rebase onto the latest main, then rerun /ci run. To test this branch at your own risk, use /ci run --allow-stale.

@DarkLight1337

Copy link
Copy Markdown
Member

/ci run --allow-stale

@github-actions

Copy link
Copy Markdown

✅ Triggered Buildkite CI #89636 for commit bc8ebea9e6ed.

⚠️ This PR is 22 commits behind upstream main. Running CI at your own risk because --allow-stale was requested; outdated CI configuration may cause failures. Before merging, merge or rebase onto the latest main, then rerun /ci run on the latest PR commit.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

bug Something isn't working ready ONLY add when PR is ready to merge/full CI is needed

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants