大模型GPU在线推理中重要的默认优化思路-算子融合
如 QKV 线性层融合、Attention 中的 Softmax + 缩放 + Mask 融合、残差连接与 LayerNorm 融合;
为什么需要融合?做了些什么呢?怎么融合的呢?
为了彻底理解算子融合,不要把它当成抽象的概念,而是看成“减少 GPU 显存搬运次数”的优化手段。
GPU 的计算速度极快(算力),但显存读写速度相对较慢(带宽瓶颈)。
算子融合的核心目的就是:让数据尽可能在 GPU 的“高速缓存(SRAM)”里处理完,不要反复往“低速显存(HBM)”里存和取。
融合一:QKV 线性层融合(减少内核启动 + 减少显存读取)
标准的 Transformer 中,我们需要将输入 X 分别乘以三个权重矩阵,得到 Q、K、V。
-
输入
X:形状[Batch, Seq_Len, Hidden],假设为[1, 4, 8](1句话,4个词,每个词8维)。 -
权重矩阵:
W_q、W_k、W_v,形状均为[8, 4](为了计算简单,假设输出维度为4)。
Q = X @ W_q # 形状 [1, 4, 4] K = X @ W_k # 形状 [1, 4, 4] V = X @ W_v # 形状 [1, 4, 4]
发生了什么(显存视角):
-
从显存读取
X和W_q,计算 Q,将 Q 写回显存。 -
再次从显存读取
X和W_k,计算 K,将 K 写回显存。 -
再次从显存读取
X和W_v,计算 V,将 V 写回显存。
问题:X 被从显存重复读取了 3 次!启动了 3 次 GPU 内核,每次都要等数据搬运。
融合写法(实际工程做法)
# 先将 W_q, W_k, W_v 在内存层面拼成一个大矩阵 [8, 12] (4*3=12) W_qkv = torch.cat([W_q, W_k, W_v], dim=-1) # 只执行 1 次矩阵乘法! QKV = X @ W_qkv # 形状变为 [1, 4, 12] # 在最后一维(12)上切成 3 块,每块 4 维(注意:这只是逻辑切分,不复制显存!) Q, K, V = QKV.chunk(3, dim=-1)
详细举例演示整个过程吧!前面的优点抽象
第一步:先看“朴素写法”(不融合)
假设随机生成了三个权重矩阵(数字是随意编的,只为展示结构):
-
W_q (4行3列):
[1, 0, 2] [0, 1, 1] [1, 1, 0] [0, 0, 1]
-
W_k (4行3列):
[2, 1, 0] [1, 0, 1] [0, 2, 1] [1, 1, 0]
-
W_v (4行3列):
[0, 1, 1] [1, 0, 2] [1, 1, 0] [0, 1, 1]
朴素写法要算 3 次矩阵乘法:
-
Q = X @ W_q(读一遍 X,读一遍 W_q) -
K = X @ W_k(读一遍 X,读一遍 W_k) -
V = X @ W_v(读一遍 X,读一遍 W_v)
问题:X 被从显存里重复读取了 3 次,耗时!
第二步:再看“融合写法”(代码里的做法)
工程上的骚操作来了:不分开存 3 个小矩阵,先把它们横向拼成 1 个大矩阵!
拼接操作:W_qkv = torch.cat([W_q, W_k, W_v], dim=-1)
按列(最后一维)拼接。W_q 有 3 列,W_k 有 3 列,W_v 有 3 列,拼完后变成 4行,9列 的大矩阵。它长这样:
W_qkv (4行 × 9列):
[ 1, 0, 2 | 2, 1, 0 | 0, 1, 1 ] <- 第1行(前3列是W_q,中间3列是W_k,后3列是W_v) [ 0, 1, 1 | 1, 0, 1 | 1, 0, 2 ] <- 第2行 [ 1, 1, 0 | 0, 2, 1 | 1, 1, 0 ] <- 第3行 [ 0, 0, 1 | 1, 1, 0 | 0, 1, 1 ] <- 第4行
(用竖线 | 标出了三块的分界线)
第三步:执行唯一的矩阵乘法
现在,只执行 1 次矩阵乘法:QKV = X @ W_qkv
假设输入 X 是 [1, 2, 4],也就是两行数据(两个单词):
-
单词1:
[a1, a2, a3, a4] -
单词2:
[b1, b2, b3, b4]
X 乘以这个 4×9 的 W_qkv,得到的结果 QKV 形状为 [1, 2, 9]。
结果矩阵 QKV 长什么样? 它每一行有 9 个数字,而这 9 个数字天然就是 前3个是 Q,中间3个是 K,最后3个是 V!
假设算出来的第一行(单词1的结果)是:
[ q1, q2, q3 | k1, k2, k3 | v1, v2, v3 ]
第四步:执行 chunk(3, dim=-1) 切分
代码里的 Q, K, V = QKV.chunk(3, dim=-1),就是在最后一维(9)上切成 3 块,每块 3 维。
物理上发生的事:
-
Q指向第 1~3 列([q1, q2, q3]) -
K指向第 4~6 列([k1, k2, k3]) -
V指向第 7~9 列([v1, v2, v3])
关键点:这个 chunk 操作没有复制任何显存! 它只是创建了三个“指针”(视图),告诉 GPU:“以后读 Q 就去读大矩阵的第 1~3 列,读 K 就去读第 4~6 列。”
第五步:对比“显存搬运”的巨大差异
| 操作方式 | 内核启动次数 | 读取 X 的次数 | 显存写入中间结果的次数 |
|---|---|---|---|
| 朴素写法 | 3 次 (matmul) |
3 次 | 写 3 次 (Q, K, V分别写) |
| 融合写法 | 1 次 (matmul) |
1 次 | 只写 1 次大矩阵 (QKV) |
总结:代码里的那两行,到底在物理上干了什么?
# 1. 内存层面:把三块小铁板(W_q, W_k, W_v)焊成一块大铁板(W_qkv) W_qkv = torch.cat([W_q, W_k, W_v], dim=-1) # 2. 计算层面:用一把大锤子(矩阵乘法)砸下去,一次性算出 Q、K、V 的原始混合体 QKV = X @ W_qkv # 3. 逻辑层面:在大铁板上画两条线,左边叫Q,中间叫K,右边叫V(不切割铁板,只是做标记!) Q, K, V = QKV.chunk(3, dim=-1)
为什么大模型必须这么干?
因为 X 可能非常大(比如 [1, 2048, 4096]),如果重复读 3 次,显存带宽就成了瓶颈。
融合后,X 只读 1 次,计算一次搞定,直接省掉 2/3 的显存读取时间。
这就是工程上“算子融合”最朴素、最暴力的物理意义。
补充:
张量变化过程:
| 步骤 | 操作 | 张量形状 | 显存访问次数 |
|---|---|---|---|
| 朴素 | 3次 matmul |
每次读 [1,4,8],写 [1,4,4] |
读3次,写3次 |
| 融合 | 1次 matmul + 1次 chunk |
读 [1,4,8],写 [1,4,12] |
读1次,写1次 |
效果:
只发起了 1 次大型矩阵乘法(GPU 利用率更高),显存读写次数直接降至原来的 1/3。
这就是为什么 vLLM 和 Hugging Face 的 from_pretrained 都默认开启 QKV 融合。
=========================================================================
融合二:Attention 中的 Softmax + 缩放(Scale)+ Mask 融合(消除中间变量)
第 0 步:回顾前面讲的(我们手里有什么?)
假设上一讲中,输入 X(2个单词)经过融合后的 W_qkv 大矩阵相乘并 chunk 切分后,得到了如下 Q、K、V(数字是为方便计算精心设计的):
-
Q(2行 × 3列,代表 2 个 Query):Q = [ 1, 2, 1 ] <- 第1个单词产生的 Query [ 3, 1, 2 ] <- 第2个单词产生的 Query
-
K(2行 × 3列,代表 2 个 Key):K = [ 0, 1, 1 ] <- 第1个单词产生的 Key [ 2, 1, 0 ] <- 第2个单词产生的 Key
第 1 步:计算原始注意力分数 Scores = Q @ K^T(矩阵乘法)
要把 Q 的每一行,分别与 K 的每一行做点积(内积)。注意,K 需要转置(K^T),形状变成 [3, 2]。
计算第 1 行(Q1 对所有 Key 的打分):
-
Q1 =
[1, 2, 1]与 K1 =[0, 1, 1]点积:1×0 + 2×1 + 1×1 = **3** -
Q1 =
[1, 2, 1]与 K2 =[2, 1, 0]点积:1×2 + 2×1 + 1×0 = **4**
计算第 2 行(Q2 对所有 Key 的打分):
-
Q2 =
[3, 1, 2]与 K1 =[0, 1, 1]点积:3×0 + 1×1 + 2×1 = **3** -
Q2 =
[3, 1, 2]与 K2 =[2, 1, 0]点积:3×2 + 1×1 + 2×0 = **7**
最终得到的原始分数矩阵 Scores(形状 2×2):
Scores = [ 3, 4 ] <- 第1个Query对第1、2个Key的打分 [ 3, 7 ] <- 第2个Query对第1、2个Key的打分
注意:这里的数值越大,代表两个单词的“相关性”越强。
第 2 步:缩放(Scale)—— 除以 √d
假设模型超参数 √d = 1(为了演示简单,缩放因子为 1,数值不变)。
如果 √d = 2,那就是把每个数字除以 2。这里我们假设除以 1,所以 Scores 保持不变。
第 3 步:应用因果掩码(Causal Mask)
在自回归语言模型中,第 1 个单词(位置0)不能看到第 2 个单词(位置1),所以矩阵右上角(第0行第1列)必须被遮住,设为 -inf(负无穷)。
Mask 后的矩阵(记作 S_masked):
S_masked = [ 3, -inf ] <- 第1个Query只看得到自己(3),看不到未来(屏蔽) [ 3, 7 ] <- 第2个Query能看到过去和自己(3和7)
第 4 步:Softmax 归一化
现在对每一行单独做 Softmax(公式:e^x / sum(e^x))。
第 0 行:[3, -inf]
-
e^3 = 20.085,e^(-inf) = 0 -
分母 =
20.085 + 0 = 20.085 -
结果 =
[20.085/20.085, 0/20.085]=[1.0, 0.0]
第 1 行:[3, 7]
-
e^3 = 20.085,e^7 = 1096.63 -
分母 =
20.085 + 1096.63 = 1116.715 -
结果 =
[20.085/1116.715, 1096.63/1116.715]≈[0.018, 0.982]
最终注意力矩阵 Attn:
Attn = [ 1.000, 0.000 ] <- 第1个单词 100% 注意自己 [ 0.018, 0.982 ] <- 第2个单词 1.8% 注意第1个,98.2% 注意自己
关键:朴素写法 VS 融合写法(显存里到底发生了什么?)
朴素写法(如果不做融合):
GPU 会傻乎乎地在显存里按顺序创建这些中间大矩阵:
-
写出
Scores(原始分数)占一份显存。 -
除以
√d,写出S_scaled(缩放后)占第二份显存。 -
应用 Mask,写出
S_masked(掩码后)占第三份显存。 -
算 Softmax,写出最终的
Attn占第四份显存。
显存峰值:4 份矩阵同时存在!读写显存次数高达 6 次(读Scores、写A、读A、写B、读B、写C)。
融合写法(FlashAttention / torch.compile 的做法):
GPU 内部只开一个内核,在 寄存器(超快缓存) 里完成所有操作,只把最后的 Attn 写回显存。
伪代码执行流程(在 GPU 核心内部):
-
第一趟遍历(求分母):
-
读取
Scores第 0 行[3, 4],发现有 Mask(右上角),无视 4,记录max=3,累加器sum_exp = e^3 = 20.085。 -
读取
Scores第 1 行[3, 7],记录max=7,累加器sum_exp = e^3 + e^7 = 20.085 + 1096.63 = 1116.715。
-
-
第二趟遍历(写结果):
-
再次读取
Scores第 0 行[3, 4],计算e^3 / 20.085 = 1.0,有 Mask 的位置写0.0,直接写入显存中的最终Attn矩阵。 -
再次读取
Scores第 1 行[3, 7],计算e^3/1116.715 ≈ 0.018,e^7/1116.715 ≈ 0.982,直接追加写入显存中的Attn。
-
显存峰值:只有原始的 Scores 和最终的 Attn 两份!读写显存次数仅 2 次(读一次 Scores,写一次 Attn)。
总览全景图:从输入到最终注意力(完整串联)
| 阶段 | 操作 | 输入形状 | 输出形状 | 显存中间变量(朴素) | 显存中间变量(融合) |
|---|---|---|---|---|---|
| QKV 融合 | X @ W_qkv + chunk |
[2,4] |
Q,K,V [2,3] |
无(一步到位) | 无(一步到位) |
| 算分数 | Q @ K^T |
[2,3] + [3,2] |
Scores [2,2] |
存在 | 存在 |
| 缩放 | / √d |
[2,2] |
S_scaled [2,2] |
多出一份 | 不存在! |
| Mask | 右上角置 -inf |
[2,2] |
S_masked [2,2] |
再多出一份 | 不存在! |
| Softmax | 指数归一化 | [2,2] |
Attn [2,2] |
最终结果 | 最终结果 |
最终结论
融合优化的本质,就是砍掉了上表中标注的“缩放”和“Mask”两步的显存写入。
在真实的 70B 大模型中,Scores 矩阵可能是 [1, 32, 4096, 4096](约 5 亿个元素,2GB 显存)。
如果朴素写法,中间会额外产生 2 个这样的 2GB 临时张量,直接导致显存溢出(OOM)。
而融合写法(比如调用的 F.scaled_dot_product_attention 或开启了 @torch.compile),把这 2GB 的中间变量永远地消灭在了寄存器里,这就是它能让你在同样的显卡上跑起来且速度快几倍的终极奥秘。
专业说法:“算子融合优化的本质是 Kernel 融合。它将多个依赖的数学操作(如 Scale、Mask、Softmax)合并到一个 GPU 内核中执行。
数据从全局显存(HBM)加载到片上高速缓存(SRAM/寄存器)后,所有中间结果都不再写回全局显存,而是在寄存器中流水线式完成计算,最终只将唯一的结果张量写回显存。
这极大减少了全局内存的访问次数(IO 带宽瓶颈)和显存峰值占用。”
融合三:残差连接与 LayerNorm(RMSNorm)融合(减少往返读写)
在标准 Pre-Norm 架构中,公式是:Output = RMSNorm( X + Residual )。
-
X:当前子层(如 Attention)的输出,形状[1, 4, 8]。 -
Residual:残差连接(原始输入),形状[1, 4, 8]。
朴素写法(未融合):
# 步骤1:残差相加
temp = X + Residual # 生成中间张量 temp,写回显存 (100MB)
# 步骤2:计算 RMSNorm(需要读 temp)
variance = temp.pow(2).mean(dim=-1, keepdim=True)
output = temp / sqrt(variance) * weight
发生了什么:
-
temp是一个 100MB 的中间张量,被完整地写入显存,又被立刻读出来做归一化。 -
浪费了一次巨大的显存写入和读取。
融合写法(代码已经包含了,但可以更深一层):
在 1 个内核中完成:
# 伪代码:读取 X 和 Residual 的块
for each element i:
# 计算加和(直接在寄存器中,不写回显存)
tmp_i = X[i] + Residual[i]
# 计算临时累加和(用于求 mean)
sum_sq += tmp_i ** 2
# 计算完 mean 后,再次遍历(或流式处理):
for each element i:
tmp_i = X[i] + Residual[i]
output[i] = tmp_i / sqrt(sum_sq/N + eps) * weight[i]
# 只把最终的 output 写回显存
张量变化过程(显存足迹):
| 模式 | 中间变量 temp | 显存读写次数 |
|---|---|---|
| 朴素(先加后归一化) | 存在(100MB) | 读2次+写1次(写 temp,读 temp) |
| 融合(内核内相加) | 不存在 | 读1次(只读 X 和 Residual,直接输出) |
代码实现注意:
def residual_rms_forward(self, x, residual) 返回了 (归一化结果, x+residual),这在逻辑上很清晰。
但如果要极致性能融合,应该把 x = x + residual 和 rms_forward 内部的 pow、mean 放在同一个 @torch.compile 函数里,编译器会自动帮你消除 temp 的显存分配。
======================================================
一张表看懂三大融合
| 融合场景 | 朴素写法瓶颈 | 融合后的效果 | 性能提升来源 |
|---|---|---|---|
| QKV 线性层 | 3 次内核启动,X 读 3 次 |
1 次内核启动,X 读 1 次 |
减少内存带宽消耗 3 倍 |
| Softmax+Mask | 3 个中间大张量(200% 显存浪费) | 0 个中间张量 | 显存占用降低 75%,减少 2 次全局读写 |
| 残差+LN | 必须写出 100MB 的加法结果 | 加法在寄存器完成,不写显存 | 省去 1 次大张量写入 + 1 次大张量读取 |
工程落地:
-
PyTorch 2.0+ 的
@torch.compile会自动帮做后两种融合(只要不刻意写中间变量)。 -
FlashAttention 专门做第二种融合。
-
vLLM / TensorRT-LLM 在模型加载时,会把 QKV 的权重自动拼接,实现第一种融合。
现在写的这两个模块(SiluAndMul 和 RMSNorm),配上 @torch.compile,刚好踩中了算子融合的所有关键点,所以它们的运行效率会很接近底层 C++ 手写内核的水平!
浙公网安备 33010602011771号