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 训练过程的关键。