L2.4 Triton 编译器源码级实战
三维坐标
layer: L2(数据与算子层)|level: Senior|pillar: 编程与编译上一篇我们站在「数据与算子层」的入口,看清了算子如何被喂进硬件。本文把镜头推到这条链路最硬核的一段——编译器。目标不是教你「会写 Triton kernel」,而是带你穿透
@triton.jit这层语法糖,沿着 DSL → AST → TTIR → TTGIR → LLVM IR → PTX 的多级 IR 一路看到底,理解 Triton 凭什么用百行 Python 逼近手写 CUDA 的性能,以及它如何跨硬件生态(NVIDIA / 昇腾)做后端适配。
学习目标
- 前置知识:读过 L2 前序(尤其 L2.3 CUDA 编程:知道 warp、shared memory、Tensor Core,建立「算子 → 机器码」的硬件直觉 );写过基础 Python;对 GPU 执行模型有 L1.1 级别的认知,但无需深入的 CUDA 手写经验(Triton 正是来降低这道门槛的)。「算子 → IR → 机器码」的多级下降范式本篇会带你逐级建立,后续 L2.7 编译器与 MLIR 会进一步深化。
- 学完产出:① 能画出 Triton 完整编译流水线(DSL → AST → TTIR → TTGIR → LLVM IR → PTX),并说清每一级「丢掉什么抽象、加入什么硬件细节」;② 能用「块级编程模型」一句话讲清 Triton 为什么写
program_id而非threadIdx,并解释「线程级脏活被 Layout 推导自动化」是它易写又高性能的根因;③ 能说清 Triton 性能魔法集中在 TTIR→TTGIR 这一步的硬件根因,区分blocked/mma/shared三种 Layout 的用途;④ 能对比「矩阵乘 + softmax」融合与纯逐元素融合的难度差异,理解 Layout 不兼容如何决定 fusion 边界;⑤ 亲手写一个 fused softmax kernel 并逐级 dump TTIR/TTGIR/PTX,肉眼验证tl.max如何一路 lowering 成shfl.sync树形归约。 - 阅读姿势:盯住一条主线——「Triton 的全部价值,是把手写 CUDA 时最折磨人的『线程映射 + 访存布局 + 流水编排』交给编译器 Pass 自动完成」。无论是块级编程模型、Layout 编码,还是多后端适配,本质都在回答同一个问题:如何让算法工程师用「想数据块」的方式,写出贴着硬件跑的 kernel。
背景与现状
Triton 是 OpenAI 开源的一门面向 GPU 的 Python 嵌入式 DSL(领域特定语言)+ 编译器。它的核心命题只有一句话:让算法工程师用「块级(block-level)编程模型」写出接近 cuBLAS/CUTLASS 性能的 GPU kernel,而无需手动管理线程、寄存器分配与 shared memory 调度。
要理解它的定位,先看清它解决的痛点:
- CUDA 太底层:手写 CUDA 要求工程师亲自处理
threadIdx/blockIdx的线程映射、shared memory 的 bank conflict、向量化访存与 warp 调度——这是一道劝退算法同学的高墙。 - PyTorch 太上层:
a @ b一行虽优雅,但算子边界固定,无法做 fusion(融合)。一个 softmax 要读写显存好几趟,带宽被白白浪费。 - Triton 卡在中间的甜点位:你以 「一个 program 处理一个数据块(block/tile)」 的视角写代码,由编译器自动完成线程级并行、访存合并与软件流水。fusion 是天然的——把 max、exp、sum、div 写在一个 kernel 里,数据只过一遍显存。
业界信号:PyTorch 2.0 的
torch.compile默认后端 TorchInductor 生成的 GPU 代码就是 Triton。这意味着哪怕你从不直接写 Triton,只要用了torch.compile,你的模型早已跑在 Triton 编译出的 kernel 上。FlashAttention、vLLM、各类 SOTA fused kernel 也大量以 Triton 实现——Triton 已是事实上的「GPU kernel 中间语言」。
从产业演进看:
- 2021 前:定制 kernel = 手写 CUDA,门槛极高,只有少数 infra 团队玩得动。
- 2021–2023:Triton 开源并被 PyTorch 收编为 Inductor 后端,「写 kernel」的人群从 CUDA 专家扩展到了算法工程师。
- 2023 至今:Triton 底层全面 MLIR 化,多级 Dialect IR(TTIR/TTGIR) 成为跨硬件(NVIDIA GPU、AMD ROCm、华为昇腾、国产 GPU)适配的统一抽象层——这才是本文要源码级拆解的核心价值。
原理与架构
理解 Triton 编译器,最有效的方式是沿着一段 @triton.jit Python 函数被编译的「数据流」,看它如何逐级 lowering(下降)到最终的 GPU 机器码。
2.1 Triton 多级 IR 编译流水线全景
Triton 的现代实现(2023 起)全面构建在 MLIR(Multi-Level Intermediate Representation) 之上,编译过程是一条多级 Dialect 逐步下降的流水线:
逐级读这张图——每一级 IR 都在「丢掉一些抽象、加入一些硬件细节」:
| 阶段 | IR 形态 | 关键职责 | 硬件感知度 |
|---|---|---|---|
| ① DSL → AST | Python AST | @triton.jit 装饰器拦截函数,用 ast.parse 拿到语法树,不真正执行 Python | 无 |
| ② AST → TTIR | Triton IR(MLIR Dialect) | CodeGenerator 遍历 AST,把 tl.load/tl.dot/tl.reduce 发射为 tt.* 算子;表达块级语义,仍是硬件无关的张量运算 | 无 |
| ③ TTIR → TTGIR | Triton GPU IR | 编译器核心 Pass convert-triton-to-tritongpu:为每个张量选择 Layout(布局编码)——blocked(普通访存)、mma(Tensor Core)、shared(共享内存)。这一步决定了线程如何分摊数据 | 高 |
| ④ TTGIR → LLVM IR | LLVM IR | convert-tritongpu-to-llvm:把块级算子 lowering 成线程级指令,插入 barrier、shared memory 读写、向量化 load | 线程级 |
| ⑤ LLVM IR → PTX | PTX / cubin | 复用 LLVM 的 NVPTX 后端生成 PTX 汇编,再由 ptxas 编成 cubin | 机器码 |
要点在于:Triton 性能的「魔法」几乎全部发生在 ③ TTIR → TTGIR 这一步。TTIR 还只是「数学上要算什么」,而 TTGIR 通过 Layout 编码回答了「这些数据该让哪些线程、用什么访存模式、要不要走 Tensor Core 来算」——这正是手写 CUDA 时最耗心力、最易出错的部分,被编译器的 Pass 自动化了。
2.2 块级编程模型:为什么是 program_id 而不是 threadIdx
手写 CUDA 思考的是单个线程:int i = blockIdx.x * blockDim.x + threadIdx.x。Triton 反其道而行——你思考的是一个 program(≈ 一个 CUDA block)负责处理一整块(tile)数据:
pid = tl.program_id(axis=0):当前 program 的编号(类似blockIdx,但粒度是「块」)。offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE):本 program 负责的一段连续索引(一个向量,而非一个标量)。mask = offsets < n:边界保护,处理「数据量不是 BLOCK_SIZE 整数倍」的尾块。tl.load(ptr + offsets, mask=mask):向量化访存——一条语句把一整块数据从 HBM 搬进来,编译器自动做访存合并(coalescing)。
本质区别:CUDA 让你管理「线程」,Triton 让你管理「数据块」。线程级的脏活(如何把 BLOCK_SIZE 个元素摊到 warp 的 32 个线程上、要不要向量化成
ld.global.v4)全部交给 TTIR→TTGIR 的 Layout 推导。这就是 Triton「易写」与「高性能」能兼得的根因。
2.3 三大支柱在本文的落点
| 支柱 | 在 Triton 编译链路中的体现 |
|---|---|
| 编程与编译(本篇核心) | DSL → 多级 IR → PTX 的整条 lowering 流水线,就是「编程与编译」支柱最具代表性的工业实现 |
| 基础硬件架构 | TTGIR 的 Layout 编码(mma / shared)直接对应 Tensor Core 与 shared memory——不懂硬件就读不懂 TTGIR |
| 训推框架 | torch.compile 的 Inductor 后端产出的就是 Triton kernel,上层框架性能强依赖这条链路 |
动手实践:fused softmax kernel + dump 各级 IR
实验目标:亲手写一个 Triton fused softmax kernel,然后逐级 dump 它的 TTIR / TTGIR / PTX,肉眼观察「同一段块级 Python 代码如何被 lowering 成硬件感知的 GPU 机器码」。产出物:三份 IR 文本 + 一段对照说明。
3.1 环境准备
# 推荐 Python 3.11,用 uv 或 venv 隔离
python3 -m venv .venv && source .venv/bin/activate
# 路径 A(首选):有 NVIDIA GPU —— Triton 会随官方 CUDA 版 PyTorch 一起装好
pip install torch triton
# 路径 B(无 GPU,仅验证逻辑):装 CPU 版 torch + triton
# pip install torch --index-url https://download.pytorch.org/whl/cpu
# pip install triton
硬件路径说明(重要):
- NVIDIA GPU 路径:能完整跑通并 dump 出真实的 PTX。这是观察 lowering 的「标准答案」路径。
- Mac / 无 GPU 路径:Triton 的 GPU 代码生成需要 CUDA,无 GPU 时无法 dump PTX。两条替代方案:①设
TRITON_INTERPRET=1走解释模式,在 CPU 上模拟执行以验证 kernel 逻辑正确性(但不会产生 GPU IR);②用 Google Colab 免费 T4 GPU(运行时 → 更改类型 → T4 GPU),免费跑通完整 GPU 路径与 IR dump。
3.2 代码:fused softmax kernel
把 softmax 的 max → 减 max → exp → sum → 除 sum 五步融合进一个 kernel,数据只过一遍显存——这正是 fusion 的价值所在。
import torch
import triton
import triton.language as tl
@triton.jit
def softmax_kernel(
out_ptr, in_ptr,
in_row_stride, out_row_stride,
n_cols,
BLOCK_SIZE: tl.constexpr,
):
# 块级编程模型:每个 program 负责输入矩阵的「一整行」
row_idx = tl.program_id(axis=0)
row_start = in_ptr + row_idx * in_row_stride
# 一条语句向量化加载整行;mask 处理 n_cols 不是 BLOCK_SIZE 整数倍的尾 部
col_offsets = tl.arange(0, BLOCK_SIZE)
in_ptrs = row_start + col_offsets
mask = col_offsets < n_cols
row = tl.load(in_ptrs, mask=mask, other=-float("inf"))
# 数值稳定的 softmax:先减去行最大值,避免 exp 溢出
row_minus_max = row - tl.max(row, axis=0)
numerator = tl.exp(row_minus_max)
denominator = tl.sum(numerator, axis=0)
softmax_out = numerator / denominator
# 写回结果
out_row_start = out_ptr + row_idx * out_row_stride
tl.store(out_row_start + col_offsets, softmax_out, mask=mask)
def triton_softmax(x: torch.Tensor) -> torch.Tensor:
n_rows, n_cols = x.shape
# BLOCK_SIZE 取大于等于列数的最近 2 的幂
BLOCK_SIZE = triton.next_power_of_2(n_cols)
out = torch.empty_like(x)
# grid:启动 n_rows 个 program,每行一个
grid = (n_rows,)
kernel = softmax_kernel[grid](
out, x,
x.stride(0), out.stride(0),
n_cols,
BLOCK_SIZE=BLOCK_SIZE,
)
return out, kernel
if __name__ == "__main__":
dev = "cuda" if torch.cuda.is_available() else "cpu"
print(f"[device] {dev}")
x = torch.randn(1823, 781, device=dev)
y_triton, kernel = triton_softmax(x)
y_ref = torch.softmax(x, axis=1)
# 正确性校验(解释模式下也能跑这一步)
print("[correctness] max abs diff =",
(y_triton - y_ref).abs().max().item())
3.3 dump 各级 IR
方法一:编译产物对象的 .asm 字典(推荐,最直观)。kernel[grid](...) 返回的编译产物对象携带各级 IR:
# 紧接 3.2 的 __main__,在 NVIDIA GPU 路径下执行
if dev == "cuda":
print("=" * 30, "TTIR", "=" * 30)
print(kernel.asm["ttir"]) # ① Triton IR(硬件无关)
print("=" * 30, "TTGIR", "=" * 30)
print(kernel.asm["ttgir"]) # ② Triton GPU IR(带 layout 编码)
print("=" * 30, "PTX", "=" * 30)
print(kernel.asm["ptx"][:2000]) # ③ PTX(截前 2000 字符即可)
方法二:环境变量 TRITON_KERNEL_DUMP(一次性落盘所有阶段)。无需改代码,自动把每一级 IR 写到磁盘:
# 把 TTIR / TTGIR / LLVM IR / PTX 全部 dump 到指定目录
TRITON_KERNEL_DUMP=1 TRITON_DUMP_DIR=./triton_dump python softmax.py
# 看 dump 出的目录结构(用 search_files / ls 等价命令查看)
# triton_dump/<kernel_hash>/
# ├── softmax_kernel.ttir # ① 硬件无关块级 IR
# ├── softmax_kernel.ttgir # ② 硬件感知 + layout 编码
# ├── softmax_kernel.llir # ③ LLVM IR
# └── softmax_kernel.ptx # ④ PTX 汇编
无 GPU 时的验证路径:
# 解释模式:CPU 上模拟执行,验证 kernel 逻辑正确(不产生 GPU IR)
TRITON_INTERPRET=1 python softmax.py
# 期望输出 [correctness] max abs diff = 1e-6 量级,证明 softmax 逻辑正确
3.4 运行与观察:lowering 的「丢抽象、加细节」
逐级对照你 dump 出的 IR,会看到一条清晰的下降轨迹:
- TTIR:保留块级语义。你会看到
tt.load、tt.reduce(对应tl.max/tl.sum)、tt.store,类型是tensor<...xf32>——完全看不到「线程」概念,纯粹是「对一整块数据做什么运算」。 - TTGIR:每个张量类型上多了
#triton_gpu.blocked<...>之类的 Layout 编码,描述「这块数据如何摊到 warp / 线程」。tt.reduce周围出现了 跨线程归约的结构——这就是 ③ Pass 注入的硬件感知信息。 - PTX:彻底是 线程级机器码——
ld.global.v4.f32(向量化访存)、shfl.sync(warp 内 shuffle 做归约)、bar.sync(同步屏障)。你在 Python 里写的tl.max(row, axis=0),最终变成了一串shfl指令在 warp 内做树形归约。
观察结论:
tl.max(row, axis=0)这一行 Python,在 TTIR 里是一个抽象的tt.reduce,到 TTGIR 里被赋予了「跨线程如何分摊」的 layout,到 PTX 里彻底展开成shfl.sync树形归约——这就是 Triton 编译器替你完成的、手写 CUDA 时最折磨人的那部分工作。
踩坑预警 (Gotchas)
TRITON_INTERPRET=1不产生 GPU IR:解释模式是在 CPU 上「假装」执行来验证逻辑,kernel.asm里不会有 ttgir/ptx。想看真实 IR 必须有 CUDA GPU。误以为解释模式能 dump PTX 是最常见的坑。BLOCK_SIZE必须是 2 的幂且为tl.constexpr:tl.arange(0, BLOCK_SIZE)要求编译期常量,且 Triton 的 layout 推导假设 2 的幂。直接传列数(如 781)会编译报错——务必用triton.next_power_of_2。- 首次编译极慢别误判:
@triton.jit是 JIT,首次调用会触发完整的 AST→PTX 编译(可达数秒),结果按 kernel hash 缓存。测性能前务必预热,否则把编译时间算进了 kernel 耗时。 mask缺失 = 越界读写:当n_cols < BLOCK_SIZE,不加mask会读到行外的脏数据甚至非法地址。tl.load的other=-float("inf")是为了让 padding 位置在 softmax 的max/exp中不产生贡献。kernel.asm的 key 随版本变化:不同 Triton 版本字典 key 可能是'ttir'/'ttgir'/'ptx'或带前缀。若KeyError,先print(kernel.asm.keys())确认实际键名,或改用TRITON_KERNEL_DUMP=1落盘更稳。
深入思考
下面三题每题先给题干,再用
<details>折叠一份图文并茂的参考答案。建议先合上答案自己想 3 分钟,再展开对照。