Triton CUDA API 参考手册
本文档提供 Triton 核心 API 的详细参考,包括函数签名、参数说明和使用示例。
1. 内核装饰器
@triton.jit
@triton.jit
def kernel_function(...):
pass
- 作用: 将 Python 函数编译为 GPU 内核
- 约束: 函数内部不能使用
return、break、continue语句
2. 程序 ID 与网格 API
tl.program_id(axis)
pid = tl.program_id(axis) # axis: 0, 1, or 2
- 参数:
axis- 维度轴 (0, 1, 2) - 返回: 当前程序在该轴上的 ID
- 用途: 确定当前程序块处理的数据范围
tl.num_programs(axis)
num_pids = tl.num_programs(axis) # axis: 0, 1, or 2
- 参数:
axis- 维度轴 (0, 1, 2) - 返回: 该轴上的总程序数
- 用途: 计算网格大小和边界条件
triton.cdiv(a, b)
grid_size = triton.cdiv(total_elements, block_size)
- 参数:
a,b- 被除数和除数 - 返回: 向上取整的除法结果
- 用途: host 侧使用,计算启动网格大小
3. 内存操作 API
tl.load(pointer, mask=None, other=None, boundary_check=None)
data = tl.load(ptr + offsets, mask=mask, other=0.0)
- 参数:
pointer: 内存指针mask: 布尔掩码,True 表示有效位置other: 掩码为 False 时的默认值boundary_check: 边界检查维度 (0, 1) 或 None
- 返回: 加载的张量数据
- 用途: 从全局内存加载数据
tl.store(pointer, value, mask=None, boundary_check=None)
tl.store(ptr + offsets, result, mask=mask)
- 参数:
pointer: 内存指针value: 要存储的值mask: 布尔掩码,True 表示有效位置boundary_check: 边界检查维度 (0, 1) 或 None
- 用途: 将数据存储到全局内存
tl.make_block_ptr(base, shape, strides, offsets, block_shape, order)
block_ptr = tl.make_block_ptr(
base=ptr, # 基础指针
shape=(M, N), # 完整矩阵形状
strides=(stride_m, stride_n), # 步长
offsets=(start_m, start_n), # 当前块偏移
block_shape=(BLOCK_M, BLOCK_N), # 块形状
order=(1, 0) # 内存布局顺序
)
- 参数:
base: 基础内存指针shape: 完整张量的形状strides: 每个维度的步长offsets: 当前块的起始偏移block_shape: 当前块的大小order: 内存布局顺序 (1, 0) 表示行主序
- 返回: 块指针对象
- 用途: 高效访问 2D 数据块
tl.advance(ptr, offsets)
block_ptr = tl.advance(block_ptr, (BLOCK_M, 0))
- 参数:
ptr: 块指针offsets: 各维度的偏移量
- 返回: 移动后的块指针
- 用途: 移动块指针到下一个位置
4. 张量创建与操作 API
tl.arange(start, end)
offsets = tl.arange(0, BLOCK_SIZE)
- 参数:
start,end- 起始和结束值 - 返回: 连续整数序列
- 用途: 创建索引序列
tl.zeros(shape, dtype)
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
- 参数:
shape: 张量形状dtype: 数据类型
- 返回: 全零张量
tl.full(shape, value, dtype)
ones = tl.full((M, N), 1.0, dtype=tl.float32)
- 参数:
shape: 张量形状value: 填充值dtype: 数据类型
- 返回: 填充指定值的张量
tl.cast(input, dtype)
float_data = tl.cast(int_data, tl.float32)
- 参数:
input: 输入张量dtype: 目标数据类型
- 返回: 类型转换后的张量
5. 数学运算 API
tl.dot(a, b, acc=None, allow_tf32=True)
result = tl.dot(a, b, acc=accumulator)
- 参数:
a,b: 输入矩阵acc: 累加器 (可选)allow_tf32: 是否允许 TF32 精度(CUDA 特有,Ampere+ GPU)
- 返回: 矩阵乘法结果
- 用途: 核心矩阵乘法操作,利用 Tensor Core 加速
tl.sum(x, axis)
block_sum = tl.sum(data, axis=0)
- 参数:
x: 输入张量axis: 归约轴
- 返回: 归约结果
tl.max(x, axis)
max_val = tl.max(data, axis=0)
- 参数:
x: 输入张量axis: 归约轴
- 返回: 最大值
tl.min(x, axis)
min_val = tl.min(data, axis=0)
- 参数:
x: 输入张量axis: 归约轴
- 返回: 最小值
tl.where(condition, x, y)
result = tl.where(mask, data, 0.0)
- 参数:
condition: 条件张量x,y: 选择值
- 返回: 根据条件选择的值
- 用途: SIMD 友好的条件选择
tl.exp(x) / tl.log(x) / tl.sqrt(x)
exp_val = tl.exp(x)
log_val = tl.log(x)
sqrt_val = tl.sqrt(x)
- 用途: 基本数学函数,CUDA 后端直接支持
tl.sigmoid(x)
sigmoid_val = tl.sigmoid(x)
- 用途: Sigmoid 激活函数
tl.extra.cuda.libdevice.tanh(x)
tanh_val = tl.extra.cuda.libdevice.tanh(x)
- 用途: 双曲正切函数
- 注意: CUDA 后端无
tl.tanh或tl.math.tanh,须使用tl.extra.cuda.libdevice.tanh
tl.math.exp2(x) / tl.math.log2(x)
exp2_val = tl.math.exp2(x)
log2_val = tl.math.log2(x)
- 用途: 以 2 为底的指数/对数运算
tl.cumsum(input, axis=0, reverse=False, dtype=None)
cumulative_sum = tl.cumsum(data, axis=0)
reverse_cumsum = tl.cumsum(data, axis=1, reverse=True)
- 参数:
input: 输入张量axis: 累积求和的轴 (默认为 0)reverse: 是否反向累积 (默认为 False)dtype: 输出数据类型 (可选,默认与输入相同)
- 返回: 累积求和结果张量
- 用途: 计算沿指定轴的累积和,常用于前缀和计算
tl.cumprod(input, axis=0, reverse=False)
cumulative_prod = tl.cumprod(data, axis=0)
- 参数:
input: 输入张量axis: 累积乘积的轴 (默认为 0)reverse: 是否反向累积 (默认为 False)
- 返回: 累积乘积结果张量
6. 原子操作 API
tl.atomic_add(pointer, value)
tl.atomic_add(output_ptr, block_sum)
- 参数:
pointer: 目标内存指针value: 要添加的值
- 用途: 线程安全的加法操作
tl.atomic_max(pointer, value)
tl.atomic_max(max_ptr, local_max)
- 参数:
pointer: 目标内存指针value: 要比较的值
- 用途: 线程安全的最大值更新
tl.atomic_min(pointer, value)
tl.atomic_min(min_ptr, local_min)
- 参数:
pointer: 目标内存指针value: 要比较的值
- 用途: 线程安全的最小值更新
tl.atomic_cas(pointer, cmp, val)
old = tl.atomic_cas(ptr, expected, desired)
- 参数:
pointer: 目标内存指针cmp: 期望值val: 新值
- 返回: 原始值
- 用途: 比较并交换操作
tl.constexpr
BLOCK_SIZE: tl.constexpr = 1024
- 用途: 标记编译时常量参数
- 约束: 必须在函数签名中声明
7. 块分配优化 API
tl.swizzle2d(i, j, size_i, size_j, group_size)
task_i, task_j = tl.swizzle2d(block_i, block_j, NUM_BLOCKS_I, NUM_BLOCKS_J, GROUP_SIZE)
- 参数:
i,j: 原始块索引size_i,size_j: 总块数group_size: 分组大小(通常为 2/4/8)
- 返回: 重排后的块索引 (task_i, task_j)
- 用途: 2D 块重排,提升 L2 缓存局部性
- 适用场景: 矩阵乘法等多维块计算,改善数据复用
使用建议
- 内存操作: 优先使用
tl.make_block_ptr处理 2D 数据 - 边界检查: 始终使用
mask或boundary_check防止越界 - 原子操作: 仅在必要时使用,有性能开销
- 数据类型: 注意类型转换,使用
tl.cast显式转换 - 编译优化: 使用
tl.constexpr标记常量参数 - Tensor Core: MatMul 类操作使用
allow_tf32=True利用 Tensor Core