📅 创建时间:2026-06-03 🏷️ 标签:#Triton #GPU #Python #JIT #BlockedISA #FlashAttention #torch.compile #Pythonic 📚 前置知识:[[/04-ai/01-llm-engineering/07-llm-evolution]](LLM 发展脉络)[[17-kernel-primer]](Kernel 开发入门) 📚 相关知识:[[19-tvm-te]](TVM TE)[[24-torch-compile]](torch.compile)[[18-cutlass]](CUTLASS)
使用 Triton 编写高性能 GPU Kernel / Triton for High-Performance GPU Kernels in Python
┌──────────────────────────────────────────────────────────────────────────────┐ │ 📖 场景:Triton 生成的代码质量能打败手写 CUDA 吗? │ ├──────────────────────────────────────────────────────────────────────────────┤ │ 你的团队在用 Triton 写 Attention kernel。PyTorch 工程师写起来很顺手。 │ │ 但你怀疑它生成的代码质量:Python JIT 真的能打败手写 CUDA 吗? │ │ 实测 Flash Attention Triton 版本和 CUDA 官方实现速度接近, │ │ 但比手写的简单 kernel 快 30%。为什么 Triton 能做到? │ └──────────────────────────────────────────────────────────────────────────────┘
第1节 Triton 设计哲学——让 PyTorch 工程师写高性能 GPU Kernel
1.1 为什么需要 Triton
GPU Kernel 开发有两个极端:
极端 1: 手写 CUDA
- 优点: 完全控制,可以达到最高性能
- 缺点: 学习曲线陡峭,容易出错,代码难维护
极端 2: cuBLAS/cuDNN
- 优点: 开箱即用,性能好
- 缺点: 不支持自定义操作(如自定义激活函数)
Triton 的目标: 填补中间的空白
- 让你用 Python 写 kernel
- 自动生成接近手写 CUDA 的代码
- 但不需要你手动管理 shared memory / coalesced accessTriton 的核心洞察:
GPU kernel 的优化空间是有限的、规则的:
1. Tiling 策略: 把矩阵分成 tile,每个 tile 放入 shared memory
2. 向量化: SIMD 指令一次处理多个元素
3. 内存访问合并: 连续线程访问连续内存
4. 双缓冲: 计算和内存传输重叠
这些规则是通用的,可以被编译器自动推断。
Triton 的 Blocked ISA 让用户只需指定 tile shape,
编译器自动处理剩下的优化。1.2 Triton vs 其他方案的定位
┌──────────────────────────────────────────────────────────────────────┐
│ 性能上限 │
│ │
│ 手写 CUDA (cuBLAS 团队) ──────────────────────────► 理论峰值 │
│ │ │
│ │ │
│ CUTLASS │ 接近 cuBLAS,但需要大量模板知识 │
│ │ │
│ │ │
│ Triton │ 达到手写 CUDA 的 80-95% │
│ │ Pythonic 接口 │
│ │ │
│ │ │
│ TVM TE │ 可达到接近手写 CUDA │
│ │ 但 schedule 参数空间巨大 │
│ │ │
│ │ │
│ PyTorch │ 简单但性能受限(einsum, torch.mm) │
│ │ │
└──────────────────────────────────────────────────────────────────────┘
选择指南:
- 需要极致性能 + 愿意投入 → 手写 CUDA / CUTLASS
- 需要定制 + Python 优先 → Triton
- 需要自动搜索最优参数 → TVM AutoScheduler
- 标准 GEMM/Conv → cuBLAS/cuDNN第2节 @triton.jit 装饰器——Python 函数变成 GPU Kernel
2.1 JIT 编译模型
import triton
import triton.language as tl
# @triton.jit: 把 Python 函数标记为 Triton kernel
# Triton 会在运行时 (JIT) 编译这个函数到 PTX/CUDA 代码
@triton.jit
def add_kernel(x_ptr, y_ptr, output_ptr, n_elements, BLOCK_SIZE: tl.constexpr):
"""
y = x + y, 每个元素一个 thread 处理
但实际每个 thread 处理 BLOCK_SIZE 个连续元素 (向量化)
参数说明:
x_ptr, y_ptr, output_ptr: GPU 内存指针 (通过 torch.Tensor_PTR 传入)
n_elements: 元素总数
BLOCK_SIZE: compile-time 常量 (tl.constexpr)
"""
# 程序 ID: 识别当前 thread block 处理的元素范围
pid = tl.program_id(axis=0) # axis=0 表示一维 grid
# 计算当前 block 处理的元素范围
# 每个 block 处理 BLOCK_SIZE 个连续元素
block_start = pid * BLOCK_SIZE
# 元素索引: 从 block 起始位置开始
offsets = block_start + tl.arange(0, BLOCK_SIZE)
# 创建 mask: 只处理有效范围内的元素 (n_elements 以内的)
# 超出 n_elements 的位置被 mask 住,不会真正执行
mask = offsets < n_elements
# 从 GPU 内存加载数据
# tl.load: 从 x_ptr[offsets] 读取,mask 住越界的访问
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
# 写回 GPU 内存
tl.store(output_ptr + offsets, output, mask=mask)
# 调用 Triton kernel
def add(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
output = torch.empty_like(x)
n_elements = output.numel()
# launch 配置
BLOCK_SIZE = 1024
# 计算 grid 大小: 需要多少个 block
# 向上取整: ceil(n_elements / BLOCK_SIZE)
grid = (triton.cdiv(n_elements, BLOCK_SIZE),)
# 启动 kernel
add_kernel[grid](
x, y, output, # torch.Tensor 会自动转换为 device pointer
n_elements,
BLOCK_SIZE=BLOCK_SIZE,
)
return output2.2 tl.constexpr——Compile-Time 常量
# tl.constexpr: 标记在编译时求值的常量
# 作用: 允许 Triton 在编译时做常量折叠和优化
@triton.jit
def kernel_with_constexpr(
x_ptr, y_ptr, output_ptr,
BLOCK_SIZE: tl.constexpr, # 编译时常量
DOUBLE: tl.constexpr, # 另一个编译时常量
):
# 编译时可以直接计算
# 比如: 1 << DOUBLE 在编译时就是 1 << 2 = 4
stride = 1 << DOUBLE # 如果 DOUBLE=2,stride=4
offsets = ...
# DOUBLE 作为 bool 用
if DOUBLE:
# 如果 DOUBLE 是 True,编译器知道只执行这个分支
# 可以消除 dead code
...第3节 Triton Blocked ISA——BLOCK 概念
3.1 什么是 Blocked ISA
传统 CUDA 编程模型:
- 你需要手动管理:
1. Thread/block 的映射到数据
2. Shared memory 的加载/同步
3. 内存访问的 coalescing
4. Warp 级别的同步
Triton Blocked ISA:
- 你只指定 BLOCK 的大小和形状
- Triton 编译器自动推断:
1. 如何把 BLOCK 映射到 thread/block
2. 如何利用 shared memory
3. 如何合并内存访问
- BLOCK = Tile = 你要处理的最小数据单元3.2 常见的 BLOCK 形状
# BLOCK 形状的选择影响性能
# 1D BLOCK: 适合向量操作
BLOCK_M = 1
BLOCK_N = 1024 # 每个 thread 处理 1024 个元素
# 2D BLOCK: 适合矩阵操作
BLOCK_M = 128 # 行方向
BLOCK_N = 128 # 列方向
# BLOCK 形状的经验法则:
#
# GEMM:
# BLOCK_M = 64-128 (行方向 tile)
# BLOCK_N = 64-128 (列方向 tile)
# BLOCK_K = 32-64 (reduce 方向)
#
# Flash Attention:
# BLOCK_M = 64-128 (query block)
# BLOCK_N = 64-128 (key block)
# 小的 BLOCK 减少 SRAM 需求,大的 BLOCK 提高数据复用
#
# Softmax:
# BLOCK_M = 1 (每个 thread 处理一行)
# BLOCK_N = 1024-4096 (每行元素数)
# 需要 reduce 操作,但行间独立3.3 自动向量化
@triton.jit
def vectorized_load_kernel(x_ptr, output_ptr, n_elements, VECTOR_SIZE: tl.constexpr):
"""
Triton 自动把 tl.load/store 向量化
VECTOR_SIZE 指定每个 memory transaction 的大小
"""
pid = tl.program_id(axis=0)
# 注意: offsets 是 BLOCK_SIZE 个元素的索引
# Triton 会识别这是连续访问,自动生成 vectorized load
offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
# 如果 BLOCK_SIZE = 256 且 VECTOR_SIZE = 4
# Triton 会生成 256/4 = 64 个 vectorized load
# 每个 load 4 个 float = 16 bytes
# 加载向量化的数据
# x_ptr + offsets 指向连续内存
# Triton 识别这是向量化访问,生成 64 个 4-element 的 vectorized load
x = tl.load(x_ptr + offsets, mask=offsets < n_elements, other=0.0)
output = x * 2.0
tl.store(output_ptr + offsets, output, mask=offsets < n_elements)第4节 Triton DSL——tl.load, tl.store, tl.dot, tl.constexpr
4.1 tl.load / tl.store
# tl.load: 从 GPU 内存加载数据到寄存器
# tl.store: 把寄存器数据写回 GPU 内存
@triton.jit
def load_store_examples(
ptr, # GPU 内存指针
offsets, # 索引 (可以是向量)
mask, # 掩码 (越界的位置)
other, # 越界时的默认值
BLOCK_SIZE: tl.constexpr
):
# 基础用法
x = tl.load(ptr + offsets, mask=mask, other=0.0)
# 展开指针 (pointer arithmetric)
# ptr_tensor + offsets 等价于 ptr[offsets]
# 你也可以手动计算偏移
base_ptr = ptr
x = tl.load(base_ptr + offsets, mask=mask, other=0.0)
# 跨步加载 (strided load)
# 用于非连续内存访问
# 例如: 读取矩阵的一列
# ptr: 矩阵起始地址
# offsets: 行索引
# stride: 列跨度 (矩阵的列数)
row_offsets = offsets * stride
x = tl.load(ptr + row_offsets, mask=mask, other=0.0)
# 缓存策略提示
# tl.RW0, tl.RW1 指定 cache modifier
# tl.load(ptr, cache_modifier="ca") # cache as streaming
# store 用法相同
tl.store(ptr + offsets, x, mask=mask)4.2 tl.dot——矩阵乘
# tl.dot: 执行矩阵乘累加
# 这是 Triton 最强大的功能之一
@triton.jit
def matmul_kernel(
# Pointers for matrices
a_ptr, b_ptr, c_ptr,
# Matrix dimensions
M, N, K,
# Strides
stride_am, stride_ak,
stride_bk, stride_bn,
stride_cm, stride_cn,
# Meta-parameters
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
GROUP_M: tl.constexpr, # 用于 swizzle
):
"""
实现: C = A @ B + C
A: M×K (行主序)
B: K×N (行主序)
C: M×N (行主序)
"""
# 程序 ID: 二维 grid
# pid_m: 行方向 block index
# pid_n: 列方向 block index
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
# 计算当前 block 的 C 矩阵范围
# 每 block 处理 BLOCK_M × BLOCK_N 的 C
rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) # C 的行索引
rn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) # C 的列索引
# A 和 B 的索引
# A: rm (行), k (列) - k 遍历 K
# B: k (行), rn (列) - k 遍历 K
# 初始化累加器
# 用 fp32 累加,精度更高
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# K 方向循环
for k in range(0, K, BLOCK_K):
# A 的索引: rm (行), k+k_idx (列)
ram = rm[:, None] # shape: (BLOCK_M, 1)
rak = k + tl.arange(0, BLOCK_K) # shape: (BLOCK_K,)
a_offsets = ram * stride_am + rak * stride_ak # 全局索引
# B 的索引: k+k_idx (行), rn (列)
rak = k + tl.arange(0, BLOCK_K) # shape: (BLOCK_K,)
rbn = rn[None, :] # shape: (1, BLOCK_N)
b_offsets = rak * stride_bk + rbn * stride_bn
# 加载 A 和 B 的 tile
# shape: (BLOCK_M, BLOCK_K) 和 (BLOCK_K, BLOCK_N)
# 自动处理 alignment 和 masking
a_mask = (rm < M)[:, None] & (rak < K)[None, :]
b_mask = (rak < K)[:, None] & (rn < N)[None, :]
a = tl.load(a_ptr + a_offsets, mask=a_mask, other=0.0)
b = tl.load(b_ptr + b_offsets, mask=b_mask, other=0.0)
# 执行矩阵乘
# accumulator += a @ b
# a: (BLOCK_M, BLOCK_K)
# b: (BLOCK_K, BLOCK_N)
# result: (BLOCK_M, BLOCK_N)
accumulator += tl.dot(a, b)
# 把累加器 (fp32) 转换回 output dtype (通常是 fp16)
c = accumulator.to(tl.float16)
# 写回 C
cm = rm[:, None]
cn = rn[None, :]
c_offsets = cm * stride_cm + cn * stride_cn
c_mask = (rm < M)[:, None] & (rn < N)[None, :]
tl.store(c_ptr + c_offsets, c, mask=c_mask)4.3 tl.constexpr 的高级用法
@triton.jit
def advanced_constexpr(
x_ptr, y_ptr,
N: tl.constexpr,
ACTIVATION: tl.constexpr, # 编译时选择激活函数
SCALE_BIAS: tl.constexpr, # 编译时决定是否做 scale+bias
):
x = tl.load(x_ptr)
# 编译时分支消除
if ACTIVATION == "relu":
x = tl.where(x > 0, x, 0.0)
elif ACTIVATION == "gelu":
x = 0.5 * x * (1.0 + tl.tanh(0.797884 * (x + 0.044715 * x * x * x)))
elif ACTIVATION == "silu":
x = x * tl.sigmoid(x)
# 编译时展开循环
if SCALE_BIAS:
scale = tl.load(y_ptr)
bias = tl.load(y_ptr + 1)
x = x * scale + bias
return x第5节 Triton 的自动优化
5.1 自动选择 Memory Access Pattern
@triton.jit
def autotuned_matmul(...):
"""
Triton 自动处理:
1. Coalesced memory access
2. Bank conflict 避免 (通过 swizzle)
3. L2 cache 友好的访问模式
"""
...Triton 的自动 coalescing:
场景: 加载 A 矩阵的 tile (行 rm, 列 k 范围)
rm = [0, 1, 2, ..., 127]
k_range = [0, 1, 2, ..., 31]
问题: 如果用简单的索引 a[rm, k_range]
每个 thread 访问的内存地址不连续
Triton 编译器识别:
原始索引: a[0,0], a[1,0], a[2,0], ... (列优先)
连续访问: a[0,0], a[0,1], a[0,2], ... (行优先)
编译器自动重排:
1. 交换 axis 顺序
2. 生成 coalesced 的 load 指令
3. 或生成 transpose + coalesced load5.2 自动选择 Tiling Factor
# Triton 2.0 引入了 autotuning
# 可以自动搜索最优的 BLOCK_SIZE, num_warps 等参数
@triton.autotune(
configs=[
# configs 是一个配置列表,Triton 会自动尝试
# 每个 config 是 {key: value} 的字典
triton.Config({'BLOCK_M': 64, 'BLOCK_N': 64, 'BLOCK_K': 32}, num_warps=4),
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 64, 'BLOCK_K': 32}, num_warps=8),
triton.Config({'BLOCK_M': 64, 'BLOCK_N': 128, 'BLOCK_K': 32}, num_warps=8),
triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 64}, num_warps=8),
],
# key: 用于分组的参数(相同 key 共享编译结果)
key=['M', 'N', 'K'],
)
@triton.jit
def autotuned_gemm(...):
...5.3 自动管理 Shared Memory
# Triton 2.0 之前: 用户需要手动申请 shared memory
# Triton 2.0+ : 自动管理
@triton.jit
def auto_smem_matmul(...):
# Triton 编译器自动决定:
# 1. 需要多少 shared memory
# 2. 如何安排 tile 的 layout
# 3. 何时 flush shared memory
# 加载 A tile 到 "自动分配的 shared memory"
# triton 的 load/store 指令隐式使用 shared memory 缓存
a = tl.load(a_ptr + a_offsets, mask=a_mask, other=0.0, eviction_policy="evict_last")
# eviction_policy="evict_last": L2 cache 友好的策略
# 先加载的数据最后被 evict,提高 cache 利用率
# ... 计算 ...
return第6节 Flash Attention 在 Triton 中的实现
6.1 Flash Attention 的核心思想
标准 Attention 实现的问题:
1. 需要完整的 S×S attention matrix (S=序列长度)
2. 内存复杂度 O(S²),无法处理长序列
Flash Attention 的核心思想:
1. 不存储完整的 attention matrix
2. 分块计算 online softmax
3. 边算边归一化,只保留必要的统计量 (m, l)
Online Softmax:
m_i = max(m_{i-1}, x_i) # max 值
l_i = l_{i-1} * exp(m_{i-1} - m_i) + exp(x_i - m_i) # sum
softmax_i = exp(x_i - m_i) / l_i
只需要 O(1) 额外内存!6.2 Triton Flash Attention 实现
@triton.jit
def flash_attention_kernel(
Q, K, V, Out, # 指针
stride_qb, stride_qh, stride_qm, stride_qk, # Q 的 stride
stride_kb, stride_kh, stride_kn, stride_kk,
stride_vb, stride_vh, stride_vn, stride_vk,
stride_ob, stride_oh, stride_om, stride_ok,
# 维度
B, H, N, D, # batch, heads, seq_len, head_dim
# Meta-parameters
BLOCK_M: tl.constexpr, # Q block 大小
BLOCK_N: tl.constexpr, # K/V block 大小
BLOCK_D: tl.constexpr, # head_dim (必须 <= 128)
):
"""
Flash Attention 的 Triton 实现
公式:
O_i = softmax(Q_i @ K^T / sqrt(D)) @ V
= sum_j (exp(q_i · k_j - m_i) / l_i) * v_j
"""
# 获取程序 ID
bid_h = tl.program_id(0) # head index
bid_b = tl.program_id(1) # batch index
start_m = tl.program_id(2) # Q block index
# Q_i: 当前 query block
# shape: (BLOCK_M, D)
rm = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
rn = tl.arange(0, BLOCK_N) # 用于加载 K
rd = tl.arange(0, BLOCK_D)
# 计算 Q block 的内存偏移
q_base = Q + bid_b * stride_qb + bid_h * stride_qh
# 加载 Q
q_offsets = rm[:, None] * stride_qm + rd[None, :] * stride_qk
q_mask = (rm[:, None] < N) & (rd[None, :] < D)
q = tl.load(q_base + q_offsets, mask=q_mask, other=0.0)
# 初始化累加器
# m_i: max 值 (用于 softmax)
# l_i: sum 值 (用于归一化)
m_i = tl.zeros((BLOCK_M,), dtype=tl.float32) + float("-inf") # 初始化为 -inf
l_i = tl.zeros((BLOCK_M,), dtype=tl.float32) + 1e-6 # 初始化很小的值
# accumulator: 存放 O_i 的累加结果
acc = tl.zeros((BLOCK_M, BLOCK_D), dtype=tl.float32)
# 遍历 K, V 的 block
for start_n in range(0, N, BLOCK_N):
# 加载 K block: shape (BLOCK_N, D)
kn = start_n + rn
k_offsets = kn[:, None] * stride_kn + rd[None, :] * stride_kk
k_mask = (kn[:, None] < N) & (rd[None, :] < D)
k = tl.load(K + bid_b * stride_kb + bid_h * stride_kh + k_offsets,
mask=k_mask, other=0.0)
# 计算 Q @ K^T
# q: (BLOCK_M, D), k: (BLOCK_N, D)
# q @ k^T: (BLOCK_M, BLOCK_N)
qk = tl.dot(q, tl.trans(k)) # transpose k
# 加上 causal mask (如果需要)
# mask: rm[:, None] >= start_n + rn[None, :]
# 即: 当前位置 q_i 只能 attend 到 k_j,其中 j <= i
# Scale by sqrt(D)
scale = 1.0 / tl.sqrt(D.to(tl.float32))
qk = qk * scale
# 计算 block 内的 max 和 sum (online softmax)
m_ij = tl.max(qk, axis=1) # max across K
m_i_new = tl.maximum(m_i, m_ij) # 更新 global max
# 计算 exp(qk - m_i_new)
qk_minus_m = qk - m_i_new[:, None]
p = tl.exp(qk_minus_m)
# 计算 block 内的 l
l_ij = tl.sum(p, axis=1)
l_i_new = tl.exp(m_i - m_i_new) * l_i + l_ij
# 计算 softmax 归一化后的 attention
p_scaled = p / l_ij[:, None]
# 加载 V block: shape (BLOCK_N, D)
v_offsets = kn[:, None] * stride_vn + rd[None, :] * stride_vk
v_mask = (kn[:, None] < N) & (rd[None, :] < D)
v = tl.load(V + bid_b * stride_vb + bid_h * stride_vh + v_offsets,
mask=v_mask, other=0.0)
# 计算 attention @ V
# p_scaled: (BLOCK_M, BLOCK_N), v: (BLOCK_N, D)
# result: (BLOCK_M, D)
pv = tl.dot(p_scaled.to(tl.float16), v)
# 累加到 acc
# 需要 scale 旧的 acc
acc_scale = tl.exp(m_i - m_i_new) * l_i / l_i_new
acc = acc * acc_scale[:, None] + pv / l_i_new[:, None]
# 更新 m 和 l
m_i = m_i_new
l_i = l_i_new
# 写回输出
out = acc.to(tl.float16)
o_base = Out + bid_b * stride_ob + bid_h * stride_oh
o_offsets = rm[:, None] * stride_om + rd[None, :] * stride_ok
o_mask = (rm[:, None] < N) & (rd[None, :] < D)
tl.store(o_base + o_offsets, out, mask=o_mask)第7节 Triton vs CUDA C++——什么时候选哪个
7.1 选 Triton 的场景
| 场景 | 原因 |
|---|---|
| PyTorch 项目 | 和 PyTorch 无缝集成,不需要单独的编译流程 |
| 快速原型 | Python 开发速度快,JIT 编译即跑 |
| 中等复杂度算子 | GEMM、Attention、Softmax、LayerNorm 等 |
| 需要 fusion | Triton 容易 fusion 多个操作 |
| 研究实验 | 容易尝试不同的算法变体 |
7.2 选手写 CUDA 的场景
| 场景 | 原因 |
|---|---|
| 极致性能 | 手写 CUDA 可以达到 100% 峰值 |
| 复杂控制流 | Triton 对复杂的分支/循环支持有限 |
| 动态 shape | Triton 的动态 shape 支持是实验性的 |
| 特殊硬件特性 | 需要手动利用特殊的硬件指令 |
| 生产级库 | 需要长期维护和优化的核心库 |
7.3 Triton vs 其他方案对比表
| 维度 | Triton | TVM TE | CUTLASS | 手写 CUDA |
|---|---|---|---|---|
| 接口易用性 | ⭐⭐⭐⭐⭐ Pythonic | ⭐⭐⭐ 需要写 schedule | ⭐ 需要 C++ 模板 | ⭐ 需要手写底层 |
| 性能 | ⭐⭐⭐⭐ 手写 CUDA 的 80-95% | ⭐⭐⭐⭐ 取决于 schedule | ⭐⭐⭐⭐⭐ 接近峰值 | ⭐⭐⭐⭐⭐ 理论峰值 |
| 灵活性 | ⭐⭐⭐⭐ 足够大多数场景 | ⭐⭐⭐⭐⭐ 高度可定制 | ⭐⭐⭐⭐ 模板化 | ⭐⭐⭐⭐⭐ 完全控制 |
| 学习曲线 | 低 | 高 | 很高 | 很高 |
| 编译时间 | 中等 (JIT) | 高 | 中等 | 无 (预编译) |
| 动态 shape | 有限 | ✅ 支持 | ❌ 不支持 | ❌ 不支持 |
| 适用场景 | PyTorch 集成、研究 | 自动搜索优化 | 生产级 GEMM | 极致优化库 |
第8节 Triton 的局限性
8.1 动态 Shape 支持
# Triton 对动态 shape 的支持有限
# 场景 1: 编译时知道 shape
@triton.jit
def static_shape_kernel(ptr, N: tl.constexpr): # N 是编译时常量
# 可以在编译时展开循环,优化代码
for i in range(N): # 编译时展开
...
# 场景 2: 运行时才知道 shape
@triton.jit
def dynamic_shape_kernel(ptr, N): # N 是运行时值
# Triton 无法在编译时优化
# 循环保持为循环,代码质量下降
for i in range(N): # 运行时循环,无法展开
...
# 解决方案: 用 autotune 覆盖常见 shape
@triton.autotune(configs=[
triton.Config({'BLOCK_SIZE': 512}),
triton.Config({'BLOCK_SIZE': 1024}),
triton.Config({'BLOCK_SIZE': 2048}),
], key=['n_elements'])
@triton.jit
def tuned_kernel(ptr, n_elements, BLOCK_SIZE: tl.constexpr):
...8.2 复杂 Control Flow
# Triton 不擅长复杂的控制流
# ❌ 不好的例子: 复杂分支
@triton.jit
def bad_control_flow(x, mode: tl.constexpr):
if mode == "relu":
x = tl.where(x > 0, x, 0.0)
elif mode == "gelu":
x = 0.5 * x * (1.0 + tl.tanh(...))
elif mode == "silu":
x = x * tl.sigmoid(x)
elif mode == "identity":
pass
# 更好的做法: 用 lookup table 或分开实现
# 或者接受编译出来的代码包含所有分支
# ✅ 更好的例子: 简化控制流
@triton.jit
def simple_activation(x, do_gelu: tl.constexpr, do_silu: tl.constexpr):
# 编译时消除不需要的分支
if do_gelu:
x = gelu(x)
if do_silu:
x = x * tl.sigmoid(x)升华
┌────────────────────────────────────────────────────────────────────────────┐
│ 🚀 Triton 核心原则 │
├────────────────────────────────────────────────────────────────────────────┤
│ │
│ ① 选 Triton 而不是手写 CUDA 的条件: │
│ 你需要定制 + 性能要求在 80-95% 内 + Python 优先 │
│ 否则应该用 cuBLAS/cuDNN │
│ │
│ ② @triton.jit 把 Python 变成 GPU 代码: │
│ 编译器自动处理 coalesced access、shared memory、vectorization │
│ 你只需要指定 BLOCK 大小 │
│ │
│ ③ Flash Attention 是 Triton 的杀手级应用: │
│ 分块 online softmax 只用 O(1) 额外内存 │
│ Triton 实现比手写 CUDA 简洁得多 │
│ │
│ ④ autotune 是性能关键: │
│ 不要手动猜 BLOCK_SIZE,用 autotune 搜索最优配置 │
│ │
│ ⑤ 动态 shape 是弱点: │
│ 编译时 shape 能达到最佳性能 │
│ 运行时 shape 需要用 autotune 覆盖常见情况 │
│ │
└────────────────────────────────────────────────────────────────────────────┘"AI 可查 vs 必须理解"清单
必须理解(不理解就等于不会):
- 🔴 Blocked ISA 的概念:为什么指定 BLOCK 大小就能自动优化
- 🔴 @triton.jit 的 JIT 编译模型:Python 函数如何变成 GPU 代码
- 🔴 tl.load/tl.store 的 mask 机制:为什么需要 mask,越界访问会发生什么
- 🔴 Flash Attention 的 online softmax:m 和 l 的递推公式,为什么能省内存
- 🔴 autotune 的机制:Triton 如何在运行时选择最优配置
AI 可查(知道去哪查就行):
- ✅ Triton 具体的 API 签名(官方文档)
- ✅ 特定 GPU 架构的 optimal BLOCK_SIZE(benchmark 数据)
- ✅ Triton 的底层代码生成(PTX/SASS 输出)
- ✅
triton.Config的完整参数(官方文档) - ✅ Triton 和 torch.compile 的集成细节(PyTorch 文档)
学习状态:🟡 开始学习