Triton IR 剖析:从 Python Kernel 到 GPU 指令,TTIR/TTGIR 全流程拆解

从头梳理 Triton 编译器从 Python DSL 一路 lowering 到 LLVM IR 的完整链路。上半部分讲 kernel 怎么变成 TTIR、MLIR 语法长什么样;下半部分拆四个关键优化 Pass——CSE、Canonicalizer、Inliner、RewriteTensorPointer,看完你就能对着 triton-opt 的输出不再发怵。


一、整体框架:Triton 编译器到底在干什么

Triton 是一个用 Python 写 GPU Kernel 的框架。你写一段类似 NumPy 的代码,它帮你编译成能在 NVIDIA/AMD GPU 上跑的高性能机器码。这件事听起来简单,中间其实藏了一条很长的 lowering 管线:

Triton DSL (Python)
    ↓  AST 解析 + 代码生成
  TTIR (Triton IR)         ← 高层张量语义,硬件无关
    ↓  Pass 优化 + lowering
  TTGIR (Triton GPU IR)    ← 引入线程/block/warp,绑定GPU执行模型
    ↓  lowering
  LLVM IR                  ← 通用底层IR
    ↓  NVPTX backend
  PTX / GPU Binary         ← 最终可执行指令

这条管线里,TTIR 描述"算什么",TTGIR 描述"怎么在 GPU 上算",LLVM IR 负责"最终怎么执行"。三层 IR 各司其职,背后依托的正是 MLIR 这套可扩展编译器基础设施。


二、上篇:从 Triton Kernel 到 TTIR

2.1 Kernel 长什么样 → TTIR 怎么来

在这里插入图片描述

以一个最简的 vector_add 为例:

@triton.jit
def add_kernel(x_ptr, y_ptr, output_ptr, n_elements, BLOCK_SIZE: tl.constexpr):
    pid = tl.program_id(axis=0)
    block_start = pid * BLOCK_SIZE
    offsets = block_start + tl.arange(0, BLOCK_SIZE)
    mask = offsets < n_elements
    x = tl.load(x_ptr + offsets, mask=mask)
    y = tl.load(y_ptr + offsets, mask=mask)
    output = x + y
    tl.store(output_ptr + offsets, output, mask=mask)

这段 Python 代码经过 Triton 前端处理后,会生成对应的 TTIR。TTIR 是一种 MLIR Dialect,里面的 op 直接对应你写的 tl.loadtl.storetl.arange 等操作,但完全不涉及线程模型、shared memory、warp 调度等 GPU 细节。它只关心张量级别的语义:从哪里加载、做什么运算、写到哪里去。

2.2 Triton 源码里的关键模块

在这里插入图片描述

Triton 的 Python 前端不是黑盒。几个核心模块的职责:

模块作用
triton/language/core.pyTriton 核心操作实现(load/store/dot/reshape 等)
triton/language/semantic.py语法层,负责把 Python AST 翻译成内部 IR
triton/language/standard.py常用算子的标准实现
triton/language/math.py数学运算(exp/sin/cos 等)
triton/language/random.py随机数生成
triton/language/extra/cuda/libdevice.pyCUDA 硬件相关的底层操作

如果你要 debug “为什么我的 kernel 生成的 IR 不对”,大概率要从 semantic.pycore.py 入手追溯。

2.3 Kernel 与 Triton IR 的映射关系

在这里插入图片描述

这一页 PPT 展示了 Python kernel 代码与生成的 TTIR 之间的对应关系。核心映射规律:

  • tl.loadtt.load
  • tl.storett.store
  • tl.arangett.make_range
  • tl.program_idtt.get_program_id
  • 算术运算 → arith.addf / arith.mulf 等 MLIR 标准 Dialect

这里有个关键认知:TTIR 不是 Triton 自己从零造的,而是基于 MLIR 的 Dialect 机制扩展出来的。算术运算用的是 MLIR 内置的 arith Dialect,内存操作用的是 tt(Triton)Dialect。这种"混搭"正是 MLIR 设计的精髓。


三、插播:MLIR 是什么,为什么 Triton 要用它

3.1 MLIR 的核心概念

在这里插入图片描述

在这里插入图片描述

MLIR(Multi-Level Intermediate Representation) 是 LLVM 项目下的可扩展编译器框架。它不像传统编译器只有一种 IR,而是允许你定义多种"方言"(Dialect),每种方言有自己的一套操作(Operation)、类型(Type)和属性(Attribute),然后通过 Pass 在不同方言之间逐层 lowering。

三个核心概念一句话:

  • Dialect:IR 的"方言",比如 arith(算术)、scf(结构化控制流)、tt(Triton 自定义)
  • Operation(op):IR 的基本计算单元,类似汇编里的指令
  • Pass:对 IR 做变换的模块,可以优化、重写、lowering

3.2 为什么 GPU 编译器非 MLIR 不可

GPU Kernel 编译涉及张量计算、线程/warp/block 调度、shared memory 分配、向量化、bank conflict 消除……这些概念横跨多个抽象层级。传统做法是为每一层手写一套 IR,结果就是 IR 之间信息割裂,Pass 复用困难。

MLIR 的方案是:所有层级共用同一套 IR 基础设施,通过 Dialect 表达不同层级的语义,通过 Pass 在不同 Dialect 之间渐进式 lowering。Triton 的 Dialect 栈长这样:

Dialect层级作用
tt (TTIR)高层张量语义,硬件无关
ttg (TTGIR)中层GPU 线程/block/warp/shared memory
arith / math通用算术与数学运算
scf / cf通用控制流
llvm底层对接 LLVM IR
nvvm底层NVIDIA PTX 专项

这套分层设计让每个 Pass 只需关注自己层级的优化,不用操心跨层问题。


四、下篇:Triton IR 中的优化 Pass

下篇聚焦编译器真正"干活"的部分——优化 Pass。Triton 在 TTIR 阶段会跑一系列 Pass,把 naive 的初始 IR 逐步转化成更适合 GPU 执行的形式。

4.1 Pass Pipeline 总览

在这里插入图片描述

在这里插入图片描述

NVIDIA GPU 对应的 make_ttir 优化 Pass pipeline 定义在 triton/backends/nvidia/compiler.py 中,由 python/triton/compiler/compiler.py 调用。整个 pipeline 的典型执行顺序:

TTIR 初始 IR
  ↓
Canonicalizer   ← 先把 IR 规范化
  ↓
CSE             ← 消除重复计算
  ↓
Inliner         ← 展开所有函数调用
  ↓
RewriteTensorPointer ← 拆解张量指针
  ↓
TTGIR (后续还有更多 Pass...)

下面逐个拆解。

4.2 CSE:公共子表达式消除

在这里插入图片描述

CSE(Common Subexpression Elimination) 做的事情非常朴素:如果同一个表达式被计算了多次,只算一次,后面直接复用结果。

比如 Triton kernel 里经常出现大量索引计算:

%off0 = arith.addi %base, %stride
%off1 = arith.addi %base, %stride   // 和上一行完全一样

CSE 会把 %off1 直接替换成 %off0,后续所有引用 %off1 的地方都改成引用 %off0

在 GPU kernel 场景下,CSE 的价值特别大。因为 GPU 的寄存器数量是硬限制,重复计算不仅浪费 ALU 指令,还会增加寄存器压力,间接导致 occupancy 下降。消除一次重复计算,可能就多塞一个 warp。

Triton 使用的是 MLIR 内置的 CSE Pass,核心逻辑是基于 op 类型 + 操作数 + 属性的哈希匹配。但 Triton 做了定制化处理:tt.loadtt.storett.reduce 这些有副作用的 op 不会被 CSE,防止把"读两次内存"错误优化成"读一次"。

执行命令:

triton-opt cse_before.ttir -cse > cse_after.ttir

4.3 Canonicalizer:规范化 Pass

在这里插入图片描述

Canonicalizer 的目标是把 IR 转成一种"规范形式"。它不是一个单独的大规则,而是大量小规则的集合——常量折叠、代数化简、冗余 op 消除等等:

  • x + 0x
  • x * 1x
  • 2 * 36
  • 空的 scf.if 块直接移除

Canonicalizer 通常放在 CSE 之前执行。原因很简单:规范化的 IR 让更多表达式长得一样,CSE 就能命中更多匹配。

Triton 的自定义 op(如 tt.expand_dimstt.broadcast 等)也都注册了自己的 canonicalization pattern,确保 Triton 特有的 IR 结构也能被规范化。

执行命令:

triton-opt canonicalizer_before.ttir -canonicalize > canonicalizer_after.ttir

4.4 Inliner:内联 Pass

在这里插入图片描述

Inliner 直接把函数调用替换成函数体代码。这是 GPU 编译里必须做的一步,不是什么"可选优化"——因为 GPU 硬件根本不支持通用 function call。所有逻辑必须在同一个 kernel 里展开,否则 NVPTX backend 根本没法生成合法代码。

Triton 会生成很多辅助函数(数学运算、向量化操作等),Inliner 把它们全部展开到 kernel 主函数里。展开之后通常会再跑一轮 Canonicalizer + CSE,因为内联后经常暴露出新的优化机会。

MLIR 的 Inliner 还有一个成本模型:如果函数体太大或者调用次数太多,它会根据启发式策略决定是否内联。但在 Triton 的 GPU kernel 场景下,基本是全量内联——没办法,硬件不答应。

执行命令:

triton-opt inliner_before.ttir -inline > inliner_after.ttir

4.5 RewriteTensorPointer:张量指针重写

在这里插入图片描述

这是 Triton 最核心、最特有的 Pass 之一。

Triton 引入了一个叫 “Tensor Pointer” 的概念——tt.ptr<tensor<128xf32>>。它不是传统的指针,而是一个携带了形状、布局、swizzle 等元信息的"智能指针"。在早期 IR 中,一个 tt.load 配合 Tensor Pointer 就能表达复杂的向量化加载,非常简洁。

但问题来了:LLVM IR 不认识什么 Tensor Pointer。它只认最原始的 i8* 加上偏移量。所以 RewriteTensorPointer 的任务就是把高级的 Tensor Pointer 拆解为:基地址 + 偏移计算 + 向量化 load/store + 内存 layout 变换

举个例子,Before:

%ptr = tt.make_tensor_ptr %base, [%stride], [%shape]
%val = tt.load %ptr : tensor<128xf32>

After RewriteTensorPointer:

  • %ptr 被拆成 base address + block offset + swizzle offset
  • 生成多个 llvm.load 或 PTX 向量化加载指令(如 ld.global.v4.f32
  • 可能会引入 shared memory promotion 和 coalescing 优化

这个 Pass 之所以叫 “Rewrite” 而不是 “Lower”,是因为它仍然运行在 TTIR/TTGIR 层面,产出的是"语义等价但更低级"的 Tensor Pointer 表示,还没到 LLVM Dialect。

执行命令:

triton-opt rewrite_before.ttir -triton-rewrite-tensor-pointer > rewrite_after.ttir

4.6 四个 Pass 总结

Pass干啥为什么重要
Canonicalizer规范化 IR,常量折叠 + 代数化简让后续 Pass 更容易命中模式
CSE消除重复计算省寄存器、省 ALU,提升 occupancy
Inliner展开所有函数调用GPU 不支持 call,必须内联
RewriteTensorPointer拆解张量指针为底层地址计算连接高层语义与 LLVM IR 的关键桥梁

五、学完能做什么

搞懂这套流程之后,你至少能做三件事:

  1. Debug IR 输出TRITON_PRINT_IR=1 打印各阶段 IR 后,你能看懂每一层在干什么,定位问题是出在 TTIR 生成还是某个 Pass 优化
  2. 手写 Triton Pass:如果要加自定义优化(比如针对某种 memory layout 做特殊处理),你知道应该在哪个 Dialect 层面介入
  3. 跨框架对比:理解了 Triton 的 lowering 管线,再看 TVM、XLA、IREE 的 IR 栈,会发现思路一脉相承——都是 Dialect + Pass 的分层 lowering

如果你想把某个环节彻底吃透,建议直接去读 triton/backends/nvidia/compiler.py 里的 make_ttir 函数,它是这条管线的"总控台",所有 Pass 的注册和顺序一目了然。


本文基于董贞汝老师《Triton-IR剖析》分享内容整理,融入 MLIR 基础概念与各 Pass 原理补充,截图均来自原始 PPT。
(内容由AI生成,仅供参考)
(内容由AI生成,仅供参考)

Logo

小龙虾开发者社区是 CSDN 旗下专注 OpenClaw 生态的官方阵地,聚焦技能开发、插件实践与部署教程,为开发者提供可直接落地的方案、工具与交流平台,助力高效构建与落地 AI 应用

更多推荐