torch_functional.py — 这个文件实现了 FusedLinearForPPO -- 一个为 PPO 训练优化的融合线性...¶
模块路径: verl.utils.experimental.torch_functional
文件概述¶
这个文件实现了 FusedLinearForPPO -- 一个为 PPO 训练优化的融合线性层。它将 "hidden_states -> logits -> log_probs + entropy" 的计算融合在一起,通过分块(chunking)处理来大幅减少 GPU 显存占用。
在 PPO 训练中,计算 log_probs 和 entropy 需要对整个词表做 softmax,产生巨大的中间 logits 张量。这个融合层避免了同时在内存中保存完整的 logits。
关键代码讲解¶
1. 前向计算(分块融合)¶
def _fused_linear_for_ppo_fwd(hidden_states, vocab_weights, input_ids, temperature=1.0):
logits = (hidden_states @ vocab_weights.t()) / temperature
logits = logits.to(torch.float32) # 数值稳定性
# 用 log_softmax 而非 probs.log()(更稳定)
probs = logits.softmax(dim=-1)
log_probs = logits.log_softmax(dim=-1)
# 只取对应 token 的 log_prob
token_log_probs = log_probs.gather(-1, input_ids.unsqueeze(-1)).squeeze(-1)
# entropy = logsumexp(logits) - sum(probs * logits)
entropy = torch.logsumexp(logits, dim=-1) - torch.sum(probs * logits, dim=-1)
return token_log_probs, entropy
2. 自定义 autograd Function(分块处理)¶
class FusedLinearForPPOFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, hidden_states, vocab_weights, input_ids, temperature=1.0, chunk_size=512):
# 将 3D 输入展平为 2D
hidden_states = hidden_states.flatten(0, 1) # [B*T, D]
# 分块计算,每次只处理 chunk_size 个 token
for chunk_start in range(0, T, chunk_size):
chunk_end = min(chunk_start + chunk_size, T)
chunk_log_probs, chunk_entropy = _fused_linear_for_ppo_fwd(
hidden_states[chunk_start:chunk_end],
vocab_weights,
input_ids[chunk_start:chunk_end],
temperature,
)
log_probs[chunk_start:chunk_end] = chunk_log_probs
entropy[chunk_start:chunk_end] = chunk_entropy
ctx.save_for_backward(hidden_states, vocab_weights, input_ids)
return log_probs, entropy
分块的关键好处:假设词表大小 V=128K,chunk_size=512,则每个 chunk 的 logits 大小为 [512, 128K],而不是 [B*T, 128K]。显存节省巨大。
3. 反向传播(也分块)¶
@staticmethod
def backward(ctx, dlog_probs, dentropy):
# 分块反向传播
for chunk_start in range(0, T, chunk_size):
h, v = _fused_linear_for_ppo_bwd(
dlog_probs=chunk_dlog_probs,
dentropy=chunk_dentropy,
hidden_states=hidden_states[chunk_start:chunk_end],
vocab_weights=vocab_weights,
input_ids=input_ids[chunk_start:chunk_end],
)
dhidden_states[chunk_start:chunk_end] += h
dvocab_weights += v # 累加
反向传播的梯度计算:
def _fused_linear_for_ppo_bwd(dlog_probs, dentropy, hidden_states, vocab_weights, input_ids, temperature):
probs = logits.softmax(dim=-1)
# log_probs 的梯度
if dlog_probs is not None:
one_hot = torch.zeros_like(logits).scatter_(-1, input_ids.unsqueeze(-1), 1)
dlogits += dlog_probs.unsqueeze(-1) * (one_hot - probs)
# entropy 的梯度
if dentropy is not None:
dlogits += probs * (log_probs + entropy.unsqueeze(-1)) * (-dentropy.unsqueeze(-1))
dhidden_states = dlogits @ vocab_weights
dvocab_weights = dlogits.t() @ hidden_states
return dhidden_states, dvocab_weights
4. Module 封装¶
class FusedLinearForPPO(torch.nn.Module):
def __init__(self, chunk_size=512):
self.chunk_size = chunk_size
def forward(self, hidden_states, vocab_weights, input_ids, temperature=1.0):
return FusedLinearForPPOFunction.apply(
hidden_states, vocab_weights, input_ids, temperature, self.chunk_size
)
核心类/函数列表¶
| 类/函数名 | 作用 |
|---|---|
FusedLinearForPPO |
nn.Module 封装 |
FusedLinearForPPOFunction |
自定义 autograd Function |
_fused_linear_for_ppo_fwd |
前向计算核心 |
_fused_linear_for_ppo_bwd |
反向传播核心 |
与其他模块的关系¶
- 可以替换 PPO 训练中的
model.lm_head+ softmax + log_prob 计算 - 通过
torch.autograd.Function自定义前向和反向传播 - 属于实验性功能,可能在未来稳定后迁移
小结¶
这是一个高度优化的计算融合模块。通过将 linear → softmax → log_prob/entropy 融合为一个操作,并使用分块策略,可以在不降低精度的情况下显著减少 PPO 训练的显存占用。自定义 autograd Function 确保梯度计算也采用同样的分块策略。这在训练大词表模型时尤其重要。