跳转至

torch_functional.py — 核心 Torch 计算函数

文件路径: verl/utils/torch_functional.py

文件概述

torch_functional.py(1000+ 行)是 verl 中计算逻辑最核心的文件之一。它包含了 RLHF 训练中需要的关键数学函数:log 概率计算、KL 散度、GAE 优势估计、学习率调度器等。

背景知识

在 RLHF 训练中: - log-prob: 模型对每个 token 的对数概率,是策略梯度的基础 - KL 散度: 衡量当前策略与参考策略的差异,用于正则化 - GAE (Generalized Advantage Estimation): 估计每个 token 的优势值,平衡偏差和方差 - PPO clip: PPO 算法的核心,通过裁剪概率比来稳定训练

核心函数详解

1. Log 概率计算

def log_probs_from_logits_response_rmpad(input_ids, attention_mask, logits_rmpad, response_length):
    """
    从去除 padding 的 logits 中计算 response 部分的 log 概率。
    Args:
        input_ids: [batch_size, seqlen]
        attention_mask: [batch_size, seqlen]
        logits_rmpad: [total_nnz, vocab_size] (去除 padding 后的紧凑表示)
        response_length: 回复的长度
    Returns:
        [batch_size, response_length] 的 log 概率
    """

rmpad 表示 "remove padding"。为了节省计算,先去掉 padding token,只对有效 token 计算 logits。

2. KL 散度

def kl_penalty(logprob, ref_logprob, kl_penalty_type="kl"):
    """
    计算 KL 惩罚项。
    支持多种变体:
    - "kl": 标准 KL 散度 ref_logprob - logprob
    - "abs": |logprob - ref_logprob|
    - "mse": (logprob - ref_logprob)^2
    - "full": 完整 KL 散度
    """

KL 惩罚防止模型偏离参考模型太远,是 RLHF 稳定训练的关键。

3. GAE 优势估计

def compute_gae_advantage_return(token_level_rewards, values, response_length, gamma, lam):
    """
    使用 GAE 计算优势值和回报。
    Args:
        token_level_rewards: [batch, response_len] 每个 token 的奖励
        values: [batch, response_len] Critic 估计的价值
        gamma: 折扣因子(通常 0.99 或 1.0)
        lam: GAE lambda 参数(通常 0.95)
    Returns:
        advantages: [batch, response_len]
        returns: [batch, response_len]
    """

GAE 通过 TD(lambda) 方法,在 MC 估计(高方差低偏差)和 TD 估计(低方差高偏差)之间取平衡。

4. 学习率调度器

def get_constant_schedule_with_warmup(optimizer, num_warmup_steps):
    """常数学习率 + warmup"""

def get_cosine_schedule_with_warmup(optimizer, num_warmup_steps, num_training_steps):
    """余弦退火学习率 + warmup"""

核心函数列表

函数 说明
log_probs_from_logits() 从 logits 计算 log 概率
log_probs_from_logits_response_rmpad() 去 padding 版本
kl_penalty() 计算 KL 惩罚
compute_gae_advantage_return() GAE 优势估计
compute_rewards() 计算带 KL 惩罚的奖励
get_cosine_schedule_with_warmup() 余弦 LR 调度
get_constant_schedule_with_warmup() 常数 LR 调度
masked_mean() 带 mask 的均值
entropy_from_logits() 从 logits 计算熵

与其他模块的关系

  • 被 PPO trainer 在训练循环中频繁调用
  • 被 Actor worker 用来计算 log-prob
  • 被 Critic worker 用来计算 GAE
  • 依赖 flash_attn 的 pad/unpad 工具处理变长序列

小结

torch_functional.py 是 RLHF 算法的"数学引擎",集中了所有与策略优化相关的计算函数。理解这些函数是理解 PPO 训练过程的关键。