缓存时间:
2026/04/20 14:55
# 介绍 Triton:用于神经网络的开源 GPU 编程
来源:https://openai.com/index/triton/
我们发布了 Triton 1.0,这是一种开源的类似 Python 的编程语言,使得没有 CUDA 经验的研究人员能够编写高效的 GPU 代码——在大多数情况下性能与专家编写的代码相当。
## 为什么这很重要
Triton 使得以相对较少的工作量达到峰值硬件性能成为可能;例如,它可以用于编写 FP16 矩阵乘法内核,在不到 25 行代码的情况下匹配 cuBLAS 的性能——这是许多 GPU 程序员无法做到的。我们的研究人员已经用它生产了效率比同等 Torch 实现高达 2 倍的内核,我们很高兴能与社区合作,使 GPU 编程对所有人更加易得。
深度学习领域的新研究思想通常使用原生框架操作符的组合来实现。虽然方便,但这种方法通常需要创建和/或移动许多临时张量,这会影响大规模神经网络的性能。这些问题可以通过编写专门的 GPU 内核来缓解,但由于 GPU 编程的许多复杂之处,这样做可能出乎意料地困难。虽然最近出现了各种系统来简化这个过程,但我们发现它们要么过于冗长,要么缺乏灵活性,要么生成的代码明显慢于我们手工调整的基线。这促使我们扩展和改进 Triton,这是一种最近的语言和编译器,其原始创建者现在在 OpenAI 工作。
## GPU 编程的挑战
现代 GPU 的架构可以大致分为三个主要组件——DRAM、SRAM 和 ALU——优化 CUDA 代码时必须考虑每一个:
- 来自 DRAM 的内存传输必须**合并**成大事务,以利用现代内存接口的大总线宽度。
- 数据必须在被重新使用之前手动存储到 SRAM,并进行管理以最小化检索时的共享内存库冲突。
- 计算必须在流式多处理器 (SM) 之间和之内仔细分割和调度,以促进指令/线程级别的并行性并利用专用 ALU(例如张量核心)。
GPU 架构图:GPU 的基本架构。
即使对于拥有多年经验的资深 CUDA 程序员来说,考虑所有这些因素也可能很有挑战性。Triton 的目的是完全自动化这些优化,使开发人员能够更好地专注于其并行代码的高级逻辑。
Triton 旨在广泛适用,因此不会自动跨 SM 调度工作——留给开发人员一些重要的算法考虑因素(例如分块、SM 间同步)的处理权。
| | CUDA | TRITON |
|---|---|---|
| 内存合并 | 手动 | 自动 |
| 共享内存管理 | 手动 | 自动 |
| 调度(SM 内) | 手动 | 自动 |
| 调度(跨 SM) | 手动 | 手动 |
CUDA 与 Triton 的编译器优化比较。
## 编程模型
在所有可用的领域特定语言和 JIT 编译器中,Triton 或许最类似于 Numba:内核定义为装饰的 Python 函数,并在所谓**实例**的网格上使用不同的 `program_id` 并发启动。但是,如下面的代码片段所示,相似之处到此为止:Triton 通过对**块**(小数组,其维度是 2 的幂)的操作来公开实例内并行性,而不是单指令多线程 (SIMT) 执行模型。这样做,Triton 有效地抽象了与 CUDA 线程块**内**并发相关的所有问题(例如,内存合并、共享内存同步/冲突、张量核心调度)。虽然这对于易于并行化的(即,按元素)计算可能不是特别有帮助,但它可以大大简化更复杂的 GPU 程序的开发。
考虑一个融合 softmax 内核的例子(下文),其中每个实例规范化给定输入张量 X ∈ R_{M×N} 的不同行。这种并行化策略的标准 CUDA 实现可能很难编写,需要显式的线程同步,因为它们同时归约 X 的同一行。使用 Triton,大部分复杂性都消失了,其中每个内核实例加载相关行并使用类似 NumPy 的原语按顺序规范化它。
#### Python
```python
import triton
import triton.language as tl
@triton.jit
def softmax(Y, stride_ym, stride_yn, X, stride_xm, stride_xn, M, N):
# row index
m = tl.program_id(0)
# col indices
# this specific kernel only works for matrices that
# have less than BLOCK_SIZE columns
BLOCK_SIZE = 1024
n = tl.arange(0, BLOCK_SIZE)
# the memory address of all the elements
# that we want to load can be computed as follows
X = X + m * stride_xm + n * stride_xn
# load input data; pad out-of-bounds elements with 0
x = tl.load(X, mask=n < N, other=-float('inf'))
# compute numerically-stable softmax
z = x - tl.max(x, axis=0)
num = tl.exp(z)
denom = tl.sum(num, axis=0)
y = num / denom
# write back to Y
Y = Y + m * stride_ym + n * stride_yn
tl.store(Y, y, mask=n < N)
import torch
# Allocate input/output tensors
X = torch.normal(0, 1, size=(583, 931), device='cuda')
Y = torch.empty_like(X)
# SPMD launch grid
grid = (X.shape[0], )
# enqueue GPU kernel
softmax[grid](Y, Y.stride(0), Y.stride(1),
X, X.stride(0), X.stride(1),
X.shape[0] , X.shape[1])
```
注意 Triton JIT 将 X 和 Y 视为**指针**而非张量;我们认为保留对内存访问的低级控制对于处理更复杂的数据结构(例如块稀疏张量)很重要。重要的是,softmax 的这个特定实现在整个规范化过程中将 X 的行保留在 SRAM 中,这在适用时最大化数据重用(~<32K 列)。这不同于 PyTorch 内部 CUDA 代码,其对临时内存的使用使其更通用但明显更慢(下文)。关键点不是 Triton 本身更好,而是它简化了专门内核的开发,这些内核可以比通用库中的内核快得多。Torch (v1.9) JIT 的较低性能突显了从高级张量操作序列自动生成 CUDA 代码的困难。
#### Python
```python
@torch.jit.script
def softmax(x):
x_max = x.max(dim=1)[0]
z = x - x_max[:, None]
numerator = torch.exp(x)
denominator = numerator.sum(dim=1)
return numerator / denominator[:, None]
```
使用 Torch JIT 的融合 softmax。
## 矩阵乘法
能够为按元素操作和归约编写融合内核很重要,但鉴于神经网络中矩阵乘法任务的突出地位,这还不够。事实证明,Triton 也非常适合这些任务,仅用约 25 行 Python 代码就能达到峰值性能。另一方面,用 CUDA 实现类似的东西需要付出更多努力,甚至可能达到更低的性能。
#### Python
```python
@triton.jit
def matmul(A, B, C, M, N, K, stride_am, stride_ak,
stride_bk, stride_bn, stride_cm, stride_cn,
**META):
# extract metaparameters
BLOCK_M, GROUP_M = META['BLOCK_M'], META['GROUP_M']
BLOCK_N = META['BLOCK_N']
BLOCK_K = META['BLOCK_K']
# programs are grouped together to improve L2 hit rate
_pid_m = tl.program_id(0)
_pid_n = tl.program_id(1)
pid_m = _pid_m // GROUP_M
pid_n = (_pid_n * GROUP_M) + (_pid_m % GROUP_M)
# rm (resp. rn) denotes a range of indices
# for rows (resp. col) of C
rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
rn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
# rk denotes a range of indices for columns
# (resp. rows) of A (resp. B)
rk = tl.arange(0, BLOCK_K)
# the memory addresses of elements in the first block of
# A and B can be computed using numpy-style broadcasting
A = A + (rm[:, None] * stride_am + rk[None, :] * stride_ak)
B = B + (rk [:, None] * stride_bk + rn[None, :] * stride_bn)
# initialize and iteratively update accumulator
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k in range(K, 0, -BLOCK_K):
a = tl.load(A)
b = tl.load(B)
# block level matrix multiplication
acc += tl.dot(a, b)
# increment pointers so that the next blocks of A and B
# are loaded during the next iteration
A += BLOCK_K * stride_ak
B += BLOCK_K * stride_bk
# fuse leaky ReLU if desired
# acc = tl.where(acc >= 0, acc, alpha * acc)
# write back result
C = C + (rm[:, None] * stride_cm + rn[None, :] * stride_cn)
mask = (rm[:, None] < M) & (rn[None, :] < N)
tl.store(C, acc, mask=mask)
```
Triton 中的矩阵乘法。
手写矩阵乘法内核的一个重要优势是它们可以根据需要自定义,以适应其输入的融合变换(例如,切片)和输出(例如,Leaky ReLU)。没有像 Triton 这样的系统,对矩阵乘法内核进行非平凡修改对于没有非凡 GPU 编程专业知识的开发人员来说将是遥不可及的。
## 高级系统架构
Triton 的良好性能来自于以 Triton-IR 为中心的模块化系统架构,这是一个基于 LLVM 的中间表示,其中多维值块是一等公民。`@triton.jit` 装饰器通过遍历提供的 Python 函数的抽象语法树 (AST) 来工作,使用通用 SSA 构建算法即时生成 Triton-IR。生成的 IR 代码随后被我们的编译器后端简化、优化和自动并行化,然后被转换为高质量的 LLVM-IR——最终转换为 PTX——以在最近的 NVIDIA GPU 上执行。目前不支持 CPU 和 AMD GPU,但我们欢迎社区贡献以解决这一限制。
## 编译器后端
我们发现通过 Triton-IR 使用块化程序表示允许我们的编译器自动执行各种重要的程序优化。例如,通过查看计算密集型块级操作(例如 `tl.dot`)的操作数,数据可以自动存储到共享内存中——并使用标准活性分析技术分配/同步。另一方面,Triton 程序可以高效地自动并行化,既可以通过并发执行不同内核实例跨 SM,也可以通过分析每个块级操作的迭代空间并跨不同 SIMD 单元适当分割它来在 SM 内进行,如下所示。
## 贡献
*如果你有兴趣加入我们的团队并从事 Triton 和 GPU 内核工作,**我们在招聘*!
- 社区与协作
- 软件与工程
## 参考资料
1. Gray, S. (2017). SGEMM Walkthrough. URL https://github.com/NervanaSystems/maxas/wiki/SGEMM
2. Kerr, A. (2020). Developing CUDA kernels to push Tensor Cores to the Absolute Limit on NVIDIA A100. URL https://developer.nvidia.com/gtc/2020/video/s21745-vid
3. Yan, D., Wang, W., & Chu, X. (2020, May). Demystifying tensor cores to optimize half-precision matrix multiply. 在 2020 IEEE International Parallel and Distributed Processing Symposium (IPDPS)。IEEE。
4. NVIDIA CUTLASS
5. Apache TVM
6. Tillet, P., Kung, H. T., & Cox, D. (2019, June). Triton: an intermediate language and compiler for tiled neural network computations. 在 Proceedings of the 3rd ACM SIGPLAN International Workshop on Machine Learning and Programming Languages (pp. 10-19)。
7. Lin, Y. & Grover, V. (2018). Using CUDA Warp-Level Primitives. URL https://developer.nvidia.com/blog/using-cuda-warp-level-primitives/
8. Braun, M., Buchwald, S., Hack, S., Leißa, R., Mallon, C., & Zwinkau, A. (2013, March). Simple and efficient construction of static single assignment form. 在 International Conference on Compiler Construction (pp. 102-122)。Springer, Berlin, Heidelberg。
## 致谢
Da Yan (HKUST)、DeepSpeed (Microsoft)、Anthropic