GPU Kernel 编程:Triton vs TileLang vs CUDA¶
更新日期:2026-04-15
一、Kernel 编程框架全景¶
| 框架 | 抽象层级 | 学习曲线 | 性能上限 | 典型场景 |
|---|---|---|---|---|
| CUDA C++ | 最底层 | 高 | 100% | 库作者 / 极致优化 |
| CUTLASS | 中-底层 | 中-高 | 95-100% | matmul / FlashAttention 内核 |
| TileLang | 中层(Tile DSL) | 中 | 95% | 自定义 Attention / MoE kernel |
| Triton | 中-高层(Python) | 低 | 90-95% | 快速原型 / 科研 |
| torch.compile | 高层 | 极低 | 80-90% | 业务代码自动加速 |
| ThunderKittens | 中层(C++ DSL) | 中 | 95% | H100 异构资源调度 |
实践笔记:如果你的目标是写自定义 Attention 变体(如 MLA、Sliding Window),建议从 Triton 入手验证正确性,再用 TileLang 优化性能。直接写 CUDA 的 ROI 很少值得,除非你在做基础设施级别的工作。
二、Triton 深入¶
Triton 是 OpenAI 开源的 Python-like DSL,专门用于写 GPU kernel。参考 Triton (Tillet et al., 2019)。
2.1 Triton 编程模型¶
import triton
import triton.language as tl
@triton.jit
def add_kernel(x_ptr, y_ptr, out_ptr, n, BLOCK: tl.constexpr):
# Triton 的核心概念: "program" = 一个 block
# 你不直接管理 threads, Triton 编译器自动并行化
pid = tl.program_id(0) # 当前 block ID
# 计算这个 block 负责的数据范围
offsets = pid * BLOCK + tl.arange(0, BLOCK)
mask = offsets < n # 边界处理
# 加载 + 计算 + 存储
x = tl.load(x_ptr + offsets, mask=mask)
y = tl.load(y_ptr + offsets, mask=mask)
tl.store(out_ptr + offsets, x + y, mask=mask)
# 启动:
def add(x, y):
n = x.numel()
out = torch.empty_like(x)
grid = lambda meta: (triton.cdiv(n, meta['BLOCK']),)
add_kernel[grid](x, y, out, n, BLOCK=1024)
return out
2.2 Triton 的抽象层次¶
2.3 Triton Flash Attention (核心)¶
@triton.jit
def flash_attention_kernel(
Q, K, V, Out,
stride_qm, stride_kn, stride_vn, stride_om,
sm_scale,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, D: tl.constexpr,
):
# 每个 block 处理 Q 的一段
block_idx = tl.program_id(0)
# 加载 Q 的 block
q_offs = block_idx * BLOCK_M + tl.arange(0, BLOCK_M)
q = tl.load(Q + q_offs[:, None] * stride_qm + tl.arange(0, D)[None, :])
# 初始化 online softmax
m = tl.full([BLOCK_M], -float('inf'), dtype=tl.float32)
l = tl.zeros([BLOCK_M], dtype=tl.float32)
o = tl.zeros([BLOCK_M, D], dtype=tl.float32)
# 遍历所有 KV blocks
for kv_start in range(0, N, BLOCK_N):
# 加载 K, V
k = tl.load(K + ...)
v = tl.load(V + ...)
# 计算注意力分数
s = tl.dot(q, tl.trans(k)) * sm_scale
# Online softmax 更新
m_new = tl.maximum(m, tl.max(s, axis=1))
alpha = tl.exp(m - m_new)
p = tl.exp(s - m_new[:, None])
l = alpha * l + tl.sum(p, axis=1)
o = alpha[:, None] * o + tl.dot(p, v)
m = m_new
# 归一化并存储
o = o / l[:, None]
tl.store(Out + ..., o)
三、TileLang 深入¶
TileLang 是新兴的 tile-based DSL,抽象层次比 Triton 更高。DeepSeek 已在生产中使用。参考 TileLang (2024)。
3.1 TileLang Flash Attention¶
from tilelang import Kernel, Tile
@Kernel
def flash_attention(
Q: Tile[M, D],
K: Tile[N, D],
V: Tile[N, D]
) -> Tile[M, D]:
# TileLang 以 Tile 为核心操作单元
# 自动管理 shared memory / registers / tensor cores
O = Tile.zeros(M, D)
m = Tile.full(M, -float('inf'))
l = Tile.zeros(M)
for kv_tile in K.tiles(BLOCK_N): # 自动分块迭代
v_tile = V[kv_tile.index]
# 矩阵乘法自动映射到 tensor core
s = Q @ kv_tile.T * scale
m_new = max(m, s.row_max())
alpha = exp(m - m_new)
p = exp(s - m_new.broadcast())
l = alpha * l + p.row_sum()
O = alpha.broadcast() * O + p @ v_tile
m = m_new
return O / l.broadcast()
# 约 30 行 vs Triton 约 100 行
# 性能: 与手写 CUDA 持平 (TileLang 论文数据)
3.2 TileLang 优势¶
3.3 TileLang FlashMLA¶
# MLA 的高性能 Kernel (DeepSeek-V3 使用)
@Kernel
def flash_mla(
Q: Tile[M, D_q],
C_kv: Tile[N, D_c], # 压缩的 KV cache
W_uk: Tile[D_c, D_k], # 上投影矩阵
W_uv: Tile[D_c, D_v]
) -> Tile[M, D_v]:
O = Tile.zeros(M, D_v)
m = Tile.full(M, -float('inf'))
l = Tile.zeros(M)
for c_tile in C_kv.tiles(BLOCK_N):
# 关键: 在 shared memory 中即时恢复 K, V
K_tile = c_tile @ W_uk
V_tile = c_tile @ W_uv
s = Q @ K_tile.T * scale
# online softmax (同上)
...
return O / l.broadcast()
四、何时手写 CUDA¶
| 场景 | 推荐 | 为什么 |
|---|---|---|
| 通用矩阵乘法 | cuBLAS (不用写) | NVIDIA 工程师针对每代 GPU 手动调优到极致,自己写不可能超过 |
| 标准 Attention | Flash Attention library | Tri Dao 团队持续优化(FA-⅔),直接调库最省事 |
| 自定义融合算子 | Triton 或 TileLang | 融合多个操作减少 HBM 读写是性能关键;库不支持的融合模式需要自己写 |
| 极致性能 + 熟悉 CUDA | CUDA | 当 Triton 编译器生成的代码离理论峰值差距 >20% 时,手写 CUDA 是最后手段 |
| 新硬件 / 新特性 | CUDA (先行支持) | Hopper 的 TMA、Warp Specialization 等新特性通常 CUDA 先支持,Triton 滞后 6-12 个月 |
| 快速原型 | Triton | 从想法到可运行 kernel 最快几小时,CUDA 可能要几天 |
| 生产部署 | TileLang 或 CUDA | TileLang 有跨平台优势;纯 NVIDIA 环境用 CUDA 更成熟 |
五、Warp Specialization (Hopper+)¶
H100 和 Blackwell 支持 Warp Specialization: 同一 block 内不同 warp 做不同工作。FlashAttention-3 的核心优化。
# 传统: 所有 warp 都做相同工作
# Warp Specialization:
# - Producer warp: 从 HBM 加载数据 (用 TMA 指令)
# - Consumer warp: 计算 (用 Tensor Core)
# 两者 pipeline 并行执行
# Triton 3.x 支持:
@triton.jit
def flash_attn_hopper(Q, K, V, Out,
num_consumer_groups: tl.constexpr = 2,
num_buffers_warp_spec: tl.constexpr = 3,
):
# Triton 编译器自动生成 producer-consumer warp 代码
# 用户只需指定参数
...
参考 FlashAttention-3 (Shah et al., 2024)。
六、工具链¶
| 工具 | 用途 | 何时用 / 备注 |
|---|---|---|
| Triton | 通用 GPU kernel DSL,Python 编写 | 自定义算子首选入口;PyTorch 2.x 的 torch.compile 后端之一 |
| TileLang | 高级 tile DSL,比 Triton 更抽象 | 需要跨平台或追求极致性能时用;DeepSeek 生产验证 |
| CUTLASS | NVIDIA C++ 高性能 GEMM/Attention 模板库 | 需要在 C++ 层面集成高性能 GEMM 时用;学习成本高但灵活 |
| cuBLAS/cuDNN | NVIDIA 官方黑盒库 | GEMM/Conv 的默认选择,不用自己写;torch.mm 底层就是 cuBLAS |
| ROCm/HIP | AMD GPU 编程栈(API 类似 CUDA) | 用 AMD MI250/MI300 时必须;HIP 代码可用 hipify 从 CUDA 半自动转换 |
| TVM | 多后端编译器(CPU/GPU/NPU) | 端侧部署和异构硬件编译优化;Apache 项目,社区活跃 |
| XLA | Google 编译器 (TPU 为主,也支持 GPU) | 用 JAX/TPU 训练时的默认编译器;GPU 上不如 Triton 灵活 |
| Nsight Compute | NVIDIA Kernel profiling 工具 | 分析 kernel 的 occupancy、memory throughput、warp stall 原因的必备工具 |
参考文献¶
-
[1] Tillet et al. Triton. 2019. 博客
-
[2] TileLang. 2024. 论文
-
[3] Dao et al. FlashAttention-2. 2023. 论文
-
[4] Shah et al. FlashAttention-3. 2024. 论文
-
[7] ML-Triton: Multi-level compilation. 2024. 论文
↑ 上级 · C. 分布式训练基础设施