跳转至

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)优化和分布式训练的正确归一化。