首页
/ vLLM 权重传输系统基类设计:训练器与推理引擎之间的可插拔同步架构与自定义引擎开发指南

vLLM 权重传输系统基类设计:训练器与推理引擎之间的可插拔同步架构与自定义引擎开发指南

2026-09-07 09:00:39作者:韦蓉瑛

导读

在线强化学习(RLHF、GRPO 等)中,训练进程需要不断把更新后的策略模型权重同步到推理引擎,用于 rollout 采样。本文深入解析 vLLM 的 weight transfer(权重传输)系统在 vllm/distributed/weight_transfer/base.py 中定义的四大可独立替换的抽象——训练侧的 WeightSourceVLLMWeightSyncClientTrainerWeightTransferEngine 与推理侧的 WeightTransferEngine,并配套介绍两套互不耦合的工厂注册机制。读完本文,你将能够:理解权重从 FSDP/TP/PP/EP 分片模型流向推理 worker 的完整调用链;在自研 RL 框架中通过四个方法即接入控制面;以及按官方约定从零实现并注册一个自定义 weight transfer 后端。

一、整体架构:两个进程、一套对称协议

vLLM 权重传输系统遵循"每侧一个引擎、两端对称"的设计。下表(对应 docs/training/weight_transfer/README.mdbase.py 的类定义)说明了角色分工:

抽象 所在侧 回答的问题
WeightSource 训练器 发什么——如何从你的训练框架中抽出权重
VLLMWeightSyncClient 训练器 怎么联系推理引擎——适配你的 RL 框架自带的 vLLM wrapper
TrainerWeightTransferEngine 训练器 怎么传字节(传输状态与收发调度)
WeightTransferEngine 推理侧 怎么收并加载进模型

两个引擎分别注册在 WeightTransferTrainerFactoryWeightTransferEngineFactory 两个工厂里。它们仅靠 backend 字符串命名约定共享,但注册表相互独立——训练进程永远不会实例化 worker 引擎,反之亦然(源码注释明确说明这样做是为了不耦合 import graph)。

四阶段协议

每一次同步在底层都是同一个四阶段协议,由训练器引擎替你驱动:

  1. 初始化init_weight_transfer_engine):在训练循环开始前,由 trainer_init 调用一次,建立训练器与推理 worker 之间的通信信道。
  2. 开始start_weight_update):让推理引擎为一个权重更新做准备。
  3. 权重更新update_weights):传输更新后的权重,可调用一次或多次(如分块传输)。
  4. 结束finish_weight_update):在所有权重传完后调用一次,做收尾(如 checkpoint 格式权重的后处理)。

内置后端与配置入口

当前工厂内置注册了 ncclipcsparse_ncclsharded_rdt 四种后端,注册集中在 factory.py

backend 传输方式 适用场景
nccl NCCL broadcast 训练与推理分占不同 GPU
ipc CUDA IPC handle 训练与推理同机/同卡共置
sparse_nccl NCCL broadcast checkpoint 坐标稀疏权重补丁(delta)
sharded_rdt NIXL / Ray Direct Transport(拉取式) 超大模型、每个 worker 只需自己分片(MoE + EP)

推理侧只认识一个配置项 WeightTransferConfigbackend 字符串,且默认值就是 "nccl"。其余传输参数全部由训练器决定并在 init 握手时"空运"过去,推理侧从不自行猜测。注意,配置里刻意不携带 packed、缓冲大小等 wire 参数,就是为了让"两端对 wire 参数理解不一致"这种情况在类型层面不可表达。

从代码结构看,worker 启动路径(vllm/v1/worker/gpu_worker.py)会在 self.vllm_config.weight_transfer_config 非空时调用 WeightTransferEngineFactory.create_engine(...),随后 HTTP/Ray 层暴露的四个方法(gpu_worker.py)逐一转发到引擎:init_weight_transfer_engineparse_init_info + init_transfer_enginestart_weight_updatestart_weight_updateupdate_weightsupdate_weightsfinish_weight_updatefinish_weight_update,并带 session 状态检查(如未 start 就 update 会报错)。

二、训练侧:把权重变成可发送的流

WeightSource:训练框架形状的适配器

WeightSource 是"你的训练器权重长什么样"的适配器,负责为特定框架抽取权重。它的核心约束是:推理引擎拿到的永远是同一形态——HF 格式的参数名,以及已经物化为完整(未分片)形状的张量。任何 gather、重新融合(re-fusing)、反量化、重命名的脏活都必须封装在 source 内部。

一个 WeightSource可重复迭代的,且必须提供两条通道(base.py):

  • metadata() -> list[ParamMeta]:只声明每个参数的名称、线上 dtype 与完整形状,不做任何数据传输。当形状本地可知时开销很低(FSDP DTensor 自带全局形状);对那些必须先物化才能知道形状的生产方(如 Megatron-Bridge export)首次调用可能昂贵,此时应做缓存。
  • 迭代通道:逐个产出已完全物化的 (name, tensor) 对。

ParamMeta 是三条数据的冻结 dataclass(base.py):

@dataclass(frozen=True)
class ParamMeta:
    name: str
    dtype: torch.dtype
    shape: tuple[int, ...]

!!! warning "两条通道必须逐元素一致" metadata() 声明的必须与迭代将产出的完全一致:同样的参数、同样的顺序、同样的 dtype 与 shape。这是 ABC 的不变量,而非某一后端的局部约定。一个在两条通道之间重排、漏掉或改 dtype 的 source 即使在你测试的那个后端上侥幸不报错,也是坏掉的。后端可以自由地同时读取两条通道并信任它们一致——密集 NCCL 就这么做,并在传输时强制执行该约定(详见 nccl.md)。

物化通常本身是一个集合操作,因此每个训练 rank 都必须按相同顺序、步调一致地迭代同一个 source,否则会死锁。metadata() 对自定义 producer 而言也可能是集合操作,所以它也在每个 rank 上执行——只有 sender 才把结果送出去。iter(source) 每轮必须产出一次全新的遍历。

ModuleSource:最常见的现成实现

ModuleSource(module) 覆盖 module.named_parameters(),是普通稠密模块与 FSDP 分片模块的通用解(base.py)。它不做特殊分支处理:迭代时对每个 DTensor 通过 full_tensor() 做 all-gather(内部是集合操作),普通张量则原样放行;metadata() 只读全局 .shape/.dtype,因此永远不会触发 gather。

from vllm.distributed.weight_transfer import ModuleSource

source = ModuleSource(model)

其底层依赖 materialize_full_tensor:优先调用张量的 full_tensor(),没有该方法的普通张量原样返回。这样"昂贵的 gather 恰好只发生一次"——且只在发送时刻发生,读元数据不触发。

自定义 WeightSource:把框架导出变成 HF 格式

当权重需要经过处理才能到达 HF 格式(框架特定导出、重新融合、dtype 转换)时,子类化 WeightSource。文档给出的 MegatronBridgeSource 是一个完整范式:桥接器在内部完成 TP/PP/EP gather 并让每个 rank 都拿到完整张量,metadata()_export() 结果做缓存(因为它是昂贵通道),__iter__ 产出与 metadata() 声明完全一致的流,并在发送前做 to(dtype).detach().contiguous()

from vllm.distributed.weight_transfer import ParamMeta, WeightSource


class MegatronBridgeSource(WeightSource):
    """Megatron model -> HF names, via a bridge that gathers TP/PP/EP internally
    and returns full tensors on every rank."""

    def __init__(self, bridge, module, dtype):
        self._bridge, self._module, self._dtype = bridge, module, dtype
        self._meta: list[ParamMeta] | None = None

    def _export(self):
        return self._bridge.export_hf_weights(self._module)

    def metadata(self) -> list[ParamMeta]:
        # Cache: for producers that must materialize to learn shapes, this is
        # the expensive channel. Runs on every rank (it may be a collective).
        if self._meta is None:
            self._meta = [
                ParamMeta(name, self._dtype, tuple(t.shape))
                for name, t in self._export()
            ]
        return self._meta

    def __iter__(self):
        # Must yield exactly what metadata() declared, in the same order.
        for name, tensor in self._export():
            yield name, tensor.to(self._dtype).detach().contiguous()

held_names():部分所有权

默认假设每个 rank 都能产出全部参数(上面的 bridge source 就是这样——它在 yield 之前已跨所有并行维度 gather)。这简单且永远正确,但意味着 gather 代价在每个 rank 上被全额支付一次。

当 rank 被拆分、各自只持有模型的一部分时,可以覆写 held_names()。它返回本 rank 持有的参数名集合(或返回默认值 None 表示持有全部):

def held_names(self):
    # This pipeline stage's layers, and within them only this EP rank's experts.
    return self._my_stage_names - self._foreign_expert_names

这覆盖了多种训练布局——流水线并行(一个 rank 持有部分层)、专家并行(一个 rank 持有部分专家)、二者叠加,或两者都不匹配的形状。能够按参数路由的后端(见 sharded_rdt.md)随后会从真正持有该名字的 rank 拉取。

覆写它伴随三条硬性要求(base.py 的 docstring 同样声明):

  • metadata() 在每个 rank 上仍必须描述整个模型。只有 sender 的 metadata 能到达推理侧,所以如果某个 rank 只上报自己的份额,其余部分会在无人知晓的情况下漏传。sharded RDT 在 init 时跨 rank 交叉校验这一点。
  • 每个名字必须至少被一个 rank 持有,否则永远无法被提供。引擎在 init 时抛错并点名第一个孤儿参数。
  • 迭代对未持有的名字产出 None。名字仍要出现(保持 metadata 顺序),以便顺序检查跨 rank 对齐——只是数据缺失。声明持有某名字却又对它产出 None 是错误,引擎会按名字报告。

!!! warning "部分所有权只对 sharded RDT 生效" 只有按参数路由的后端才会尊重 held_names()。广播类后端会忽略它、从每个 rank 发出每个名字,因此在那些后端上声明部分所有权不会改变任何行为。

Gather groups:按"层"分组传输

某些后端按层传输而不是按整个模型传输,因此需要把 metadata() 划分成 gather groupslayerwise_groups 用每个名字包含的最外层数字索引段做 key(base.py),因此一个 group 就是一个 decoder layer,无索引的名字段(embeddings、末尾 norm、lm_head)各自成组:

group 0     model.embed_tokens.weight
group 1     model.layers.0.*          <- 一个 decoder layer
group 2     model.layers.1.*
...
group N+1   model.norm.weight, lm_head.weight

按索引而非字面前缀做 key,意味着无需逐架构维护命名表model.layers.0.model.language_model.layers.0.(近期 Qwen 文本权重)、transformer.h.0.(GPT-2、Falcon)、backbone.layers.0.(Mamba)、视觉塔的 visual.blocks.0. 都会以同样方式切分。取的是最外层索引,这让一个 MoE 层保持完整——如 model.layers.3.mlp.experts.7.w1 这类 per-expert 名字按 layer 而非 expert 归组。实现上 _stack_key 逐段扫描名字、取第一个数字段及其前缀;无数字段的名字按"相对第一个有索引名字的位置"归入 pre 块(embeddings)或 post 块(最终 norm、lm_head 与跨栈投影),其中 post 块无论多早到达都排在最后——这正是 pipeline 并行 source 需要的(Megatron-Bridge 会在 layers 之前先流式输出最后一个 stage 的输出块)。

Group index g 在每个 rank 与每个消费方都指向同一层,因为它由某一个 rank 的 metadata() 顺序推导而来。这种一致性正是后端能把缓冲绑定到单层、并在所有人都处理完一层后立即释放该层的前提。

!!! warning "一个叶子模块的全部名字必须落在同一组" sharded-RDT 引擎在最后一个 chunk 落地后立即释放该组,因此一个跨组拆分的模块会把一次 pull 挂起直到 stall 看门狗触发。默认划分保证这一性质;groups() 覆写必须维持它。

分组之上有两个带可用默认实现的钩子:

  • groups()——本 rank 的组,按 metadata 顺序排列。默认是 layerwise_groups(metadata()) 过滤掉不包含任何 held name 的组;某组在这里没有持有的名字则整体跳过。
  • iter_groups()——同一流,按组批量产出。默认实现驱动 __iter__ 并把输出分批,边分边检查名字是否按 metadata 顺序到达。当你的框架能一步产出整个组时覆写它:物化通常是集合操作,按组而非按张量驱动,在 per-expert MoE 模型上能把约 3.7 万次生成器 resume 降到约 95 次(源码注释给出每次同步约省 0.9s)。

由于 metadata() 顺序决定划分,共享同一 layer 索引的所有名字在 metadata 中必须连续。自然导出顺序会穿插不同层(比如把所有 MoE 专家聚在一起)的 source,必须在返回前重排。

三、训练侧:控制面客户端 VLLMWeightSyncClient

为什么需要这一层

很多 RL 框架把推理引擎包进自己的抽象,每个框架访问 vLLM 的方式各不相同。VLLMWeightSyncClient 就是把这套"各自的形状"统一适配进去的唯一接缝,让权重同步引擎保持控制面无关。

契约只有一条:无论 wrapper 长什么样,它最终都要落到底层同样的四个调用——setup 时调用一次 init_weight_transfer_engine,随后每轮调用 start_weight_update → 一次或多次 update_weightsfinish_weight_update。训练器引擎需要来自推理侧的一切都经由这四个调用完成。

class VLLMWeightSyncClient(Protocol):
    def init_weight_transfer_engine(self, init_info: dict[str, Any]) -> None: ...
    def start_weight_update(self) -> None: ...
    def update_weights(self, update_info: dict[str, Any]) -> None: ...
    def finish_weight_update(self, weight_version: str | None = None) -> None: ...

它声明为 @runtime_checkable结构化 Protocol(PEP 544)base.py),这是让适配变廉价的关键:任何具备这四个方法的对象都已经满足该协议。你框架里现成的 wrapper 通常只需要补四个转发方法就能变成一个 client。

两个随 vLLM 发布的内置客户端

源码位于 vllm/distributed/weight_transfer/clients.py

Client 对话对象
RayVLLMWeightSyncClient(handle) 一个或多个 AsyncLLM/LLM Ray actor。接受列表并对每个 handle 扇出每次调用、阻塞等待全部完成,因此多 actor(如 multi-DP)部署被当作一个整体驱动
HTTPVLLMWeightSyncClient(base_url, timeout=300) 通过 RLHF HTTP 路由对接 vLLM server

RayVLLMWeightSyncClientray.get([...]) 对全部 handle 扇出并阻塞(clients.py);HTTPVLLMWeightSyncClient 则向 /init_weight_transfer_engine/start_weight_update/update_weights/finish_weight_update 四个路由 POST(clients.py)。

!!! note "HTTP 与 CUDA IPC 句柄" HTTP 无法承载原始 CUDA IPC handle,因此 HTTPVLLMWeightSyncClient 会把它们 pickle 并 base64 编码进 ipc_handles_pickled 字段(见 _json_safe_update_info);worker 仅在 VLLM_ALLOW_INSECURE_SERIALIZATION=1 时才反序列化。载荷本就是 JSON 原生的后端(NCCL)则原样透传。

自定义 client:适配你自己的 rollout pool

文档给出的自定义范式如下——要点是把你 RL 框架已有的 rollout pool 广播到所有副本:

class MyFrameworkWeightSyncClient:
    """Adapts one RL framework's rollout pool to the four weight-sync calls."""

    def __init__(self, rollout_pool):
        self.pool = rollout_pool          # whatever your stack already has

    def init_weight_transfer_engine(self, init_info):
        # Fan out to every replica and block: all of them receive weights.
        self.pool.broadcast_rpc("init_weight_transfer_engine", init_info=init_info)

    def start_weight_update(self):
        self.pool.broadcast_rpc("start_weight_update")

    def update_weights(self, update_info):
        self.pool.broadcast_rpc("update_weights", update_info=update_info)

    def finish_weight_update(self, weight_version=None):
        self.pool.broadcast_rpc("finish_weight_update")
        if weight_version is not None:
            self.pool.broadcast_rpc("update_weight_version", weight_version)

任何适配器都有两件事必须做对:

  • 覆盖每一个副本,并阻塞直到全部完成。 权重更新不是负载均衡请求:持有模型副本的每个 worker 都必须收到。在它们全部完成前返回,会让训练器冲到仍在加载的 worker 前面。(两个内置 client 都这样做——Ray 通过对 handles 扇出,HTTP 则因为服务端的 DP client 内部广播。)
  • 失败必须抛出。 训练器引擎依赖异常来暴露推理侧错误;一个吞掉异常的 client 会把一次失败的同步变成静默过期的权重,或对于与 worker 做传输汇合的后端变成死锁。

四、训练侧引擎:TrainerWeightTransferEngine 与 TrainerInitInfo

职责与方法

训练器侧引擎负责持有传输状态(NCCL communicator、IPC 设备信息、传输计划)、从 WeightSource 抽取权重、并通过 VLLMWeightSyncClient 驱动推理侧。它对自己的 init info 类型是泛型的,通过 trainer_init classmethod 工厂构造,由 send_weights() 驱动:

方法 说明
trainer_init(init_info, *, client, source=None) Classmethod。与推理侧会合(rendezvous)并返回就绪实例
send_weights() 推送权重并驱动完整的更新往返
shutdown() 拆除 communicator / process group。默认为 no-op

trainer_initsend_weights 会在每一个训练 rank 上被调用。 is_sendertrainer_init 时根据 init_info.rank 一次性解析。每个引擎在每个 rank 上都持有真实 client,但把控制面 RPC 和发送都放在 self.is_sender 守卫之后,因此只有 sender 触达线路;非 sender rank 仍会运行每个集合操作以保持组对齐。

训练侧不接收 WeightTransferConfig。后端来自 init info 上的 backend ClassVar,wire 参数同样由 init info 携带。这一点与推理侧对称地构成"谁配置、谁说了算"的单一事实来源。

TrainerInitInfo:显式 rank 与 ClassVar backend

trainer_init 收到的 init_info 就是 TrainerInitInfo,调用者通过它配置一次传输:选择后端、声明本进程是哪个 rank、携带 wire 参数。每个后端子类化它;基类只保留所有后端都需要的那个字段:

@dataclass
class TrainerInitInfo:
    backend: ClassVar[str]        # factory dispatch key
    rank: int = field(kw_only=True)

    @property
    def is_sender(self) -> bool:
        return self.rank == 0
  • rank 是本训练进程的 rank,由调用者显式提供。引擎不从一个全局 process group 里读它——一旦同时存在多个组(FSDP / TP / PP / EP),那是有歧义的。Rank 0 永远是 sender,这正是 trainer_init 解析成 is_sender 的东西。它是 keyword-only 字段,因此后端子类可以自由新增位置参数。
  • backendClassVar 而非 __init__ 字段:它是工厂读取以分发路由的、每个后端固定不变的常量,所以调用者从不传 backend= 参数。每个子类必须设置它——__init_subclass__ 在缺失时会抛出 TypeErrorbase.py 中的校验)。

子类还携带传输的 wire 参数packed、缓冲大小等)。sender 在 trainer_init 内部把它们传播给 worker,因此两端不可能各执一词。具体字段见 nccl.mdNCCLTrainerInitInfoipc.mdIPCTrainerInitInfo

Full-Resync 后端 vs Delta 后端

source 是可选参数,这把后端分成两种形态:

  • Full resync(NCCL、IPC)——稳定的 WeightSourcetrainer_init 时固定下来,每轮重新迭代;send_weights() 不接收参数。这类后端自己校验 source 非空。
  • Delta(sparse NCCL)——每轮载荷都不同,没有稳定 source。引擎不接收 source,每轮载荷直接传给 send_weights(patches)

实现自定义训练器引擎

文档给出的完整范式(要点:后端用 ClassVar 声明、wire 参数在 trainer_init 的握手里空运给 worker、非 sender rank 也要排空 source 以留在集合操作里):

from dataclasses import dataclass
from typing import ClassVar

from typing_extensions import Self

from vllm.distributed.weight_transfer.base import (
    TrainerInitInfo,
    TrainerWeightTransferEngine,
    VLLMWeightSyncClient,
    WeightSource,
)


@dataclass
class MyTrainerInitInfo(TrainerInitInfo):
    backend: ClassVar[str] = "my_backend"

    endpoint: str
    chunk_size_bytes: int = 256 * 1024 * 1024   # a wire param: shipped to the worker


class MyTrainerWeightTransferEngine(TrainerWeightTransferEngine[MyTrainerInitInfo]):
    init_info_cls = MyTrainerInitInfo

    def __init__(self, *, client, source, is_sender=True, chunk_size_bytes=0):
        super().__init__(client=client, source=source, is_sender=is_sender)
        self.chunk_size_bytes = chunk_size_bytes

    @classmethod
    def trainer_init(
        cls,
        init_info: MyTrainerInitInfo,
        *,
        client: VLLMWeightSyncClient,
        source: WeightSource | None = None,
    ) -> Self:
        if source is None:
            raise ValueError("my_backend requires a WeightSource.")
        engine = cls(
            client=client,
            source=source,
            is_sender=init_info.is_sender,
            chunk_size_bytes=init_info.chunk_size_bytes,
        )
        if engine.is_sender:
            # Ship the must-agree wire params so the worker decodes exactly as
            # this trainer encodes, then open the trainer-side endpoint.
            engine.client.init_weight_transfer_engine(
                {"chunk_size_bytes": init_info.chunk_size_bytes}
            )
        return engine

    def send_weights(self) -> None:
        assert self.source is not None
        meta = self.source.metadata()      # every rank: may be a collective
        if not self.is_sender:
            for _ in self.source:          # stay in the trainer-side collective
                pass
            return

        self.client.start_weight_update()
        self.client.update_weights(
            {
                "names": [m.name for m in meta],
                "dtype_names": [str(m.dtype).split(".")[-1] for m in meta],
                "shapes": [list(m.shape) for m in meta],
            }
        )
        for name, tensor in self.source:
            ...                            # transmit
        self.client.finish_weight_update()

两件事必须做对,而它们都咬过内置后端:

  • 返回前排空(Drain before returning)。 send_weights 不能在仍有传输在途时返回。任何让发送缓冲存活的东西都会随栈帧消亡,而推理侧的 finish_weight_update 后处理可能把尚未落地的权重误判为已完成。
  • 错误路径上绝不 join 控制面线程。 如果你像 NCCL 那样在与传输并发的侧线程上跑 update_weights,而传输抛错,worker 仍阻塞在配对的集合操作里永远不返回。此时应该不等待地关闭 executor,让真正的异常浮出水面而不是挂死。

五、训练侧工厂注册:WeightTransferTrainerFactory

WeightTransferTrainerFactory 与推理侧工厂保持平行结构。它支持两种注册方式(懒加载按名字导入模块,或直接注册类):

from vllm.distributed.weight_transfer import WeightTransferTrainerFactory

# Lazy loading (recommended): the module is imported only when the backend is used
WeightTransferTrainerFactory.register_engine(
    "my_backend",
    "my_package.my_module",
    "MyTrainerWeightTransferEngine",
)

# Or register the class directly
WeightTransferTrainerFactory.register_engine("my_backend", MyTrainerWeightTransferEngine)

engine = WeightTransferTrainerFactory.trainer_init(
    init_info=MyTrainerInitInfo(rank=0, endpoint="..."),  # `backend` selects the engine
    client=client,
    source=source,
)

注意训练侧没有 backend= 参数:MyTrainerInitInfo.backend 声明自己,工厂依据它分发。任何与两端"必须一致"的静态参数都放在 init info 上,由 sender 在 init 握手时传给 worker。

六、推理侧引擎:WeightTransferEngine

类型参数与五个抽象方法

推理侧基类是一个泛型抽象类,由两个 dataclass 类型参数化(base.py):

  • TInitInfo(继承 WeightTransferInitInfo):后端特定的初始化参数。
  • TUpdateInfo(继承 WeightTransferUpdateInfo):后端特定的权重更新元数据。

子类必须实现五个方法:

方法 说明
init_transfer_engine(init_info) 在每个推理 worker 上初始化通信信道,并记录训练器送来的 wire 参数
start_weight_update() 为一次更新做准备(如开始逐层重载 layerwise reload);in-place 引擎为 no-op
finish_weight_update() 结束更新(如收尾逐层重载);in-place 引擎为 no-op
receive_weights(update_info) 接收权重并加载进 self.model
shutdown() 清理资源

基类提供四件套:

  1. __init__,接收 configWeightTransferConfig)、vllm_configVllmConfig)、devicetorch.device)、modelnn.Module)。
  2. update_weights(update_info_dict),是 receive_weights 的薄包装:把 dict 解析成类型化 dataclass、调用 receive_weights、再做设备同步——除非引擎设置了下面的 defers_processing
  3. parse_init_info / parse_update_info,把 API 层 dict 转成类型化 dataclass,坏载荷抛 ValueError(实际是捕获 TypeError 后包装重抛)。
  4. set_weight_update_target / reset_weight_update_target,用于把一次更新重定向到投机解码的 draft 模型上(base.py 保存/恢复默认 model 与 model_config)。

gpu_worker.py 的接线中可以看到推理侧引擎之上的 API 约束:start_weight_update 处于活动状态时再次调用会报错(必须先 finish_weight_update);update_weights 之前必须已 start_weight_update;每个 chunk 都会加载进该次 start_weight_update 会话所指向的模型。

!!! note "wire 参数从握手读取,而不是从载荷读取" 任何两端必须一致的参数——packed、缓冲几何——都随 init info 到达,应该在 init_transfer_engine 里存到 self 上、然后在 receive_weights 里从 self 读取。每轮的 update info 只携带每轮元数据。这正是让"训练器/worker 不匹配"在类型上不可表达的原因。

!!! note "defers_processing:当一次 update 返回意味着"已入队"而非"已生效"" 一个把 GPU 后处理流水化到后台线程的引擎,不能让 update_weights 同步设备——那会阻塞在那些线程上并把流水线串行化。这类引擎设置类属性 defers_processing = True,省略每次更新的同步,改为在 finish_weight_update 里保证完成。

经由 `finish_weight_update` 的调用方无需任何动作——引擎在那里排空。但自己驱动收尾的调用方(例如自己跑 `finalize_layerwise_reload`)必须先检查该标志并调用 `drain_pending()`,因为置位后一次返回的 `update_weights` 意味着*已入队*而非*已生效*。`drain_pending()` 是幂等的,对同步处理的引擎是 no-op,因此永远可以安全调用。

[sharded_rdt.md](https://gitcode.com/GitHub_Trending/vl/vllm/blob/dd0760165c17fdf750b1aa8e2a26ac885eafd598/docs/training/weight_transfer/sharded_rdt.md?utm_source=gitcode_repo_files) 中介绍的内置引擎正是设置它的那个:它在拥有自己 CUDA stream 的后台线程上 scatter 与量化,因此其 `drain_pending()` 会在 `finalize_layerwise_reload` 运行前把两条队列都 join、两条 stream 都同步。

请求类与 draft 模型更新

API 层请求类用纯 dict 提供与后端无关的序列化(base.py):

from vllm.distributed.weight_transfer.base import (
    WeightTransferInitRequest,
    WeightTransferUpdateRequest,
)

# Init request (dict is converted to backend-specific TInitInfo)
init_request = WeightTransferInitRequest(
    init_info={"master_address": "10.0.0.1", "master_port": 29500, ...}
)

# Update request (dict is converted to backend-specific TUpdateInfo)
update_request = WeightTransferUpdateRequest(
    update_info={"names": [...], "dtype_names": [...], "shapes": [...]}
)

使用内置 client 时你从不需要手工构造它们——RayVLLMWeightSyncClient 替你包装 dict(封装成 WeightTransferInitRequest/WeightTransferUpdateRequest 再 remote 调用),HTTPVLLMWeightSyncClient 则把它们作为 JSON 发布。

在 LLM/API 层,调用 start_draft_weight_update() 而不是 start_weight_update() 可把更新目标指向投机 draft 模型;update_weights / finish_weight_update 不变。不支持的引擎设置 supports_draft_weight_update = Falsebase.pygpu_worker.py 中会用该标志拒绝并给出明确报错)。worker 侧的 update_weights 收到带 draft 会话的 chunk 时,会先用 set_weight_update_target 重定向、结束后用 reset_weight_update_target 还原(gpu_worker.py 附近的逻辑)。

七、实现并注册一个自定义推理引擎

1. 定义信息 dataclass

from dataclasses import dataclass
from vllm.distributed.weight_transfer.base import (
    WeightTransferEngine,
    WeightTransferInitInfo,
    WeightTransferUpdateInfo,
)

@dataclass
class MyInitInfo(WeightTransferInitInfo):
    endpoint: str
    chunk_size_bytes: int = 256 * 1024 * 1024   # must-agree wire param

@dataclass
class MyUpdateInfo(WeightTransferUpdateInfo):
    names: list[str]
    dtype_names: list[str]
    shapes: list[list[int]]
    # Per-round metadata only.

2. 实现引擎

注意 start_weight_update/finish_weight_update 的注释区分了两类引擎:checkpoint 格式引擎在此调用 initialize_layerwise_reload/finalize_layerwise_reload,in-place 引擎则为 no-op。receive_weights 从 update info 重建张量并交给 self.model.load_weights

class MyWeightTransferEngine(WeightTransferEngine[MyInitInfo, MyUpdateInfo]):
    init_info_cls = MyInitInfo
    update_info_cls = MyUpdateInfo

    def init_transfer_engine(self, init_info: MyInitInfo) -> None:
        # Record the trainer's wire params, then set up the connection.
        self.chunk_size_bytes = init_info.chunk_size_bytes
        ...

    def start_weight_update(self) -> None:
        # Checkpoint-format engines: run initialize_layerwise_reload(self.model).
        # In-place engines: no-op
        ...

    def finish_weight_update(self) -> None:
        # Checkpoint-format engines: run finalize_layerwise_reload(...).
        # In-place engines: no-op
        ...

    def receive_weights(self, update_info: MyUpdateInfo) -> None:
        weights = []
        for name, dtype_name, shape in zip(
            update_info.names, update_info.dtype_names, update_info.shapes
        ):
            dtype = getattr(torch, dtype_name)
            weight = self._fetch_weight(name, shape, dtype)
            weights.append((name, weight))
        self.model.load_weights(weights)

    def shutdown(self) -> None:
        # Clean up resources
        ...

3. 注册到工厂

from vllm.distributed.weight_transfer import WeightTransferEngineFactory

# Option 1: Lazy loading (recommended for built-in engines)
WeightTransferEngineFactory.register_engine(
    "my_backend",
    "my_package.my_module",
    "MyWeightTransferEngine",
)

# Option 2: Direct class registration
WeightTransferEngineFactory.register_engine(
    "my_backend",
    MyWeightTransferEngine,
)

注册完成后,用户通过 WeightTransferConfig(backend="my_backend") 即可选中你的后端。

WeightTransferEngineFactory 的懒加载注册表

推理侧工厂采用带懒加载的注册表模式(factory.py)。内置引擎(ncclipcsparse_ncclsharded_rdt)在 import 时注册,但其模块只在实际请求该后端时才被 import——这避免了在不需要时导入重依赖(如 NCCL communicator)。重复注册同名后端会抛 ValueError;未注册的后端在 create_engine/trainer_init 时会得到带可用列表的 ValueError 提示:

from vllm.distributed.weight_transfer import WeightTransferEngineFactory

# Create an engine from config
engine = WeightTransferEngineFactory.create_engine(
    config=weight_transfer_config,
    vllm_config=vllm_config,
    device=device,
    model=model,
)

vLLM 会在 worker 启动期间替你调用它(gpu_worker.py);只有当你需要把引擎嵌入自己的 worker 时才需要直接调用。

八、验证与测试佐证

从测试代码可以看出这套 ABC 的用法约定确实可按文档书写。仓库中的 tests/entrypoints/weight_transfer/test_weight_transfer_llm.py 定义了 MockInitInfo(WeightTransferInitInfo)MockUpdateInfo(WeightTransferUpdateInfo) 与一个实现 init_transfer_engine/start_weight_update/receive_weights/finish_weight_update/shutdown 五个方法的 MockWeightTransferEngine,并测试了:get_world_size(TP1 下 WeightTransferConfig(backend="nccl"))、init_weight_transfer_engine 确实触达引擎(通过 mock gpu_worker.WeightTransferEngineFactory.create_engine 断言 init_transfer_engine_called)、以 WeightTransferInitRequest(init_info={...}) 形式下发参数等路径。这验证了"推理侧从 WeightTransferConfig 建引擎 → dict 请求 → parse_* → 五个抽象方法"这条完整链路的可运行性。

九、常见误区与设计要点速查

把本文反复强调的约束汇总成一张自查表,写自定义 source / client / 引擎前值得逐条过一遍:

关注点 正确做法 违反后果
metadata() 与迭代一致 同参数、同顺序、同 dtype/shape 两端对数据流切分不一致,数据错乱
多 rank 迭代同步 所有 rank 按相同顺序、步调一致地迭代 source rank 间死锁
send_weights 返回前 确保传输全部落地再返回 推理侧 finalize 掉未落地的权重
client 失败处理 必须向上抛异常 静默过期权重或死锁
client 覆盖全部副本 广播所有持有模型副本的 worker 并阻塞 部分 worker 权重过期、训练器超前
wire 参数的位置 init info(握手),不是 per-round update info 两端各执一词、不可排查的不一致
defers_processing=True 的引擎 直接驱动收尾前先调 drain_pending() 用到未生效的权重
部分所有权声明 仅当后端按参数路由(sharded RDT) 广播后端上声明无效;孤儿参数、漏传

想继续深入某条传输路径的细节,可继续阅读本目录下的后端专章:NCCL 后端与 sparse NCCLCUDA IPC 后端sharded RDT 后端,以及总览文档 README(含 HTTP 控制面端点表与 RLHF HTTP 路由 /init_weight_transfer_engine/start_weight_update/update_weights/finish_weight_update/update_weight_version/weight_info/pause/resume/get_world_size 的完整说明)。

登录后查看全文
热门项目推荐
相关项目推荐