Prhub

#51729 [Docs][RL] Rewrite weight-transfer docs; standardize examples

原始 PR 作者 hao-aaron 合并时间 2026-08-12 09:53 文件变更 16 提交数 3 评论 4 代码增减 +1226 / -1157

执行摘要

重写权重传输文档并统一 RL 示例到 vllm serve + HTTP 模式

PR 描述明确这是系列收尾:"Closes out the two follow-ups that PR 3 of the trainer-side weight-transfer split (#48042, #48981, and the NCCL/sparse PR) deferred"。两个遗留问题:一是文档仍描述 PR 2/3 已删除的静态 trainer_send_weights / TrainerSendWeightsArgs 路径("Every trainer-side snippet documented the static trainer_send_weights / TrainerSendWeightsArgs path that PR 2 and PR 3 removed"),文档与代码脱节;二是 worker ABC 上残留的抽象方法迫使三个后端保留只抛 NotImplementedError 的过渡 stub("the three transitional NotImplementedError stubs that existed only to satisfy it")。示例层面,Ray 版 rlhf_nccl.py / rlhf_ipc.py 与 HTTP 版重复("Beyond using Ray in place of vllm serve they demonstrated nothing their HTTP counterparts do not"),且原 HTTP 示例要求用户手动先启动 server,无法无人值守运行,导致唯一接入 CI 的示例就是将被删除的两个 Ray 文件。

建议精读,尤其是三类读者:做 RLHF/RL 集成的工程师(文档 + 示例直接决定接入方式)、计划自定义 weight-transfer 后端的人(base.md 的 WeightSource 双通道契约与 VLLMWeightSyncClient 结构型 Protocol 是核心设计)、以及维护示例库的团队(自启 server + sleep/wake 编排是值得复用的模式)。值得关注的设计决策:用"传输即客户端"(the transport is the client you pass)消解 IPC/NCCL 的 send_mode 概念;用单个 vllm serve --data-parallel-size N 取代多 actor fan-out;用分级 sleep/wake 在传输前腾出显存。

讨论亮点

本 PR 没有实质性的 review 讨论:review 评论列表为空,仅 claude[bot] 自动提示 "This pull request is from a fork — automated review is disabled. A repository maintainer can comment @claude review to run a one-time review.",未触发人工 bot 审查。核心设计权衡记录在 PR body 中,作者给出了三个关键决策理由:

  • 删除 Ray 示例:"Beyond using Ray in place of vllm serve they demonstrated nothing their HTTP counterparts do not; their one unique behaviour (fp8 quantization) moves to rlhf_http_nccl.py."
  • 四客户端收敛为一个:"one HTTPVLLMWeightSyncClient replaces the four-handle fan-out (the API server's DP client already broadcasts each weight-transfer RPC to every engine core)."
  • 示例可无人值守:"The HTTP examples now launch and tear down their own vllm serve (printing the exact command), so they are runnable unattended. This also preserves CI coverage."

这些决策未收到相反意见即被合并。

实现拆解

按 5 步完成:

  1. 移除过渡 API,收敛 worker ABC:在 vllm/distributed/weight_transfer/base.py 中删除 WeightTransferEngine.trainer_send_weights 抽象静态方法(旧无状态设计的关键残留),并删除 ipc_engine.py、nccl_engine.py、sparse_nccl_engine.py 中三个仅用于满足该抽象方法的 NotImplementedError 过渡 stub,同时清理三个文件中不再使用的 typing.Any 导入。worker ABC 收敛为五个抽象方法(init_transfer_engine / start_weight_update / update_weights / receive_weights / shutdown)。配套测试 tests/entrypoints/weight_transfer/test_weight_transfer_llm.py 移除对 trainer_send_weights 的引用。同一提交还在 WeightSource.metadata() 的 docstring 中补齐"metadata 与迭代必须逐元素一致"的双通道契约说明,并说明 NCCL 依赖该契约(接收缓冲区大小 + packed 分块边界)。

  2. 重写四页文档:docs/training/weight_transfer/ 下 base.md 从描述旧单一 ABC 改为"四个可独立替换的抽象"(WeightSource / VLLMWeightSyncClient / TrainerWeightTransferEngine / WeightTransferEngine),新增自定义 WeightSource(MyExportSource 示例,含 metadata 缓存)与自定义 client 示例,以及两个工厂的注册关系说明;nccl.md 覆盖 NCCLTrainerInitInfo 字段语义,并把 sparse NCCL 重新定义为基于 baseline 的 delta 后端;ipc.md 删除旧 send_mode 框架,改为"传输方式就是你传入的 client";README.md 新增 'Where Each Setting Lives' 配置对照表与多 rank 章节。

  3. 统一 HTTP 示例为自启 server 模式:rlhf_http_nccl.py 与 rlhf_http_ipc.py 新增 start_vllm_server()(subprocess.Popen 启动 vllm serve,轮询 /health 就绪,默认 900 秒超时,打印精确启动命令),main() 用 try/finally 保证 server 随脚本退出(terminate → wait 30s → kill)。rlhf_http_nccl.py 固定 3-GPU 布局:server 用 --device-ids 0,1 + TP=2 + --quantization fp8,训练侧固定 cuda:2,替代原先基于 cuda:{world_size} 的设备猜测;fp8 行为由此从被删的 rlhf_nccl.py 迁入。

  4. FSDP+EP 示例迁移到 HTTP:rlhf_ipc_fsdp_ep.py 删除 MyLLM 子类、VLLM_DP_* 环境变量 SPMD 协调、DataParallelInferenceEngine 与 get_weight_metadata,改为:先启动 4 个 FSDP Ray worker(ray.get_gpu_ids() 查询实际 GPU),再用 --device-ids 把单个 vllm serve --data-parallel-size 4 --enable-expert-parallel 钉到同一批 GPU;每个 FSDP rank 仍构建 IPCTrainerWeightTransferEngine,但 client 从四 handle fan-out 的 RayVLLMWeightSyncClient 换成单个 HTTPVLLMWeightSyncClient(API server 的 DP client 会把每个 RPC 广播到所有 engine core)。传输间隙引入分级睡眠:/sleep?level=1(卸载权重 + 释放 KV cache)→ /wake_up?tags=weights → 传输 → /wake_up?tags=kv_cache&tags=scheduling。

  5. CI 配置与验证:.buildkite/test_areas/distributed.yaml 把原指向已删除 Ray 示例的两个条目改为 HTTP 示例(rlhf_http_*),.buildkite/test-amd.yaml 相应调整。验证:CPU-only 主机 pytest tests/distributed/test_weight_transfer.py tests/distributed/test_packed_tensor.py 结果为 61 passed / 45 skipped;8xH100 上 6 个示例全部端到端通过(同步前为 dummy 乱码、同步后为正常文本),IPC FSDP 示例与 NCCL FSDP 示例输出逐字符一致,rlhf_sparse_nccl.py 保持 after_equal / patch_digest_equal = True。

文件 模块 状态 重要度
vllm/distributed/weight_transfer/base.py 权重传输 modified 7.01
vllm/distributed/weight_transfer/nccl_engine.py 权重传输 modified 6.69
vllm/distributed/weight_transfer/ipc_engine.py 权重传输 modified 6.65
vllm/distributed/weight_transfer/sparse_nccl_engine.py 权重传输 modified 6.71
examples/rl/rlhf_http_nccl.py 示例脚本 modified 8.14
examples/rl/rlhf_http_ipc.py 示例脚本 modified 8.49
examples/rl/rlhf_ipc_fsdp_ep.py 示例脚本 modified 8.69
examples/rl/rlhf_nccl.py 示例脚本 removed 8.4
examples/rl/rlhf_ipc.py 示例脚本 removed 8.29
docs/training/weight_transfer/base.md 文档 modified 6.28
docs/training/weight_transfer/nccl.md 文档 modified 4.93
docs/training/weight_transfer/ipc.md 文档 modified 4.62
docs/training/weight_transfer/README.md 文档 modified 4.38
tests/entrypoints/weight_transfer/test_weight_transfer_llm.py 测试 modified 3.89
.buildkite/test_areas/distributed.yaml CI 配置 modified 3.2
.buildkite/test-amd.yaml CI 配置 modified 2.68

关键符号

trainer_send_weights(移除) start_vllm_server print_generations get_gpu_ids setup_engine gather_and_broadcast_weights_ipc WeightTransferEngine.shutdown

关键源码片段

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

唯一 core-logic 变更:从 worker ABC 删除 trainer_send_weights 抽象方法,并强化 WeightSource 双通道契约文档,是系列重构的收尾动作。

class WeightSource(ABC):
    """可重复迭代的权重来源,供 trainer 引擎消费,包含两条通道:    * metadata() —— 声明 (name, wire dtype, full shape),不传输数据。
    * 迭代通道 —— 逐个产出物化后的 (name, tensor)。
    两条通道必须逐元素一致:NCCL 后端用 metadata 计算接收缓冲区与
    packed 分块边界,再按迭代顺序发字节;两通道不一致会导致收发两侧
    流切分错位。
    """
​
    @abstractmethod
    def metadata(self) -> list[ParamMeta]:
        # 对 FSDP DTensor 而言全局 shape / dtype 本地可知,天然廉价;
        # 对必须物化才能得知 shape 的自定义 producer(如 Megatron 导出)
        # 首次调用可能昂贵,应缓存。
        raise NotImplementedError
​
    @abstractmethod
    def __iter__(self) -> Iterator[tuple[str, torch.Tensor]]:
        # 物化通常是集合操作(FSDP full_tensor()),所有 trainer rank
        # 必须以相同顺序同步迭代同一 source,否则死锁。
        raise NotImplementedError
​
​
class ModuleSource(WeightSource):
    """覆盖 module.named_parameters() 的通用实现,对普通稠密与
    FSDP 分片模块无需特判。
    """
​
    def __init__(self, module: torch.nn.Module) -> None:
        self._module = module
​
    def metadata(self) -> list[ParamMeta]:
        return [
            ParamMeta(name, p.dtype, tuple(p.shape))
            for name, p in self._module.named_parameters()
        ]
​
    def __iter__(self) -> Iterator[tuple[str, torch.Tensor]]:
        for name, param in self._module.named_parameters():
            yield name, materialize_full_tensor(param)
​
​
# 迁移全部完成后,worker ABC 从六个抽象方法收敛为五个:
# init_transfer_engine / start_weight_update / update_weights /
# receive_weights / shutdown。此处展示其收尾形态与紧随其后的
# 结构化传输客户端协议。
class WeightTransferEngine(ABC, Generic[TInitInfo, TUpdateInfo]):
    @abstractmethod
    def shutdown(self) -> None:
        """释放通信资源;worker 退出时应调用。"""
        raise NotImplementedError
​
​
@runtime_checkable
class VLLMWeightSyncClient(Protocol):
    """Trainer 侧操作 inference 引擎的传输无关控制面。    结构化 Protocol(PEP 544):任何实现这四个方法的对象都自动满足,
    RL 框架自带的 vLLM 包装器只需补四个转发方法即可作为 client 使用。
    """
​
    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: ...
examples/rl/rlhf_http_nccl.py dependency-wiring

标准模式的样板示例:自启 vllm serve(TP=2 + fp8 + --device-ids),并承接被删 rlhf_nccl.py 的唯一独有行为 fp8 量化。

def start_vllm_server() -> subprocess.Popen:
    """自行拉起 vllm serve 并阻塞等待 /health 就绪。    --device-ids 0,1 显式固定 server 占用的物理 GPU,训练侧固定取下一张
    卡 cuda:2,替代原先基于 cuda:{world_size} 的脆弱设备猜测;fp8 量化
    在 server 加载 dummy 权重时完成,训练侧只需发送 bf16 权重。
    """
    serve_args = [
        "vllm", "serve", MODEL_NAME,
        "--tensor-parallel-size", str(INFERENCE_TP_SIZE),
        "--device-ids", SERVER_DEVICE_IDS,
        "--quantization", "fp8",
        "--enforce-eager", "--load-format", "dummy",
        "--port", str(SERVER_PORT),
        "--weight-transfer-config", '{"backend": "nccl"}',
    ]
    env = os.environ.copy()
    # 该开关暴露 weight-transfer 控制面与 pause/resume 端点
    env["VLLM_SERVER_DEV_MODE"] = "1"
    print(f"[server] Launching: {' '.join(serve_args)}")
    proc = subprocess.Popen(serve_args, env=env, stdout=sys.stdout,
                            stderr=sys.stderr, start_new_session=True)
​
    # 轮询健康检查,900 秒超时;进程提前退出视为启动失败
    deadline = time.monotonic() + 900
    while True:
        if proc.poll() is not None:
            raise RuntimeError("vLLM server exited before becoming ready.")
        try:
            if requests.get(f"{BASE_URL}/health", timeout=5).status_code == 200:
                break
        except requests.RequestException:
            pass
        if time.monotonic() > deadline:
            raise RuntimeError("vLLM server failed to start in time.")
        time.sleep(2)
    print("[server] Ready.")
    return proc
​
​
def main():
    server_proc = start_vllm_server()
    try:
        # 训练侧位于 server 之后的那张 GPU,并作为 NCCL 通信组的 rank 0
        torch.accelerator.set_device_index(TRAINER_DEVICE)
        train_model = AutoModelForCausalLM.from_pretrained(
            MODEL_NAME, dtype=torch.bfloat16).to(TRAINER_DEVICE)
        client = OpenAI(base_url=f"{BASE_URL}/v1", api_key="EMPTY")
​
        # 同步前:dummy 权重应产生乱码
        outputs = generate_completions(client, MODEL_NAME, PROMPTS)
        print_generations("BEFORE weight sync (dummy weights):", PROMPTS, outputs)
​
        # NCCL 通信组 = trainer + 全部 inference worker
        world_size = get_world_size(BASE_URL) + 1
        engine = WeightTransferTrainerFactory.trainer_init(
            init_info=NCCLTrainerInitInfo(
                master_address=get_ip(), master_port=get_open_port(),
                world_size=world_size, rank=0, packed=True,
            ),
            client=HTTPVLLMWeightSyncClient(BASE_URL),
            source=ModuleSource(train_model),
        )
        pause_generation(BASE_URL)
        # 一次调用驱动 inference 侧 start/update/finish,与 NCCL 广播并发
        engine.send_weights()
        resume_generation(BASE_URL)
​
        # 同步后:应产生正常文本
        outputs_updated = generate_completions(client, MODEL_NAME, PROMPTS)
        print_generations("AFTER weight sync (real weights):", PROMPTS, outputs_updated)
    finally:
        server_proc.terminate()
        try:
            server_proc.wait(timeout=30)
        except subprocess.TimeoutExpired:
            server_proc.kill()

评论区精华

Fork PR 自动审查被禁用 other

claude[bot] 评论指出该 PR 来自 fork,自动 review 被禁用;维护者可评论 @claude review 触发一次性人工审查。整个 PR 无其他 review 评论。

结论:未进行 bot 审查;PR 经 CI(Buildkite #83457)验证后由 ywang96 合并。 · 已解决

风险与影响

  1. 公共 ABC 成员移除(breaking change):WeightTransferEngine.trainer_send_weights 从抽象方法变为不存在,任何仍按旧文档/旧示例调用该静态方法的第三方代码会立即 AttributeError。破坏面已被 PR 2/3 的 stub 提前收窄(stub 早已抛 NotImplementedError),但外部资料无法同步。
  2. 示例删除的迁移成本:rlhf_nccl.py / rlhf_ipc.py 是唯一接入 Buildkite 的示例,也是 Ray 集成的主要参照。删除后 Ray 路径只剩 rlhf_sparse_nccl.py 与 rlhf_async_new_apis.py 覆盖;多 rank Ray 编排的参照消失(HTTP 版覆盖同能力,但传输路径不同)。
  3. 自启 server 的稳定性:两个 HTTP 示例固定端口 8000,不能并行运行;900 秒(单卡)与 1800 秒(FSDP EP 大模型)启动超时在冷启动 + 大模型下载场景可能偏紧;start_new_session=True 子进程若脚本被 SIGKILL 可能残留孤儿进程,依赖 finally 的 terminate 兜底。
  4. FSDP 示例对睡眠/唤醒语义的依赖:/sleep?level=1 与 /wake_up?tags=weights|kv_cache|scheduling 的分级语义是示例正确性的前提,sleep mode 或 weight-transfer 端点行为变化时示例可能静默失败(无断言,只有输出对比)。
  5. fp8 行为迁移的一致性:rlhf_http_nccl.py 的 fp8 由 server 加载时量化实现,与删除的 rlhf_nccl.py(LLM 构造参数 quantization='fp8')路径不同,行为等价性依赖训练侧 bf16 → server fp8 的量化一致性,作者仅以端到端输出验证。
  6. 文档-代码同步风险:+1226/-1157 的大规模文档重写,存在与后续代码演进脱节的编辑维护成本。

用户/开发者:所有接触 RLHF 权重同步的开发者都将面对新的文档与示例结构;示例从"手动启动 server + 复杂 Ray 编排"变为一行命令可跑,极大降低上手成本;但使用 Ray 同步 API(LLM 实例)的集成者失去两个现成参照。
系统:运行时行为无变化(仅删除死代码与 stub,净约 90 行),不触碰推理/训练热路径;对 vLLM 包体积与启动无影响。
团队:完成 4-PR 架构迁移的最后一块拼图,weight-transfer 的架构叙事统一到 trainer/worker 双引擎 + client 适配层;CI 的示例覆盖从被删文件平滑切到 HTTP 版本。
影响程度:中。范围集中在 RL 示例与文档,但 API 移除属于对外的静默 breaking change。

公共 ABC 方法移除 示例删除与迁移 固定端口与启动超时 sleep/wake 语义依赖 文档量大

关联 Issue

#48042 [rl] Stateful Trainer Send: New Abstractions [1/N]
#48981 [rl] Stateful Trainer Send: IPC [2/N]

完整报告

参与讨论