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 机制确保多步训练的理论正确性。陈旧样本统计帮助监控异步训练中的数据新鲜度。