
如果你写过几行 Triton 算子肯定绕不开tl.reduce。很多教程会把 reduce 讲成“向量求和的一行 API”却没讲它背后那个叫 combine 的东西。我第一次啃 Triton 源码时就是被这个名字勾住的源码里没有 magic只有一层层把用户想法变成机器的过程。要理解 Triton 的高性能 reduce第一步就是搞清楚 combine 到底被编译成了什么所以这篇“Triton 源码阅读系列”的第一篇我就先拿它开刀。这篇文章不会逐行贴代码因为 Triton 版本迭代很快今天的主分支和半年前可能就差出好几个 pass。我会按主线逻辑拆开讲combine 在用户 API 里长什么样、在 MLIR 里长什么样、又是怎么一步步落到 LLVM 和 PTX 的。适合正在学 Triton 怎么写高效算子、想进一步读源码、或者已经被 reduce 的 layout 问题折磨过的人。1. 为什么要从 combine 开始读 Triton 源码1.1 用户眼里的 tl.reduce一个函数参数而已先看一个最常见的例子import triton import triton.language as tl triton.jit def sum_kernel(x_ptr, out_ptr, BLOCK: tl.constexpr): offs tl.arange(0, BLOCK) x tl.load(x_ptr offs) s tl.reduce(x, axis0, combine_fnlambda a, b: a b) tl.store(out_ptr, s)在这个 kernel 里combine_fnlambda a, b: a b就是我们要读的 combine。从用户视角看它只是一个“把两个元素合并成一个元素”的二元函数Triton 负责把 128 个元素归约成一个。问题在于这个“两个元素”是哪两个是相邻元素吗是同一个线程内的两个元素吗跨线程的归约又怎么处理如果不读源码你只能靠猜。我一开始甚至以为 Triton 会在运行时反复调用这个 lambda后来读源码才发现完全不是这么回事。1.2 combine 对应的源码位置不止一处Triton 的编译链路大致是Python AST - Triton IRTTIR- TritonGPU IRTTGIR- LLVM IR - PTX - 机器码combine 在这个链路里至少要经过三个关口Python API 入口python/triton/language/core.py里的tl.reduce语义分析python/triton/language/semantic.py里处理 combine_fn 的逻辑中间表示include/triton/Dialect/Triton/IR/TritonOps.td里的ReduceOp定义再往下走到 TritonGPU还会遇到 layout 分布、warp shuffle、共享内存归约等一堆问题。所以“读 combine”不是读一个文件而是读一条完整的数据流。抓住这条线后面再读 load、dot、convert_layout 都会轻松很多。2. combine 的编译链路从 Python 函数到 MLIR 区域2.1 前端不会真的“调用”你的 lambda阅读源码时最颠覆我认知的一点是Triton 前端不会像普通 Python 那样在运行时把你的combine_fn函数一遍遍调用。它做的是语法捕获。combine_fn通常是一个 lambda也可能是一个普通函数。编译器会把这个函数的 AST 解析出来把它转换成一个 MLIR 区域region。也就是说你写的是 Python 语法但编译器看到的是一个“可以被内联的表达式模板”。这也是为什么combine_fn里不能随便写带副作用的 Python 代码。你可以在里面用tl.maximum、tl.where、算术运算但不应该在里面修改外部变量、打印、调用任意 Python 库。那些操作没法被静态转换成 MLIR region。2.2 semantic.py 里到底做了什么我把semantic.py里 reduce 的主线逻辑整理成了伪代码# 伪代码只保留核心动作 def reduce(self, input, axis, combine_fn, keep_dimsFalse): input self.tensor(input) # 1. 检查 input 是不是 block tensor不是就报错 # 2. 把 axis 归一化成合法的非负整数 # 3. 根据输入 shape 和 keep_dims 计算输出 shape # 4. 从 combine_fn 的 AST 构建出 combine region combine_region build_combine_region(combine_fn) # 5. 创建 tt.reduce op reduce_op ReduceOp([input], axis, keep_dims, combine_region) return reduce_op.result这里有三个容易忽略的细节。第一个axis必须归一化。用户可能传-1也可能传一个 Python intsource 里会把它映射成实际维度避免后面 layout 转换时出现负数索引。第二个keep_dims会直接影响输出 shape。keep_dimsTrue时reduce 后的结果还保留那个长度为 1 的维度这对广播很重要如果忘记这个参数后续和其他 tensor 做运算时很容易出现 shape mismatch。第三个combine_fn的 AST 构建并不是简单把函数体抄一遍。它要处理参数名绑定lambda 的两个参数a、b必须映射到 region 块的两个入口参数。有时候你会看到lambda x, y: x y和lambda u, v: u v生成完全一样的 region就是因为参数名在 AST 阶段就被替换成了块参数。2.3 MLIR 里的 combine 区域长什么样tl.reduce落到 TTIR 之后大概长这样%sum tt.reduce(%x) { axis 0 : i32 } : (tensor128xf32) - f32 { ^bb0(%lhs: f32, %rhs: f32): %0 arith.addf %lhs, %rhs : f32 tt.reduce.return %0 : f32 }注意看这个 region块参数%lhs和%rhs的类型是f32不是tensor128xf32。说明 combine 操作的是“元素”而不是“整个 block”。这个设计很重要它意味着编译器可以把 region 里的算术逻辑当作标量指令嵌入到任意归约位置不管这个位置在线程内、线程束内还是跨 block。tt.reduce.return是这个区域里的 terminator表示“本次 combine 的返回值”。在 MLIR 里一个 op 可以带 regionregion 里可以有多个 block但 reduce 的 combine 区域一般只有一个 block。你如果去读TritonOps.td会发现ReduceOp的 region 定义非常克制这保证了后续 lowering 时不至于被复杂控制流绊住。3. 手把手读一遍核心实现3.1 先从 Python 入口的 reduce 入手打开python/triton/language/core.py搜索def reduce你会看到一个很薄的封装。它真正做的事情是转发到semantic.reduce同时处理一些参数兼容问题。我自己读源码时不太建议一上来就深挖 AST 解析细节容易陷进去。更高效的做法是先跑一个最小 kernel把kernel.asm[ttir]打出来看中间表示。比如刚才那个 sum kernel编译后可以看到tt.reduce的 axis 和 region 是不是符合预期。这一步能帮你把“Python API”和“MLIR 结构”对应起来。compiled_kernel sum_kernel[(1,)](x, out, BLOCK128) print(compiled_kernel.asm[ttir])Triton 编译后的对象里会带多个中间表示ttir是最接近前端的那一层。如果你看到tt.reduce的 region 里已经变成了arith.addf说明前端 AST 捕获成功。如果 region 还残留 Python 对象那大概率是版本问题或者 combine_fn 写得过于复杂。3.2 ReduceOp 的 MLIR 定义要点回到 IR 定义ReduceOp的几个关键字段值得记一下字段作用我的理解operands参与归约的输入 tensor可以是一个也可以是多个multi-reduce 时会用到axis在哪个维度归约必须是常量属性编译期确定keep_dims是否保留长度为 1 的维度影响 shape 推导和后续广播regioncombine 逻辑二输入单输出主体是一堆纯算术 op这里有个容易踩坑的点axis是I32Attr也就是说它必须在编译期固定。你不能让 axis 由某个 tensor 的值决定也不能在运行时用变量改变归约方向。如果你想做动态维度的归约得用循环加掩码模拟性能会差不少。3.3 往下走到 TritonGPUlayout 是主角一旦tt.reduce进入 TritonGPU 阶段问题就从“怎么合并两个元素”变成“合并哪两个元素”。这取决于 tensor 的 layout。举个最简单的例子一个tensor128xf32如果 layout 是blocked且每个线程持有连续的 4 个元素那么归约可以分两步每个线程先把本地持有的 4 个元素用 combine 区域内的逻辑归约成 1 个值。再把 32/64/128 个线程之间的部分结果用 warp shuffle 或 shared memory 归约成最终值。你写的arith.addf在这个过程中会被复用多次。它先出现在本地寄存器归约又出现在 shuffle 后的合并逻辑里。source 中的ReduceOpToLLVM这类 lowering 主要就是在做这件事把同一份 combine 逻辑插入到不同归约阶段。这也是了解 combine 价值的核心你只需要写一次“如何合并两个元素”编译器负责把它变成一台高效的归约机器。但反过来你的 combine 逻辑越复杂它在每个归约阶段被重复执行的代价就越高。比如lambda a, b: a b和lambda a, b: tl.max(a, b)都很便宜但如果在 combine 里塞一个除法或指数运算性能会肉眼可见地变差。3.4 从 LLVM 到 PTXshuffle 与 shared memory再往下LLVM IR 会出现类似__shfl_xor_sync或 shared memory load/store 的调用。这些是 GPU 上跨线程归约的底层原语。我在读代码时特别留意到Triton 并不会盲目使用 warp shuffle。它先判断当前 axis 对应的数据分布是否足够“对齐”。如果对齐到 warp 内部走 shuffle 是最快的如果数据分散到多个 warp或者跨 block就必须用 shared memory 中转。这个决策过程的很多细节藏在 layout 分析和 conversion pass 里combine 本身的逻辑反而很简单。所以如果你发现 reduce 性能不好不要第一时间怀疑 combine 写错了先去看kernel.asm[ttgir]里插了多少个convert_layout。有一次我的 kernel 比手写 CUDA 慢三倍问题不在 reduce而在前后各插了一次convert_layout数据来回倒了好几个来回。4. 常见问题与排查技巧实录4.1 combine_fn 里写了 if 怎么编不过很多刚接触 Triton 的人会在 combine_fn 里写def combine_fn(a, b): if a b: return a else: return b如果a、b是标量 tensor这个写法并不安全。因为合并逻辑最终会被放进 region 内联执行Python 的if对 tensor 值的判断语义会和tl.where不一样。源码里对这类控制流的处理非常保守与其纠结能不能写不如直接改成def combine_fn(a, b): return tl.where(a b, a, b)tl.where会生成真正的select指令两个分支都会被计算但 GPU 的 select 开销很低语义也更清晰。4.2 axis、keep_dims 不一致导致的 shape 问题常见报错是Cannot broadcast或者shape mismatch。比如二维 tensor 沿 axis0 归约得到的结果是一维的这时候再想和原来的二维 tensor 做广播必须把keep_dimsTrue打开# 错误示范输出 shape 是 (128,) 无法直接和 (64, 128) 广播 s tl.reduce(x, axis0, combine_fnlambda a, b: a b) # 正确示范输出 shape 是 (1, 128) s tl.reduce(x, axis0, keep_dimsTrue, combine_fnlambda a, b: a b)这个和我前面说的一样keep_dims不是装饰性参数它会影响整个 shape 推导链路。4.3 浮点归约顺序不确定tl.reduce默认不保证元素按从左到右的顺序归约。combine_fn 如果只是a b你得到的结果可能和 CPU 上顺序求和的结果有微小差异这是正常的。但如果你写的 combine 不具备结合律比如lambda a, b: a * 0.5 b * 0.5那结果会随着归约树的形状变化而波动。这种 combine 在语义上就不适合并行归约。读源码后你也会发现Triton 的并行归约基于一个明确的假设combine 必须满足结合律绝大多数情况下最好还满足交换律。如果你的业务确实需要严格顺序比如做某种状态累积应该用for循环而不是tl.reduce。4.4 如何确认 reduce 的性能瓶颈我建议固定三个检查点检查点方法目标TTIRkernel.asm[ttir]确认 reduce 的 axis 和 region 正确TTGIRkernel.asm[ttgir]统计convert_layout数量排查多余 layout 转换PTXkernel.asm[ptx]确认有没有生成shfl.sync或共享内存访问我见过很多性能问题最后都落在convert_layout上。reduce 本身并不贵贵的是为了归约而把数据从一种分布变成另一种分布。如果 TTGIR 里出现连续两次convert_layout十有八九是 layout 选择策略和你的num_warps、BLOCK尺寸不匹配。5. 实操心得与后续阅读建议5.1 我踩过的三个坑第一个坑是闭包变量。有段时间我在 combine_fn 里引用了外层 Python 变量作为缩放系数编译能过但运行时改这个系数并不会生效因为它在编译时就被固化成常量了。后来我把所有可变参数都改成 kernel 参数从tl.load或标量输入里取再也没有这种诡异行为。第二个坑是 axis 传值。我在做 softmax 时想对最后一维归约结果误传了axis1而输入恰好是二维确实也能编译但归约方向完全错了。排查了很久才发现是 axis 语义理解错了。建议写 kernel 时先打印一下x.shape再对照tl.reduce的 axis 文档确认方向。第三个坑是过度优化。我一开始觉得combine_fn写得越紧凑越好于是把一堆tl.where塞进去。结果代码可读性极差MPI 集群上报错时根本看不出问题。后来我改成一个命名函数并在函数里写了几个中间变量生成的中间表示反而更容易读定位问题快得多。5.2 下一步建议读什么combine 只是 Triton 源码阅读的一个入口。读完 reduce 之后我建议按这个顺序继续先读tl.load和tl.store的 lowering理解内存访问是怎么变成指针算术和掩码的。再读tt.dot和矩阵乘法相关 pass这是 Triton 能高效跑 FlashAttention、MLA 这些算子的核心。最后读convert_layout因为大多数性能问题都能在 layout 转换里找到答案。读源码不需要从第一个文件开始读。找到一个像 combine 这样小而完整的入口跟着它走一遍编译链路比单纯按目录读更不容易迷路。这也是我接下来这篇系列要继续用的方式。最后分享一个小技巧读 Triton 源码时手上最好备一个能跑 GPU 的环境每看完一个 pass就写一个最小 kernel 把中间结果打出来对照。源码里的类型推导和 region 处理光看很难有体感跑一次kernel.asm[ttgir]往往比读十段注释都管用。