CUDA 算子工程:手写 FlashAttention v2 之路

附录 B · CUDA C++ vs Triton

作者 杨艺韬 · 1,800 字 · 发布于 · 更新于

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 的"甜区":

  1. 原型验证:新算法快速实现,看效果。
  2. 中等复杂度算子:fused softmax、dropout、新激活函数等。
  3. PyTorch 生态:torch.compile 内部就用 Triton 生成 kernel。
  4. 快速迭代:调 BLOCK 大小不需要重编译几分钟。

B.5 什么时候坚持 CUDA C++

CUDA C++ 的"甜区":

  1. 生产 GEMM / Attention:CUTLASS、FA3 都是 CUDA C++,cuBLAS 是 NVIDIA 闭源的原生库。目前的极致性能仍在这一侧。
  2. 复杂 epilogue fusion:CUTLASS 的 EVT 比 Triton 灵活。
  3. 跨硬件代际:CUTLASS 同一套模板覆盖 Volta 到 Blackwell(cutlass-4.7.0/include/cutlass/arch/mma_sm70.h 一直到 cutlass-4.7.0/include/cutlass/arch/mma_sm100.h);Triton 的后端跟进有滞后。
  4. TMA / WGMMA / Warp Specialization 极致优化:Triton 还在追赶这些特性。
  5. 库的作者: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 是互补关系:

  1. Triton 是高效的 DSL:开发速度快得多,性能差距随算子形态而定——element-wise / reduce 持平,大 GEMM 与 attention 明显落后。
  2. CUDA C++ 是极致性能的最后一公里:CUTLASS、FA3 都是 C++,cuBLAS 是闭源原生库。
  3. 工业上两者并存:核心 kernel C++,外围 kernel Triton。
  4. 学习路径:先学 CUDA C++(理解硬件),再学 Triton(提升生产力)。读完本专栏,再去看 Triton 会非常顺。