在 OpenAI 开发的 Triton 语言中,Program、Tile和Mask是构建其“以 Block 为中心(Block-Level)”编程模型的三个核心基石。
理解这三者的关系,可以用一句话概括:
- Program决定了“当前任务在哪个块(Grid)上跑”;
- Tile决定了“当前块要处理哪一整片数据”;
- Mask决定了“这片数据里哪些是安全的有效数据”。
1. Program(程序实例 / 任务网格)
在 CUDA 中,你写的 Kernel 是基于单个 Thread(线程)的视角;而在 Triton 中,代码是基于单个 Program(线程块)的视角编写的。
- 对应概念:类似于 CUDA 中的
Thread Block。 - 核心 API:
tl.program_id(axis) - 机制与作用:
- 当你 Launch 一个 Triton Kernel 时,启动的是一个三维网格(Grid:
x, y, z)。每一个独立的处理单元就是一个Program。 tl.program_id(axis=0)用于获取当前 Program 在网格中的唯一索引(等价于 CUDA 的blockIdx.x)。- 开发者利用
pid来计算当前 Program 负责处理的全局数据偏移量。
2. Tile(分块 / 数据张量块)
Tile(在 Triton 官方文档中常称为Block或Tensor Tile)是 Triton 最具特色的设计。Triton 摒弃了手写 Thread 循环,直接对一维/二维的矢量矩阵块(Tile)进行操作。
- 对应概念:一整块连续或有规律步长的数据集合(如128128128个元素,或64×6464 \times 6464×64的矩阵块)。
- 核心 API:
tl.arange(start, stop),广播操作[:, None]/[None, :] - 机制与作用:
- 隐式并行:在 Triton 中,
tl.arange(0, 128)会自动生成一个长度为 128 的索引 Tile。对 Tile 执行+、*或tl.dot(),编译器会自动映射到底层硬件多线程去并行计算。 - 二维矩阵块构建:通过 numpy 风格的增加轴(Broadcasting),可以极其轻松地构建出二维 Tile 指针:
offs_m=pid_m*BLOCK_M+tl.arange(0,BLOCK_M)# 1D 行索引 [BLOCK_M]offs_n=pid_n*BLOCK_N+tl.arange(0,BLOCK_N)# 1D 列索引 [BLOCK_N]# 构建 2D Tile 指针:(BLOCK_M, 1) + (1, BLOCK_N) -> (BLOCK_M, BLOCK_N)a_ptrs=a_ptr+(offs_m[:,None]*stride_a_m+offs_n[None,:]*stride_a_n)3. Mask(掩码 / 边界保护)
在 GPU 硬件中,Tile 的尺寸(BLOCK_SIZE)通常要求是2 的幂次方(如 32, 64, 128, 256),以便充分利用硬件对齐与 Warp 调度。然而,实际传入的动态矩阵尺寸(如N=100N=100N=100)往往不能被BLOCK_SIZE整除。
- 对应概念:布尔张量(Boolean Tile),指示 Tile 中每个位置是否合法。
- 核心 API:
mask = offsets < N,结合tl.load()与tl.store() - 机制与作用:
- 防止内存越界:在加载(
tl.load)或写回(tl.store)Tile 时传入mask。对于mask为False的位置,load会自动填充为安全值(如0.0),store则不执行物理写入。 - 消除分支分化:开发者不需要手写
if-else条件判断,硬件底层会通过条件指令/Predicate Mask 高效执行,避免了 GPU Warp Divergence。
完整代码协同示例(矢量加法 Kernel)
下面是一个将三者结合的典型 Triton 代码片段:
importtritonimporttriton.languageastl@triton.jitdefadd_kernel(x_ptr,# 输入向量 X 的指针y_ptr,# 输入向量 Y 的指针output_ptr,# 输出向量的指针n_elements,# 向量总长度 (例如 100)BLOCK_SIZE:tl.constexpr,# Tile 的固定尺寸 (例如 128)):# -------------------------------------------------------------# 1. Program: 获取当前 Program ID,确定任务分工# -------------------------------------------------------------pid=tl.program_id(axis=0)# -------------------------------------------------------------# 2. Tile: 构建当前 Program 负责的数据块索引 (Tile)# -------------------------------------------------------------block_start=pid*BLOCK_SIZE offsets=block_start+tl.arange(0,BLOCK_SIZE)# [BLOCK_SIZE] 维度的 Tile# -------------------------------------------------------------# 3. Mask: 计算边界保护掩码,防止尾部 Block 访问溢出# -------------------------------------------------------------mask=offsets<n_elements# 布尔 Tile,越界元素为 False# -------------------------------------------------------------# 协作执行:使用 Mask 安全地从指针加载与存储 Tile# -------------------------------------------------------------# mask=mask 保证越界位置不报错,other=0.0 填补越界处的默认值x=tl.load(x_ptr+offsets,mask=mask,other=0.0)y=tl.load(y_ptr+offsets,mask=mask,other=0.0)output=x+y# 使用 mask 保证不写穿非法内存区tl.store(output_ptr+offsets,output,mask=mask)三者核心关系总结
| 概念 | 物理/逻辑映射 | 解决的问题 | 代码常用模式 |
|---|---|---|---|
| Program | GPU Thread Block / Grid 坐标 | 决定任务分发与并行块定位 | pid = tl.program_id(0) |
| Tile | 片上寄存器/共享内存中的数据小块 | 提供向量化/张量化计算表达,代替单线程指针运算 | offs = pid * B + tl.arange(0, B) |
| Mask | 硬件 Predicate 掩码寄存器 | 解决尾块对齐与越界防护,避免分支惩罚 | tl.load(ptr, mask=mask, other=0) |