Prhub

#31985 [Perf] Fold dspark dense draft embedding into the draft graph via forward_embed

原始 PR 作者 hnyls2002 合并时间 2026-07-22 08:00 文件变更 2 提交数 1 评论 5 代码增减 +14 / -4

执行摘要

将 DSpark 密集草稿嵌入查询移入草稿图

密集 DSpark 草稿模型在 attach_shared_modules 中丢弃了共享嵌入(del embed_tokens),导致运行器每次草稿图回放前都必须急切地准备 input_embeds。而 deepseek_v4_dspark 已经使用了 forward_embed 机制(保留 embed_tokens,暴露 forward_embed),使得 CUDA 图运行器和草稿提议者可以跳过 eager 阶段的复制。本 PR 让密集 DSpark 模型采用相同方式,将嵌入查找移入草稿图内,消除不必要的 eager 阶段拷贝。

这是一个小而精准的性能优化 PR,值得相关开发者阅读了解如何通过 forward_embed 机制消除 eager 阶段拷贝。建议精读 dspark.py 和 dflash.py 中的改动以理解设计模式。

讨论亮点

无 review 讨论。

实现拆解

  1. dspark.py – 保留 embed_tokens 并添加 forward_embed 方法:将 DSparkDraftMixin.attach_shared_modules 中的 del embed_tokens 改为 self.embed_tokens = embed_tokens,新增 forward_embed(self, input_ids) 方法,通过调用 self.embed_tokens(input_ids) 完成嵌入查找,并附带注释说明这样可以让运行器跳过 eager 阶段的 input_embeds 准备。
  2. dflash.py – 为 DFlashDraftModel.forward 添加 forward_embed 回退逻辑:当 input_embeds 为 None 时,先检查模型是否有 forward_embed 方法(通过 hasattr),若有则调用 self.forward_embed(input_ids) 获取嵌入;否则保持原有行为,抛出 ValueError。这样 base 类不会破坏已有模型,而实现了 forward_embed 的模型可以自动利用该优化。
  3. 无需修改运行器代码:CUDA 图运行器和草稿提议者已经支持 forward_embed 机制,只需模型实现该接口即可自动生效。
文件 模块 状态 重要度
python/sglang/srt/models/dspark.py 模型定义 modified 6.85
python/sglang/srt/models/dflash.py 模型定义 modified 6.58

关键符号

DSparkDraftMixin.attach_shared_modules DSparkDraftMixin.forward_embed DFlashDraftModel.forward

关键源码片段

python/sglang/srt/models/dspark.py data-contract

核心变更文件:保留 embed_tokens 并新增 forward_embed 方法,将嵌入查找移入草稿图。

# python/sglang/srt/models/dspark.py 中 DSparkDraftMixin 的修改def attach_shared_modules(
    self, *, embed_tokens: nn.Module, lm_head: nn.Module
) -> None:
    self.embed_tokens = embed_tokens # 之前是 del embed_tokens,现在保留
    self.lm_head = lm_headdef forward_embed(self, input_ids: torch.Tensor) -> torch.Tensor:
    # 使用共享的目标嵌入在草稿图内部进行嵌入查找
    # 当草稿模型暴露 forward_embed 时,运行器会跳过 eager 阶段的 input_embeds 准备
    return self.embed_tokens(input_ids)
python/sglang/srt/models/dflash.py data-contract

修改 forward 方法以支持 forward_embed 回退,使得基类可以自动利用优化。

# python/sglang/srt/models/dflash.py 中 DFlashDraftModel.forward 的修改@torch.no_grad()
def forward(
    self,
    input_ids: torch.Tensor,
    positions: torch.Tensor,
    forward_batch: ForwardBatch,
    input_embeds: Optional[torch.Tensor] = None,
    get_embedding: bool = False,
    pp_proxy_tensors=None,
) -> LogitsProcessorOutput:
    if input_embeds is None:
        if hasattr(self, "forward_embed"):
            # 如果模型实现了 forward_embed,则直接在图中获取嵌入
            input_embeds = self.forward_embed(input_ids)
        else:
            raise ValueError(
                "DFlashDraftModel requires `input_embeds` (use the target "
                "embedding)."
            )
    hidden_states = input_embeds
    # ... 后续层处理不变

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

本 PR 风险极低。核心改动仅涉及两处:

1) DSparkDraftMixin 保留 embed_tokens 而非丢弃,添加 forward_embed 方法;
2) DFlashDraftModel.forward 增加 forward_embed 回退分支。不会影响现有行为,因为只有模型存在 forward_embed 时才会走新路径。未看到直接对应的测试变更,但运行了 DSPark 和 DFlash 的 sanity 测试并通过。潜在的兼容性风险:若其他子类依赖 attach_shared_modules 中丢弃 embed_tokens 的行为,可能受影响,但当前仅 DSparkDraftMixin 使用该混入。

影响范围局限于 DSPark 密集草稿模型的推理流程。用户无需任何配置更改即可自动获得性能提升(减少一次 eager 阶段的数据拷贝,将嵌入计算纳入 CUDA 图)。对系统其他模块无影响。

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论