CUDA 算子工程:手写 FlashAttention v2 之路
附录 B · CUDA C++ vs Triton
B.1 Triton 是什么
Triton 是一门 GPU 编程 DSL(Domain-Specific Language):Philippe Tillet 等人 2019 年在 MAPL 上发表同名论文,2021 年由 OpenAI 正式开源并推广。核心理念是用 Python 语法、tile-level 抽象、编译器自动调优写 GPU kernel。本附录的 Triton 代码按 Triton 3.6.0 核对,这是 PyTorch v2.11.0 钉的版本(pytorch-v2.11.0/.ci/docker/triton_version.txt)。
一个 Triton 风格的 GEMM:
import triton
import triton.language as tl
@triton.jit
def gemm_kernel(
a_ptr, b_ptr, c_ptr,
M, N, K,
stride_am, stride_ak,
stride_bk, stride_bn,
stride_cm, stride_cn,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
offs_k = tl.arange(0, BLOCK_K)
a_ptrs = a_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak
b_ptrs = b_ptr + offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k in range(0, K, BLOCK_K):
# 别漏 other=0.0:tl.load 的 other 默认是 None,被 mask 掉的元素
# 取到的是未定义值,直接喂进 tl.dot 累加会算错 K 不是 BLOCK_K 整数倍的尾块。
# M、N 方向同理要 mask,否则 M/N 不是 BLOCK 整数倍时会越界读写。
a = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & (offs_k[None, :] < K - k), other=0.0)
b = tl.load(b_ptrs, mask=(offs_k[:, None] < K - k) & (offs_n[None, :] < N), other=0.0)
acc += tl.dot(a, b)
a_ptrs += BLOCK_K * stride_ak
b_ptrs += BLOCK_K * stride_bk
c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
tl.store(c_ptrs, acc.to(tl.float16), mask=(offs_m[:, None] < M) & (offs_n[None, :] < N))
短得惊人——三十多行 Python 写出一个 GEMM。要在 CUDA C++ 里拿到同等的 Tensor Core 与分块效果,代码要长得多(对照第 11、12 章)。
B.2 Triton 的设计哲学
Triton 提供 tile-level(瓦片级)抽象:
- 程序员不需要管 thread——你只看到 BLOCK_M × BLOCK_N 大小的 tile。
- Triton 编译器自动决定 tile 内每个元素分配给哪个 thread/warp。
- Triton 编译器自动选择 SMEM 布局、是否用 cp.async 与 ldmatrix、是否用 Tensor Core。
- Triton 有
triton.autotune装饰器自动搜索最优 tile size。
这种"声明式"抽象比 CUDA C++ 高一层——程序员描述意图("我要一个 BLOCK_M×BLOCK_N 的 GEMM tile"),编译器负责实现(thread 分配、SMEM 布局、寄存器分配)。
B.3 性能对比
两者的性能差距强烈依赖算子形态、尺寸和 Triton 版本,任何一张"百分比对照表"都会很快过期。本专栏没有条件实测,这里只给定性的分层(读者要具体数字,请用 triton.testing.do_bench 在自己的卡和形状上量):
| 算子形态 | Triton 相对手写 CUDA C++ / CUTLASS |
|---|---|
| Element-wise、逐行 reduce(LayerNorm、Softmax、RoPE) | 持平或更好。这类 kernel 的最优解就是"合并访存 + 一次 warp reduce",Triton 编译器生成的就是这个,手写也很难再有明显提升。 |
| 中等复杂度的融合(dequant + GEMV、fused MoE 的路由段) | 接近。差在寄存器分配和边界处理的细节上。 |
| 大 GEMM、FlashAttention | 明显落后。CUTLASS / FA3 用满了 Hopper 的 TMA、WGMMA、warp specialization 和手工流水;Triton 对这些特性的支持来得晚,且不是所有形状都能触发。 |
规律很清楚:Triton 的抽象层次越贴合算子的本质结构,它就越接近手写。element-wise 和 reduce 的"最优实现"本来就没什么自由度,编译器一定找得到;而 GEMM 和 attention 的最优实现是一整套跨层次的协同设计(tile 大小 × 流水深度 × warp 角色 × 寄存器预算),这是编译器目前搜不动的空间。
Triton 真正稳定的优势是开发速度——几十行 Python 对几百行 C++,改 tile 大小不用重编译几分钟。很多场景下这个 trade-off 非常划算。
B.4 什么时候用 Triton
Triton 的"甜区":
- 原型验证:新算法快速实现,看效果。
- 中等复杂度算子:fused softmax、dropout、新激活函数等。
- PyTorch 生态:
torch.compile内部就用 Triton 生成 kernel。 - 快速迭代:调 BLOCK 大小不需要重编译几分钟。
B.5 什么时候坚持 CUDA C++
CUDA C++ 的"甜区":
- 生产 GEMM / Attention:CUTLASS、FA3 都是 CUDA C++,cuBLAS 是 NVIDIA 闭源的原生库。目前的极致性能仍在这一侧。
- 复杂 epilogue fusion:CUTLASS 的 EVT 比 Triton 灵活。
- 跨硬件代际:CUTLASS 同一套模板覆盖 Volta 到 Blackwell(
cutlass-4.7.0/include/cutlass/arch/mma_sm70.h一直到cutlass-4.7.0/include/cutlass/arch/mma_sm100.h);Triton 的后端跟进有滞后。 - TMA / WGMMA / Warp Specialization 极致优化:Triton 还在追赶这些特性。
- 库的作者:cuBLAS、Megatron-LM 这种"被无数下游用"的库不能容忍哪怕几个百分点的性能损失。
B.6 工业上怎么共存
vLLM 是个有趣的案例:
按 v0.8.5 的源码目录实际分布(csrc/ 是 CUDA C++,vllm/**/triton_* 与 fused_moe/*.py 是 Triton):
vLLM v0.8.5 的 kernel 分布(路径均相对 vllm-0.8.5/):
├── PagedAttention: CUDA C++ csrc/attention/paged_attention_v{1,2}.cu
├── RMSNorm / LayerNorm: CUDA C++ csrc/layernorm_kernels.cu
├── RoPE: CUDA C++ csrc/pos_encoding_kernels.cu
├── Activation: CUDA C++ csrc/activation_kernels.cu
├── Marlin INT4 GEMM: CUDA C++ csrc/quantization/gptq_marlin/
├── MoE: 两条路都有 csrc/moe/{topk_softmax_kernels.cu, moe_wna16.cu,
│ marlin_moe_ops.cu, moe_align_sum_kernels.cu}
│ + vllm/model_executor/layers/fused_moe/fused_moe.py(Triton)
├── LoRA: Triton vllm/lora/ops/triton_ops/{lora_shrink.py, lora_expand.py}
└── 部分 attention 后端: Triton vllm/attention/ops/{triton_decode_attention.py,
triton_flash_attention.py, triton_merge_attn_states.py}
规律和 B.3 的分层大体对得上:形状固定、值得为它手工调一个月的核心算子用 CUDA C++;形状多变、需要跟着模型结构快速改的用 Triton。但要注意,上面是源码目录的分布,不等于默认运行时走哪条路:v0.8.5 默认的 V1 引擎在未开 enforce_eager 时把 custom_ops 设成 ["none"](vllm-0.8.5/vllm/config.py:3884-3892,注释说分段 CUDA Graph 与自定义 CUDA kernel 配合不好),RMSNorm、RoPE、SiluAndMul 这些 CustomOp 于是走 forward_native 的 PyTorch 写法,再由 torch.compile(Inductor)生成 Triton kernel;csrc 里的 CUDA 版本要在 --enforce-eager 或 V0 下才用得上。也就是说,element-wise、逐行 reduce 这类算子在 vLLM 里已经交给编译器生成的 Triton,和 B.3 的分层一致。
B.7 PyTorch 2.x 的 TorchInductor
PyTorch 2.x 的 torch.compile 内部用 TorchInductor 把 PyTorch 计算图编译成 Triton kernel(GPU 上;CPU 上生成 C++,见《PyTorch 训练框架内核深度解析》专栏第 14 章)。这意味着普通 PyTorch 代码加一行 model = torch.compile(model) 就能享受 Triton 的优化。
但 TorchInductor 不能替代手写 kernel——对 attention 和大 GEMM 这些核心算子,它默认调用预编译的 cuBLAS / cuDNN / FlashAttention。开 max-autotune 后它确实会生成 Triton 的 GEMM 模板并与 cuBLAS 打擂台择优,但"生成的模板赢了 cuBLAS"和"Triton 能替代 CUTLASS"是两回事:前者只发生在特定形状上。
B.8 这个附录的小结
CUDA C++ 和 Triton 是互补关系:
- Triton 是高效的 DSL:开发速度快得多,性能差距随算子形态而定——element-wise / reduce 持平,大 GEMM 与 attention 明显落后。
- CUDA C++ 是极致性能的最后一公里:CUTLASS、FA3 都是 C++,cuBLAS 是闭源原生库。
- 工业上两者并存:核心 kernel C++,外围 kernel Triton。
- 学习路径:先学 CUDA C++(理解硬件),再学 Triton(提升生产力)。读完本专栏,再去看 Triton 会非常顺。