Prhub

#44353 Weight sync refactor + move sparse nccl engine

原始 PR 作者 hao-aaron 合并时间 2026-07-01 16:25 文件变更 32 提交数 14 评论 17 代码增减 +953 / -1046

执行摘要

权重同步重构,抽取稀疏 NCCL 独立引擎

参考 PR body 描述:当前权重同步路径要求非权重传输引擎组件过度了解内部细节,start_weight_update(is_checkpoint_format=...) 仅用于切换是否执行逐层重载,实际上所有密集流程都使用 checkpoint 格式,唯一的 kernel 格式使用者(稀疏更新)现已独立为单独的引擎。因此该标志成为泄露引擎内部决策到公共 API 的死负载。

建议精读该 PR,特别是引擎抽象设计、NCCL 初始化共享化以及 Worker 瘦身的模式。对于需要深度定制权重传输的开发者,理解 WeightTransferEngine 的新生命周期至关重要。该 PR 值得关注的设计决策包括:将 start_weight_update/finish_weight_update 声明为抽象方法以保证一致性,以及通过 nccl_common 避免代码复用时的继承耦合。

讨论亮点

1. IPC 引擎设备索引问题(bedeks → hao-aaron → SumanthRH)

bedeks 指出在 IPCWeightTransferEngine.start_weight_update 中调用 torch.accelerator.current_device_index() 可能不匹配 worker 的设备,因为 worker.update_weights 不再使用 with torch.device(self.device) 包装。作者确认该问题,SumanthRH 解释 init_device 已预先设置设备索引,因此当前调用是安全的。

2. 文档与 docstring 改进(SumanthRH)

SumanthRH 建议在 base.pyipc_engine.py 中提供更友好的 docstring,并添加指向逐层重载文档的链接。作者采纳建议。

3. 移除废弃代码(aoshen02)

aoshen02 指出 gpu_model_runner 中的 apply_sparse_weight_patches 方法已是死代码,可以移除。作者在后续提交中进行了清理。

实现拆解

  1. 移除格式区分:删除 start_weight_update(is_checkpoint_format=...) 中的 is_checkpoint_format 参数,将该决策下沉到引擎内部。涉及文件 gpu_worker.pyasync_llm.pybase.py 等。
  2. 显式生命周期:将 start_weight_update/finish_weight_update 声明为 WeightTransferEngine 的抽象方法,每个后端必须实现自己的准备和收尾逻辑。密集引擎在 start/finish 中执行逐层重载,而稀疏引擎实现为 no-op。
  3. 抽取稀疏 NCCL 引擎:新建 SparseNCCLWeightTransferEngine(不继承自密集引擎),共享的进程组初始化移至 nccl_common。同时移除了模型运行器中的 apply_sparse_weight_patches 方法。
  4. 简化 WorkerGPUWorker 不再处理格式判断、加载器构造等,仅转发 start/update/finish 到引擎,维护一个 _weight_update_active 守卫。

此外,更新了文档 docs/training/weight_transfer/base.md 以反映新架构,并调整了所有相关测试文件以适应接口变更。

文件 模块 状态 重要度
vllm/distributed/weight_transfer/sparse_nccl_engine.py 稀疏引擎 added 9.26
vllm/distributed/weight_transfer/nccl_common.py 公共组件 added 8.83
vllm/distributed/weight_transfer/nccl_engine.py 密集引擎 modified 8.86
vllm/distributed/weight_transfer/base.py 基类 modified 8.44

关键符号

start_weight_update finish_weight_update receive_weights init_transfer_engine worker_init_process_group trainer_init create_engine update_weights

关键源码片段

vllm/distributed/weight_transfer/sparse_nccl_engine.py core-logic

新增稀疏 NCCL 权重传输引擎,从密集引擎解耦,展示独立生命周期和原地补丁模式。

# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Sparse NCCL weight transfer engine.A standalone engine (not a subclass of `NCCLWeightTransferEngine`) for applying
sparse, flat-index weight patches in place. It shares only NCCL process-group
initialization with the dense engine (via `nccl_common`); the update path
applies index/value patches directly to existing model parameters and never runs
layerwise reload.MVP limitations:
* TP=1 and PP=1 only
* uses runtime/kernel-format parameter names
* not composable with checkpoint-format or packed updates
"""from dataclasses import dataclass
from typing import TYPE_CHECKINGimport torchif TYPE_CHECKING:
    from vllm.config import VllmConfigfrom vllm.distributed.weight_transfer.base import (
    WeightTransferEngine,
    WeightTransferUpdateInfo,
)
from vllm.distributed.weight_transfer.nccl_common import (
    NCCLWeightTransferInitInfo,
    worker_init_process_group,
)
​
​
@dataclass
class SparseWeightPatch:
    """A sparse in-place patch for one existing parameter."""
    name: str
    indices: torch.Tensor
    values: torch.Tensor
​
​
@dataclass
class SparseNCCLWeightTransferUpdateInfo(WeightTransferUpdateInfo):
    """Update info for the sparse NCCL weight transfer backend."""
    names: list[str]
    dtype_names: list[str]
    shapes: list[list[int]]
    num_updates_list: list[int]
​
    def __post_init__(self) -> None:
        num_params = len(self.names)
        if len(self.dtype_names) != num_params:
            raise ValueError(
                f"`dtype_names` should be of the same size as `names`: "
                f"got {len(self.dtype_names)} and {len(self.names)}"
            )
        if len(self.shapes) != num_params:
            raise ValueError(
                f"`shapes` should be of the same size as `names`: "
                f"got {len(self.shapes)} and {len(self.names)}"
            )
        if len(self.num_updates_list) == 0:
            raise ValueError("`num_updates_list` cannot be empty for sparse updates")
        if len(self.num_updates_list) != num_params:
            raise ValueError(
                f"`num_updates_list` should be of the same size as `names`: "
                f"got {len(self.num_updates_list)} and {len(self.names)}"
            )
        if any(num_updates < 0 for num_updates in self.num_updates_list):
            raise ValueError("Sparse `num_updates_list` entries must be non-negative")
​
​
class SparseNCCLWeightTransferEngine(
    WeightTransferEngine[NCCLWeightTransferInitInfo, SparseNCCLWeightTransferUpdateInfo]
):
    """
    Sparse weight transfer engine using NCCL.
    Receives flat-index (indices, values) patches broadcast from the trainer
    and applies them in place to existing model parameters. Weights are
    applied directly without layerwise reload, so `start_weight_update` and
    `finish_weight_update` are no-ops.
    """
​
    def start_weight_update(self) -> None:
        if self.parallel_config.world_size != 1:
            raise NotImplementedError(
                "Sparse weight updates currently require TP=1 and PP=1"
            )
​
    def finish_weight_update(self) -> None:
        pass
​
    def receive_weights(self, update_info: SparseNCCLWeightTransferUpdateInfo) -> None:
        ...
vllm/distributed/weight_transfer/nccl_common.py core-logic

新增共享 NCCL 初始化模块,密集和稀疏引擎共用的进程组创建工作迁移至此,避免代码重复。

# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Shared NCCL initialization helpers for weight transfer engines.The dense (`NCCLWeightTransferEngine`) and sparse
(`SparseNCCLWeightTransferEngine`) backends are independent engines that share
*only* their process-group initialization. That common logic lives here so the
sparse engine does not have to subclass the dense one.
"""from dataclasses import dataclass
from typing import TYPE_CHECKINGimport torchif TYPE_CHECKING:
    from vllm.config.parallel import ParallelConfig
    from vllm.distributed.device_communicators.pynccl import PyNcclCommunicatorfrom vllm.distributed.weight_transfer.base import WeightTransferInitInfo
​
​
@dataclass
class NCCLWeightTransferInitInfo(WeightTransferInitInfo):
    """Initialization info for NCCL-based weight transfer backends."""
    master_address: str
    master_port: int
    rank_offset: int
    world_size: int
​
​
def stateless_init_process_group(
    master_address: str, master_port: int, rank: int, world_size: int, device,
) -> "PyNcclCommunicator":
    from vllm.distributed.device_communicators.pynccl import PyNcclCommunicator
    from vllm.distributed.utils import StatelessProcessGroup
    pg = StatelessProcessGroup.create(
        host=master_address, port=master_port, rank=rank, world_size=world_size
    )
    return PyNcclCommunicator(pg, device=device)
​
​
def worker_init_process_group(
    init_info: NCCLWeightTransferInitInfo,
    parallel_config: "ParallelConfig",
) -> "PyNcclCommunicator":
    dp_rank = parallel_config.data_parallel_index
    world_size_per_dp = parallel_config.world_size
    rank_within_dp = parallel_config.rank
    worker_rank = dp_rank * world_size_per_dp + rank_within_dp
    rank = worker_rank + init_info.rank_offset
    device = torch.accelerator.current_device_index()
    return stateless_init_process_group(
        init_info.master_address, init_info.master_port,
        rank, init_info.world_size, device=device,
    )
​
​
def trainer_init(
    init_info: NCCLWeightTransferInitInfo | dict,
) -> "PyNcclCommunicator":
    if isinstance(init_info, dict):
        master_address = init_info["master_address"]
        master_port = init_info["master_port"]
        world_size = init_info["world_size"]
    else:
        master_address = init_info.master_address
        master_port = init_info.master_port
        world_size = init_info.world_size
    device = torch.accelerator.current_device_index()
    return stateless_init_process_group(
        master_address, master_port, 0, world_size, device,
    )

评论区精华

IPC 引擎设备索引问题 正确性

bedeks 指出在 IPC 引擎的 start_weight_update 中调用 torch.accelerator.current_device_index() 可能不匹配 worker 的设备,因为 worker.update_weights 不再使用 with torch.device(self.device) 包装。

结论:SumanthRH 指出 init_device 已预先设置设备索引,因此当前调用是安全的。作者确认感谢。 · 已解决

文档和 docstring 改进 documentation

SumanthRH 建议在 base.py 和 ipc_engine.py 中提供更友好的 docstring,并添加指向逐层重载文档的链接。

结论:作者采纳建议,改进 docstring。 · 已解决

移除废弃代码 apply_sparse_weight_patches other

aoshen02 指出 gpu_model_runner 中的 apply_sparse_weight_patches 方法已是死代码,可以移除。

结论:作者在提交 8ce70be 中移除了该方法。 · 已解决

风险与影响

  1. IPC 引擎设备索引假设:尽管当前通过 init_device 保证设备索引一致,但若未来调用流程变化,可能引入多 GPU 下设备不匹配的风险。
  2. 稀疏引擎局限性:新引擎仅支持 TP=1 和 PP=1,且使用 runtime/kernel-format 参数名,不兼容 checkpoint-format 或打包更新。用户需谨慎评估适用场景。
  3. API 兼容性变更start_weight_update 参数移除,WeightTransferEngine 构造函数增加 vllm_configdevice 参数,所有自定义引擎需要适配新接口。
  4. 测试覆盖:虽然测试文件大量更新,但稀疏引擎与密集引擎组合的复杂场景可能测试不足。
  1. 用户/开发者:使用权重同步 API(如 RLHF 示例)需移除 is_checkpoint_format 参数;使用稀疏更新的用户需配置 backend="sparse_nccl"
  2. 系统架构:新引擎降低了模块耦合,密集引擎不再包含稀疏分支,代码清晰度提升;共享初始化模块便于后续扩展。
  3. 团队维护:未来增加新后端(CUDA IPC、RDMA)时,可参考 SparseNCCLWeightTransferEngine 模式,实现独立生命周期。
IPC 多 GPU 设备索引风险 稀疏引擎 TP=1 PP=1 限制 API 构造函数参数变更

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论