返回
查看原链接原链接
Bilibili45分23秒 · —

Triton 与 TileLang:Tile-Level Programming 思想与实现解析

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 启动后,硬件按以下层次展开并行任务:

  1. Grid:kernel 启动后生成的逻辑上的全部任务集合。
  2. Block:grid 被划分为多个可独立调度的 block,它是真正被调度到 SM(流多处理器) 上的单位。
  3. Thread:每个 block 内包含多个线程,线程是最小的工作单位
  4. Warp:线程以 32 个为一组(一个 warp)执行,同一 warp 内的 32 个线程是同步执行的,且执行同一条指令

内存层次结构

CUDA 内存分为三个主要层次,从下到上速度递增、容量递减、权限越发私有:

内存层次说明特点
Global Memory最大的内存容量大,延迟较高,所有 block 均可访问
Shared Memory每个 block 内共享延迟较低,一般用于数据复用
Register每个线程单独拥有的寄存器访问最快,只能存当前正在算的值

CUDA 编程的问题

在编写 CUDA kernel 时,开发者必须在代码中手动完成以下工作:

  • 使用 blockIdxthreadIdx 等 ID 手动计算每个线程负责的数据下标。
  • 手动指定数据在三个内存层次中(global → shared → register)的搬运方式。
  • 手动处理同步问题(如 __syncthreads())和边界问题(如 grid 界限判断)。

总结:数据分块到线程执行的映射全部需要手写。CUDA 的优势是把所有可控制性能的东西都交给开发者,但也因此无形中增加了工作量和代码复杂度。


Tile-Level Programming 的引入

核心思想

回顾前面提到的 CUDA kernel 的常见操作模式:加载一部分数据 → 对其进行计算 → 写回原来的内存。这些操作本质上都是围绕一整块数据进行的。因此,引入 tile-level programming 的出发点便是:

  • 不再从每个线程的角度思考问题,而是将 block 内的线程工作抽象为一个对整个分块(tile) 的整体操作。
  • 省去手动管理线程细节的部分,由编译器负责把 tile 的操作自动映射到 block 内的具体线程上。
  • 开发者只需要对整个分块进行“块级别”的操作,从而简化实现难度。

Tile-Level Programming 的通用流程

在写一个基于 tile 的 kernel 时,通常包含以下固定步骤:

  1. 确定当前 tile 的位置:知道当前 program(或 tile)在整体数据中的位置。
  2. 从 global memory 加载数据:把当前要计算的这块数据加载到当前计算单元。
  3. 执行计算
  4. 写回结果:将计算结果存回 global memory。
  5. 处理边界(注意:不一定是最后一步,大概率夹杂在计算过程中):处理分块不整除时超出范围的部分,需要根据具体代码调整位置。

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 的 blockIdxthread_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 Z
CPU 入口部分解析
  1. 定义函数与变量:定义函数把 XY 相加,Z 存放结果,N 是向量长度。
  2. 确定块长BLOCK_SIZE):由开发者手动指定(此处为 1024),可以根据不同机器和实际效果灵活调整。
  3. 确定 grid 大小N / BLOCK_SIZE 得到分块数量。
  4. 调用 kernel:Triton 的特点是用方括号 [grid] 传入分块参数,表示将函数分成 grid 个部分并行调用。
GPU Kernel 部分解析

三步固定流程,也是写其他 kernel 的通用模板:

  1. 获取关键值
  • 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 的部分(处理最后一块不是整块的情况)。
  1. 加载数据:使用 tl.loadXY 在当前下标位置的值加载到本地临时变量。mask=mask 是语法,表示按照 mask 对数组下标进行屏蔽;other=0.0 作为 padding,对屏蔽掉的部分以 0.0 填充。
  1. 计算与写回:直接计算加法,再用 tl.store 写回全局内存。

核心差异:Triton 操作完全建立在整块基础上,把 XY 在当前块的整个向量取出来就直接计算,避免了繁琐的手写过程。这是“对向量进行操作”,而非“对逐个元素操作”。

矩阵乘法代码详解

分块思想推导

矩阵乘法的分块方法不像向量加法那么直观。考虑 C = A × B:

  • 如果先对 A 或 B 分块,会发现对 C 的影响比较模糊,某一 A 块可能影响 C 的多个位置,不容易拆出互不相关的独立部分。
  • 因此,从结果矩阵 C 入手:C 的每一位与其他位置无关,可以对 C 进行拆分,然后推导 C 的每一块由 A 和 B 的哪些贡献计算出来。

三次分块过程

  1. 第一次分块:对 C 分块(如深蓝色的一块),对应由 A 的若干行和 B 的若干列矩阵乘得到。
  2. 第二次分块:问题在于 A 的行向量和 B 的列向量可能仍然很大,需要继续对这两个矩阵进行分块。
  3. 第三次分块:对中间的循环变量 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 的简洁性来源于对部分控制权的放弃,这同时构成了它的局限:

  1. 编译器是黑盒:无法手动管理内存访问、数据布局等细节,这些全部由编译器决定。在 CUDA 中可逐行控制的操作在 Triton 中无法指定,只能通过调整 BLOCK_SIZE 等参数间接影响结果。无法指定如将两块数据错开放置等高级布局。
  2. 生成的汇编结果不可见:如果某个 kernel 性能很低,无法直观判断 Triton 代码与生成汇编之间的关系。profiling 时只能看到编译后的内容,难以定位原始代码的问题出处。
  3. 性能天花板受限:编译器自动生成代码意味着很多优化空间无法手动挖掘,性能与手写 CUDA 相比存在差距。
  4. 表达力不足:比较复杂的流水线(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 中的不同表达(不是严格逐行翻译,而是对应关系):

任务阶段CUDATileLang
任务划分blockIdx 确定当前块负责的输出位置,用 threadIdx 确定每个线程负责的元素,手动计算行列kernel 直接定义 Cgrid_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 代码时的主要组成部分:

  1. 定义与启动
  • 声明 kernel(类似 @triton.jit 的方式)。
  • 定义一个 lang 函数,创建 grid;一个实例对应一个 CTA(可理解为 block)。
  1. 内存相关的部分
  • TensorT.Tensor):从 global memory 拿数据。
  • sharedT.alloc_shared):shared memory 分配。
  • fragment:寄存器中的计算结构(每个线程各自寄存器的计算结果)。
  1. 数据流部分
  • 从 global memory 到 shared memory:使用 T.copy
  • 计算后的结果存储在寄存器的 fragment 中。
  • 最后用 T.copy 将寄存器内容写回 global memory。
  1. 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)为后续深入学习提供了进一步的阅读材料。