kernels.py — Triton 内核实现¶
文件路径: verl/utils/kernel/kernels.py
文件概述¶
使用 Triton 编写的高性能 GPU 内核,实现了 Linear Cross Entropy 的前向和反向传播。这是计算密集型操作的最底层实现。
背景知识¶
Triton 是 OpenAI 开发的 GPU 编程语言,比 CUDA C 更易用,同时能生成接近手写 CUDA 的高性能代码。它的编程模型基于 block(每个 program instance 处理一个 tile)。
核心函数¶
1. 前向传播¶
def efficient_entropy_forward(hidden, weight, labels, reduction, temperature, dist_process_group):
"""
计算 log-prob 和 entropy,不需要存储完整 logits。
内部使用 Triton kernel 分 tile 计算:
1. 每个 tile 计算局部的 hidden @ weight[tile].T
2. 用 online softmax 技巧累积全局统计量
3. 最终得到 log-prob 和 entropy
"""
2. 反向传播¶
def efficient_entropy_backward(dlogprobs, dentropy, hidden, weight, labels, ...):
"""
计算 hidden 和 weight 的梯度。
同样使用 tiling 策略,避免存储完整 logits。
"""
3. Triton Kernel 示例¶
@triton.jit
def _forward_kernel(
HIDDEN, # [N, D] 隐藏状态
WEIGHT, # [V, D] 词表权重
LABELS, # [N] 标签
LOGPROBS, # [N] 输出 log 概率
ENTROPY, # [N] 输出 entropy
...
BLOCK_D: tl.constexpr,
BLOCK_V: tl.constexpr,
):
# 每个 program instance 处理一个 token
pid = tl.program_id(0)
# 分 tile 遍历词表,累积 max 和 sum_exp
for v_start in range(0, V, BLOCK_V):
# logit = hidden[pid] @ weight[v_start:v_start+BLOCK_V].T
# 更新 running max 和 sum_exp (online softmax)
# 计算最终的 log-prob 和 entropy
性能优化技巧¶
- Tiling: 将大矩阵分成小块,适应 GPU shared memory
- Online Softmax: 不需要两遍扫描,一遍完成 max + sum
- Fused Operations: 将线性层和损失函数合并为单个 kernel
- 张量并行支持: 通过 all_reduce 合并分片 logits 的统计量
与其他模块的关系¶
- 被
linear_cross_entropy.py调用 - 依赖 Triton 库(
import triton) - 在没有 Triton 时会 fallback 到 PyTorch 实现
小结¶
kernels.py 是 verl 中计算密度最高的文件,用 Triton 实现了内存高效的 log-prob 和 entropy 计算。理解它需要 GPU 编程基础,但使用时只需调用 linear_cross_entropy() 即可。