跳转至

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

性能优化技巧

  1. Tiling: 将大矩阵分成小块,适应 GPU shared memory
  2. Online Softmax: 不需要两遍扫描,一遍完成 max + sum
  3. Fused Operations: 将线性层和损失函数合并为单个 kernel
  4. 张量并行支持: 通过 all_reduce 合并分片 logits 的统计量

与其他模块的关系

  • 被 linear_cross_entropy.py 调用
  • 依赖 Triton 库(import triton)
  • 在没有 Triton 时会 fallback 到 PyTorch 实现

小结

kernels.py 是 verl 中计算密度最高的文件,用 Triton 实现了内存高效的 log-prob 和 entropy 计算。理解它需要 GPU 编程基础,但使用时只需调用 linear_cross_entropy() 即可。