ray_trainer.py — 这是 整个 verl 框架最核心的文件¶
文件概述¶
模块路径: verl.trainer.ppo.ray_trainer
这是 整个 verl 框架最核心的文件,实现了基于 Ray 的分布式 PPO 训练器 RayPPOTrainer。它作为"单控制器"(single controller),在一个驱动进程上协调多个分布式 Worker Group,完成 PPO 训练的完整流程。
该文件约 1600 行代码,包含:
- KL 惩罚计算
- 优势估计
- 完整的分布式训练循环(fit())
- Worker Group 初始化
- Checkpoint 保存/加载
- 验证和指标收集
在训练流程中的位置¶
这是训练流程的"心脏",由 main_ppo.py 的 TaskRunner.run() 创建并调用 fit() 启动训练。
关键代码讲解¶
1. 顶层辅助函数¶
KL 惩罚¶
def apply_kl_penalty(data: DataProto, kl_ctrl, kl_penalty="kl"):
"""在 token 级别的奖励上施加 KL 惩罚"""
# 计算当前策略和参考策略之间的 KL 散度
kld = core_algos.kl_penalty(
data.batch["old_log_probs"],
data.batch["ref_log_prob"],
kl_penalty=kl_penalty
)
kld = kld * response_mask
beta = kl_ctrl.value # KL 系数
# 原始奖励 - KL 惩罚
token_level_rewards = token_level_scores - beta * kld
# 更新 KL 控制器(自适应模式下会调整 beta)
kl_ctrl.update(current_kl=current_kl, n_steps=batch_size)
data.batch["token_level_rewards"] = token_level_rewards
KL 惩罚的作用是防止策略偏离参考策略太远,是 RLHF 中的关键技术。
优势估计入口¶
def compute_advantage(data, adv_estimator, gamma=1.0, lam=1.0, ...):
"""根据选择的优势估计器计算优势值"""
if adv_estimator == AdvantageEstimator.GAE:
advantages, returns = core_algos.compute_gae_advantage_return(
token_level_rewards=data.batch["token_level_rewards"],
values=data.batch["values"], # 来自 Critic
response_mask=data.batch["response_mask"],
gamma=gamma, lam=lam,
)
elif adv_estimator == AdvantageEstimator.GRPO:
advantages, returns = core_algos.compute_grpo_outcome_advantage(
token_level_rewards=data.batch["token_level_rewards"],
response_mask=grpo_calculation_mask,
index=data.non_tensor_batch["uid"], # 按 prompt 分组
...
)
else:
# 通过注册表动态获取(REINFORCE++, RLOO, OPO 等)
adv_estimator_fn = core_algos.get_adv_estimator_fn(adv_estimator)
advantages, returns = adv_estimator_fn(**adv_kwargs)
data.batch["advantages"] = advantages
data.batch["returns"] = returns
这个函数根据配置选择不同的优势估计方法。优势估计是 RL 训练质量的关键。
2. RayPPOTrainer 类¶
初始化¶
class RayPPOTrainer:
def __init__(self, config, tokenizer, role_worker_mapping, resource_pool_manager, ...):
self.tokenizer = tokenizer
self.config = config
self.hybrid_engine = config.actor_rollout_ref.hybrid_engine
assert self.hybrid_engine, "Currently, only support hybrid engine"
# 判断需要哪些组件
self.use_reference_policy = need_reference_policy(self.config)
self.use_rm = need_reward_model(self.config)
self.use_critic = need_critic(self.config)
# LoRA 模式下,Ref Policy 就是没有 LoRA 的 Actor
self.ref_in_actor = lora_rank > 0 or config.actor_rollout_ref.model.get("lora_adapter_path") is not None
# KL 控制器
if self.config.algorithm.use_kl_in_reward:
self.kl_ctrl_in_reward = core_algos.get_kl_controller(self.config.algorithm.kl_ctrl)
self._create_dataloader(train_dataset, val_dataset, collate_fn, train_sampler)
init_workers() - Worker Group 创建¶
def init_workers(self):
"""创建分布式训练的 Worker Group"""
# 1. 创建资源池
self.resource_pool_manager.create_resource_pool()
# 2. 为每个角色创建 RayClassWithInitArgs
actor_rollout_cls = RayClassWithInitArgs(
cls=self.role_worker_mapping[actor_role],
config=self.config.actor_rollout_ref,
role=str(actor_role),
)
# 3. 创建 Critic Worker(如果需要)
if self.use_critic:
critic_cls = RayClassWithInitArgs(cls=self.role_worker_mapping[Role.Critic], ...)
# 4. 使用 create_colocated_worker_cls 创建共置 Worker
# (多个角色共享同一组 GPU)
for resource_pool, class_dict in self.resource_pool_to_cls.items():
worker_dict_cls = create_colocated_worker_cls(class_dict=class_dict)
wg_dict = self.ray_worker_group_cls(
resource_pool=resource_pool,
ray_cls_with_init=worker_dict_cls,
)
spawn_wg = wg_dict.spawn(prefix_set=class_dict.keys())
all_wg.update(spawn_wg)
# 5. 提取各角色的 Worker Group
self.actor_rollout_wg = all_wg[str(actor_role)]
self.actor_rollout_wg.init_model()
if self.use_critic:
self.critic_wg = all_wg[str(Role.Critic)]
self.critic_wg.init_model()
# 6. 创建 Reward Loop Manager 和 Checkpoint Manager
self.reward_loop_manager = RewardLoopManager(config=self.config, ...)
self.async_rollout_manager = AgentLoopManager.create(...)
self.checkpoint_manager = CheckpointEngineManager(...)
这里的 create_colocated_worker_cls 是关键设计:它允许多个角色(如 Actor 和 Critic)共享同一组 GPU,在需要时切换模型,以节省 GPU 内存。
3. fit() - 核心训练循环¶
def fit(self):
"""PPO 训练主循环"""
logger = Tracking(...)
self.global_steps = 0
# 加载 checkpoint 并同步权重
self._load_checkpoint()
self.checkpoint_manager.update_weights(self.global_steps)
# 训练前验证
if self.config.trainer.get("val_before_train", True):
val_metrics = self._validate()
logger.log(data=val_metrics, step=self.global_steps)
for epoch in range(current_epoch, self.config.trainer.total_epochs):
for batch_dict in self.train_dataloader:
metrics = {}
timing_raw = {}
batch = DataProto.from_single_dict(batch_dict)
Step 1: 生成回复¶
with marked_timer("gen", timing_raw):
gen_batch_output = self.async_rollout_manager.generate_sequences(gen_batch_output)
self.checkpoint_manager.sleep_replicas()
调用推理引擎生成回复。生成完成后,让推理副本休眠以释放 GPU 显存。
Step 2: 计算奖励¶
with marked_timer("reward", timing_raw):
if self.use_rm and "rm_scores" not in batch.batch.keys():
batch_reward = self._compute_reward_colocate(batch)
batch = batch.union(batch_reward)
reward_tensor, reward_extra_infos_dict = extract_reward(batch)
Step 3: 计算旧策略概率¶
with marked_timer("old_log_prob", timing_raw):
old_log_prob, old_log_prob_mfu = self._compute_old_log_prob(batch)
# 计算熵(用于正则化)
entropys = old_log_prob.batch["entropys"]
batch = batch.union(old_log_prob)
Step 4: 计算参考策略概率(如果需要 KL 惩罚)¶
if self.use_reference_policy:
with marked_timer(str(Role.RefPolicy), timing_raw):
ref_log_prob = self._compute_ref_log_prob(batch)
batch = batch.union(ref_log_prob)
Step 5: 计算价值函数(如果使用 Critic)¶
if self.use_critic:
with marked_timer("values", timing_raw):
values = self._compute_values(batch)
batch = batch.union(values)
Step 6: 计算优势和奖励¶
with marked_timer("adv", timing_raw):
batch.batch["token_level_scores"] = reward_tensor
if self.config.algorithm.use_kl_in_reward:
batch, kl_metrics = apply_kl_penalty(batch, ...)
else:
batch.batch["token_level_rewards"] = batch.batch["token_level_scores"]
# Rollout 校正(处理 off-policy 问题)
if rollout_corr_config is not None and "rollout_log_probs" in batch.batch:
batch, is_metrics = compute_rollout_correction_and_add_to_batch(batch, ...)
# 计算优势
batch = compute_advantage(batch, adv_estimator=..., gamma=..., lam=...)
Step 7: 更新 Critic 和 Actor¶
if self.use_critic:
with marked_timer("update_critic", timing_raw):
critic_output = self._update_critic(batch)
if self.config.trainer.critic_warmup <= self.global_steps:
with marked_timer("update_actor", timing_raw):
actor_output = self._update_actor(batch)
critic_warmup 允许 Critic 先训练几步再开始更新 Actor,使得价值估计更稳定。
Step 8: 同步权重、保存和验证¶
# 保存 checkpoint
if self.config.trainer.save_freq > 0 and (is_last_step or ...):
self._save_checkpoint()
# 同步权重到推理引擎
with marked_timer("update_weights", timing_raw):
self.checkpoint_manager.update_weights(self.global_steps)
# 验证
if self.config.trainer.test_freq > 0 and (is_last_step or ...):
val_metrics = self._validate()
metrics.update(val_metrics)
# 记录所有指标
metrics.update(compute_data_metrics(batch=batch, use_critic=self.use_critic))
metrics.update(compute_timing_metrics(batch=batch, timing_raw=timing_raw))
metrics.update(compute_throughout_metrics(batch=batch, timing_raw=timing_raw, n_gpus=n_gpus))
logger.log(data=metrics, step=self.global_steps)
4. 序列长度均衡¶
def _balance_batch(self, batch, metrics, ...):
"""重新排列数据,使每个 DP rank 获得相似的总 token 数"""
global_seqlen_lst = batch.batch["attention_mask"].view(batch_size, -1).sum(-1)
dp_size = self._get_dp_size(self.actor_rollout_wg, "actor")
global_partition_lst = get_seqlen_balanced_partitions(workload_lst, k_partitions=dp_size, ...)
batch.reorder(global_idx)
这是性能优化的关键:如果不均衡,某些 rank 的序列特别长,其他 rank 就要等待,造成 GPU 空闲。
核心类/函数列表¶
| 名称 | 类型 | 作用 |
|---|---|---|
apply_kl_penalty() |
function | 在奖励上施加 KL 惩罚 |
compute_response_mask() |
function | 计算回复部分的 attention mask |
compute_advantage() |
function | 根据配置选择优势估计方法 |
RayPPOTrainer |
class | 核心训练器类 |
RayPPOTrainer.__init__() |
method | 初始化配置、数据加载器等 |
RayPPOTrainer.init_workers() |
method | 创建分布式 Worker Group |
RayPPOTrainer.fit() |
method | 执行完整训练循环 |
RayPPOTrainer._compute_values() |
method | 调用 Critic 计算价值函数 |
RayPPOTrainer._compute_ref_log_prob() |
method | 调用 Ref Policy 计算参考概率 |
RayPPOTrainer._compute_old_log_prob() |
method | 调用 Actor 计算旧策略概率 |
RayPPOTrainer._update_actor() |
method | 更新 Actor 策略网络 |
RayPPOTrainer._update_critic() |
method | 更新 Critic 价值网络 |
RayPPOTrainer._save_checkpoint() |
method | 保存 checkpoint |
RayPPOTrainer._load_checkpoint() |
method | 加载 checkpoint |
RayPPOTrainer._validate() |
method | 执行验证 |
RayPPOTrainer._balance_batch() |
method | 序列长度均衡 |
数据流和调用关系¶
fit()
|
+-- _load_checkpoint()
+-- checkpoint_manager.update_weights()
+-- _validate() (训练前验证)
|
+-- for epoch:
for batch in train_dataloader:
|
+-- generate_sequences() --> gen_batch_output (含 responses, rollout_log_probs)
|
+-- _compute_reward() --> reward_tensor (rm_scores)
|
+-- _compute_old_log_prob() --> old_log_probs, entropys
|
+-- _compute_ref_log_prob() --> ref_log_prob (如果需要 KL)
|
+-- _compute_values() --> values (如果使用 Critic)
|
+-- apply_kl_penalty() --> token_level_rewards
|
+-- compute_advantage() --> advantages, returns
|
+-- _update_critic() --> 更新 Critic
|
+-- _update_actor() --> 更新 Actor
|
+-- checkpoint_manager.update_weights() --> 同步到推理引擎
|
+-- _save_checkpoint() (定期)
+-- _validate() (定期)
+-- logger.log(metrics)
小结¶
ray_trainer.py 是 verl 框架的核心,它实现了:
- 完整的 PPO 训练循环:从生成回复到计算奖励、优势估计、策略更新
- 分布式协调:通过 Ray Worker Group 管理多个分布式组件
- 灵活的算法支持:支持 GAE、GRPO、REINFORCE++ 等多种优势估计
- 高效的资源管理:共置模型、序列长度均衡、休眠/唤醒机制
- 完善的工程特性:checkpoint、断点恢复、指标收集、性能分析
理解这个文件是理解整个 verl 训练系统的关键。