首页
/ vLLM Sharded RDT 权重传输引擎:面向大模型 MoE 的按需切片拉取式权重同步

vLLM Sharded RDT 权重传输引擎:面向大模型 MoE 的按需切片拉取式权重同步

2026-09-07 13:53:08作者:庞队千Virginia

导读

sharded_rdt 是 vLLM 权重传输系统中面向超大模型 / MoE + 专家并行场景的后端,它通过 Ray Direct Transport(RDT)在训练进程与推理 worker 之间做点对点(pull-based)拉取式权重同步:每个推理 worker 只拉取自己在张量并行与专家并行下真正消费的那一小片(slice)权重,而非广播整份权重。本文以 docs/training/weight_transfer/sharded_rdt.md 为主线,结合 vllm/distributed/weight_transfer/ 源码与 examples/rl/rlhf_sharded_rdt_small_ep.py 示例,讲清它的适用场景、切片追踪原理、分组建 gather 的流水线、所有权路由,以及推理侧 / 训练侧完整的接入与调参方式。读完你可以在自己的 RL(RLHF / GRPO)训练栈中独立配置并验证一次 sharded RDT 权重同步。

适用场景与前提条件

从仓库文档与源码看,sharded RDT 专为“worker 不需要整份权重”的架构设计,适用场景集中在三类:

  • 超大模型,广播整份参数本身成为瓶颈——典型是配合专家并行(EP)服务的 MoE,每个 worker 只拥有专家的很小一部分;
  • 训练与推理分处于不同 GPU,两者通过 NIXL 支持的互联(InfiniBand、RoCE、EFA)传输;
  • 训练方本身就是分片的,包括**流水线并行(PP)**训练方——某个 rank 只持有模型的一部分,此时全量广播根本不可行。

作为对照,同一套 docs/training/weight_transfer/README.md 框架下的 NCCL(广播)、IPC(CUDA IPC 句柄、同机同卡)、sparse_nccl(稀疏 patch)后端都不具备“按名字路由到持有者”的能力——这一点在 docs/training/weight_transfer/base.md 中有明确警告:只有按参数路由的后端才尊重 held_names() 的部分所有权声明。

启用 sharded RDT 需要同时满足以下硬性前提(来自原文档与 sharded_rdt_common.py):

  • distributed_executor_backend="ray"——worker 必须是 Ray actor;
  • 训练方与 worker 两侧都需 Ray >= 2.56.0(源码中 RDT_MIN_RAY_VERSION = "2.56.0",并附有 check_ray_rdt_version() 启动期检查:RDT 依赖 ray.experimentalregister_nixl_memory / set_target_for_ref,最低 2.55 才能导入,2.56 是实际被测试的版本);
  • 训练方与 worker 共享的环境中安装 nixl
  • 权重加载器必须停留在支持的操作集合之内(下文详解);
  • 不允许 EPLB(enable_eplb=true),因为它会在运行时重排专家,使训练期记录下来的计划失效。

工作原理:三段式架构

1. 切片通过 vLLM 自己的权重加载器被追踪(BAKE,一次性)

权重加载器通常收到一个完整的 HF 格式张量,再切出本 worker 需要的部分。sharded RDT 反其道而行之:引擎递给加载器一个 FakeRDTTensor——一个零存储的张量子类(实现见 sharded_rdt_fake.py),它只回答 .shape / .dtype / .size(),不持有任何数据。加载器对它做的每一个 view 或 slice 操作,都会返回一个把该操作追加进**记录链(op chain)**的新 fake;copy_ 是终结整条链的“汇”(sink)。

这条链就是线上格式。例如

("model.layers.0.w", (("narrow", (0, 512, 512), ()), ("t", (), ())))

意思是告诉训练方:“取出这个张量,做一次 narrow,再转置,把结果发给我”。训练方回放时执行 getattr(tensor, op)(*args, **kwargs)。源码中链的每个元素是 ("op_name", 位置参数, 排好序的 kwargs 项),可哈希,从而可以作为去重键 FetchKey = (name, op_chain)

允许哪些操作由 SUPPORTED_OPS 表唯一决定(sharded_rdt_common.py):narrowviewreshape__getitem__unsqueezesqueezetransposetpermuteflattencontiguouschunkunbind——全部是纯 view 操作。表中注释特别强调:to 永远不会被加入,因为它是 bake 想要拒绝的 dtype/device“逃生门”。ALLOWED_OPS 由同一张表派生,因此记录方(recorder)与回放方(replayer)不可能漂移。

加载器只要做了需要真实数据的操作——算术、.to().float().item().data 访问、bool mask 索引——都会在 init 时抛错。源码里 FakeRDTTensor.__torch_dispatch__ 会明确点名不支持的 op 和当前链并列出允许集合。原文档说得很清楚:在 setup 阶段大声失败,好过静默传输错误的字节

关键设计是发现(discovery)很昂贵,所以只做一次:在 init_transfer_engine 时,让每个参数都位于 meta 设备上、对 model.load_weights 做一次空跑(dry run),不传输任何数据,只是按叶子模块记录“哪一片喂给哪个目的区域”。此后每次同步都是纯回放:不再有 load_weights、不再有 FakeRDTTensor 分发、不再有任何发现动作。这也正是引擎模块 docstring(sharded_rdt_engine.py)中 “BAKE / REPLAY” 两阶段划分的含义——每个落地的目的区域被记录为参数的 as_strided 描述(offset/shape/stride,见 _Scatter)。

2. 到达的切片直接落入 layerwise reload 缓冲区

引擎自己驱动 layerwise 重载(源码在 start_weight_update / finish_weight_update 中分别调用该模块的 initialize_layerwise_reload / finalize_layerwise_reload,对应的上层抽象约定见 docs/training/weight_transfer/base.mdWeightTransferEngine 一节)。由于空跑已经把每个目的区域记录成参数的 as_strided 区域,到达的切片被直接拷进正在重载的那一层——worker 上永远不会实例化完整的 HF 张量,也不会对 load_weights 跑第二遍。每一层在其最后一片落地后立即被量化并拷入其持久化的 kernel 存储。

一个值得注意的配套机制是 defers_processing:这个引擎在后台线程用自己的 CUDA stream 做 scatter 与量化,所以其 update_weights 返回只代表“已入队”而非“已应用”,真正的完成保证落在 finish_weight_update——引擎提供幂等的 drain_pending(),它先合并两个队列并同步两条 stream,再执行 finalize_layerwise_reload

3. gather 与拉取被流水线化:gather_lookahead 的语义

训练方通常不能直接以“原样”参数对外服务:FSDP 分片了它们,即便是 EP 切分的训练方也得先把一个完整 expert 拼起来。所以每次同步仍然要跑 gather 集合通信——但一次只 gather 一层,而不是一次 gather 整个模型

一个 gather 组 = 一个 decoder layer。 参数列表按每个名字最外层索引段(outermost index segment)分 key,于是那些没有索引的连续名字——第一层之前的 embeddings、最后一层之后的 final norm 与 lm_head——各自成组:

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

这一分组的权威实现在 base.pylayerwise_groups(names)。对索引而非字面前缀分 key,意味着不需要逐架构建表:model.layers.0.、VLM 的 model.language_model.layers.0.、GPT-2 / Falcon 的 transformer.h.0.、Mamba 的 backbone.layers.0.、视觉塔的 visual.blocks.0. 都用同一条规则切分。取的是最外层索引,这使一个 MoE 层保持完整——model.layers.3.mlp.experts.7.w1 按层而非按 expert 分 key。注意控制组的语义一致性:原文档警告“一个叶子模块的所有源必须落在同一组”,因为引擎在某层最后一片落地时就释放该组,跨组的模块会卡住一次拉取直到 stall 看门狗触发。

“层”是其后一切流程的单元:训练方 gather 完一层就发布它(立即可拉取),随即去 gather 下一层,同时消费者拉取刚发布的那层;当所有消费者都发信号表示用完某层后,训练方才丢弃它并获得一个 gather 的信用额度gather_lookahead 就是这个循环允许领先消费者的组数,因此训练方上同时驻留不超过 gather_lookahead + 1。默认值 1 表示“当前层正被拉取时,下一层已 gather 好并可拉”,足以藏住交接延迟又不必让训练方显存翻倍。只有在“某一层的 gather 比它的拉取更慢”时才需要调大它。

层同时是消费者释放的单位和接收缓冲区定长的单位,因此它是两侧内存都有界的保证:没有它,整个模型就是一次传输,两侧就得同时持有各自完整份额。源码中 ShardedRDTTrainerInitInfo 的注释也印证:gather lookahead “bounds trainer-resident memory at gather_lookahead + 1 groups”。

4. 所有权(Ownership)与路由

一个训练方 rank 不必持有整个模型。每个 rank 通过 WeightSource.held_names() 声明自己持有什么,全舰队在 trainer_init 时 all-gather 这些声明,消费者把每次拉取路由到真正持有该名字的 rank。流水线阶段、专家并行、以及两者的组合,都只是同一种声明。消费者会把拉取分摊到持有该名字的各 rank 上,因此不会出现某个训练方 NIC 成为瓶颈。

路由表在消费者侧构建(sharded_rdt_common.pyRdtRouter)。两个引擎从同一份线上数据推导出相同的表,所以它们对“谁服务什么”意见一致——不一致不是“答案错误”而是“挂死或大声误路由”:发给从未 gather 过该名字的 producer 的拉取会触发其 served-names 守卫。带部分所有权有几个硬性要求(见 base.md):每个 rank 的 metadata() 仍须描述整个模型(只有发送方的 metadata 会到达推理侧);每个名字至少被一个 rank 持有,否则引擎在 init 时指名道姓地抛错;迭代对未持有的名字产出 None 以保持跨 rank 的顺序一致。

推理侧接入

推理侧只需要一个后端名,其余一切——producer 有哪些、模型如何分成 layer 组、所有权表——都在 init 握手时由训练方送达。

from vllm import LLM
from vllm.config import WeightTransferConfig

llm = LLM(
    model="my-model",
    weight_transfer_config=WeightTransferConfig(backend="sharded_rdt"),
    distributed_executor_backend="ray",
)

对应 vllm serve 的命令行形式为:

vllm serve my-model \
  --distributed-executor-backend ray \
  --weight-transfer-config '{"backend": "sharded_rdt"}'

WeightTransferConfigvllm/config/weight_transfer.py)的 backend 字段只是一个字符串选择器,在引擎创建时对照 WeightTransferEngineFactory 注册表校验,合法值为 "nccl" | "ipc" | "sparse_nccl" | "sharded_rdt"

!!! warning "在选择 gpu_memory_utilization 之前先估好接收缓冲区尺寸" 每个 worker 持有 num_rdt_buffers 个接收缓冲区,每个都要足够容纳它拉取的最大单次切片批次。和 NCCL、NIXL 内部缓冲区一样,它们不计入 gpu_memory_utilization——所以一个不留余量的利用率配置会让引擎“健康启动”却在第一次同步时 OOM。缓冲区尺寸由最大的原子切片驱动——对一个未被切分持有未拼接词表矩阵(untied vocab matrix)的 worker 而言,那就是整个 embedding。源码侧,缓冲区槽位通过 buffer_alloc_bytes() 一次性按“请求大小、presize 下限、256MB 粗粒度向上取整”三者取最大来分配且永不重新增长——重新增长在注释中被定性为正确性隐患而非单纯性能问题(sharded_rdt_common.py)。

训练侧接入

训练侧在每个 rank 上都要执行 trainer_initsend_weights:每个 rank 都拥有一个 serve actor、都参与 gather;只有 rank 0 驱动推理侧的握手。任何满足 VLLMWeightSyncClient 协议的对象都可以作为 client(内置的 HTTPVLLMWeightSyncClient 走 RLHF HTTP 路由,RayVLLMWeightSyncClient 直连 Ray actor)。

from vllm.distributed.weight_transfer import (
    ModuleSource,
    HTTPVLLMWeightSyncClient,
    WeightTransferTrainerFactory,
)
from vllm.distributed.weight_transfer.sharded_rdt_trainer import (
    ShardedRDTTrainerInitInfo,
)

engine = WeightTransferTrainerFactory.trainer_init(
    init_info=ShardedRDTTrainerInitInfo(
        rank=rank,                                # rank 0 is the sender
        num_consumers=8,                          # inference workers, fleet-wide
        trainer_actor_namespace="my_namespace",   # must be visible to the workers
    ),
    client=HTTPVLLMWeightSyncClient("http://localhost:8000"),
    source=ModuleSource(model),
)

engine.send_weights()   # once per sync, on every trainer rank

ModuleSource 覆盖普通与 FSDP 分片模块:迭代时通过 full_tensor() all-gather 每个 DTensor,而 metadata() 只读全局 .shape / .dtype,绝不触发 gather。若训练方不是普通的 nn.Module——例如 Megatron 导出、原始分片 checkpoint——就子类化 WeightSource

需要为自定义 client 适配时注意两点(base.md 反复强调):必须触达每个副本并阻塞到全部完成(权重更新不是负载均衡请求),以及失败必须抛异常(吞掉异常会把一次失败的同步变成静默过期权重或挂死)。

ShardedRDTTrainerInitInfo 参数一览

字段定义与文档位于 sharded_rdt_trainer.py

字段 默认值 说明
rank keyword-only。本训练方 rank;0 是发送方。显式传入而非读全局进程组——FSDP/TP/PP/EP 多进程组并存时全局 rank 有歧义
num_consumers 全舰队推理 worker 数(DP × TP × PP × PCP),用于 M:N 块分配与 free 引用计数
workers_per_replica 0 每个推理 DEPLOYMENT 的消费者数(num_consumers // num_replicas)。用来形成 slot-sharing 组:id 相差该值整数倍的消费者是不同 deployment 的同一号 worker,bake 出相同计划、可共享一个 serve 槽位。0 或等于 num_consumers 则关闭共享
trainer_actor_namespace None 引擎在其中 spawn serve actor 的 Ray namespace;worker(在自己的 EngineCore 子进程里 ray.init)按名字在此解析 actor,因此必须对 worker 可见
num_rdt_buffers 2 两侧的环形缓冲深度(ring depth K,须与 worker 一致)
buffer_presize_gb 0.0 每个缓冲区槽位的预分配下限,单位 GiB;建议设为能覆盖最大原子切片的值,避免 NIXL 描述缓存抖动
gather_lookahead 1 gather 循环可领先于消费者的“已 gather 未释放”层数,训练方常驻内存上界为 gather_lookahead + 1 个组
stall_timeout_s 300.0 无任何 publish/serve/free 进展即判同步失败的时间(秒)。这是消费者中途死亡的活性兜底,不是延迟目标

这些“双方必须一致”的 wire 参数都由训练方在 trainer_init 握手时原样转发给 worker,因此两侧不可能出现不一致——这正是 base.md 强调的“wire 参数不上 WeightTransferConfig、不上每轮 update info”的设计原因。

实战示例:4 卡 MoE 权重同步

原文档给出的参考示例是 examples/rl/rlhf_sharded_rdt_small_ep.py——2 个 FSDP2 训练 rank → 2 个 vLLM DP rank(开专家并行),单节点 4 卡,训练舰队与推理舰队分离(这是本后端唯一支持的布局)。这个示例刻意保持训练方极小——FSDP2 只用来让权重真实化——使文件聚焦于权重同步本身,可在 CI 无人值守运行。

值得注意的实现细节:

  • 默认模型是 Qwen/Qwen1.5-MoE-A2.7B(可用环境变量 RDT_MODEL 换成更小的 MoE),但权重从不被下载:训练方自己构造 config 建模型,server 用 --load-format dummy 启动,因此被测试的是同步本身。默认配置每张卡约需 29 GiB(bf16)显存,适合 40 GiB 以上显卡,因为 fully_shard 切分前每个 rank 都会先构建整份模型;
  • 训练方发布的是 checkpoint 名字而非自身模块名:Transformers 会把每层 expert 融合成 [E, ...] 张量(mlp.experts.gate_up_proj / down_proj),而 checkpoint 按单 expert 存储,vLLM 的 MoE 加载器读的是单 expert 形式。CheckpointNameSource(同目录的 examples/rl/rdt_weight_source.py)会把融合张量再拆回单 expert——这正是真实训练方把内部布局映射到 checkpoint 名字时要做的事;
  • 推理侧以 --enable-expert-parallel + DP2 获得真正的 EP;
  • 主流程(对应 examples/rl/rdt_vllm_serve.py 提供的 HTTP 辅助函数):先占住训练方 GPU 的 placement group → 启动 vLLM serve → 生成基线输出 → 两次 send_weights() 同步(分别通过 HTTP 端点的 pause / resume 包住)→ 断言。两条断言构成冒烟测试:第一次同步后生成必须变化(证明权重真的动了,静默跳过的同步会失败而非输出看似合理的文本);第二次同步后生成必须不变(证明回放稳定)。

何时使用、何时不用的速查

应当使用 sharded RDT:超大 MoE + 专家并行、训练与推理分 GPU 集群、训练方自身 PP/EP 分片且每个 worker 只需自己的切片、total_bytes / num_workers 的移动量显著低于 total_bytes 广播。

不应使用或无法使用:训练与推理共享同卡(改用 IPC);需要广播整份权重的小模型或纯张量并行(改用 NCCL);加载器含 to() / 算术 / bool mask 索引等数据依赖操作;开了 enable_eplb;或环境中 Ray 版本低于 2.56、未安装 nixl。

更进一步

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

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.13 K
2.75 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
857
1.35 K
docsdocs
暂无描述
Markdown
897
5.8 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
529
593
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
915
1.83 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.58 K
1.01 K
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.35 K
1.46 K
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.01 K
515
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
547
388