跳转至

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 确保梯度计算也采用同样的分块策略。这在训练大词表模型时尤其重要。