TIRx:面向演进中的前沿机器学习(ML)内核的开放编译器栈
TIRx:面向演进中的前沿机器学习(ML)内核的开放编译器栈
今天,我们推出 TIRx——一个基于 Apache TVM 构建的开源、硬件原生 DSL(领域特定语言)与编译器,专为 ML 内核设计。它瞄准 AI 软件栈中快速演进的内核与快速迭代的硬件相交的部分:TIRx 目前编译到 GPU 和专用 AI 加速器,并设计为能够伴随后续各代硬件一同成长。同一套设计服务于专家编写的内核、智能体生成的内核以及巨型内核(megakernel)系统。
我们与更广泛的社区合作,在发布时提供了以下材料:
- PyPI wheel 和 Python 前端。一个嵌入 Python 的硬件原生内核 DSL,支持
@T.jit/@T.prim_func风格的编写、解析器工具,以及用于构造 TIRx 程序的 Python API。 - TIRx 内核库与基准测试。端到端示例,涵盖在 Blackwell GPU 上的 GEMM、注意力风格内核以及低精度算子。
- 现代 GPU 编程开放课程。这门精心策划的在线课程是卡内基梅隆大学机器学习系统课程的一部分,使用 TIRx 教授学生 面向机器学习系统的现代 GPU 编程。
你可以找到以下资源:
- GitHub:https://github.com/apache/tvm
- 文档:https://tvm.apache.org/docs/tirx/overview.html
- PyPI wheel:https://pypi.org/project/apache-tvm/0.25.0/
pip install apache-tvm==0.25.0 - 社区 TIRx 内核库:https://github.com/mlc-ai/tirx-kernels
- 面向机器学习系统的现代 GPU 编程:https://mlc.ai/modern-gpu-programming-for-mlsys/index.html
动机
内核 DSL(领域特定语言)在选择程序员与机器之间的正确边界时最为有效。对于成熟的内核和成熟的硬件,该边界可以是高层级的:编译器将线程分配、内存移动、布局细节和指令选择隐藏在紧凑的张量或图块(tile)抽象背后。Triton 是典型的例子,它的普及显示出这在既定内核模式上效果有多好。在前沿领域,同一边界承受着更大的压力。新的指令、内存空间、协作模式和内核算法往往在编译器具备自动化它们的内置机制之前就已出现。当这种情况发生时,高层级编译器通常会隐藏的部分,恰恰是专家仍然需要手动控制的部分。
TIRx(发音为“tier-ex”)通过选择一条更低、更显式的边界来回应,其组织围绕三个决策:
- 编排保持在硬件原生源码中。 流水线结构、同步、角色分配、内存放置和后端内建函数(intrinsic)是前沿领域最常需要专家控制的部分,因此 TIRx 将它们保留在源码中,而不是隐藏在可能尚未建模新功能的抽象之后。
- 重复出现的图块原语暴露给编译器。 执行作用域(execution scope)、张量布局和图块原语分发使得常见操作可以跨后端保持可复用、可分析和可移植,而无需将整个内核强制通过固定的编译器流水线。硬件原生控制的代价是工程工作量:为每个内核和每个后端手动编写每个操作是繁琐的。将重复出现的操作暴露为图块原语缓解了这一问题,这样作者可以复用分发的实现,而不是每次重写相同的数据移动或矩阵乘法。
- 新硬件首先以内建函数形式引入,随后再升级为图块原语。 一个新特性可以立即作为原生内建函数使用——一个薄薄的、后端特定的单硬件操作封装。一旦跨内核的使用模式稳定下来,它可以提升为图块原语:一个感知布局的操作,能够跨作用域、操作数和后端进行分发。核心抽象保持小巧,而为一个新特性添加内建函数永远不会破坏已有的内建函数。
结果是一个能够随硬件一起成长的 DSL 和编译器栈。这就是 TIRx 背后的核心设计哲学:保持基础小巧且显式,让后端库随着新一代加速器的到来而演进。
这使得 TIRx 位于像 TileLang 这样的系统之下。TileLang 也通过暴露内存作用域和流水线来相对于 Triton 降低边界,同时仍然将布局推断和线程绑定留给编译器。TIRx 有意将这些更高层次的问题留在其核心之外,并提供一个最小的基础,使此类系统可以在其上构建;我们正在与 TileLang 社区合作,将 TIRx 作为新的最小基础来支持 TileLang 编译。
同样的小巧、显式基础使得一种设计能够服务于几种追求峰值性能同时尽可能减少工程工作量的用户:专家编写的生产内核、智能体生成的内核,以及巨型内核系统——每一种都需要在原生层面的控制以及编译器可见的重复出现的操作。
本篇博文的其余部分将先介绍编程模型,然后依次讨论每个方向。
TIRx 编程模型
以下是该边界在实际中的样子。一个 TIRx 程序读起来就像结构化的原生内核:循环、分支、张量、同步、流水线状态和后端内建函数都是直接编写的。图块原语出现在需要使重复出现的硬件操作变得可复用和可分发的场合。三个要素承载了大部分模型。
执行作用域决定谁执行某个操作以及以何种粒度执行。两个东西选择它:控制流(选择进入某个区域的硬件角色)和原语命名空间(设置调用的粒度)。未限定的 Tx.* 调用在线程级别运行;Tx.wg.* 在 warpgroup 级别运行。像 T.ptx.elect_sync() 这样的谓词可以进一步将线程级调用缩小到单个发起线程。
张量布局通过存储优先的接口描述逻辑张量位于何处。一个图块可以位于全局内存、共享内存、寄存器、张量内存或加速器 SRAM 中。用户声明每个图块驻留在哪里以及其元素如何分布在通道(lane)、warp 和寄存器中;该声明保持附加到图块上。当调用一个原语时,编译器读取这些声明以选择实现。布局是一种存储描述,而不是循环变换工具:用户可以构造图块的布局,但绝不用布局来变换循环。
图块原语分发将一个调用转化为原生 IR(中间表示)。根据操作数布局、执行作用域和目标,或者显式的 dispatch= 提示,它选择匹配的实现:从全局到共享的拷贝解析为 TMA,从共享到寄存器的拷贝解析为 ldmatrix,从张量内存到寄存器的拷贝解析为 tcgen05.ld;矩阵乘法解析为 WGMMA、tcgen05 或脉动阵列指令。然后分发生成在整个图块上应用该指令所需的循环和寻址。
这些要素在作用域重要的地方结合起来。在下面的 GEMM 尾声(epilogue)中,warpgroup 作用域和线程作用域的原语位于同一区域:Tx.wg.* 调用跨 warpgroup 移动和转换一个图块,而最后通过显式发起线程谓词保护的线程作用域 Tx.copy_async 执行 TMA 存储。
以上摘录是简化的。要了解全貌,这里有一个完整的 FP16/BF16 GEMM 内核中的两个角色——TMA 生产者和张量内存写回。你不需要逐行阅读。关键是所有与编排相关的内容(流水线状态、屏障协议、角色选择、低级同步内建函数如 tcgen05.wait 和 cp_async.bulk)都保留在源代码中(此处应有具体代码示例,但原文未给出完整代码,我们保持结构)。