modelscope / modelscope/ms-swift

[RFC] 实现可扩展、解耦的分布式PPO训练

Open
#6,361 1 comment 2 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

stale
Dominant language
Python
Stars
15.7k
Forks
1.7k
Avg merge
1d 16h
Merged PRs (30d)
136

Description

描述
实现一个分布式PPO训练模块,具备解耦性、可扩展性和灵活性,可以扩展为其他RLHF算法的训练。

一、总体架构设计

1.1 架构图

+-------------------------------------------------------------------+
|                           Controller                              |
| (驱动进程)                                                          |
| - 管理配置与资源                                                     |
| - 编排训练循环                                                      |
+-------------------------------------------------------------------+
       | 1. sync()      ^ 7. 训练状态          | 4. collect()     ^ 6. 返回数据句柄
       |                |                    |                   | (DataHandle)
       v                |                    v                   |
+--------------------------------+       +---------------------------------------+
|       WeightSynchronizer       |       |                 Data Bus              |
| (位于'transfer'模块的工具库)       |       | (位于'transfer'模块, 底层为Ray对象存储)  |
| - API: sync()                  |       | - API: put(), get()                   |
+--------------------------------+       +---------------------------------------+
       | 2. get_weights()  | 3. set_weights()        ^ 5. 训练批次句柄
       |                   |_____________________    | (DataHandle)
       v                                        v    |
+-------------------------------------------+  +-------------------------------------------+
|             Training Engine               |  |             Rollout Engine                |
| (Ray Actor组 - GPU密集型)                   |  | (Ray Actor组 - GPU密集型)                  |
| API: `update_policy()`, `get_weights()`   |  |-------------------------------------------|
|                                           |  | +---------- RolloutWorker 1 ----------+   |
| +--------- TrainingBackend -----------+   |  | | (Ray Actor)                         |   |
| | (例如 FSDPBackend)                   |   |  | | API: `collect()`, `set_weights()`   |   |
| |                                     |   |  | |                                     |   |
| |   +-----------------------------+   |   |  | | +------- InferenceBackend -------+  |   |
| |   |      权威的ActorCritic角色    |   |   |  | | (例如 HFBackend, vLLMBackend)     |  |   |
| |   +-----------------------------+   |   |  | | |                                |  |   |    
| |                                     |   |  | | |   +-------------------------+  |  |   |
| +-------------------------------------+   |  | | |   | ActorCritic角色 (副本)    |  |  |   |
|                                           |  | | |   +-------------------------+  |  |   |
+-------------------------------------------+  | | |   +-------------------------+  |  |   |
                                               | | |   |     Reference角色        |  |  |   |
                                               | | |   +-------------------------+  |  |   |
                                               | | |   +-------------------------+  |  |   |
                                               | | |   |       Reward角色         |  |  |   |
                                               | | |   +-------------------------+  |  |   |
                                               | | |                                |  |   |
                                               | | +-------------------------------+|  |   |
                                               | |                                     |   |
                                               | +-------------------------------------+   |
                                               | |         (... 更多 Workers ...)       |   |
                                               +-------------------------------------------+

1.2 数据流

sequenceDiagram
    participant C as Controller
    participant WS as WeightSynchronizer
    participant TE as Training Engine
    participant RE as Rollout Engine
    participant DB as Data Bus

    C->>WS: 1. sync() 指令 (启动同步)
    WS->>TE: 2. get_weights() (请求权重句柄)
    TE-->>WS: 权重句柄 (DataHandle)
    WS->>RE: 3. set_weights(Handle) (广播更新)

    Note over WS, RE: WS等待所有Rollout Worker完成set_weights()

    C->>RE: 4. collect(Prompts Handle) (开始采样)
    Note over RE: Rollout Worker使用最新的Actor/Critic进行推理
    RE->>DB: DataBus.put(rollout数据)
    RE-->>C: 5. rollout数据句柄 (DataHandle)

    C->>C: 5b. ray.get(Handles) & 计算 GAE/Returns 
    C->>DB: DataBus.put(训练批次数据)
    C->>TE: 6. update_policy(Batch Handle) (启动训练)

    Note over TE: Training Engine Backend开始多轮PPO优化 (GPU密集型)

    TE-->>C: 7. 训练统计信息 (Losses, KL, etc.)
    C->>C: 8. 记录日志, 开始下一轮迭代

1.2 代码结构

swift/
│
├── distributed_rlhf/               
│   │
│   ├── controller/
│   │   └── ppo_controller.py # PPO算法的中央控制器实现
│   │
│   ├── engines/              # 服务化引擎的定义
│   │   ├── training_engine.py   # TrainingEngine基类
│   │   ├── rollout_engine.py # RolloutWorker基类
│   │   ├── training_backends/   # TrainingEngine的后端实现
│   │   │   ├── base.py       # 后端抽象基类
│   │   │   ├── fsdp_backend.py
│   │   │   └── megatron_backend.py
│   │   └── rollout_backends/ # RolloutEngine的后端实现
│   │       ├── base.py       # 后端抽象基类
│   │       ├── hf_backend.py
│   │       ├── vllm_backend.py
│   │       └── sglang_backend.py
│   │
│   ├── roles/                # 模型角色抽象层
│   │   ├── base.py           # ActorRole, CriticRole等接口
│   │   ├── actor_critic.py   # ActorCritic 的标准实现
│   │   ├── reference.py
│   │   └── reward.py
│   │
│   ├── transfer/             # 框架的核心传输抽象与工具
│   │   ├── data_transfer.py
│   │   └── weight_transfer.py
│   │
│   ├── algorithms/           # 算法相关的特定计算模块
│   │   └── ppo_utils.py      # PPO特有的计算逻辑, 如GAE
│   │
│   └── utils/                # 通用工具模块

二、Controller设计

2.1 角色与职责

Controller作为中央协调器,其核心职责是编排分布式组件来执行RL算法。它不存储模型权重或大规模经验数据,而是管理这些资源在专用服务间的流动。

  • 算法逻辑实现: PPO的 "采集-处理-训练" 循环逻辑在此实现。
  • 任务分发: 向 Rollout Engine 发起数据采集任务,向 Training Engine 发起训练任务。
  • 资源管理: 负责根据配置文件,动态地分配Ray集群资源给各个服务。
  • 监控与管理: 负责日志记录、指标收集和模型检查点的触发。

2.2 PPO主训练循环

Controller是PPO主训练循环的实现者,其逻辑遵循标准的PPO流程:

  1. 同步权重: 调用WeightSynchronizerTraining EngineActorCritic的最新权重同步给所有RolloutWorker
  2. 并行采集: 向所有RolloutWorker下发collect任务,并异步等待它们返回经验数据句柄。
  3. 数据后处理: 收集所有经验数据,在CPU上计算GAE和Returns。
  4. 模型训练: 将处理好的训练批次发送给Training Engine执行update_policy
  5. 日志记录: 收集训练统计信息并记录。
  6. 循环: 重复以上步骤。

2.3 Ray资源管理

Ray的资源管理体现在声明、配置、分配三个层面:

  1. 资源声明: 在TrainingEngineRolloutWorker@ray.remote装饰器中,声明其需要GPU资源。

  2. 资源配置: 用户可以灵活指定Training Engine需要多少个GPU,以及Rollout Engine需要启动多少个worker。

  3. 资源分配: 动态地为TrainingEngineRolloutWorker实例请求具体的GPU数量,从而触发Ray调度器进行物理资源分配。

三、Training Engine 设计

3.1 角色与职责

Training Engine 是框架的训练中心,封装了所有与模型训练相关的复杂性,类似Trainer的概念。

  • 权威权重存储: 持有一个可训练的ActorCritic角色对象的权威副本。
  • 分布式训练执行: 在多GPU环境下执行模型的训练步骤。
  • 后端抽象: 对上层隐藏了FSDP、Megatron-LM等分布式训练后端的实现细节。
  • 状态管理: 负责模型的加载与保存。

3.2 API 定义

TrainingEngine Class
  • __init__(self, role_config: Dict, train_config: Dict): 初始化函数,根据配置选择并实例化相应的训练后端和角色对象(如ActorCriticMegatronActorCritic)。
TrainingEngine.get_weights()
  • 函数签名: get_weights(self) -> DataHandle
  • 描述: 获取当前ActorCritic角色对象的权重句柄,用于分发给Rollout Engine。这是一个非阻塞调用。
TrainingEngine.update_policy()
  • 函数签名: update_policy(self, batch_handle: DataHandle, algorithm_config: Dict) -> Dict[str, float]
  • 描述: 通用策略更新接口。使用一批经验数据执行一次完整的训练步骤,更新模型,并返回训练统计信息。algorithm_config 传入算法相关的超参数(如PPO的clip_epsilon或DPO的beta),使得此API可被多种算法复用。
TrainingEngine.save_checkpoint()
  • 函数签名: save_checkpoint(self, path: str) -> None
  • 描述: 将模型的当前状态(包括权重和优化器状态)持久化到存储。
TrainingEngine.load_checkpoint()
  • 函数签名: load_checkpoint(self, path: str) -> None
  • 描述: 从一个检查点加载模型状态。

3.3 内部架构:后端抽象

Training Engine 内部持有一个训练后端实例(如FSDPBackend),所有API调用都被委托给该后端执行。后端负责管理具体的分布式训练逻辑和它所持有的Role对象(如ActorCritic)。这种双层结构实现了关注点分离Role负责定义计算逻辑,而Backend负责管理分布式环境和训练流程

  • 对于Megatron: MegatronBackend会管理一个MegatronActorCritic实例。该角色对象的构建和训练步骤完全遵循Megatron的范式,但其接口与标准ActorCritic保持一致,从而对上层透明。

四、Roles (模型角色) 设计

4.1 角色与职责

Roles模块将算法中涉及的各个模型组件,从底层的神经网络实现中抽象为具有明确职责的逻辑角色。它专注于定义纯粹的前向传播计算逻辑。而“怎么更新”的问题则由TrainingEngine Backend负责。

  • Roles 模块职责:

    • 定义神经网络架构。
    • 实现无状态的前向传播方法,如forward(用于计算log_prob)和compute_values
    • 不包含任何优化器、损失函数计算、梯度反向传播或分布式包裹(如FSDP)的代码。
  • TrainingEngine Backend 模块职责:

    • 管理分布式环境。
    • 持有并包裹Role对象。
    • 实现update_policy方法,该方法内部会调用Role对象进行前向计算,然后计算PPO损失,并执行完整的训练步骤(backward, optimizer.step)。

4.2 API 定义

ActorRole Interface
  • forward(self, data: Dict) -> Dict:
    • 描述: 核心计算方法。接收批次数据,计算并返回包含log_probs等信息的字典。
  • generate(self, data: Dict) -> torch.Tensor:
    • 描述: 根据给定的prompts生成文本序列。
CriticRole Interface
  • compute_values(self, data: Dict) -> torch.Tensor:
    • 描述: 接收批次数据,计算并返回每个状态的价值(values)张量。
核心角色实现
  • ActorCritic(ActorRole, CriticRole):
    • 描述: 同时实现了ActorRoleCriticRole的接口,适合PPO中共享骨干网络的情况。

五、Rollout Engine 设计

5.1 角色与职责

Rollout Engine 是系统的分布式数据采集前端,由一组并行的RolloutWorker Ray Actor构成。

  • 并行数据采集: 通过横向扩展RolloutWorker数量来提升数据采集吞吐量。
  • 模型推理: 每个worker内部署了ActorCriticReferenceReward等角色对象的推理副本。
  • PPO数据生成: 执行完整的PPO采样流程。
  • 后端抽象: 对上层屏蔽了Hugging Face Transformers、vLLM等推理后端的差异。

5.2 API 定义

RolloutWorker Class
  • __init__(self, role_config: Dict, generation_config: Dict): 初始化函数,根据配置选择并实例化相应的推理后端和所有角色对象。
RolloutWorker.set_weights()
  • 函数签名: set_weights(self, weights_handle: DataHandle) -> None
  • 描述: 接收最新的权重句柄,并更新内部的ActorCritic角色副本。
RolloutWorker.collect()
  • 函数签名: collect(self, prompts_handle: DataHandle) -> DataHandle
  • 描述: 根据给定的prompts,执行一次完整的PPO数据采集流程,并返回采集到的数据句柄。

5.3 内部架构:后端抽象

RolloutWorker 内部持有一个推理后端实例(如vLLMBackend),collect方法的核心逻辑委托给该后端执行。后端负责协调其内部持有的ActorCriticReferenceReward角色对象来完成数据生成。

六、Data Bus 设计

6.1 角色与职责

Data Bus 是一个逻辑概念层,充当框架中所有服务之间数据交换的数据总线。它通过传递轻量级的数据句柄(DataHandle)而非数据本身,来解耦数据存储与计算。

6.2 API 定义

DataBus.put()
  • 函数签名: put(data: Any) -> DataHandle
  • 描述: 将一个Python对象放入共享存储中,并立即返回一个非阻塞的句柄。
DataBus.get()
  • 函数签名: get(handle: DataHandle) -> Any
  • 描述: 通过句柄从共享存储中取回实际对象,这是一个阻塞操作。
DataBus.wait()
  • 函数签名: wait(handles: List[DataHandle], num_returns: int = 1) -> (List[DataHandle], List[DataHandle])
  • 描述: 阻塞等待,直到列表中至少num_returns个句柄就绪。

七、WeightSynchronizer 设计

7.1 角色与职责

WeightSynchronizer 是一个由Controller调用的高级工具,负责在Training EngineRollout Engine之间高效、可靠地同步模型权重,将同步的复杂性从算法逻辑中抽象出来。

7.2 API 定义

WeightSynchronizer.sync()
  • 函数签名: sync(self, source: ActorHandle, destinations: List[ActorHandle]) -> None
  • 参数:
    • source: 源服务句柄(Training Engine),必须实现get_weights() -> DataHandle
    • destinations: 目标服务句柄列表(RolloutWorkers),每个都必须实现set_weights(weights_handle: DataHandle) -> None
  • 描述: 执行一次完整的权重同步操作。这是一个阻塞调用,确保所有RolloutWorker都更新完成后才返回。

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start by reviewing the proposed entry points in distributed_rlhf/controller/ppo_controller.py, engines/training_engine.py, engines/rollout_engine.py, transfer/data_transfer.py, and transfer/weight_transfer.py. Trace the controller's sync, collect, and update_policy flow, then define completion against the listed interfaces and the distributed PPO architecture; no test path is named in the issue.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
backend, distributed-systems, machine-learning
Issue type
Feature
Difficulty
5/5
Estimated time
Over a week
Activity status
Stale
Clarity
Needs clarification
Newbie friendliness
20/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.