在 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 的改造步骤
- 识别融合边界:QK^T 与 Softmax 与 PV 三者共享数据流,且逐 token 依赖(softmax 归一化),必须在线处理
- 分块设计:Q、K、V 都按
(block_m, block_n)切块,块大小受共享内存与寄存器约束 - 在线 Softmax:维护
m_i(行 max)与l_i(行 sum),新块到来时重新归一化 - 消除中间矩阵: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 | 解决的问题 | 收益 |
|---|---|---|
| FlashDecoding | Decode 阶段 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 的典型步骤:
- 写出算子链的数学公式,标出可复用/可消除的中间量
- 设计并行映射:每个线程/块负责哪部分数据(对齐 Warp 粒度)
- 选择显存层级:全局 → 共享内存 → 寄存器,逐层搬运
- 处理边界与归约:余数线程、块内归约
- 用 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/ 中的方法在真实负载下验证。
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。