跳转至

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 是全异步训练中训练和推理之间的"桥梁"。它协调了暂停推理 -> 同步权重 -> 恢复推理的完整流程,确保推理服务器始终使用最新的模型参数。异步验证支持让权重同步和模型验证可以重叠执行,减少等待时间。