跳转至

fully_async_trainer.py — 实现了全异步 PPO 训练器——全异步架构中的"消费者"

文件路径: verl/experimental/fully_async_policy/fully_async_trainer.py

文件概述

实现了全异步 PPO 训练器——全异步架构中的"消费者"。它从 MessageQueue 中获取 Rollouter 生成的样本,执行完整的 PPO 训练流程(计算奖励、优势估计、更新 Critic 和 Actor),并定期触发参数同步。

关键代码讲解

1. 类定义

@ray.remote(num_cpus=10)
class FullyAsyncTrainer(SeparateRayPPOTrainer):
    def __init__(self, config, tokenizer, ...):
        # 训练状态
        self.global_steps = 1
        self.current_param_version = 0
        self.local_trigger_step = 1  # 本地训练步数(用于控制参数同步频率)

        # 所需样本数 = ppo_mini_batch_size * require_batches
        self.required_samples = config.actor_rollout_ref.actor.ppo_mini_batch_size * self.require_batches

        # 指标聚合器(因为全异步模式下每步指标需要聚合多个训练步的结果)
        self.metrics_aggregator = MetricsAggregator(total_gpus=total_gpus)

2. 训练主循环

async def fit(self):
    self.logger = Tracking(...)

    # 先获取初始验证数据
    self._log_validation_data()

    # 持续训练,直到收到终止信号
    while True:
        try:
            await self.fit_step()
        except TrainingStopException:
            break

    # 最终参数同步和验证
    ray.get(self.param_synchronizer.wait_last_valid.remote())
    self._log_validation_data()
    self.progress_bar.close()
    self._fit_save_checkpoint()

3. 单步训练 - fit_step

async def fit_step(self, batch_dict=None):
    with marked_timer("step", self.timing_raw):
        batch = self._fit_generate(None)        # 从队列获取数据
        batch = self._fit_compute_reward(batch)   # 计算奖励
        batch = self._fit_compute_log_prob(batch)  # 计算 log prob
        batch = self._fit_compute_ref_log_prob(batch)  # 计算参考策略 log prob
        batch = self._fit_compute_critic(batch)   # Critic 前向
        batch = self._fit_compute_advantage(batch) # 计算优势
        batch = self._fit_update_critic(batch)    # 更新 Critic
        batch = self._fit_update_actor(batch)     # 更新 Actor
        await self._fit_update_weights()          # 触发参数同步
        self._fit_dump_data(batch)                # 保存数据

    self._fit_save_checkpoint()
    self._fit_collect_metrics(batch)

4. 从队列获取样本

def _get_samples_from_queue(self):
    queue_samples = []
    while len(queue_samples) < self.required_samples:
        sample, queue_len = self.message_queue_client.get_sample_sync()
        if sample is None:  # 终止信号
            break
        queue_samples.append(sample)

    # 反序列化
    queue_samples = [ray.cloudpickle.loads(x) for x in queue_samples]

    # 组装成训练 batch
    batch = assemble_batch_from_rollout_samples(queue_samples, self.tokenizer, self.config, ...)
    return 0, batch

5. 参数同步控制

每训练 trigger_parameter_sync_step 步后触发一次参数同步:

async def _trigger_parameter_sync_after_step(self, validate=False):
    if self.local_trigger_step < self.trigger_parameter_sync_step and not validate:
        self.local_trigger_step += 1
        return  # 还不到同步时间

    # 触发参数同步
    self.current_param_version += 1
    self.local_trigger_step = 1

    # 等待上一轮同步完成
    ray.get(self.param_synchronizer.wait_last_valid.remote())
    # 启动新的同步
    ray.get(self.param_synchronizer.sync_weights.remote(self.current_param_version, ...))

6. MIS(Multiple Importance Sampling)支持

当 trigger_parameter_sync_step > 1 时,一个参数版本可能执行多步训练。为保证理论正确性,需要使用第一步的参数计算 old_log_prob:

def _compute_old_log_prob(self, batch):
    if self.local_trigger_step == 1:
        # 第一步:保存当前参数到 CPU
        self.actor_rollout_wg.save_model_to_cpu(1)
        old_log_prob, mfu = super()._compute_old_log_prob(batch)
    else:
        # 后续步:恢复第一步的参数计算 old_log_prob,再恢复当前参数
        self.actor_rollout_wg.save_model_to_cpu(self.local_trigger_step)
        self.actor_rollout_wg.restore_model_from_cpu(1)
        old_log_prob, mfu = super()._compute_old_log_prob(batch)
        self.actor_rollout_wg.restore_model_from_cpu(self.local_trigger_step)
    return old_log_prob, mfu

7. 陈旧样本统计

def _collect_metrics_from_samples(self, batch, metrics):
    samples_param_versions = batch.meta_info["rollout_param_versions"]
    # 统计使用旧版本参数生成的样本数
    stale_count = sum(1 for v in samples_param_versions if self.current_param_version - v >= 1)
    self.stale_samples_processed += stale_count

核心类/函数列表

名称 类型 说明
FullyAsyncTrainer Ray Remote 类 全异步 PPO 训练器
fit() 方法 训练主循环
fit_step() 方法 单步训练(获取数据 -> PPO更新 -> 参数同步)
_get_samples_from_queue() 方法 从 MessageQueue 获取样本
_trigger_parameter_sync_after_step() 方法 触发参数同步
_compute_old_log_prob() 方法 MIS 支持的 old_log_prob 计算
TrainingStopException 异常 训练终止信号

与其他模块的关系

  • 继承自 SeparateRayPPOTrainer(separation/ray_trainer.py)
  • 使用 MessageQueueClient(message_queue.py)获取样本
  • 使用 ParameterSynchronizer(param_sync.py)同步参数
  • 使用 MetricsAggregator(detach_utils.py)聚合指标
  • 使用 assemble_batch_from_rollout_samples(detach_utils.py)组装 batch

小结

FullyAsyncTrainer 是全异步训练的"消费者"。它持续从 MessageQueue 获取样本进行 PPO 训练,通过 trigger_parameter_sync_step 控制参数同步频率,通过 MIS 机制确保多步训练的理论正确性。陈旧样本统计帮助监控异步训练中的数据新鲜度。