跳转至

linear_cross_entropy.py — 融合线性交叉熵

文件路径: verl/utils/kernel/linear_cross_entropy.py

文件概述

实现融合的 "线性层 + 交叉熵损失" 计算(Linear Cross Entropy)。传统方法先算 logits (hidden x weight),再算 cross entropy,需要存储完整的 logits 矩阵。融合方法将两步合一,大幅减少内存使用。

背景知识

语言模型的最后一步是: 1. logits = hidden @ weight.T — 将隐藏状态映射到词表空间 2. loss = cross_entropy(logits, labels) — 计算损失

logits 的形状是 [batch * seq_len, vocab_size],对于大词表(如 128K)和长序列,这个矩阵非常大。融合内核避免了存储完整 logits。

核心类

LinearCrossEntropy

class LinearCrossEntropy(torch.autograd.Function):
    @staticmethod
    def forward(ctx, hidden, weight, labels, temperature=1.0, reduction="none", dist_process_group=None):
        """
        Args:
            hidden: [batch_size * num_tokens, hidden_size]
            weight: [vocab_size, hidden_size]
            labels: [batch_size * num_tokens]
        Returns:
            logprobs: log 概率
            entropy: token 级别的熵
        """
        logprobs, entropy, _maximum, _accumulate, _entropy_b = kernels.efficient_entropy_forward(
            hidden, weight, labels, REDUCTION, temperature, dist_process_group
        )
        ctx.save_for_backward(hidden, weight, labels, _maximum, _accumulate, _entropy_b)
        return logprobs, entropy

    @staticmethod
    def backward(ctx, dlogprobs, dentropy):
        d_hidden, d_weight = kernels.efficient_entropy_backward(...)
        return (d_hidden, d_weight, None, None, None, None)

linear_cross_entropy = LinearCrossEntropy.apply

优势

方面 传统方法 融合方法
内存 O(batch * seq * vocab) O(batch * seq)
速度 两次 kernel launch 一次 kernel launch
数值 标准 同样精确

与其他模块的关系

  • 依赖 kernels.py 中的 Triton 内核实现
  • 被 Actor worker 在计算 log-prob 和 entropy 时使用
  • 支持张量并行(dist_process_group 参数)

小结

linear_cross_entropy.py 通过 autograd Function 封装融合内核,让上层代码可以像普通 PyTorch 操作一样使用,同时获得巨大的内存节省。