三年前第一次用CUDA写矩阵乘,对着 __shared__warp divergencebank conflict 调了一整个通宵,运行时间从 12ms 降到 8ms。隔壁实习生用 Triton 花了 30 分钟写出 30 行 Python,跑出来 7ms。那一刻想摔键盘。

这是很多人接触 Triton 时的真实心路历程。它不是另一个深度学习框架——而是一把专门针对 GPU Kernel 开发的"自动挡"螺丝刀。


1. Triton 到底是什么

简单说:Triton = Python DSL + 分块编译器

2019 年 Philippe Tillet 在 MAPL workshop 上发布了这篇开创性论文——Triton: an intermediate language and compiler for tiled neural network computations。拆开三个关键词:

关键词含义
Intermediate Language基于 Python 的领域特定语言(DSL),不是直接写 C++/CUDA
Tiled Neural Network Compute自动分析神经网络计算模式,实施分块(Tiling)策略
Compiler编译器负责把 Python DSL 变成高效 GPU 机器码

在这里插入图片描述

核心思路很"反直觉":手动管理 GPU 的 thread block 和 shared memory 是反人类的,这件事应该交给编译器做。


2. Triton 在整个生态中的位置

AI 编译器赛道一直很拥挤。XLA、TVM、MLIR 各自画了一块地盘,而 CUDA 作为 NVIDIA 的亲儿子长期占据底层话语权。Triton 的定位很精准:

在这里插入图片描述

和深度学习编译器(XLA/TVM)比:XLA 和 TVM 做的是"把整张计算图编译优化",Triton 只做单算子 Kernel 级别的开发和编译。一个是"整栋楼的施工图",一个是"单个零件的锻造工艺"。

和 CUDA 比:CUDA 让你手动管理一切——线程束、共享内存、寄存器、内存合并——灵活度拉满但门槛高得离谱。Triton 用 Python 级别的抽象让开发者专注算法逻辑,编译器自动搞定分块和内存优化。

换句话说:CUDA = 手动挡赛车,Triton = 带拨片换挡的性能车——你仍然能跑很快,但不用踩离合了。


3. Triton 语言层:Python 语法,GPU 灵魂

手写 CUDA Kernel 开发是个什么体验?

  • Float32 / BF16 / Int8 各种精度都来一套
  • Linear / Convolution / Normalization / Pooling 每种算子都写一遍
  • 每个 Kernel 都要手动调 block size、grid size、shared memory 分配

高性能和低灵活性这道选择题,Trade-off 了几十年。

在这里插入图片描述

Triton 把解题思路改了:

@triton.jit
def matmul_kernel(a_ptr, b_ptr, c_ptr,
                  M, N, K,
                  BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr):
    pid = tl.program_id(axis=0)
    # 你只管写分块逻辑,剩下的线程调度交给编译器
    ...

你不需要知道 GPU 上有多少个 SM、每个 warp 有多少个线程。你只需要定义 BLOCK_SIZE,告诉 Triton “把数据切成多大的块”。编译器会自动把这块逻辑映射到 GPU 硬件上。

编程模型简化带来的红利:一个 Linear 算子同时支持 FP32/BF16/INT8,不再需要为每种精度各写一套 Kernel。


4. Triton 编译器:MLIR 驱动的多层优化管线

这才是 Triton 的真功夫。

在这里插入图片描述

Triton 编译器构建在 MLIR(Multi-Level Intermediate Representation) 之上,分三层逐级下降:

第一层:Triton Dialect(TTIR)——“干净的算法图”

这一层保留用户写的原始语义。操作是张量级别的(tt.dottt.load),不涉及任何硬件细节。编译器先在这里做一轮通用优化:公共子表达式消除(CSE)、死代码消除(DCE)、函数内联(Inlining)。

第二层:Triton GPU Dialect(TTGIR)——“魔法发生的地方”

关键创新:Layout 被编码进类型系统。

一个 Tensor 不只是 tensor<128xf32>,而是 tensor<128xf32, #blocked_layout>#blocked_layout 告诉编译器数据如何在 128 个线程之间分布。如果这个 Tensor 要喂给 tt.dot(矩阵乘),编译器会自动插一条 convert_layout 把它变成 #mma_layout(针对 Tensor Core 优化的交错排布)。

这一层的优化 pass 包括:

  • Pipeline / Prefetch:软件流水线,把数据搬运和计算重叠起来
  • Coalesce:确保连续线程访问连续内存地址,榨干显存带宽
  • Remove Layout:消除不必要的 Layout 转换开销

第三层:LLVM — “落地成 PTX”

Triton 把优化好的 TTGIR 转成标准 LLVM IR,然后借 LLVM 后端生成 PTX 代码——NVIDIA GPU 真正执行的指令。这层涉及最硬核的指针运算:BlockedLayout 里每个线程到底该读哪些偏移量,全在这层算清楚。

三个层次一个管线跑完,Python 进去,PTX 出来。


5. 为什么这很重要

传统深度学习开发有个割裂的痛点:

  1. 算法研究员用 PyTorch 写 Python,开心又愉快
  2. 性能不够,需要写 CUDA Kernel
  3. 把 PyTorch 丢给系统工程师,手动翻译成 CUDA,来回对齐精度

Triton 直接让研究员在 Python 层面就能写出接近手写 CUDA 性能的 Kernel。不用等人翻译,不用跨团队沟通。

而且 Triton 天生兼容 PyTorch——Kernel 可以直接挂到 torch.autograd.Function 里:

class TritonMatMul(torch.autograd.Function):
    @staticmethod
    def forward(ctx, a, b):
        output = torch.empty(...)
        matmul_kernel[(grid,)](
            a, b, output,
            M, N, K,
            BLOCK_M=128, BLOCK_N=128, BLOCK_K=32
        )
        return output

6. 参考资料

写这篇博客时主要参考了以下资料,建议深入阅读:

  1. Triton 原始论文Philippe Tillet et al., “Triton: an intermediate language and compiler for tiled neural network computations,” MAPL 2019.
  2. Triton 官方文档triton-lang.org
  3. MLIR 相关演讲jokeren.tech/slides
  4. OpenAI Triton 发布博客openai.com/index/triton
  5. Triton MLIR 发布解读superjomn.github.io

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传


封面页嘉宾信息:董贞汝,先进编译实验室(Advanced Compiler Lab)。本博客以 PPT 主线为基础,结合公开资料对 Triton 编译器架构、MLIR 优化管线等概念做了补充展开。

更多推荐