通过一个简单的向量加法例子,学习 Triton 的基本编程模型。
配套代码
Triton 是一种用 Python 语法编写 GPU kernel 的语言。与 CUDA 相比,Triton 自动处理了很多底层细节(如 shared memory 管理、线程同步等),让我们可以专注于算法本身。
本章通过一个最简单的例子——向量加法——来了解 Triton 的核心编程模型。
在开始写代码前,我们需要理解 Triton 的核心思想:SPMD(Single Program, Multiple Data)。
简单说:同一份 kernel 代码会被多个 “program” 并行执行,每个 program 处理数据的不同部分。
假设我们要对两个长度为 256 的向量做加法,BLOCK_SIZE 设为 64。Triton 会启动 4 个 program:
输入向量 (N=256, BLOCK_SIZE=64):
┌────────────┬────────────┬────────────┬────────────┐
│ 0 ... 63 │ 64 ... 127 │ 128 .. 191 │ 192 .. 255 │
├────────────┼────────────┼────────────┼────────────┤
│ Program 0 │ Program 1 │ Program 2 │ Program 3 │
└────────────┴────────────┴────────────┴────────────┘
每个 program 只负责处理自己那一块数据。那么问题来了:每个 program 怎么知道自己应该处理哪一块?
答案是 tl.program_id()——它返回当前 program 的编号(0、1、2、3…)。
我们不一次性展示完整代码,而是逐步构建,理解每一部分的作用。
systems/triton_basics/01_vector_add.py
pid = tl.program_id(axis=0) # 获取当前 program 的编号
block_start = pid * BLOCK_SIZE # 计算这个 program 负责的数据起始位置
如果 pid = 2,BLOCK_SIZE = 64,那么 block_start = 128。这个 program 负责处理从索引 128 开始的元素。
systems/triton_basics/01_vector_add.py