triton中的progran/tile/mask

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