Triton 与 TileLang:Tile-Level Programming 思想与实现解析
Tile-level programming 的核心思想是从“每个线程做什么”转向“对每个数据分块做什么”,通过将线程级细节交给编译器处理来简化 GPU 编程,但这也带来了控制力与性能优化空间上的取舍。
核心要点
- 背景与动机:CUDA 编程需要手动管理从数据分块到线程执行的完整映射,逻辑繁琐且难以优化。Tile-level programming 提供了一个更抽象的视角,让开发者专注于对数据块的整体操作,而非单个线程。
- Triton(简写为 `triton`,代码中常用 `tl` 指代):强调从块(program)出发的自上而下的编程思想,保留数据分块决策权,但把线程级映射完全交给编译器。它在高级抽象下变得简洁,但存在“黑盒”问题,性能可控性弱。
- TileLang:核心主张是在 Triton 的简单与 CUDA 的复杂之间取得平衡,显式地表达内存层次结构和数据流,指导数据在 global/shared/register 之间的搬运,同时仍将处理线程细节的事务交给编译器。它在保留性能优化空间的同时提升开发效率,比 Triton 更可控,比 CUDA 更简洁。
- 选型建议:选择工具取决于性能瓶颈和需求——追求快速开发选 Triton;追求极致性能或硬件特化选 CUDA;希望在性能与开发效率之间取平衡、特别是优化数据流相关瓶颈时选 TileLang。
CUDA 编程回顾:从线程到内存的映射
理解 Tile-level programming 的前提,是先理解传统 CUDA 是如何将任务映射到 GPU 硬件上的。
硬件层级结构(任务划分)
在 CUDA 中,一个 kernel 启动后,硬件按以下层次展开并行任务:
- Grid:kernel 启动后生成的逻辑上的全部任务集合。
- Block:grid 被划分为多个可独立调度的
block,它是真正被调度到 SM(流多处理器) 上的单位。 - Thread:每个 block 内包含多个线程,线程是最小的工作单位。
- Warp:线程以 32 个为一组(一个 warp)执行,同一 warp 内的 32 个线程是同步执行的,且执行同一条指令。
内存层次结构
CUDA 内存分为三个主要层次,从下到上速度递增、容量递减、权限越发私有:
| 内存层次 | 说明 | 特点 |
|---|---|---|
| Global Memory | 最大的内存 | 容量大,延迟较高,所有 block 均可访问 |
| Shared Memory | 每个 block 内共享 | 延迟较低,一般用于数据复用 |
| Register | 每个线程单独拥有的寄存器 | 访问最快,只能存当前正在算的值 |
CUDA 编程的问题
在编写 CUDA kernel 时,开发者必须在代码中手动完成以下工作:
- 使用
blockIdx和threadIdx等 ID 手动计算每个线程负责的数据下标。 - 手动指定数据在三个内存层次中(global → shared → register)的搬运方式。
- 手动处理同步问题(如
__syncthreads())和边界问题(如 grid 界限判断)。
总结:数据分块到线程执行的映射全部需要手写。CUDA 的优势是把所有可控制性能的东西都交给开发者,但也因此无形中增加了工作量和代码复杂度。
Tile-Level Programming 的引入
核心思想
回顾前面提到的 CUDA kernel 的常见操作模式:加载一部分数据 → 对其进行计算 → 写回原来的内存。这些操作本质上都是围绕一整块数据进行的。因此,引入 tile-level programming 的出发点便是:
- 不再从每个线程的角度思考问题,而是将 block 内的线程工作抽象为一个对整个分块(tile) 的整体操作。
- 省去手动管理线程细节的部分,由编译器负责把 tile 的操作自动映射到 block 内的具体线程上。
- 开发者只需要对整个分块进行“块级别”的操作,从而简化实现难度。
Tile-Level Programming 的通用流程
在写一个基于 tile 的 kernel 时,通常包含以下固定步骤:
- 确定当前 tile 的位置:知道当前 program(或 tile)在整体数据中的位置。
- 从 global memory 加载数据:把当前要计算的这块数据加载到当前计算单元。
- 执行计算。
- 写回结果:将计算结果存回 global memory。
- 处理边界(注意:不一定是最后一步,大概率夹杂在计算过程中):处理分块不整除时超出范围的部分,需要根据具体代码调整位置。
Triton:简洁至上的 Tile-Level 抽象
基本定位与思维方式
Triton 与 CUDA 最根本的区别在于编写 kernel 的出发点:
- CUDA:从 thread 开始,自下而上构建。先决定每个 thread 做什么,再决定每个 block 做什么,最后决定整体做什么。
- Triton:基本单位是 program,可理解为一个分块过程。从 program(块)出发,自上而下考虑问题,只需描述每个 tile 的位置和形状,再对 tile 的整体进行操作,由编译器将 tile 内的计算映射到 thread 上。
关键点:Triton 保留了数据如何分块的决策权,但把线程级映射的权利交给了编译器。
program近似对应一个 CUDA block(CTA),program_id对应 CUDA 的blockIdx。thread_id不再显式出现。
向量加法代码详解
以下是一个简单的向量加法示例,用于说明 Triton 的具体写法:
import triton
import triton.language as tl
@triton.jit
def add_kernel(X_ptr, Y_ptr, Z_ptr, N, BLOCK_SIZE: tl.constexpr):
# —— GPU kernel 部分 ——
pid = tl.program_id(axis=0) # 第一个关键值:program ID
offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) # 第二个关键值:offsets
mask = offsets < N # 第三个关键值:mask(边界判断)
# 加载数据
x = tl.load(X_ptr + offsets, mask=mask, other=0.0)
y = tl.load(Y_ptr + offsets, mask=mask, other=0.0)
# 计算
z = x + y
# 写回
tl.store(Z_ptr + offsets, z, mask=mask)
# —— CPU 入口部分 ——
def add(x, y):
Z = torch.empty_like(x)
N = x.shape[0]
BLOCK_SIZE = 1024 # 块长,灵活可调
grid = (N // BLOCK_SIZE,) # 分块数量
add_kernel[grid](x, y, Z, N, BLOCK_SIZE=BLOCK_SIZE)
return ZCPU 入口部分解析
- 定义函数与变量:定义函数把
X和Y相加,Z存放结果,N是向量长度。 - 确定块长(
BLOCK_SIZE):由开发者手动指定(此处为 1024),可以根据不同机器和实际效果灵活调整。 - 确定 grid 大小:
N / BLOCK_SIZE得到分块数量。 - 调用 kernel:Triton 的特点是用方括号
[grid]传入分块参数,表示将函数分成grid个部分并行调用。
GPU Kernel 部分解析
三步固定流程,也是写其他 kernel 的通用模板:
- 获取关键值:
pid = tl.program_id(axis=0):当前 program 的 ID(第几个块)。axis=0表示维度——向量加法只有一维;对于二维矩阵乘法,axis=0是横坐标,axis=1是纵坐标。offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE):当前块在原始数组中的下标集合,从开头位置到块内偏移。mask = offsets < N:判断下标的合法性状态,屏蔽超出N的部分(处理最后一块不是整块的情况)。
- 加载数据:使用
tl.load将X和Y在当前下标位置的值加载到本地临时变量。mask=mask是语法,表示按照 mask 对数组下标进行屏蔽;other=0.0作为 padding,对屏蔽掉的部分以 0.0 填充。
- 计算与写回:直接计算加法,再用
tl.store写回全局内存。
核心差异:Triton 操作完全建立在整块基础上,把
X和Y在当前块的整个向量取出来就直接计算,避免了繁琐的手写过程。这是“对向量进行操作”,而非“对逐个元素操作”。
矩阵乘法代码详解
分块思想推导
矩阵乘法的分块方法不像向量加法那么直观。考虑 C = A × B:
- 如果先对 A 或 B 分块,会发现对 C 的影响比较模糊,某一 A 块可能影响 C 的多个位置,不容易拆出互不相关的独立部分。
- 因此,从结果矩阵 C 入手:C 的每一位与其他位置无关,可以对 C 进行拆分,然后推导 C 的每一块由 A 和 B 的哪些贡献计算出来。
三次分块过程:
- 第一次分块:对 C 分块(如深蓝色的一块),对应由 A 的若干行和 B 的若干列矩阵乘得到。
- 第二次分块:问题在于 A 的行向量和 B 的列向量可能仍然很大,需要继续对这两个矩阵进行分块。
- 第三次分块:对中间的循环变量 K 进行分块。最终 C 的答案是由 A 的每一块乘以 B 的每一块值的贡献加在一起得到的(累加)。A 和 B 在 K 维度上的分块是在 kernel 内部循环中完成的。
代码要点解析
# 假设已经有 A, B, C 矩阵,维度为 M×K, K×N, M×N
# BLOCK_M, BLOCK_N, BLOCK_K 是三次分块的块长
# 关键:stride 变量(步长)
# 矩阵在内存中是展开成一维数组存储的(行主序),需要用步长访问不同位置
stride_am = A.stride(0) # A 每增加一行所增加的步长(通常等于 K)
stride_ak = A.stride(1) # A 每增加一列所增加的步长(通常等于 1)
# 同理有 stride_bn, stride_bk, stride_cm, stride_cn
# 在 kernel 内:
pid_m = tl.program_id(axis=0) # 块的行编号
pid_n = tl.program_id(axis=1) # 块的列编号
# offset 计算:用一维坐标拼出二维位置集合
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
offs_k = tl.arange(0, BLOCK_K)
# 累加器
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# K 维循环(第三次分块在 kernel 内完成)
for k in range(0, K, BLOCK_K):
# 当前 K 块的位置
offs_kt = k + offs_k
# 计算 A 和 B 的下标位置(利用 stride 将一维与二维坐标对应)
a_ptrs = A_ptr + offs_m[:, None] * stride_am + offs_kt[None, :] * stride_ak
b_ptrs = B_ptr + offs_kt[:, None] * stride_bk + offs_n[None, :] * stride_bn
# 加载 A 和 B(带 mask 判断合法性)
a = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & (offs_kt[None, :] < K), other=0.0)
b = tl.load(b_ptrs, mask=(offs_kt[:, None] < K) & (offs_n[None, :] < N), other=0.0)
# 计算并累加(直接对块进行矩阵乘,细节由编译器处理)
acc += tl.dot(a, b)
# 计算 C 的位置,判断合法性,存储结果
c_ptrs = C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
tl.store(c_ptrs, acc, mask=(offs_m[:, None] < M) & (offs_n[None, :] < N))关键点说明:
- stride(步长):矩阵是二维数组,但在内存中按一维向量存储。步长用于便捷地访问二维位置。例如
stride_am是 A 竖着走一行需要的下标增量,stride_ak是横着走一列的增量。Triton 提供.stride()方法自动求解步长,无需手动计算。 - 维度坐标:program 的块标号是坐标形式(第几行第几列),不是一维编号;
axis=0对应竖着的行维度,axis=1对应横着的列维度。 - 累加器:K 循环过程中,每次 K 分块的乘积结果都累加到同一个
acc上(因为 C 的每一块是 A 对应行块与 B 对应列块在 K 维内积的结果)。 - 整体流程与向量加法完全相同,只是更复杂一些——但拿出来的关键量和流程是固定的,差异只在中间计算部分。整体思想依然是操作整块,而不是拆分内部元素。
性能说明:此代码仅作为教学示例讲解大致过程。要真正写出高性能代码还需要许多细节优化,不能直接端上去用。
Triton 的局限性与适用边界
Triton 的简洁性来源于对部分控制权的放弃,这同时构成了它的局限:
- 编译器是黑盒:无法手动管理内存访问、数据布局等细节,这些全部由编译器决定。在 CUDA 中可逐行控制的操作在 Triton 中无法指定,只能通过调整
BLOCK_SIZE等参数间接影响结果。无法指定如将两块数据错开放置等高级布局。 - 生成的汇编结果不可见:如果某个 kernel 性能很低,无法直观判断 Triton 代码与生成汇编之间的关系。profiling 时只能看到编译后的内容,难以定位原始代码的问题出处。
- 性能天花板受限:编译器自动生成代码意味着很多优化空间无法手动挖掘,性能与手写 CUDA 相比存在差距。
- 表达力不足:比较复杂的流水线(pipeline)操作在 Triton 中难以表达,复杂写起来会比较别扭。
小结:Triton 放弃了一部分人的控制权,换取了较高的开发效率。当细节控制成为瓶颈,或追求极高性能时,需要使用其他工具(如 TileLang 或回归 CUDA)。
TileLang:平衡控制与简洁的编程模型
设计动机
TileLang 的设计旨在避开 Triton 和 CUDA 各自的缺陷:
- Triton 太简单、太黑盒,看不见内存层次结构。
- CUDA 太复杂,需要手动控制一切。
TileLang 的核心主张是在两者之间取平均:
- 更强调显式地把层次结构写出来,将数据在哪一层、如何流动等影响性能的关键因素明确写进代码。
- 但线程级别的事情仍交给编译器处理。
- 代码更像一个数据流描述:显式定义 kernel、global memory、shared memory、按维度分块的循环,形成一条清晰的数据流水线(global → shared → 计算 → 写回)。
三类抽象层次
TileLang 提供三个层次的编程抽象,每个层次对应不同的控制力度和适用场景:
| 层次 | 抽象程度 | 特点 | 适用场景 |
|---|---|---|---|
| 第一层 | 较简略 | 内存和并行均为隐式控制,大部分由自动调度 | 较少使用,优于简洁但不显式 |
| 第二层(TILE 层) | 常用层 | 显式写出并行和内存控制,获得较好性能 | 最常用的 tile-level 视角 |
| 第三层(更细粒度) | 较细 | 可写入类似 thread 的细节(类似 CUDA) | 不常用,失去简洁性 |
一般使用的都是中间层次。
与 Triton 的对比
- Triton:更接近一个向量表达式。写一个 tile,至于它落在 shared memory 还是寄存器,直接由编译器决定。
- TileLang:更接近手写 CUDA kernel 的结构,需要明确写出把数据放进 shared memory、从 shared 开始算等流程,但优势是不需要考虑线程和同步的问题。
换而言之:Triton 更直接地让你描述“算什么”,TileLang 让你描述“怎么在内存层次里把数据串起来”。对于高性能 kernel,数据的流动方式往往决定性能本身。
CUDA 与 TileLang 的代码结构对比
以下是对比同一矩阵乘法在 CUDA 和 TileLang 中的不同表达(不是严格逐行翻译,而是对应关系):
| 任务阶段 | CUDA | TileLang |
|---|---|---|
| 任务划分 | 用 blockIdx 确定当前块负责的输出位置,用 threadIdx 确定每个线程负责的元素,手动计算行列 | kernel 直接定义 C 的 grid_dim,其中 bx, by 表示当前 tile 负责的输出位置;线程 ID 不直接对应 |
| 存储分配 | 用 __shared__ 显式分配 shared memory,缓存当前 K 分块中的 A 和 B | 使用对应函数(如 T.alloc_shared)显式分配 shared memory,做相同的事 |
| K 维循环 | 普通 for 循环,每次处理块宽为 BK 的数据 | 使用 T.Pipelined 表达式表达同样的分块过程;包含额外调度信息,迭代计算与数据加载流水化(提前加载下一块),避免机器空闲 |
| 数据搬运 | 手动对每个线程计算地址、判断越界、将元素逐个写入 shared memory,最后同步保证整个 tile 加载完成 | 用 T.copy 直接将数组和线程级操作概括为 copy,描述“从 global memory 哪个 tile 搬到哪个 shared memory tile”;合并访存和同步交给编译器 |
| 计算阶段 | 使用内层循环,每个线程读取 shared memory 更新累加器 | 直接使用 T.gemm(或类似)这个函数完成,提供更概括的语义;编译器进一步映射到硬件指令 |
| 循环末尾同步 | 第二个 __syncthreads() 保证当前块已用完,避免下一次覆盖 memory | 这类操作由编译器与 pipeline 一起组织,无需单独写出 |
| 结果写回 | 手动计算地址并写回 | 直接使用 T.copy 完成 |
总结:CUDA 将线程索引和其他线程级细节逐步展开;TileLang 只保留任务划分、内存层次、数据搬运和计算结构,将大量重复的线程级操作交给 primitive 和编译器完成。TileLang 写起来比 CUDA 简洁,但不至于像 Triton 那样不可控。
TileLang 常见 Primitive 说明
以下是实际编写 TileLang 代码时的主要组成部分:
- 定义与启动:
- 声明 kernel(类似
@triton.jit的方式)。 - 定义一个
lang函数,创建 grid;一个实例对应一个 CTA(可理解为 block)。
- 内存相关的部分:
Tensor(T.Tensor):从 global memory 拿数据。shared(T.alloc_shared):shared memory 分配。fragment:寄存器中的计算结构(每个线程各自寄存器的计算结果)。
- 数据流部分:
- 从 global memory 到 shared memory:使用
T.copy。 - 计算后的结果存储在寄存器的 fragment 中。
- 最后用
T.copy将寄存器内容写回 global memory。
- Pipeline:
T.Pipelined让访存和计算重叠,隐藏 global memory 的高延迟。T.copy负责数据搬运与对齐。
完整代码流程示例
# 以矩阵乘法为例
# 1. 定义 kernel,声明为 no 函数
# 2. 将 A, B 从 global memory 取出,开始循环
# 3. 循环中:计算块的行/列坐标,确定线程数
# 4. 从 shared memory 中加载 A 和 B 的结果
# 5. 放入 C 的 local(fragment,即寄存器)——结果是放在单线程的寄存器中的
# 6. 清空累加器,沿 K 维按 pipeline 进行计算
# 7. 最后用 copy 将答案写回 global memory三者对比与选型总结
三种 GPU Kernel 写法的定位
CUDA、Triton、TileLang 三者不是包含关系,而是同一条发展线上侧重点不同的三个点:
| 工具 | 控制力度 | 优点 | 缺点 | 适用场景 |
|---|---|---|---|---|
| CUDA(C++) | 最细 | 可完全控制性能;灵活实现特殊机制 | 代码繁琐、开发效率低 | 极端性能调优、特殊硬件机制、硬件特化 |
| Triton | 最粗 | 快速简洁地写出简单 kernel | 黑盒、不可控、性能可能受限 | 向量相加、逐元素相乘等简单 kernel |
| TileLang | 中间 | 显式表达内存层次和数据复用,保留 CUDA 结构,同时简洁 | 比 Triton 控制强但仍不如 CUDA 精细 | 需要理解高性能 kernel 的内存结构、数据流优化 |
选择依据:瓶颈在哪里
选择工具主要取决于当前面临的主要瓶颈:
- 追求写作容易 & 算子不复杂:直接用 Triton。
- 对特定硬件做特化:用 CUDA 更合适。但极端手调优化的泛化性略逊——换一张卡性能可能受影响。
- 瓶颈是性能,且性能取决于内存层次的数据搬运和 shared memory:大部分 kernel 的性能瓶颈都卡在这里,使用 TileLang 把数据流显式写出来更利于优化。TileLang 的可移植性更强。
核心竞争力总结
- Triton 的核心优势:开发效率最高,思考方式简洁(从块出发)。
- TileLang 的核心优势:在保留显示数据流控制的同时提升开发效率,解决“Triton 太黑盒、CUDA 太复杂”的问题。
- CUDA 的核心优势:控制力最细,对于需要硬件特化的场景不可替代。
总结
Tile-level programming(以 Triton 与 TileLang 为代表)是对传统 CUDA 编程范式的重要补充。其根本思路是将编程视角从线程级提升到数据块级,在开发效率和性能控制之间建立不同的平衡点:
- 逻辑映射:CUDA 需要开发者手动完成每个 block/thread 与数据分块之间的映射;Triton 保留了块级操作但自动化了线程映射;TileLang 则在中间显式控制内存层次与数据流,同时把线程级细节交给编译器。
- 思想路径:CUDA 是自下而上(从线程构建 block/grid),Triton 是自上而下(从 program/块出发),TileLang 是在中间层显式控制数据流。
- 性能与开发的取舍:控制力越细,性能优化空间越大,但开发复杂度越高;抽象层级越高,开发越高效,但性能天花板受限。三者各有优势区间,开发者应根据具体需求(瓶颈是性能、可移植性、还是开发速度)在三条路径中做出选择。
参考资料:文中提到的官方文档(Triton 与 TileLang)为后续深入学习提供了进一步的阅读材料。