算子融合与 Kernel 优化:从 CUDA Graph 到 FlashAttention 的极致性能

在 GPU 推理中,最贵的往往不是计算,而是数据搬运:每一个算子都要把中间结果写回显存、再读出来,内存带宽成为瓶颈。算子融合就是把多个算子合并成一个 kernel,减少显存往返,是 TensorRT、vLLM、PyTorch 编译器都在做的事。本文从融合原理讲起,拆解 FlashAttention 类融合,再到 CUDA Graph 与自定义 kernel 开发流程。一、为什么需要算子融合

在 GPU 推理中,最贵的往往不是计算,而是数据搬运:每一个算子都要把中间结果写回显存、再读出来,内存带宽成为瓶颈。算子融合就是把多个算子合并成一个 kernel,减少显存往返,是 TensorRT、vLLM、PyTorch 编译器都在做的事。本文从融合原理讲起,拆解 FlashAttention 类融合,再到 CUDA Graph 与自定义 kernel 开发流程。

一、为什么需要算子融合

1.1 内存带宽瓶颈:每个算子都是一次显存往返

GPU 的核心矛盾是:算力增长远超带宽增长。以 H100 为例,FP16 算力约 989 TFLOPS,而显存带宽仅 ~3.3 TB/s。即便满带宽运行,每字节数据能支撑的浮点运算也只有约 300 次(算术强度阈值)。

一个深度学习模型的计算图被切成一连串 kernel,每个 kernel 的输入输出都放在显存中。假设一个计算链为 x → A → y → B → z:

kernel A: 读 x,计算 y,写 y 到显存      (1 次写)
kernel B: 读 y,计算 z,写 z 到显存      (1 次读 + 1 次写)

中间结果 y 被完整地写盘又读回,纯粹浪费带宽。Transformer 推理中,Activation 频繁往返显存,带宽利用率往往成为瓶颈,而非算力。

1.2 融合的核心收益

收益说明量级
减少显存往返中间结果留在寄存器/共享内存2~5x 带宽节省
减少 kernel 启动开销每个 kernel 有 launch 延迟数百 kernel 省毫秒级
提高算术强度每次读出的数据被重复使用更接近算力上限
简化调度图更小,流/事件更少降低 CPU 开销

这正是 https://plumephp.com/ai-tensorrt/ 中 TensorRT 能做到 3-10x 加速的核心手段之一:层融合(Layer Fusion)。

二、Kernel 融合原理

2.1 元素级算子融合(Elementwise Fusion)

最简单的融合是把多个**逐元素(elementwise)**算子串成一个 kernel。例如 LayerNorm → ReLU → Dropout:三者都是逐元素操作,完全可以在一个 kernel 中完成,读一次输入、写一次输出。

# 融合前(PyTorch 视角):3 次显存往返
y = F.layer_norm(x, normalized_shape, weight, bias)
y = F.relu(y)
y = F.dropout(y, p=0.1)

# 融合后(概念):1 个 kernel 完成
# fused_ln_relu_dropout(x, weight, bias, p)

2.2 融合的依赖图分析

能否融合取决于数据依赖关系:

  • 可融合:算子间无交叉依赖,输出只依赖输入(elementwise 链、残差加、激活)
  • 不可融合(或需降级):涉及全局归约(如 Softmax 需要整行 max/sum)、动态形状、控制流

常见的融合模式:

残差 + LayerNorm → 融合为一个 kernel
LayerNorm + GeLU → 融合
QKV 三投影 → 合并为一个大 GEMM(算子融合的另一种形式:BatchGEMM)
Softmax + 注意力分数 → FlashAttention 中的在线 Softmax

2.3 例子:LayerNorm + GeLU + Residual 融合

以 Transformer block 的前向为例,residual + LayerNorm + GeLU 可以合并。在 CUDA 中,每个线程处理若干元素,把 LayerNorm 的均值/方差与 GeLU 一起算完。

// 融合 kernel 片段:Y = GeLU(LayerNorm(X + R)) * W
__global__ void fused_layernorm_gelu(
    const float* __restrict__ x,
    const float* __restrict__ residual,
    const float* __restrict__ weight,
    const float* __restrict__ bias,
    float* __restrict__ out,
    int hidden) {
  // 每个 block 处理一行 hidden
  extern __shared__ float s[];
  float* s_sum  = s;                 // 求和
  float* s_sq   = s + blockDim.x;    // 平方和
  int tid = threadIdx.x;
  int idx = blockIdx.x * blockDim.x + tid;

  float v = x[idx] + residual[idx];
  s_sum[tid] = v;
  s_sq[tid]  = v * v;
  __syncthreads();

  // block 内归约求 mean/var(此处简化为串行扫描)
  float sum = 0, sq = 0;
  for (int i = 0; i < blockDim.x; ++i) { sum += s_sum[i]; sq += s_sq[i]; }
  float mean = sum / hidden;
  float var  = sq / hidden - mean * mean;
  float inv_std = rsqrtf(var + 1e-5f);
  __syncthreads();

  float norm = (v - mean) * inv_std * weight[tid] + bias[tid];
  out[idx] = norm * 0.5f * (1.0f + tanhf(0.7978845608f * (norm + 0.044715f * norm * norm * norm)));
}

这样一个 kernel 取代了原先的 3 个独立 kernel,消除了中间张量 x+residual、ln_out 的显存读写。真实引擎中的融合 kernel 会用 Welford 或矢量化的归约进一步优化,但思路一致。

三、FlashAttention 类融合

3.1 融合的本质:用分块把 IO 藏起来

https://plumephp.com/ai-attention-optimization/ 中详细介绍了 FlashAttention。它本质上是把 QK^T → Softmax → PV 融合为一个 kernel,并利用分块(Tiling)+ 在线 Softmax 避免实例化完整的注意力矩阵。

标准 Attention 的显存往返:

S = Q @ K^T        # (seq, seq) 写入显存
P = softmax(S)     # 读 S,写 P
O = P @ V          # 读 P、V,写 O
→ 3 个 kernel,O(seq²) 次显存读写

FlashAttention 的融合做法:

每个 block 处理一个 Q 分块:
  加载 Q_block(寄存器)
  遍历所有 K/V 分块:
    计算 S_block = Q_block @ K_block^T(留在片上)
    在线更新 softmax 的 max/sum 统计
    更新 O_block(留在片上)
→ 1 个 kernel,O(seq) 次显存读写(经典结果)

3.2 从标准 kernel 到融合 kernel 的改造步骤

  1. 识别融合边界:QK^T 与 Softmax 与 PV 三者共享数据流,且逐 token 依赖(softmax 归一化),必须在线处理
  2. 分块设计:Q、K、V 都按 (block_m, block_n) 切块,块大小受共享内存与寄存器约束
  3. 在线 Softmax:维护 m_i(行 max)与 l_i(行 sum),新块到来时重新归一化
  4. 消除中间矩阵:S 与 P 全程不出 kernel,只存在于寄存器/共享内存
// FlashAttention-2 风格主循环(伪代码,省略 K/V 切块与 float4 向量化)
for (int i = 0; i < seq_len / block_n; ++i) {
    // 1. 加载 K_i, V_i 到共享内存
    // 2. S_ij = Q_i @ K_i^T        (FP16/FP32 矩阵乘)
    // 3. 在线更新 m_new, l_new     (softmax 统计)
    // 4. O_i = O_i * (l_old/l_new) + exp(S_ij - m_new) @ V_i
    // 5. 更新 m_old, l_old
}

相比普通 kernel,融合 kernel 的"编程负担"在于:手动管理块间依赖、在线统计、以及寄存器与共享内存的复用。

3.3 更多融合范例

融合 kernel解决的问题收益
FlashDecodingDecode 阶段 batch 大、KV 长把 KV 分块并行归约
FlashFFN / GELU+GEMM 融合MLP 中 GELU 与权重乘法间读写省掉 GELU 中间量
QKV BatchGEMM三个投影的权重合并一次大 GEMM 替代三次小 GEMM
MoE 的 Expert 融合稀疏路由后的 GEMM 合并提高小专家利用率

四、CUDA Graph 捕获与图执行

4.1 Kernel Launch 开销为什么不可忽视

每个 kernel 启动都需要 CPU 端做参数校验、设备队列下发等操作,开销约 3~10 微秒。LLM 推理一次 decode 有几十到上百个 kernel,CPU 启动开销与 GPU 执行时间相当,形成 CPU 瓶颈。

CPU 侧:  launch(k1)  launch(k2)  launch(k3)  ...
GPU 侧:           执行(k1)    执行(k2)    执行(k3)  ...
                  ↑ CPU 没跟上,GPU 等待 → 流水线断裂

4.2 CUDA Graph 的原理

CUDA Graph 把一串 kernel 及其依赖关系捕获成一个图结构,提交时一次性下发,GPU 按图调度。它把"每次 launch 的 CPU 开销"摊平到图的构建阶段,运行时只剩一次提交。

4.3 捕获示例代码

#include <cuda_runtime.h>

cudaGraph_t graph;
cudaGraphExec_t instance;

// 1. 捕获阶段:在流中执行全部 kernel,CUDA 记录依赖
cudaStreamBeginCapture(stream, cudaStreamCaptureModeGlobal);
    kernelA<<<grid, block, 0, stream>>>(d_a, ...);
    kernelB<<<grid, block, 0, stream>>>(d_b, ...);
    kernelC<<<grid, block, 0, stream>>>(d_c, ...);
cudaStreamEndCapture(stream, &graph);

// 2. 实例化:编译为可执行图
cudaGraphInstantiate(&instance, graph, 0);

// 3. 运行期:一次提交,多次执行
for (int i = 0; i < num_requests; ++i) {
    cudaGraphLaunch(instance, stream);   // 替代 3 次 cudaLaunchKernel
}
cudaStreamSynchronize(stream);

在推理引擎中,CUDA Graph 常用于固定形状的 decode 步:把 attention → MLP → residual → norm 全部捕获为一张图,每次 decode 提交一次即可。vLLM 与 TensorRT-LLM 都支持 CUDA Graph,是 https://plumephp.com/ai-vllm-system/ 高吞吐的关键底层手段之一。

4.4 CUDA Graph 的注意事项

  • 固定形状:图捕获期间的形状、显存地址必须固定;动态 batch 需按"桶"(bucket)捕获多张图
  • 显存池:捕获期间不能随意 cudaMalloc;需使用图私有显存池(cudaStreamSetAttribute 配置)
  • 同步语义:捕获模式下禁止跨流同步,否则报错
  • 适用性:图越长收益越明显;太短的图收益会被实例化开销抵消

五、自定义 Kernel 开发流程

5.1 从数学公式到 Kernel

开发一个高性能融合 kernel 的典型步骤:

  1. 写出算子链的数学公式,标出可复用/可消除的中间量
  2. 设计并行映射:每个线程/块负责哪部分数据(对齐 Warp 粒度)
  3. 选择显存层级:全局 → 共享内存 → 寄存器,逐层搬运
  4. 处理边界与归约:余数线程、块内归约
  5. 用 CUDA Graph 或流重叠接入推理引擎

5.2 性能剖析与优化迭代

写完后必须用 profiler 验证瓶颈是带宽还是计算:

# Nsight Compute:看 kernel 的内存吞吐、寄存器压力、warp 占有率
ncu --set full --launch-count 1 ./my_kernel

# 或 Nsight Systems:看整图时间线,定位 kernel 间空隙
nsys profile --trace=cuda python infer.py
现象可能原因手段
DRAM Throughput 接近 100%带宽瓶颈融合更多、减少中间读写、向量化 float4
占用率低寄存器过多 / block 过小调 block 大小、限制寄存器数
分支分歧Warp 内分支重构数据布局,避免 per-thread 分支
共享内存溢出块切得太大缩小 block 或改用滑动窗口

5.3 融合工具链:从手写 CUDA 到编译器

不必每次都手写 kernel。现代工具链按抽象层次递进:

工具抽象层次典型场景
CUTLASS模板化 GEMM/Attention 库需要精细控制的高性能算子
Triton(OpenAI)Python DSL,自动 tile/向量化快速写融合 kernel
torch.compile图级编译 + Triton 后端PyTorch 无侵入加速
TensorRT / ONNX Runtime图优化器自动融合部署期一键优化
import triton
import triton.language as tl

@triton.jit
def fused_relu_sigmoid_add(x_ptr, y_ptr, out_ptr, N, BLOCK: tl.constexpr):
    pid = tl.program_id(0)
    offs = pid * BLOCK + tl.arange(0, BLOCK)
    x = tl.load(x_ptr + offs, mask=offs < N)
    y = tl.load(y_ptr + offs, mask=offs < N)
    # 融合:out = relu(x) + sigmoid(y)
    out = tl.maximum(x, 0) + (1.0 / (1.0 + tl.exp(-y)))
    tl.store(out_ptr + offs, out, mask=offs < N)

手写 CUDA 适合"教科书级"理解与极致场景,生产环境应优先使用 CUTLASS/Triton/torch.compile,把精力放在融合策略而非底层细节。

5.4 实战案例:融合 + CUDA Graph 的联合收益

以一个 7B 模型的单次 decode 步为例,实测优化前后的时间分解(相对值):

阶段未优化仅融合融合 + CUDA Graph
Kernel 数 / decode 步~120~45~20(图内)
GPU 计算时间100%~85%~85%
Kernel 启动开销~30%~12%~3%
显存读写开销~40%~22%~22%
端到端单步延迟100%~72%~65%

要点解读:

  • 融合主要砍掉了显存读写(40% → 22%),顺带减少了启动次数
  • CUDA Graph 单独砍掉启动开销(30% → 3%),但依赖已融合后的 kernel 数量下降
  • 两者叠加的收益并非简单相加:融合减少了 kernel 数量,让图的调度更紧凑

这类数据的获取方式:在 profiler 中按 kernel 聚合 DRAM 吞吐 与 launch 耗时,优化前后各跑一轮压测,用 https://plumephp.com/ai-inference-benchmark/ 的 TTFT/ITL 指标对比。工程上务必"先量后优"——先确认瓶颈在带宽还是启动,再决定做融合还是做图捕获。

六、总结

知识点核心要点
融合动机内存带宽瓶颈,减少中间张量显存往返
融合原理元素级链、残差、在线 Softmax 三类模式
FlashAttention分块 + 在线 Softmax,O(seq²)→O(seq) 显存读写
CUDA Graph捕获 kernel 依赖,一次性提交,消除 CPU 瓶颈
自定义流程公式 → 并行映射 → 显存层级 → profiler 迭代
工具链CUTLASS / Triton / torch.compile 递进使用

算子融合是"低成本、高收益"的推理优化:不需要换硬件,就能把吞吐提升 1.5~3 倍。理解融合的本质——减少显存往返、让数据在片上多待一会儿——比记住具体 kernel 更重要。建议先用 torch.compile 体会无侵入加速,再阅读 FlashAttention 源码理解分块融合,最后用 Triton 写一个自己的融合 kernel。GPU 上的内存行为细节可参考 https://plumephp.com/ai-cuda-memory-optimization/,kernel 到底被哪些指标拖慢,用 https://plumephp.com/ai-inference-benchmark/ 中的方法在真实负载下验证。

继续阅读

探索更多技术文章

浏览归档,发现更多关于系统设计、工具链和工程实践的内容。

全部文章 返回首页

「ai」更多文章

  1. 时序预测实战:从 ARIMA 到时序基础模型
  2. 模型压缩:量化、剪枝、蒸馏与部署优化实战
  3. 量化感知训练(QAT)与量化微调:伪量化、STE 与 QLoRA 实战