我从零开始用 Triton 编写了一个 GPU 矩阵乘法内核:这是我学到的一切
原文:I Wrote a GPU Matmul Kernel From Scratch in Triton. Here's Everything I Learned
原文更新于:2026 年 6 月 23 日
翻译:GPT 5.6 Terra
我最近开始学习 Triton。这是 OpenAI 推出的、基于 Python 的 GPU 内核编程语言。我的项目是:从零开始,一步步构建一个矩阵乘法(matrix multiplication,简称 matmul)内核,直到它能与 PyTorch 内置的 torch.matmul() 一较高下。
本文完整记录了这段过程。如果你对 GPU 编程感到好奇,或曾想知道神经网络执行一次前向传播时到底发生了什么,这篇文章适合你。
最终结果在文末。
为什么 matmul 如此重要
矩阵乘法是深度学习中最重要的单个算子。每个线性层、每个注意力头,以及每个卷积(经降维转换后)最终都会归结为 matmul。人们谈论“推理优化”时,绝大多数情况下说的都是怎样让 matmul 更快。
一个拥有 70 亿参数的 LLM 在执行一次前向传播时,会进行数千次 matmul。若能把 matmul 时间削减哪怕 10%,在大规模部署下也能节省数百万美元的 GPU 成本。理解这个算子在硬件层面的工作方式,彻底改变了我对机器学习系统的看法。
Triton 是什么,为什么要用它?
通常,编写 GPU 内核意味着使用 CUDA,也就是某种意义上的“面向 GPU 的 C++”。它很强大,但也很痛苦:线程、共享内存、同步屏障等,全都需要手动管理。
Triton 位于更高一层。你写的是看起来像 Python 的代码,但它会被编译为 GPU 机器码。你不必逐个考虑线程,而是以数据块(block)为单位思考;Triton 会替你处理线程管理。
代价是放弃一部分底层控制能力;作为回报,你可以用 30 行代码写出一个可运行的内核,而不是 300 行。
GPU 的内存层级:必须先理解这一点
在写任何代码之前,先要理解 GPU 为什么在某些任务上很快、在另一些任务上却很慢。
GPU 有两种内存:
- DRAM(HBM):容量大(A100 上为 80 GB),但速度较慢(带宽约 2 TB/s)。当你调用
tensor.cuda()时,张量就存放在这里。 - SRAM(共享内存):容量很小(所有 SM 合计约 20 MB),但速度极快(19 TB/s)。它位于芯片上,紧挨着计算单元。
计算单元本身的速度极高。现代 GPU 每秒可完成数万亿次乘加运算。瓶颈几乎从来不是“GPU 算得够不够快”,而是“能否足够快地向计算单元供给数据”。
本文的每一项优化都围绕同一个目标:减少访问慢速 DRAM 的次数,并尽可能复用快速 SRAM 中的数据。
第 1 步:朴素 matmul(三层嵌套循环)
先从最简单的 matmul 开始。两个矩阵:形状为 (M, K) 的 A 和形状为 (K, N) 的 B,输出形状为 (M, N) 的 C。每个元素 C[i][j] 都是 A 的第 i 行与 B 的第 j 列的点积。
def naive_matmul(a, b):
M, K, N = a.shape[0], a.shape[1], b.shape[1]
c = torch.zeros((M, N), device=a.device)
for i in range(M):
for j in range(N):
for k in range(K):
c[i][j] += a[i][k] * b[k][j]
return c
三层循环,进行 M*N*K 次乘法。对于一个 26x26 矩阵,这大约要花 0.77 秒。
torch.matmul 完成同样的计算只需 0.03 秒。计算公式完全相同,速度却快了 25 倍。
这里有两个问题。第一,Python 循环会被解释执行,因此每次迭代都要付出 Python 字节码解释器的开销。第二,也是更根本的问题,是数据完全没有被复用。计算 C[0][0] 时,我们加载 A 的第 0 行;计算 C[0][1] 时,又重新加载完全相同的那一行。A 的第 0 行总共会被加载 N 次,浪费了大量内存带宽。
第 2 步:分块 matmul(核心思路)
如果不再一次只计算 C 的一个元素,而是同时计算一整块元素,会怎样?
将 C 切分为较小的 tile(例如 4x4)。对每个 tile,加载 A 和 B 中对应的条带,进行乘法并累加结果。
def blockwise_matmul(a, b, block_size=4):
M, K, N = a.shape[0], a.shape[1], b.shape[1]
c = torch.zeros((M, N), device=a.device)
for i in range(0, M, block_size):
for j in range(0, N, block_size):
for k in range(0, K, block_size):
for ii in range(i, min(i + block_size, M)):
for jj in range(j, min(j + block_size, N)):
s = 0.0
for kk in range(k, min(k + block_size, K)):
s += a[ii][kk] * b[kk][jj]
c[ii][jj] += s
return c
循环从三层变成了六层,看起来更复杂;乘法总数仍然完全相同,都是 M*N*K。在 Python 里,由于循环开销增加,它实际上还会稍慢一点。
那为什么还要这样做?
因为这种循环结构与 GPU 的工作方式高度契合。外层三个循环(i、j、k)遍历 tile,内层三个循环(ii、jj、kk)负责计算一个 tile 内的元素。在 GPU 上:
- 外层的
i和j循环会变成并行的 program,每个 program 在不同的 SM 上运行; - 内层循环完全在高速 SRAM 内进行;
- A 的一个 tile 从 DRAM 加载一次后,会被 C tile 中的每一列复用;
- B 的一个 tile 从 DRAM 加载一次后,会被 C tile 中的每一行复用。
当 block_size=128 时,A 的每个元素从 DRAM 加载的次数会从 N 次降至 N/128 次,内存流量减少 128 倍。由于 GPU 上的 matmul 受内存带宽限制,这几乎能直接转化为 128 倍加速。
min() 调用用于处理矩阵维度不能被 block size 整除的边界情况。剩余的 tile 会更小,但计算方法相同。
第 3 步:我的第一个 Triton 内核
现在进入最令人兴奋的部分:把分块思路转换为真正的 Triton GPU 代码。
在 Triton 中,使用 @triton.jit 装饰的、运行在 GPU 上的函数称为 kernel(内核)。启动内核时,会并行运行许多份副本,每一份称为一个 program。每个 program 处理输出的一块。
@triton.jit
def matmul_kernel(a_ptr, b_ptr, c_ptr, M, K, N, BLOCK_SIZE: tl.constexpr):
row_start = tl.program_id(axis=0) # C 的第几个块行
col_start = tl.program_id(axis=1) # C 的第几个块列
row_offsets = row_start * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
col_offsets = col_start * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
c_block = tl.zeros((BLOCK_SIZE, BLOCK_SIZE), dtype=tl.float32)
for k in range(0, K, BLOCK_SIZE):
k_offsets = k + tl.arange(0, BLOCK_SIZE)
a_block = tl.load(
a_ptr + row_offsets[:, None] * K + k_offsets[None, :],
mask=(row_offsets[:, None] < M) & (k_offsets[None, :] < K),
other=0.0
)
b_block = tl.load(
b_ptr + k_offsets[:, None] * N + col_offsets[None, :],
mask=(k_offsets[:, None] < K) & (col_offsets[None, :] < N),
other=0.0
)
c_block += tl.dot(a_block, b_block)
tl.store(
c_ptr + row_offsets[:, None] * N + col_offsets[None, :],
c_block,
mask=(row_offsets[:, None] < M) & (col_offsets[None, :] < N)
)
下面解释几个关键部分。
tl.program_id(axis=0) 和 tl.program_id(axis=1):每个 program 都有一个唯一的二维 ID。假设 C 有 250 个块行和 250 个块列,我们就启动 62,500 个 program。Program (3, 7) 计算块网格中第 3 行、第 7 列的那个块。
tl.arange(0, BLOCK_SIZE):生成数组 [0, 1, 2, ..., 63]。将它与 row_start * BLOCK_SIZE 结合,就得到该 program 的 tile 所对应的实际行索引。
row_offsets[:, None] 和 k_offsets[None, :]:这是二维指针运算中的广播技巧。[:, None] 将 [0,1,2,3] 变为列向量 [[0],[1],[2],[3]];[None, :] 则保持它为行向量。二者组合后生成一个二维内存地址网格,tile 的每个元素都有一个地址。
mask 和 other=0.0:BLOCK_SIZE 未必能整除 M、K 或 N,因此有些位置会越界。mask 能防止读取无效内存;other=0.0 会用零填充这些位置。零参与加法是安全的,因为加零不会改变结果。
tl.dot(a_block, b_block):这是 tile 级矩阵乘法,会映射到 GPU 的 Tensor Core 硬件。一次指令得到 BLOCK_SIZE^2 个输出值。
tl.load 和 tl.store:它们是仅有的 DRAM 操作。加载和存储之间的所有操作(减法、exp、除法,或这里的点积)都发生在 SRAM 中。
内核的启动方式如下:
grid = (triton.cdiv(M, BLOCK_SIZE), triton.cdiv(N, BLOCK_SIZE))
matmul_kernel[grid](a, b, c, M, K, N, BLOCK_SIZE=64)
triton.cdiv 是向上取整的除法。对于一个 1023x1023 矩阵,若 BLOCK_SIZE=64,每个维度的 cdiv(1023, 64) 都等于 16,因此总共启动 256 个 program。
这个内核在我的 V100 上达到了约 10 TFLOPS。对于仅 25 行实际内核代码而言,这已经不错了。
第 4 步:group-major PID 排序
这是最让我意外的一项优化,因为它完全不改变计算本身,只改变哪些 program 被分配到同一个 SM 上。
先说背景。GPU 有许多 SM(Streaming Multiprocessor,流式多处理器):A100 有 108 个,V100 有 80 个。Program 大致按 PID 的顺序被分配给 SM。同一 SM 上的 program 共享 SRAM。如果 Program 0 将 A 的一个 tile 加载到 SRAM,位于同一 SM 上的 Program 1 可以看到它并跳过一次 DRAM 加载。
在上面的二维网格中,PID 以行优先(row-major)顺序分配:
[ 0, 1, 2, 3]
[ 4, 5, 6, 7]
[ 8, 9, 10, 11]
[12, 13, 14, 15]
PID 0 至 3 都在计算 C 的同一行块。它们需要 A 的相同行 tile,因此能通过 SRAM 共享这些数据;但它们各自需要 B 的不同列 tile,无法共享。4 个 PID 要加载 1 份共享的 A 行数据和 4 份独立的 B 列数据,共 5 份独特数据。
如果将 PID 重排为方形分组呢?
group-major 布局:
[ 0, 2, 4, 6]
[ 1, 3, 5, 7]
[ 8, 10, 12, 14]
[ 9, 11, 13, 15]
现在 PID 0 至 3 构成一个 2x2 方块。它们共享 A 的 2 行和 B 的 2 列,也就是 4 份独特数据。输出块数相同,但从 DRAM 读取的数据量减少了 20%。
节省效果会随 group size 增大而提升。在大矩阵上使用 GROUP_SIZE=8 时,内存流量的减少会很可观。
为实现这一点,我把二维网格改为一维网格,并手动重新映射 PID:
PID = tl.program_id(axis=0) # 单个数字:0、1、2、...
# group-major 重映射
group_id = PID // (GROUP_SIZE * num_blocks_n)
first_row = group_id * GROUP_SIZE
group_size_adj = min(num_blocks_m - first_row, GROUP_SIZE)
row_start = first_row + ((PID % (GROUP_SIZE * num_blocks_n)) % group_size_adj)
col_start = (PID % (GROUP_SIZE * num_blocks_n)) // group_size_adj
% group_size_adj 会让行索引循环,例如 (0, 1, 0, 1, ...);// group_size_adj 则会在每次循环结束后推进列索引。因此,相邻 PID 会形成一条条垂直带,最终铺成方形。
内核其余部分完全相同。变化只发生在开头“我该处理哪个 block”的逻辑中。
第 5 步:Autotune(让 GPU 自己决定)
如何选择合适的 BLOCK_SIZE?应该是 32、64 还是 128?每个 program 需要多少个 warp(每个 warp 有 32 个线程)?又该使用多少个流水线 stage?
我花了不少时间推理:“128 能带来更多 SRAM 复用,但每个 program 会占用更多 SRAM,导致每个 SM 可容纳的 program 变少,所以 64 也许更好……”这种推理很脆弱。答案取决于寄存器压力、SRAM 容量、内存带宽以及十多个硬件细节。
真正的答案是:别猜,直接测。
Triton 提供了 @triton.autotune:
autotune_configs = [
triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64,
'GROUP_SIZE': 8}, num_stages=3, num_warps=8),
triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32,
'GROUP_SIZE': 8}, num_stages=4, num_warps=4),
# ... 另外 13 种配置
]
@triton.autotune(configs=autotune_configs, key=['M', 'N', 'K'])
@triton.jit
def autotuned_matmul_kernel(...):
运行时会发生以下事情:
- 首次以
M=N=K=1024调用时,Triton 会编译全部 15 个内核变体,分别进行基准测试,选出最快的一个,并缓存结果。 - 第二次以相同维度调用时,会立即使用缓存中的获胜配置。
- 以
M=N=K=4096调用时,Triton 会认为“这是新维度,需要重新测试”。这次可能是另一种配置胜出。
key=['M', 'N', 'K'] 表示只要其中任一值变化,就重新调优。适合 256x256 矩阵的最优配置,在 4096x4096 上可能表现很差。
自动调优内核还有几个重要细节:
独立的 BLOCK_SIZE_M、BLOCK_SIZE_N、BLOCK_SIZE_K:tile 不必是方形。A 的 tile 形状为 (BLOCK_SIZE_M x BLOCK_SIZE_K),B 的 tile 形状为 (BLOCK_SIZE_K x BLOCK_SIZE_N)。保持较小的 BLOCK_SIZE_K(32 至 64),可以减少每个 k 步骤的 SRAM 占用;同时允许较大的 BLOCK_SIZE_M 和 BLOCK_SIZE_N(128 至 256),以获得更多输出复用。
num_warps:更多 warp 意味着每个 program 有更多线程。更大的 tile 需要更多线程处理,因此大块配置(128x256)使用 8 个 warp,小块配置(32x64)使用 2 个。
num_stages:这是软件流水线。使用 2 个 stage 时,GPU 在计算当前 tile 的同时,会从 DRAM 加载下一个 tile;使用 3 个 stage 时,一个 stage 计算、一个加载、一个存储。更多 stage 需要更多 SRAM,以保存多个“在途”的 tile。
网格是 lambda:在 autotune 选出获胜配置前,我们不知道 block size,因此无法预先计算网格维度。lambda 会延后这一步计算:
grid = lambda meta: (triton.cdiv(M, meta['BLOCK_SIZE_M']) * triton.cdiv(N, meta['BLOCK_SIZE_N']), )
滑动指针,而不是每轮重新计算:与其在每次 k 循环中建立新的指针数组,不如直接把已有指针向前移动:
a_offsets += BLOCK_SIZE_K * stride_a_k
b_offsets += BLOCK_SIZE_K * stride_b_k
这会节省每次迭代中少量的计算工作。
使用 stride,而不是写死维度:自动调优内核使用 stride_a_m、stride_a_k 等,而不是把 * K 和 * N 硬编码进去。因此,即使面对非连续张量(转置视图、切片等),内核仍能正确工作。
测试结果
测试环境为 V100、float32,矩阵尺寸从 256 到 4096:
| 内核 | 峰值 TFLOPS | 说明 |
|---|---|---|
| PyTorch(cuBLAS) | ~13 | NVIDIA 多年工程积累的成果 |
| Triton:grouped + autotuned | ~12 | Group-major + 自动调优,共 15 种配置 |
| Triton:grouped | ~10.5 | Group-major,固定 block size = 64 |
| Triton:blockwise | ~10 | 简单二维网格,固定 block size = 64 |
自动调优内核达到了 cuBLAS 大约 90% 的性能。最后那 10%,正是 NVIDIA 内核工程师花费数月在寄存器级优化、warp shuffle 指令和 Triton 所抽象掉的架构专属技巧上的地方。
但用 50 行类 Python 代码达到 90%,确实令人印象深刻。对于许多实际场景,例如自定义融合算子、非常规形状和研究型内核,这已经完全够用。
令我意外的事
硬件调度器替你完成了很多工作。 我的第一个简单内核没有流水线、也没有 group 排序,却已达到 cuBLAS 约 75% 的性能。GPU 的设计会在某个 program 等待数据时切换到其他 program,以隐藏延迟。只要启动足够多的 program,就能获得大量“免费”的优化。
group 排序的帮助比我预期的小。 从 row-major 改为 group-major 大约只提升了 5%。理论上它应该更有效,在更大的矩阵或 L2 缓存较小的 GPU 上可能也是如此。但硬件 L2 缓存已经完成了一部分 group 排序试图显式实现的数据共享。
Autotune 比手动优化更有效。 我花时间实现 group-major PID 映射,得到 5% 提升;随后添加 autotune,只需花两分钟写配置字典,又带来 15% 提升。结论是:让硬件告诉你什么才快。
使用 float32 累加很重要。 当我尝试用 float16 累加时,512x512 矩阵的结果出现了明显数值误差。半精度下数百次加法会累积舍入误差。常见模式是:以 float16 加载、以 float32 累加、以 float16(或 matmul 输出使用 float32)存储。
BLOCK_SIZE_K 必须不小于 16。 我起初尝试 BLOCK_SIZE=4,得到了一个令人摸不着头脑的错误:“Input shapes should have K >= 16”。这是硬件约束:GPU Tensor Core 以最小宽度为 16 的块运行,不能再小。
哪些地方我会换一种做法
我会跳过 Python 的分块 matmul。 它在 Python 中并不会更快,而且六层循环可能让人困惑。它确实有助于理解概念,但我本可以从朴素 matmul 直接进入 Triton。
我会从自动调优内核开始,而不是简单内核。 加 autotune 非常容易,性能差异又很明显。先写出最简单、正确的内核,加上 autotune 后再继续迭代即可。
我会更多使用 TRITON_INTERPRET=1。 该模式会在 CPU 上借助 NumPy 运行内核,因此可以加 print 语句进行调试。我当时浪费了时间盯着错误输出,其实直接打印中间结果就行。
如果你也想学习
从 Triton 官方教程(triton-lang.org)开始,但不要只是阅读。亲手敲出来,故意把它改坏,修改 block size,看看会发生什么。移除 mask,观察哪里会出问题。
对我而言有效的学习路径是:
- 用 Python 写朴素 matmul,验证它与
torch.matmul的结果一致。 - 用 Python 写分块 matmul,面对一个
4x4矩阵,用纸笔跟踪其计算过程。 - 写简单的 Triton 内核。最难的是带广播的二维指针运算,把它画出来。
- 加入 group-major 排序。它不容易理解,但实现只要 6 行数学运算。
- 加入 autotune。它容易理解,也容易实现,能立即获得性能提升。
五个版本的完整代码及基准测试都在这里:jaygala223/tryton:通过动手尝试学习 Triton 语言。
GPU 内核编程是一种看上去难如登天、但亲自动手后会发现并非如此的技能。Triton 降低了入门门槛,让你能在一天内从零写出一个有竞争力的 matmul 内核。这里的概念,例如 tiling、内存层级和 occupancy(占用率),能够直接迁移到理解其他 GPU 内核,包括 FlashAttention 和融合优化器。
如果你已经读到这里,就去写一个内核吧。若觉得 matmul 太大,可以从向量加法开始。重要的是亲手接触 tl.load、tl.store 和 tl.program_id;其他知识都会建立在这三个基础之上。
最终结果


浙公网安备 33010602011771号