modelscope / modelscope/ms-swift
[RFC] 实现可扩展、解耦的分布式PPO训练
Nobody has claimed this yet.
- 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流程:
- 同步权重: 调用
WeightSynchronizer将Training Engine中ActorCritic的最新权重同步给所有RolloutWorker。 - 并行采集: 向所有
RolloutWorker下发collect任务,并异步等待它们返回经验数据句柄。 - 数据后处理: 收集所有经验数据,在CPU上计算GAE和Returns。
- 模型训练: 将处理好的训练批次发送给
Training Engine执行update_policy。 - 日志记录: 收集训练统计信息并记录。
- 循环: 重复以上步骤。
2.3 Ray资源管理
Ray的资源管理体现在声明、配置、分配三个层面:
-
资源声明: 在
TrainingEngine和RolloutWorker的@ray.remote装饰器中,声明其需要GPU资源。 -
资源配置: 用户可以灵活指定
Training Engine需要多少个GPU,以及Rollout Engine需要启动多少个worker。 -
资源分配: 动态地为
TrainingEngine和RolloutWorker实例请求具体的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): 初始化函数,根据配置选择并实例化相应的训练后端和角色对象(如ActorCritic或MegatronActorCritic)。
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):- 描述: 同时实现了
ActorRole和CriticRole的接口,适合PPO中共享骨干网络的情况。
- 描述: 同时实现了
五、Rollout Engine 设计
5.1 角色与职责
Rollout Engine 是系统的分布式数据采集前端,由一组并行的RolloutWorker Ray Actor构成。
- 并行数据采集: 通过横向扩展
RolloutWorker数量来提升数据采集吞吐量。 - 模型推理: 每个worker内部署了
ActorCritic、Reference和Reward等角色对象的推理副本。 - 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方法的核心逻辑委托给该后端执行。后端负责协调其内部持有的ActorCritic、Reference和Reward角色对象来完成数据生成。
六、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 Engine和Rollout 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
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- 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