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 操作一样使用,同时获得巨大的内存节省。