param_sync.py — 实现了参数同步器——负责在训练器(Actor)和推理器(Rollout)之间同步模型参数¶
文件路径: verl/experimental/fully_async_policy/param_sync.py
文件概述¶
实现了参数同步器——负责在训练器(Actor)和推理器(Rollout)之间同步模型参数。在全异步训练中,训练和推理使用不同的 GPU,训练更新的参数需要定期同步到推理服务器。
关键代码讲解¶
1. ParameterSynchronizer¶
@ray.remote
class ParameterSynchronizer:
def __init__(self, config, trainer, rollouter, mq):
self.trainer = trainer
self.rollouter = rollouter
self.mq_client = mq
self.actor_wg = ray.get(trainer.get_actor_wg.remote()) # 训练侧 Worker Group
self.rollout_wg = ray.get(rollouter.get_rollout_wg.remote()) # 推理侧 Worker Group
# 使用 CheckpointEngineManager 进行权重传输
self.checkpoint_manager = CheckpointEngineManager(
config=checkpoint_engine_config,
trainer=self.actor_wg,
replicas=replicas,
)
2. 参数同步流程 - sync_weights¶
def sync_weights(self, version, validate=False, global_steps=0, use_trainer_do_validate=False):
self.current_version = version
# 1. 暂停 Rollouter(等待当前 rollout 完成或取消)
ray.get(self.rollouter.pause.remote())
# 2. 更新 MessageQueue 的参数版本
self.mq_client.update_param_version_sync(version)
# 3. 同步权重(通过 CheckpointEngine)
self.checkpoint_manager.update_weights(global_steps)
# 4. 异步更新 Rollouter 的参数版本(可能包含验证)
self.wait_last_update = self.rollouter.update_param_version.remote(
version, validate, global_steps, use_trainer_do_validate
)
# 5. 恢复 Rollouter(依赖于 update 完成)
self.wait_last_resume = self.rollouter.resume.remote(self.wait_last_update)
3. 同步流程图¶
Trainer 训练 N 步完成
│
▼
ParameterSynchronizer.sync_weights(version)
│
├── 1. Rollouter.pause()
│ ├── 取消进行中的 rollout(如果启用 partial_rollout)
│ ├── 等待活跃任务完成
│ └── 清除 KV Cache
│
├── 2. MessageQueue.update_param_version(version)
│
├── 3. CheckpointManager.update_weights()
│ └── Actor GPU → Rollout GPU 权重传输
│
├── 4. Rollouter.update_param_version(version)(异步)
│ └── 可能触发验证
│
└── 5. Rollouter.resume()(异步,依赖步骤4)
4. 等待上一轮同步¶
def wait_last_valid(self):
"""等待上一轮同步和验证完成"""
if self.wait_last_update:
ray.get(self.wait_last_update)
if self.wait_last_resume:
ray.get(self.wait_last_resume)
if self.validate_task:
ray.get(self.validate_task)
核心类/函数列表¶
| 名称 | 类型 | 说明 |
|---|---|---|
ParameterSynchronizer |
Ray Remote Actor | 参数同步器 |
sync_weights() |
方法 | 执行完整的参数同步流程 |
wait_last_valid() |
方法 | 等待上一轮同步完成 |
rollouter_save_checkpoint() |
方法 | 触发 Rollouter 保存检查点 |
与其他模块的关系¶
- 由
FullyAsyncTaskRunner(fully_async_main.py)创建 - 被
FullyAsyncTrainer在训练步后调用sync_weights - 调用
FullyAsyncRollouter的pause/resume/update_param_version - 使用
CheckpointEngineManager进行实际的权重传输
小结¶
ParameterSynchronizer 是全异步训练中训练和推理之间的"桥梁"。它协调了暂停推理 -> 同步权重 -> 恢复推理的完整流程,确保推理服务器始终使用最新的模型参数。异步验证支持让权重同步和模型验证可以重叠执行,减少等待时间。