Prhub

#44814 [Bugfix] Fix layerwise reload dropping params after a composed weight loader

原始 PR 作者 hallerite 合并时间 2026-06-10 21:42 文件变更 2 提交数 2 评论 1 代码增减 +96 / -1

执行摘要

修复组合权重加载器导致参数丢失的 bug

修复在线权重重载(online weight reload)场景中,组合权重加载器(composed_weight_loader)导致 trailing 参数被静默丢弃的 bug。该问题源于 CopyCounter 在 composed_weight_loader 执行时对同一参数重复计数,使得层的已加载元素计数提前达到总量,层被过早终结,后续参数未被加载。此 bug 影响使用 Mamba2 混合器的模型(如 NemotronH),导致 mixer.D(SSD 跳跃连接)未被加载,推理输出 NaN logits。

此 PR 属于关键 bug 修复,逻辑精巧,值得精读。建议关注 get_numel_loaded 函数的设计模式(通过 TorchDispatch 拦截 copy_ 来计数),以及如何通过上限截断解决 composed loader 的双重复制问题。新增的测试用例是良好的回归测试范例。

讨论亮点

kylesayrs 在 review 中建议使用 args.arguments.get("param", None) 安全地获取 param 参数,以避免某些非常规加载器(unconventional loaders)中 param 缺失或类型不匹配导致的错误。该建议已被采纳,并在第二次 commit 中实现。

实现拆解

  1. 修改 vllm/model_executor/model_loader/reload/meta.py 中的 get_numel_loaded 函数:在通过 CopyCounter 获取已复制的元素数后,增加对目标参数 param 的检查。若 paramtorch.Tensor,则将计数值上限截断为 param.numel(),防止因 composed_weight_loader 的二次复制导致翻倍计数。
  2. 新增测试文件 tests/model_executor/model_loader/test_reload.py 中的测试
    • test_get_numel_loaded_caps_at_param_size:验证 composed_weight_loader 返回的计数值等于 param.numel(),而非 2×。
    • _ComposedLoaderLayer:模拟 Mamba2 混合器的参数结构(A、D、dt_bias 等大小),A 使用组合加载器,D 和 dt_bias 使用默认加载器。
    • test_layerwise_reload_composed_loader_does_not_drop_params:端到端回归测试,确保在按层重载后所有参数均保持其加载值,而非未初始化。
  3. 导入调整:在测试文件中新增对 composed_weight_loaderdefault_weight_loader 的导入。
文件 模块 状态 重要度
vllm/model_executor/model_loader/reload/meta.py 权重加载器 modified 6.85
tests/model_executor/model_loader/test_reload.py 测试 modified 7.06

关键符号

get_numel_loaded test_get_numel_loaded_caps_at_param_size test_layerwise_reload_composed_loader_does_not_drop_params

关键源码片段

vllm/model_executor/model_loader/reload/meta.py core-logic

核心修复文件。修改了 get_numel_loaded 函数,在 CopyCounter 计数后增加上限截断逻辑:若 param 为 Tensor,则计数值取 counter.copied_numel 与 param.numel() 的较小值。该一行变更即可修复 composed_weight_loader 导致的参数丢失问题。

def get_numel_loaded(
    weight_loader: Callable, args: inspect.BoundArguments
) -> tuple[int, object]:
    """
    Determine how many elements would be loaded by a weight loader call.    Args:
        weight_loader: used to load weights
        args: bound arguments to weight loader    Returns:
        number of elements loaded by the weight loader, the return value of the
        weight loader
    """
    with CopyCounter() as counter:
        return_value = weight_loader(*args.args, **args.kwargs)
​
    # 一个 weight loader 只填充一个目标参数,所以加载的元素数最多为该参数的大小。
    # 某些加载器(如 composed_weight_loader)会复制多次(初始加载 + 原地后处理变换),
    # 导致 CopyCounter 报告两倍的大小。超量计数会使层的已加载元素总数膨胀,
    # 可能在其他参数还未加载时就终结当前层,静默丢弃后续参数(如 Mamba mixer.D)。
    # 因此将计数值上限设为目标参数的大小,以保持逐层计数的正确性。
    numel = counter.copied_numel
    param = args.arguments.get("param", None)
    if isinstance(param, torch.Tensor):
        numel = min(numel, param.numel())
    return numel, return_value

评论区精华

param 参数的安全获取方式 正确性

kylesayrs 在 review 中指出,直接通过 args.arguments['param'] 访问可能因参数名不存在而抛出 KeyError,建议使用 .get() 方式以避免对非常规加载器的兼容性问题。

结论:提交者采纳建议,将代码改为 args.arguments.get("param", None) 并使用 isinstance 检查类型。 · 已解决

风险与影响

无显著风险。变更仅在一行核心逻辑(get_numel_loaded)中增加了一个上限截断操作,且该操作只有在 paramtorch.Tensor 时才会生效,不影响其他场景。新增的测试覆盖了核心回归路径,且已有测试全部通过。

影响范围限于使用 composed_weight_loader 加载参数的模型进行在线按层权重重载的场景,尤其是 Mamba2 混合器(如 NemotronH)。修复后这些模型在权重热更新后能正确加载所有参数,避免 NaN 输出。对不使用组合加载器的模型无影响。

核心路径变更 新增测试覆盖

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论