losses.py — 损失函数¶
文件概述¶
定义 PPO 训练中 Actor 和 Critic 使用的三种核心损失函数:SFT 损失、PPO 策略损失、Value 损失。这些函数是整个 RLHF 训练流程中梯度计算的核心。
核心函数¶
1. sft_loss - 监督微调损失¶
SFT(Supervised Fine-Tuning)损失是标准的语言模型交叉熵损失,仅在 response 部分计算。
def sft_loss(config: ActorConfig, model_output, data: TensorDict, dp_group=None):
log_prob = model_output["log_probs"]
if pad_mode == DatasetPadMode.NO_PADDING:
# 无填充模式:使用 nested tensor
log_prob_flatten = log_prob.values()
loss_mask_flatten = loss_mask.values()
# 左移 loss_mask 一个 token,与 log_prob 对齐
loss_mask_flatten = torch.roll(loss_mask_flatten, shifts=-1, dims=0)
# 在全局 batch 上取平均(跨 DP 组)
loss = -masked_sum(log_prob_flatten, loss_mask_flatten) / batch_num_tokens * dp_size
else:
# 有填充模式
response_mask = data["response_mask"].to(bool)
loss = -masked_sum(log_prob, response_mask) / batch_num_tokens * dp_size
return loss, {}
关键细节:
- loss_mask 需要左移一个 token,因为语言模型的 log_prob 对应的是预测下一个 token 的概率
- 损失除以 batch_num_tokens 并乘以 dp_size,实现跨数据并行组的正确归一化
2. ppo_loss - PPO 策略损失¶
PPO 损失是 RLHF 训练的核心,由三部分组成:策略梯度损失 + 熵正则 + KL 惩罚。
def ppo_loss(config: ActorConfig, model_output, data: TensorDict, dp_group=None):
# 1. 将无填充的 log_prob 转换为填充格式
log_prob = no_padding_2_padding(model_output["log_probs"], data)
# 2. 设置全局 batch 信息用于损失归一化
config.global_batch_info["dp_size"] = data["dp_size"]
config.global_batch_info["batch_num_tokens"] = data["batch_num_tokens"]
# 3. 计算策略梯度损失(PPO clip)
policy_loss_fn = get_policy_loss_fn(loss_mode)
pg_loss, pg_metrics = policy_loss_fn(
old_log_prob=old_log_prob, # rollout 时的 log_prob
log_prob=log_prob, # 当前模型的 log_prob
advantages=advantages, # GAE 计算的优势
response_mask=response_mask,
loss_agg_mode=loss_agg_mode,
config=config,
)
policy_loss = pg_loss
# 4. 熵正则(鼓励探索)
if entropy is not None:
entropy_loss = agg_loss(loss_mat=entropy, loss_mask=response_mask, ...)
policy_loss -= entropy_coeff * entropy_loss # 减号:最大化熵
# 5. KL 散度惩罚(防止偏离参考模型太远)
if config.use_kl_loss:
kld = kl_penalty(logprob=log_prob, ref_logprob=ref_log_prob, ...)
kl_loss = agg_loss(loss_mat=kld, loss_mask=response_mask, ...)
policy_loss += kl_loss * config.kl_loss_coef
return policy_loss, metrics
PPO 损失的三个组成部分:
| 组成部分 | 作用 | 系数 |
|---|---|---|
pg_loss |
策略梯度损失(带 clip) | 1.0 |
entropy_loss |
熵正则,鼓励探索 | -entropy_coeff(负号表示最大化) |
kl_loss |
KL 散度惩罚,防止偏离参考模型 | kl_loss_coef |
指标聚合策略:
- 当存在多卡 DP 时,使用 AggregationType.SUM(因为损失已经按全局 batch 归一化)
- 单卡时使用 AggregationType.MEAN
3. value_loss - 价值函数损失¶
Critic 模型的损失函数,用于训练价值估计。
def value_loss(config: CriticConfig, model_output, data: TensorDict, dp_group=None):
# 从无填充输出中切片出 response 部分的预测值
vpreds = _slice_response_from_unpad_output(model_output["values"], data)
values = data["values"] # 旧的价值预测
returns = data["returns"] # GAE 计算的回报
response_mask = data["response_mask"].to(bool)
# 计算带 clip 的价值损失
vf_loss, vf_clipfrac = compute_value_loss(
vpreds=vpreds,
values=values,
returns=returns,
response_mask=response_mask,
cliprange_value=config.cliprange_value,
loss_agg_mode=config.loss_agg_mode,
)
return vf_loss, metrics
4. _slice_response_from_unpad_output - 辅助函数¶
从无填充模型输出中提取 response 部分并恢复为填充格式。
def _slice_response_from_unpad_output(tensor, data):
"""
输入: 展平的模型输出 [total_tokens]
输出: 按样本切片并填充的 response [bsz, max_response_len]
"""
for resp_len, seq_offset in zip(response_lens, sequence_offsets):
pad_size = max_response_len - resp_len
# 左移一个 token(log_prob 对应下一个 token)
response_list.append(
F.pad(values[seq_offset - resp_len - 1 : seq_offset - 1], (0, pad_size))
)
return torch.stack(response_list, dim=0)
数据流示意¶
PPO 训练的损失计算流程:
Rollout 阶段(生成数据):
Actor(old) → old_log_probs, responses
Ref Model → ref_log_prob
Critic → values
Reward → rewards → GAE → advantages, returns
训练阶段(计算损失):
Actor(new) → log_probs, entropy
│
▼
┌─── ppo_loss ──────────────────┐
│ pg_loss = clip_ppo( │
│ old_log_prob, log_prob, │
│ advantages) │
│ │
│ entropy_loss = mean(entropy) │
│ │
│ kl_loss = KL(log_prob, │
│ ref_log_prob) │
│ │
│ total = pg + kl - entropy │
└───────────────────────────────┘
Critic(new) → vpreds
│
▼
┌─── value_loss ────────────────┐
│ vf_loss = clip_value( │
│ vpreds, values, returns) │
└───────────────────────────────┘
与其他模块的关系¶
- 被
engine/fsdp/transformer_impl.py中的FSDPEngineWithLMHead和FSDPEngineWithValueHead调用 - 被
engine/megatron/transformer_impl.py和其他引擎调用 - 依赖
verl/trainer/ppo/core_algos.py中的核心 PPO 算法实现 - 使用
padding.py中的no_padding_2_padding进行格式转换 - 使用
ActorConfig和CriticConfig获取超参数
小结¶
本文件实现了 RLHF 训练中最关键的三种损失函数。SFT 损失用于监督微调阶段,PPO 策略损失和 Value 损失用于强化学习阶段。所有损失函数都支持无填充(remove padding)优化和分布式训练的正确归一化。