Prhub

#45352 [Bugfix] Forward callable hf_overrides to the draft model config

原始 PR 作者 HumphreySun98 合并时间 2026-07-07 12:12 文件变更 3 提交数 6 评论 12 代码增减 +148 / -3

执行摘要

修复 draft 模型未接收目标可调用 hf_overrides

PR body 指明根本原因是:SpeculativeConfig.__post_init__ 内部构造 draft ModelConfig 时始终使用 hf_overrides=SpeculativeConfig.hf_config_override,静默丢弃了目标模型配置上的任何 hf_overrides。导致 dummy_hf_overrides shrink(减少层数和专家数)只应用于目标模型,而 Eagle draft 仍以完整 675B 维度实例化,CI GPU 即使 load_format="dummy" 也会 OOM。因此架构必须用 is_available_online=False 排除。需要修复 draft 配置的 override 转发路径。

建议仔细阅读该 PR 的设计决策,特别是如何区分字典和可调用覆盖,以及通过 functools.partial 保持可 pickle 性的技巧。对于维护其他 spec 方法的开发者,推荐遵循类似模式。可将其视为配置覆盖组合的参考实现。

讨论亮点

pickling 问题:作者发现初次实现返回嵌套局部闭包,导致 DFlashDraftModel 测试失败(Can't get local object)。通过改用 functools.partial 和模块级静态方法解决,并在测试中验证。
OOM 根因分析:作者定位到 draft 模型未收到 dummy_hf_overrides shrink,导致全尺寸加载 OOM,修复后验证了 EagleMistralLarge3 初始化测试通过,重新启用。
CI 失败排除:后续 CI 红项经分析均为不相关 flake 或外部 API 断裂,最终全绿。

实现拆解

  1. 新增静态方法:在 SpeculativeConfig 中添加 _apply_composed_hf_overridecompose_draft_hf_overridescompose_draft_hf_overrides 接收目标 hf_overrides,如果非可调用(dict/None)则直接返回架构映射 override;如果可调用则使用 functools.partial 将目标绑定到模块级静态方法,确保结果可 pickle。
  2. 修改 __post_init__:在非 medusa 分支中,将 draft_hf_overrides 的赋值从直接使用 SpeculativeConfig.hf_config_override 改为调用 SpeculativeConfig.compose_draft_hf_overrides(self.target_model_config.hf_overrides),从而将目标的可调用 override 组合后传递给 draft ModelConfig
  3. 解决 pickling 问题:最初实现返回嵌套局部闭包,导致 DFlashDraftModel 在 multiprocessing 反序列化时失败。改为 functools.partial 后,composed 函数成为模块级可引用对象,通过 pickle。
  4. 新增单元测试:创建 tests/config/test_speculative_draft_hf_overrides.py,包含 4 个 @pytest.mark.cpu_test 测试,覆盖 dict/None/callable override 的行为以及 composed 函数的可 pickle 性验证。
  5. 重新启用模型测试:从 tests/models/registry.py 中移除 EagleMistralLarge3ForCausalLMis_available_online=False 和关联的 TODO 注释,使该架构的初始化测试重新在 CI 中运行。
文件 模块 状态 重要度
vllm/config/speculative.py 推测解码配置 modified 7.3
tests/config/test_speculative_draft_hf_overrides.py 测试套件 added 7.18
tests/models/registry.py 模型注册表 modified 3.63

关键符号

SpeculativeConfig._apply_composed_hf_override SpeculativeConfig.compose_draft_hf_overrides SpeculativeConfig.__post_init__ test_callable_overrides_reach_the_draft_config test_arch_mapping_applies_before_callable_override test_composed_override_is_picklable

关键源码片段

vllm/config/speculative.py core-logic

核心逻辑变更:新增 compose_draft_hf_overrides 和 _apply_composed_hf_override 静态方法,修改 __post_init__ 以将目标可调用 hf_overrides 传递给 draft 配置

@staticmethod
def _apply_composed_hf_override(
    target_hf_overrides: Callable[[PretrainedConfig], PretrainedConfig],
    hf_config: PretrainedConfig,
) -> PretrainedConfig:
    # 首先应用架构映射 override(如 MTP 架构重写)
    hf_config = SpeculativeConfig.hf_config_override(hf_config)
    # 然后应用目标可调用 override(如测试 shrink)
    return target_hf_overrides(hf_config)@staticmethod
def compose_draft_hf_overrides(
    target_hf_overrides: HfOverrides | None,
) -> Callable[[PretrainedConfig], PretrainedConfig]:
    # 如果目标 override 不是可调用的(dict 或 None),则直接返回架构映射 override
    if not callable(target_hf_overrides):
        return SpeculativeConfig.hf_config_override
    # 使用 functools.partial 绑定目标 override,确保结果可 pickle(避免嵌套局部闭包)
    return functools.partial(
        SpeculativeConfig._apply_composed_hf_override, target_hf_overrides
    )# 在 __post_init__ 中(非 medusa 分支)的变更:
# 之前:
# draft_hf_overrides = SpeculativeConfig.hf_config_override
# 现在:
draft_hf_overrides = SpeculativeConfig.compose_draft_hf_overrides(
    self.target_model_config.hf_overrides
)
self.draft_model_config = ModelConfig(
    model=self.model,
    runner="draft",
    ...
    hf_overrides=draft_hf_overrides, # 现在包含组合后的覆盖
)
tests/config/test_speculative_draft_hf_overrides.py test-coverage

新增完整的单元测试,覆盖四种场景:dict/None/callable 和可 pickle 性验证,确保 compose_draft_hf_overrides 按预期工作

@pytest.mark.cpu_test
def test_callable_overrides_reach_the_draft_config():
    """验证可调用 override 会被组合并应用于 draft 配置"""
​
    def shrink(hf_config: PretrainedConfig) -> PretrainedConfig:
        hf_config.num_hidden_layers = 1
        return hf_config
​
    composed = SpeculativeConfig.compose_draft_hf_overrides(shrink)
    # 确认返回的不是原始架构映射 override
    assert composed is not SpeculativeConfig.hf_config_override
​
    out = composed(_make_hf_config(num_hidden_layers=64))
    # 确认 shrink 变换已生效
    assert out.num_hidden_layers == 1
​
​
@pytest.mark.cpu_test
def test_composed_override_is_picklable():
    """验证组合后的 override 可 pickle,避免 DFlashDraftModel 的 pickling 错误"""
    composed = SpeculativeConfig.compose_draft_hf_overrides(_module_level_shrink)
    # 必须是 functools.partial,而不是闭包
    assert isinstance(composed, functools.partial)
    assert composed.func is SpeculativeConfig._apply_composed_hf_override
​
    out = composed(_make_hf_config())
    assert out.num_hidden_layers == 1

评论区精华

compose_draft_hf_overrides 的 pickling 问题 正确性

作者发现 compose_draft_hf_overrides 最初返回嵌套局部闭包,导致 DFlashDraftModel 在 multiprocessing 反序列化时失败(Can't get local object)。作者改为使用 functools.partial 将目标 override 绑定到模块级静态方法,确保 composed 函数可 pickle。

结论:通过 functools.partial 和静态方法解决 pickling 问题,并在测试中验证。 · 已解决

EagleMistralLarge3 OOM 的根本原因 正确性

作者定位到 draft 模型未收到 dummy_hf_overrides shrink,导致全尺寸加载 OOM。修复后重新启用测试。

结论:修复 draft hf_overrides 转发,OOM 解决。 · 已解决

风险与影响

  1. 配置逻辑变更:修改了 draft_hf_overrides 的生成方式,可能影响其他 spec 方法(如 medusa 已有独立分支,不受影响)。MTP 方法因 draft_model_config = target_model_config 天然不受影响。
  2. 多进程 pickling 依赖:composed 函数必须可 pickle,设计上通过 functools.partial 和模块级静态方法保证。如果未来新增更复杂的回调,需注意约束。
  3. 测试覆盖:新增单元测试覆盖四种场景,但未测试实际模型加载路径的组合效果。CI 中 EagleMistralLarge3 初始化测试可提供端到端验证。

用户影响:普通用户无感知,仅影响使用 speculative decoding 并依赖可调用 hf_overrides 的开发者和测试场景。系统影响:修复减少 CI OOM,扩大测试覆盖率,提高稳定性。团队影响:明确了 draft 配置覆盖的设计原则,为后续类似问题提供模式。

配置逻辑变更 多进程 pickling 依赖 新增测试覆盖 回归风险

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论