Prhub

#47093 [Spec Decode] DSpark speculators checkpoint support

原始 PR 作者 mgoin 合并时间 2026-07-02 08:32 文件变更 4 提交数 1 评论 22 代码增减 +140 / -16

执行摘要

增加 DSpark speculators 格式 checkpoint 支持

支持从 HuggingFace 等源加载 speculators-format 的 DSpark 检查点(如 Qwen3-8B-speculator.dspark-reasoning),这些检查点使用与 dense DSpark 不同的 1+N fill-in block 布局,并可能具有 reduced draft vocabulary 及其到 target vocab 的映射表。需要适配 config、模型加载和采样逻辑以兼容这些检查点。

值得精读。该 PR 展示了如何在现有投机解码框架中支持不同格式的检查点,涉及配置注册、模型接口设计、运行时适配三个层面。重点关注:compute_draft_logits/map_draft_to_target 的模型-运行时合约;reduced-vocab 时的散射采样策略;block 布局的动态选择。

讨论亮点
  • Markov head vocab 维度:@benchislett 追问为何 markov_w1 使用 vocab_size 而非 draft_vocab_size。@mgoin 解释 markov_w1 嵌入的是前一个 token(来自 target vocab),需要完整词汇表;markov_w2 投影到 draft vocab,所以使用 draft_vocab_size
  • compute_draft_logits 接口:@benchislett 质疑该方法的必要性,认为其它模型未类似实现。讨论后确认这是模型与 speculator 之间的合约,统一接口未来可以改进。
  • Block 布局检测:@benchislett 指出通过 block_size 推断布局不够健壮。改为在 speculators config 翻译中直接设置 dspark_bonus_anchor 标志,运行时直接读取。
  • Reduced-vocab scatter 数值安全性:@benchislett 询问 -inf 填充是否稳定。@mgoin 确认并在 load 时用 -inf 初始化缓冲一次,后续只 scatter 特定列,避免每步 fill 重复。
  • Scatter 实现优化:@benchislett 建议用 torch.where 等原生算子融合。由于散射索引恒定,维持现有 scatter + 预初始化方案。

实现拆解

  1. 注册 speculator 配置:在 vllm/transformers_utils/configs/speculators/algos.py 中新增 update_dspark 函数,通过 @register_speculator("dspark") 注册。该函数从 config_dict 提取 draft_vocab_sizemarkov_rankblock_size 等字段写入 pre_trained_config,并设置 dspark_bonus_anchor = True 以标记使用 1+N fill-in 布局。

  2. 模型层适配:在 vllm/model_executor/models/qwen3_dspark.pyvllm/models/deepseek_v4/nvidia/dspark.py 中,修改 DSparkMarkovHead 使其接受独立的 draft_vocab_size 参数:markov_w1 嵌入 target vocab(完整大小),markov_w2 投影到 draft vocab。新增 compute_draft_logits 方法(返回 draft-vocab 的 logits,不经过 d2t 映射)和 map_draft_to_target 方法(将 draft id 映射到 target id)。权重加载时识别 d2t 张量并重命名为 draft_id_to_target_id,跳过 t2d(训练时用)。

  3. 运行时适配:在 vllm/v1/worker/gpu/spec_decode/dspark/speculator.py 中,根据 dspark_bonus_anchor 标志选择 query 布局:False 时使用原来的 anchor-first(N slots),True 时使用 1+N fill-in(1 bonus + N slots)。新增 reduced-vocab 采样路径:在 load_draft_model 中预计算散射索引 _d2t_scatter_index(恒定的 draft→target 列偏移)和用 -inf 初始化的 _draft_scatter_buf;在 _sample_sequential 的循环中,若存在 _d2t_scatter_index,则将 draft logits 散射到 target vocab 中的对应列,然后采样得到 target id,直接用于 rejection sampling(省去一次映射)。

文件 模块 状态 重要度
vllm/model_executor/models/qwen3_dspark.py 模型定义 modified 8.02
vllm/transformers_utils/configs/speculators/algos.py 配置 modified 7.09
vllm/models/deepseek_v4/nvidia/dspark.py 模型定义 modified 6.89
vllm/v1/worker/gpu/spec_decode/dspark/speculator.py 投机解码器 modified 6.79

关键符号

DSparkMarkovHead.__init__ Qwen3DSparkModel.__init__ Qwen3DSparkForCausalLM.compute_draft_logits Qwen3DSparkForCausalLM.map_draft_to_target Qwen3DSparkForCausalLM.load_weights update_dspark DeepseekDSparkForCausalLM.compute_draft_logits DeepseekDSparkForCausalLM.map_draft_to_target DSparkSpeculator.__init__ DSparkSpeculator.load_draft_model DSparkSpeculator._sample_sequential

关键源码片段

vllm/model_executor/models/qwen3_dspark.py data-contract

DSpark 模型的 Qwen3 实现,核心变更包括 DSparkMarkovHead 支持独立 draft_vocab_size、新增 compute_draft_logits/map_draft_to_target 接口、权重加载添加 d2t 映射。

class Qwen3DSparkForCausalLM(DFlashQwen3ForCausalLM):
    # ... 其他代码不变
​
    def compute_draft_logits(self, hidden_states: torch.Tensor) -> torch.Tensor:
        # Draft-vocab logits without the d2t scatter:
        # the speculator adds the Markov bias in draft space,
        # then remaps via map_draft_to_target.
        return self.logits_processor(self.lm_head, hidden_states)
​
    def map_draft_to_target(self, draft_ids: torch.Tensor) -> torch.Tensor:
        # Map draft-vocab ids to target ids (identity for full-vocab drafts).
        if self.draft_id_to_target_id is None:
            return draft_ids
        # d2t mapping is stored as an offset table:
        # target_id = draft_id + draft_id_to_target_id[draft_id]
        return draft_ids + self.draft_id_to_target_id[draft_ids]
​
    def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]):
        model_weights = {}
        includes_embed_tokens = False
        includes_lm_head = False
        includes_draft_id_mapping = False
        for name, loaded_weight in weights:
            # t2d is training-only; the draft remaps via d2t at sampling time.
            if "t2d" in name:
                continue
            if "d2t" in name:
                # Rename d2t to draft_id_to_target_id
                name = name.replace("d2t", "draft_id_to_target_id")
                includes_draft_id_mapping = True
            elif "lm_head" not in name:
                name = "model." + name
            # ... 后续加载逻辑不变
vllm/transformers_utils/configs/speculators/algos.py core-logic

注册 `update_dspark` 配置转换器,将 speculators-format 的配置字段映射为模型可用的配置,设置 `dspark_bonus_anchor` 标志。

@register_speculator("dspark")
def update_dspark(config_dict: dict, pre_trained_config: dict) -> None:
    """
    将 speculators 格式的 DSpark 配置转换为 Transformers PreTrainedConfig 字典。
    设置必要的架构字段和参数,包括 draft 词汇表大小、markov head 参数、
    block 布局标志(dspark_bonus_anchor = True)等。
    """
    pre_trained_config["architectures"] = ["Qwen3DSparkModel"]
    # Speculators DSpark 使用 1+N fill-in 布局(anchor = bonus token)
    pre_trained_config["dspark_bonus_anchor"] = True
​
    aux_layer_ids = config_dict["aux_hidden_state_layer_ids"]
    pre_trained_config["eagle_aux_hidden_state_layer_ids"] = aux_layer_ids
    # DSpark 索引 target 层的方式是 aux_id - 1(与 dense config 一致)
    pre_trained_config["target_layer_ids"] = [i - 1 for i in aux_layer_ids]
​
    for key in (
        "draft_vocab_size",
        "target_hidden_size",
        "mask_token_id",
        "markov_rank",
        "markov_head_type",
        "block_size",
        "enable_confidence_head",
        "confidence_head_with_markov",
    ):
        if config_dict.get(key) is not None:
            pre_trained_config[key] = config_dict[key]
vllm/models/deepseek_v4/nvidia/dspark.py data-contract

DeepSeek-V4 的 DSpark 模型适配,类似 Qwen3 的改动:支持 draft_vocab_size,新增 compute_draft_logits/map_draft_to_target(恒等映射)。

class DeepseekDSparkForCausalLM(DSparkPretrainedModel):
    # ... 其他代码不变
​
    def compute_draft_logits(self, hidden_states: torch.Tensor) -> torch.Tensor:
        # Full-vocab draft: 直接使用 base logits,不需要 d2t 映射
        return self.compute_logits(hidden_states)
​
    def map_draft_to_target(self, draft_ids: torch.Tensor) -> torch.Tensor:
        return draft_ids # Full-vocab: draft ids 就是 target ids

评论区精华

Markov head vocab 维度设计 设计

@benchislett 提问为何 markov_w1 使用 vocab_size 而非 draft_vocab_size。@mgoin 解释 markov_w1 嵌入 conditioning token(前一个 token),该 token 来自 target vocab,所以需要完整 vocab_size;markov_w2 投影到 draft_vocab_size。

结论:decided: markov_w1 使用 vocab_size(target vocab),markov_w2 使用 draft_vocab_size,由构造参数区分。 · 已解决

compute_draft_logits 接口必要性 设计

@benchislett 提出该方法只有一处调用,且其他模型未实现类似 helper。@mgoin 回应 speculators checkpoint 路径需要此方法,而现有 dense 路径使用 compute_logits。讨论后认为这是模型与 speculator 之间的合约,统一接口未来可以改进。

结论:accepted: 作为合约方法保留,未来可考虑统一接口。 · 已解决

Block 布局动态检测的稳健性 设计

@benchislett 指出通过 block_size 推断布局可能脆弱,不同 num_speculative_steps 设置会改变机制。@mgoin 同意并改为在 speculators config 翻译中直接设置 dspark_bonus_anchor 标志,运行时直接读取配置,不再依赖 block_size 推断。

结论:resolved: 添加 dspark_bonus_anchor 配置标志,运行时直接读取。 · 已解决

Reduced-vocab scatter 数值安全性 正确性

@benchislett 询问 -inf 填充是否数值稳定。@mgoin 确认没有问题。进一步讨论 scatter 后其他列是否会包含未定义值,@mgoin 解释在 load 时用 -inf 初始化缓冲一次,后续只 scatter 到特定列,其他列始终为 -inf,所以 softmax 不会选中。

结论:resolved: 采用一次初始化缓冲 + 运行时仅 scatter 的策略。 · 已解决

Scatter 实现优化 性能

@benchislett 询问是否有更高效的原地操作(如 torch.where)。@mgoin 指出由于散射索引恒定,已在 load 时预计算,运行时只 scatter,无需每步 fill。

结论:resolved: 维持 scatter 方案,已通过预初始化优化。 · 已解决

风险与影响

  • 回归风险:修改了 DSparkMarkovHead 构造函数签名,原有 dense DSpark 检查点加载时若未提供 draft_vocab_size,需要 fallback 到 vocab_size。代码中已用 getattr(config, "draft_vocab_size", None) or config.vocab_size 处理,但需确认两个模型(Qwen3、DeepSeek-V4)的现有检查点兼容性。
  • 性能风险:reduced-vocab 采样路径引入了 scatter 操作和额外缓冲,但已在 load 时预分配,且散射索引恒定,运行时开销可控。
  • 数值风险:scatter 时使用 -inf 填充非目标列,需要验证 softmax 数值稳定性(预期无问题)。
  • 测试缺失:本次变更未包含直接对应的测试文件,reduced-vocab 分支的覆盖完全依赖手工验证。
  • 用户影响:用户现在可以使用 speculators-format 的 DSpark 检查点(如 Qwen3-8B-speculator),扩展了投机解码灵活性和模型选择。
  • 系统影响:投机解码模块新增配置分支和采样逻辑,但整体架构保持向后兼容,现有 dense DSpark 检查点不受影响。
  • 团队影响:新增代码约 140 行,与 DFlash 框架同构,维护成本有限。
核心路径变更 缺少测试覆盖 数值稳定性 兼容性回退

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论