| name | tilelang-cuda-api |
| description | TileLang CUDA API 完整参考手册,适用于需要查阅具体 API 用法、了解函数参数含义的任意 TileLang CUDA 内核代码生成场景 |
| category | fundamental |
| version | 1.0.0 |
| metadata | {"backend":"cuda","dsl":"tilelang_cuda"} |
TileLang CUDA API 参考手册
本文档提供 TileLang 核心 API 的详细参考,包括函数签名、参数说明和使用示例。
1. 内核定义与编译
@tilelang.jit(out_idx)
@tilelang.jit(out_idx=[-1])
def my_kernel(M, N, K, block_M, block_N, block_K):
@T.prim_func
def main(
A: T.Tensor((M, K), "float16"),
B: T.Tensor((K, N), "float16"),
C: T.Tensor((M, N), "float16"),
):
pass
return main
- 作用: 将 TileLang 函数编译为 GPU 内核
- 参数:
out_idx - 指定输出张量的索引列表(如 [-1] 表示最后一个参数为输出)
- 调用: 设置
out_idx 后,运行内核时只需传入输入张量,输出由 TileLang 自动创建
tilelang.compile
kernel = tilelang.compile(my_func, out_idx=[-1])
- 作用: 编译 TileLang 函数(等价于
@tilelang.jit)
2. 内核上下文
T.Kernel(grid_x, grid_y, threads)
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (bx, by):
pass
- 参数:
grid_x: 网格 X 维度大小
grid_y: 网格 Y 维度大小(可选)
threads: 每个线程块的线程数
- 返回: 线程块索引
(bx, by)
T.ceildiv(a, b)
grid_size = T.ceildiv(N, block_N)
- 参数:
a, b - 被除数和除数
- 返回: 向上取整的除法结果
- 用途: 计算网格大小
3. 内存分配 API
T.alloc_shared(shape, dtype)
A_shared = T.alloc_shared((block_M, block_K), "float16")
- 作用: 分配共享内存(对应 GPU 共享内存)
- 参数:
- 用途: 缓存频繁访问的数据
T.alloc_fragment(shape, dtype)
C_local = T.alloc_fragment((block_M, block_N), "float")
- 作用: 分配寄存器片段(对应 GPU 寄存器文件)
- 参数:
- 用途: 累加器和临时存储
T.alloc_local(shape, dtype)
temp = T.alloc_local((1,), "float32")
- 作用: 分配线程本地内存
- 参数:
- 用途: 线程私有的临时变量
4. 数据操作 API
T.copy(src, dst)
T.copy(A[by * block_M, ko * block_K], A_shared)
T.copy(C_local, C[by * block_M, bx * block_N])
- 作用: 高效内存复制,自动合并访问
- 参数:
src: 源数据(可以是全局内存切片或寄存器片段)
dst: 目标数据
T.clear(tensor)
T.clear(C_local)
- 作用: 将张量清零
- 参数:
tensor - 要清零的张量
T.fill(tensor, value)
T.fill(buffer, -T.infinity("float"))
5. 循环控制 API
T.Pipelined(count, num_stages)
for ko in T.Pipelined(T.ceildiv(K, block_K), num_stages=3):
T.copy(A[by * block_M, ko * block_K], A_shared)
T.gemm(A_shared, B_shared, C_local)
- 作用: 软件流水线循环,重叠内存操作和计算
- 参数:
count: 循环次数
num_stages: 流水线深度(通常 2-4 效果最好)
T.Parallel(dim1, dim2, ...)
for i in T.Parallel(block_M):
pass
for i, j in T.Parallel(block_M, block_N):
C_local[i, j] = A_shared[i, j] + B_shared[i, j]
- 作用: 并行循环,自动映射到线程
- 参数: 各维度大小
- 优势: 自动处理线程映射,避免手动线程索引计算错误
T.serial(count)
for k in T.serial(block_K):
pass
- 作用: 串行循环,顺序执行
- 参数:
count - 循环次数
T.vectorized(count)
for k in T.vectorized(TILE_K):
A_local[k] = A[bk * BLOCK_K + tk * TILE_K + k]
- 作用: 向量化循环,利用向量指令
- 参数:
count - 向量化长度
6. 内置计算原语
T.gemm(A, B, C, transpose_B, policy)
T.gemm(A_shared, B_shared, C_local)
T.gemm(Q_shared, K_shared, acc_s, transpose_B=True)
T.gemm(A_shared, B_shared, C_local, policy=T.GemmWarpPolicy.FullRow)
- 作用: Tile 级别的矩阵乘法,利用 Tensor Core 加速
- 参数:
A, B: 输入矩阵(通常为共享内存)
C: 输出/累加器(通常为寄存器片段)
transpose_B: 是否转置 B 矩阵
policy: Warp 策略
T.reduce_max(input, output, dim, clear)
T.reduce_max(acc_s, scores_max, dim=1)
T.reduce_max(acc_s, scores_max, dim=1, clear=False)
- 作用: 最大值归约
- 参数:
input: 输入张量
output: 输出张量
dim: 归约维度
clear: 是否先清零输出(默认 True)
T.reduce_sum(input, output, dim)
T.reduce_sum(acc_s, scores_sum, dim=1)
- 作用: 求和归约
- 参数: 同
T.reduce_max
T.reduce_min(input, output, dim)
T.reduce_min(input_tensor, output_tensor, dim=axis)
T.reduce_mean(input, output, dim)
T.reduce_mean(input_tensor, output_tensor, dim=axis)
7. 数学函数
T.exp(x)
T.exp2(x)
T.log(x)
T.sqrt(x)
T.rsqrt(x)
T.infinity(dtype)
8. 条件与逻辑操作
T.if_then_else(condition, true_val, false_val)
result = T.if_then_else(A[idx] > 0, A[idx] + B[idx], A[idx] - B[idx])
条件分支
for i in T.Parallel(block_M):
if i < N:
pass
9. 原子操作
T.atomic_add(target, value)
T.atomic_add(C_shared[tn], C_accum[0])
- 作用: 线程安全的原子加法
- 参数:
target: 目标内存位置
value: 要添加的值
10. 线程索引获取
✅ 推荐:T.Parallel
for tn in T.Parallel(BLOCK_N):
pass
⚠️ 不推荐:T.get_thread_binding
tn = T.get_thread_binding(0)
tk = T.get_thread_binding(1)
11. 同步操作
T.sync_threads()
T.sync_threads()
- 作用: 线程块内同步
- ⚠️ 严格禁止: 在条件分支中使用(会导致死锁)
12. 高级特性
@T.macro
@T.macro
def Softmax(acc_s, acc_s_cast, scores_max, scores_sum):
T.reduce_max(acc_s, scores_max, dim=1)
for i, j in T.Parallel(block_M, block_N):
acc_s[i, j] = T.exp(acc_s[i, j] - scores_max[i])
T.reduce_sum(acc_s, scores_sum, dim=1)
T.copy(acc_s, acc_s_cast)
T.annotate_layout / T.use_swizzle
from tilelang.intrinsics import make_mma_swizzle_layout
T.annotate_layout({
A_shared: make_mma_swizzle_layout(A_shared),
B_shared: make_mma_swizzle_layout(B_shared),
})
T.use_swizzle(panel_size=10, enable=True)
类型转换
result = A[i].astype("float") * B[i].astype("float")
normalized.astype("float16")
数据类型支持
输入数据类型
float16: 半精度浮点数
float32: 单精度浮点数
bfloat16: Brain Float 16
int8: 8 位整数
int32: 32 位整数
累积数据类型
float / float32: 单精度浮点数(推荐用于累加器)
输出数据类型
- 通常与输入类型相同
- 可通过
.astype() 指定
使用建议
- 选择合适的抽象级别: Level 2 适合大多数应用,Level 3 用于极致性能优化
- 合理的内存分配: 共享内存用于频繁访问的数据,寄存器片段用于累加和临时存储
- 优化数据移动: 使用
T.Parallel 并行化数据复制,利用 T.Pipelined 重叠操作
- 选择合适的线程数: 通常为 128 或 256,考虑硬件特性和工作负载
- 利用内置原语: 使用
T.gemm、T.reduce_sum 等优化原语,避免重复实现已有功能
常见错误
- 内存分配过大: 超出硬件限制
- 流水线深度不当: 影响性能
- 线程数不匹配: 硬件利用率低
- 数据类型不匹配: 精度损失或性能下降
- ⚠️ 同步使用错误: 条件分支中的
T.sync_threads() 会导致死锁
- ⚠️ 线程索引获取错误: 使用
T.get_thread_binding() 而非 T.Parallel()