Triton Plugin Extensions: Enabling TLX and Custom Compiler Passes Out of the Box

TL;DR · AI 摘要
PyTorch 3.7引入Triton Plugin Extensions,支持动态加载自定义编译器传递,无需分叉Triton,性能匹配厂商库。
核心要点
- Triton Plugin Extensions允许动态加载自定义编译器传递,无需分叉或重新编译。
- Meta的TLX现在可直接使用,性能与NVIDIA H100和AMD MI350的厂商库相当。
- 插件通过TRITON_PLUGIN_PATHS环境变量加载,支持覆盖编译器流水线的各个阶段。
结构提纲
按章节快速跳转。
思维导图
用一张图看清主题之间的关系。
查看大纲文本(无障碍 / 无 JS 友好)
- Triton Plugin Extensions
- 问题
- 分叉维护成本高
- 解决方案
- 动态加载插件
- 无需分叉
- 实现机制
- TRITON_PLUGIN_PATHS加载
- 覆盖编译器流水线
- 应用案例
- Meta TLX性能达标
金句 / Highlights
值得收藏与分享的关键句。
PyTorch-Triton 3.7引入的Plugin Extensions系统,无需分叉即可动态加载自定义编译器传递。
Meta的TLX现在可直接使用,性能匹配NVIDIA H100和AMD MI350的厂商库。
插件通过TRITON_PLUGIN_PATHS加载,支持覆盖编译器流水线的各个阶段。
Triton插件扩展:开箱即用启用TLX和自定义编译器通道 – PyTorch
亮点项目
一句话概括
PyTorch-Triton 3.7版本引入了Triton插件扩展系统,这是一个用于在运行时动态加载自定义编译器通道、方言(包括其操作符)和DSL扩展的框架,无需分叉或重新编译。作为该系统的首个主要用户,Meta的Triton语言扩展(TLX)现已开箱即用,为标准Triton带来持久的GEMM内核和细粒度硬件控制,在NVIDIA H100和AMD MI350上性能可匹敌甚至超越厂商库。
问题:为什么需要扩展?
编写高性能GPU内核通常需要超越默认Triton编译器流水线提供的能力。自定义优化通道、硬件特定内联汇编和专用内存管理模式对于在生产负载中榨取最后一丝性能至关重要。此前,启用这些功能意味着需要维护Triton的分叉版本,而这种分叉会带来真实成本。
分叉版本很快就会落后于上游。每次上游更新都可能导致合并冲突、损坏的API和需要仔细协调的细微行为变化。使用分叉版本的团队发现自己被困在过时的发布版本中,无法利用上游的错误修复、新硬件支持和社区改进。维护负担随时间累积,分叉反而成为瓶颈而非加速器。
真正需要的是在不修改核心Triton的前提下,通过添加通道、操作符甚至完整方言来扩展Triton编译器流水线的方法。一个能在运行时动态加载扩展的插件系统,将使研究人员和工程师能够以全速迭代自定义功能,始终基于最新上游版本运行,并在无需等待更改合并到主仓库的情况下发布结果。
Triton插件扩展系统
PyTorch Triton 3.7版本正是为此而生:提供一个贯穿整个编译流水线、内置于上游Triton的通用插件扩展系统。插件是共享库(.so文件),通过TRITON_PLUGIN_PATHS环境变量在运行时被发现和加载。安装插件包时无需重新编译Triton,只需将环境变量指向插件位置,扩展即可立即可用。
可覆盖的编译器流水线
该系统的核心是嵌入在Triton后端compiler.py阶段的一组钩子。这些钩子在从高级Triton IR(TTIR)到TritonGPU IR(TTGIR),再到LLVM IR和目标特定汇编(PTX, AMDGCN)的每个降级层级上,提供对MLIR通道流水线的细粒度控制。通过这些钩子,插件可以:
- 在任意阶段的任意位置插入一个或多个自定义通道。
- 禁用某个阶段内的特定通道。
- 用专用的自定义实现替换现有通道(例如自定义的线程束优化策略)。
- 覆盖整个阶段或完整流水线。
该功能在NVIDIA和AMD后端均可用。
自定义操作符、方言和降级
插件API设计与PyBind11相辅相成,提供三个层次的可扩展性:
- 自定义转换通道:无需关联方言即可在流水线任意位置插入的单个通道。
- 自定义 MLIR 语言和转换过程:将独立编译的语言加载到 Triton 中,通过插件过程将标准 Triton IR 模式重写为自定义语言操作,实现专门的降级处理。
- 自定义顶级 DSL 操作:通过新的 Python 级语法和语义,实现全新的编程抽象,而无需修改 Triton 本身。
每个内核的控制
插件可以在内核级别动态启用或禁用。在内核代码中设置的编译器钩子会激活一个自定义流水线,所有在钩子设置后调用的内核都会使用该流水线,直到钩子被移除。自定义流水线的数量没有限制,插件需自行实现用于内核缓存管理的哈希策略——确保仅在必要时触发重新编译。这一切均由 utlx 库为用户处理。
# 启用 TLX 插件只需设置一个环境变量
import os
import sysconfig
dist_packages = sysconfig.get_paths()["purelib"]
libutlx_path = os.path.join(dist_packages, "utlx_plugin", "libutlx.so")
os.environ["TRITON_PLUGIN_PATHS"] = libutlx_pathTLX:Triton 语言扩展,现已内置
Triton 语言扩展(TLX)是 Meta 开发的一组硬件感知操作,用于显式内存管理和异步计算/加载流水线。TLX 为内核作者提供了对共享内存分配、数据移动和指令调度的直接控制——这些功能对于编写能够充分利用现代 GPU 硬件的持久内核至关重要。
核心 TLX 操作包括:
操作
描述
tlx.local_alloc(shape, dtype, num_buffers)为软件流水线分配共享内存缓冲区。
tlx.local_view(buffers, index)查看分配中的特定缓冲区。
tlx.async_load(src, dst, mask)从全局内存异步加载到共享内存。
tlx.async_load_commit_group(tokens)提交一组异步加载操作。
tlx.async_load_wait_group(n)等待异步加载组完成。
tlx.async_dot(a, b, acc)异步矩阵乘法-累加。
tlx.async_dot_wait(n, acc)等待异步点积操作完成。
tlx.local_store(dst, src)将数据存储到共享内存。
tlx.local_load(src)从共享内存加载数据到寄存器。
此前使用 TLX 需要从 Meta 的实验性 Triton 分支进行构建。通过插件扩展系统,TLX 现在作为独立的 Python 包(utlx)进行分发,可与未修改的上游 Triton 一起使用。从 PyTorch-Triton 3.7 开始,TLX 将在所有未来的 Triton 发布版本中默认启用。
跨硬件支持:NVIDIA H100 和 AMD MI350
TLX 的一个关键优势是相同的编程模型适用于不同硬件供应商,同时仍能映射到供应商特定的功能。
NVIDIA H100(Hopper)——持久 GEMM
在 Hopper GPU 上,TLX 映射到原生硬件 TMA(Tensor Memory Accelerator)异步加载和 WGMMA(Warp Group Matrix Multiply-Accumulate)指令。持久 GEMM 内核使用多阶段软件流水线与异步提交/等待组:
@triton.jit
def matmul_kernel_pipelined_hopper(
a_ptr, b_ptr, c_ptr, M, N, K,
stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_cn,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr,
NUM_STAGES: tl.constexpr,
):
# ... tile indexing ...
# 分配多阶段共享内存缓冲区
buffers_A = tlx.local_alloc((BLOCK_SIZE_M, BLOCK_SIZE_K), tlx.dtype_of(a_ptr), NUM_STAGES)
buffers_B = tlx.local_alloc((BLOCK_SIZE_K, BLOCK_SIZE_N), tlx.dtype_of(b_ptr), NUM_STAGES)
# 预取流水线序曲
for i in tl.range(0, NUM_STAGES - 1, loop_unroll_factor=NUM_STAGES - 1):
a = tlx.local_view(buffers_A, i)
b = tlx.local_view(buffers_B, i)
token_a = tlx.async_load(a_ptrs, a, mask=...)
token_b = tlx.async_load(b_ptrs, b, mask=...)
tlx.async_load_commit_group([token_a, token_b])
# 主K循环与计算和数据移动重叠
acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k in tl.range(0, tl.cdiv(K, BLOCK_SIZE_K), num_stages=0):
buf = k % NUM_STAGES
tlx.async_load_wait_group(NUM_STAGES - 2)
acc = tlx.async_dot(
tlx.local_view(buffers_A, buf),
tlx.local_view(buffers_B, buf),
acc
)
# 预取下一阶段 ...
acc = tlx.async_dot_wait(0, acc)
# 存储结果 ...NVIDIA H100 (Hopper) — 性能结果
下表显示了在NVIDIA H100上FP16 GEMM吞吐量,将标准Triton与TLX扩展插件对比cuBLAS的结果。由于插件系统生成的代码与编译内置的分支完全一致,这些结果适用于两种路径。
在主导生产LLM工作负载的大规模、计算密集型形状上,Triton + TLX在正方形GEMM上与cuBLAS性能相当,在宽大形状上则超越了cuBLAS,证实了插件加载路径不会引入任何开销,同时达到或超越厂商库性能:
128×13312×16384: cuBLAS 247.8 TFLOPS → Triton+TLX 257.0 TFLOPS (+3.7%) 16384×8192×8192: cuBLAS 549.4 TFLOPS → Triton+TLX 566.7 TFLOPS (+3.2%) 8192×16384×8192: cuBLAS 564.8 TFLOPS → Triton+TLX 575.9 TFLOPS (+2.0%) 8192×53248×8192: cuBLAS 571.3 TFLOPS → Triton+TLX 573.2 TFLOPS (+0.3%) 8192×28672×4096: cuBLAS 560.4 TFLOPS → Triton+TLX 559.8 TFLOPS (−0.1%) 8192×8192×8192: cuBLAS 582.3 TFLOPS → Triton+TLX 577.0 TFLOPS (−0.9%)
AMD MI350 — 流水线GEMM
在AMD MI350 GPU上,TLX使用显式的基于寄存器的流水线,通过local_store和local_load操作进行数据移动。相同的缓冲区管理模式适用,但数据移动路径经过寄存器而非异步硬件单元:
/think
@triton.jit
def matmul_kernel_pipelined_mi300(
a_ptr, b_ptr, c_ptr, M, N, K, ...
):
# ... tile indexing ...
# 分配共享内存缓冲区
buffers_A = tlx.local_alloc((BLOCK_SIZE_M, BLOCK_SIZE_K), tlx.dtype_of(a_ptr), NUM_STAGES - 1)
buffers_B = tlx.local_alloc((BLOCK_SIZE_K, BLOCK_SIZE_N), tlx.dtype_of(b_ptr), NUM_STAGES - 1)
# 序言:通过寄存器加载到共享内存
for i in tl.range(0, NUM_STAGES - 1, loop_unroll_factor=NUM_STAGES - 1):
a_reg = tl.load(a_ptrs, mask=...)
b_reg = tl.load(b_ptrs, mask=...)
tlx.local_store(tlx.local_view(buffers_A, i), a_reg)
tlx.local_store(tlx.local_view(buffers_B, i), b_reg)
# 主循环:重叠计算和内存操作
acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k in tl.range(NUM_STAGES - 1, K_ITERS, num_stages=0):
# 从寄存器加载下一个tile
a_reg = tl.load(a_ptrs, mask=...)
b_reg = tl.load(b_ptrs, mask=...)
# 对之前分阶段的数据进行计算
a_prev = tlx.local_load(tlx.local_view(buffers_A, buf))
b_prev = tlx.local_load(tlx.local_view(buffers_B, buf))
acc = tl.dot(a_prev, b_prev, acc)
# 将新数据存储到共享内存以供下一次迭代使用
tlx.local_store(tlx.local_view(buffers_A, ...), a_reg)
tlx.local_store(tlx.local_view(buffers_B, ...), b_reg)AMD MI350 — 性能结果
下表显示了在AMD MI350上FP16 GEMM吞吐量,将标准Triton与TLX扩展插件对比rocBLAS的结果。由于插件系统生成的代码与编译进的分支完全相同,这些结果对两种路径都适用。
Triton + TLX在所有测试矩阵尺寸上持续提供12-15%更高的TFLOPS,证实插件加载路径没有引入任何开销且超过了厂商库:
256×256×256: rocBLAS 4.4 TFLOPS → Triton+TLX 5.0 TFLOPS (+11.8%) 512×512×512: rocBLAS 29.4 TFLOPS → Triton+TLX 33.9 TFLOPS (+15.2%) 1024×1024×1024: rocBLAS 161.2 TFLOPS → Triton+TLX 180.8 TFLOPS (+12.1%) 2048×2048×2048: rocBLAS 445.1 TFLOPS → Triton+TLX 511.9 TFLOPS (+15.0%)
GPUMode Trimul 乘法更新验证
我们希望在生产环境中验证插件路径,同时在更接近实际问题的重型内核流水线中进行验证,而非单独的微基准测试。在GPU模式下,存在一个Trimul乘法更新——五个投影GEMM馈入批量矩阵乘法加上输出线性运算,封装在层规范化、sigmoid门控和置换操作中。PyTorch + torch.compile基线运行时间为19.2ms,主要由GEMM主导。通过TRITON_PLUGIN_PATHS将TLX插件加载到标准Triton中后,我们移除了 warp专用持久化GEMM(hopper_gemm_ws.py)。由于TLX将矩阵乘法暴露为另一个Triton内核,我们可以围绕它压缩流水线。最终的TLX-WS + 融合提交运行时间为12.0ms,在cuBLAS + torch.compile基线之上实现了1.61倍的加速,击败了libcuEquivariance和H100上所有其他SOTA实现。之后我们在B200上扩展了CLC流水线,进一步扩大了性能差距。最终结论是实际扩展集成只需在GPU模式下安装几个轮子文件并进行极少的设置开销即可完成。
完全相同的代码生成,无需分支
插件方法的关键验证在于生成的代码与Meta Triton分支生成的代码完全一致。TLX扩展插件会经过相同的MLIR降级流水线,唯一区别在于传递和操作是动态加载,而非编译到二进制中。
我们的演示验证了以下结果:
- 在NVIDIA H100上,持久化GEMM内核的PTX代码生成完全一致。
- 在AMD MI350上,流水线GEMM内核的AMDGCN代码生成完全一致。
- 性能相当——动态加载路径没有可测量的性能损耗。
用于验证的Colab笔记本和独立脚本如下,发布后将更新为指向官方PyTorch-Triton 3.7包:
入门指南
在上游Triton上使用TLX非常简单:
# 安装Triton(从源码)和PyPI包中的TLX扩展
git clone https://github.com/triton-lang/triton && cd triton
TRITON_EXT_ENABLED=ON pip install -e . --no-build-isolation && cd ..
# uTLX插件(发布在PyPI):
pip install triton-utlxutlx包包含预构建的扩展库。安装完成后,设置插件路径并导入扩展:
import os
import sysconfig
# 指向TLX插件
dist_packages = sysconfig.get_paths()["purelib"]
os.environ["TRITON_PLUGIN_PATHS"] = os.path.join(dist_packages, "utlx_plugin", "libutlx.so")
# 现在TLX操作符可在你的内核中使用
import triton
import triton.language as tl
import utlx_plugin as tlx要查看完整的演示和仓库:
- H100持久化GEMM笔记本:Colab
- H100独立脚本:GitHub Gist
- AMD MI350流水线GEMM:GitHub Gist
- 插件文档:triton/lib/Plugins/README.md
- 扩展仓库:triton-lang/triton-ext
- utlx PyPI项目:https://pypi.org/project/triton-utlx/
后续计划
插件扩展系统为社区开发的Triton扩展生态系统打开了大门。当前活跃开发和未来提案包括:
- 自定义后端:无需修改Triton构建系统即可动态加载Intel、CPU等目标的离树后端。
- triton-distributed:作为扩展的分布式计算原语。
- 自定义分析工具:以插件形式开发Proton和ConSan等工具的定制化版本,支持用户特定的运行时性能分析。
- 自定义优化传递:以附加组件形式提供目标特定的线程束特化、循环分割和模型特定优化。
- 专用操作符:2:4结构化稀疏性、自定义布局转换等。
我们期待社区参与,开发和实现扩展功能,为更广泛的Triton生态系统解锁新能力。如需贡献代码,请从triton-ext仓库和插件文档开始。
/post-content
/inner-wrap