大模型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_qW_kW_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]

发生了什么(显存视角)

  1. 从显存读取 XW_q,计算 Q,将 Q 写回显存。

  2. 再次从显存读取 XW_k,计算 K,将 K 写回显存。

  3. 再次从显存读取 XW_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列):

    text
    [1, 0, 2]
    [0, 1, 1]
    [1, 1, 0]
    [0, 0, 1]
  • W_k (4行3列):

    text
    [2, 1, 0]
    [1, 0, 1]
    [0, 2, 1]
    [1, 1, 0]
  • W_v (4行3列):

    text
    [0, 1, 1]
    [1, 0, 2]
    [1, 1, 0]
    [0, 1, 1]

朴素写法要算 3 次矩阵乘法

  1. Q = X @ W_q (读一遍 X,读一遍 W_q)

  2. K = X @ W_k (读一遍 X,读一遍 W_k)

  3. 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列):

text
[ 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的结果)是:

text
[ 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):

    text
    Q = 
    [ 1,  2,  1 ]    <- 第1个单词产生的 Query
    [ 3,  1,  2 ]    <- 第2个单词产生的 Query
  • K (2行 × 3列,代表 2 个 Key):

    text
    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):

text
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

text
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.085e^(-inf) = 0

  • 分母 = 20.085 + 0 = 20.085

  • 结果 = [20.085/20.085, 0/20.085] = [1.0, 0.0]

第 1 行[3, 7]

  • e^3 = 20.085e^7 = 1096.63

  • 分母 = 20.085 + 1096.63 = 1116.715

  • 结果 = [20.085/1116.715, 1096.63/1116.715][0.018, 0.982]

最终注意力矩阵 Attn

text
Attn = 
[ 1.000,  0.000 ]    <- 第1个单词 100% 注意自己
[ 0.018,  0.982 ]    <- 第2个单词 1.8% 注意第1个,98.2% 注意自己

 

关键:朴素写法 VS 融合写法(显存里到底发生了什么?)

朴素写法(如果不做融合):

GPU 会傻乎乎地在显存里按顺序创建这些中间大矩阵:

  1. 写出 Scores(原始分数)占一份显存。

  2. 除以 √d,写出 S_scaled(缩放后)占第二份显存。

  3. 应用 Mask,写出 S_masked(掩码后)占第三份显存。

  4. 算 Softmax,写出最终的 Attn 占第四份显存。

显存峰值:4 份矩阵同时存在!读写显存次数高达 6 次(读Scores、写A、读A、写B、读B、写C)。

 

融合写法(FlashAttention / torch.compile 的做法):

GPU 内部只开一个内核,在 寄存器(超快缓存) 里完成所有操作,只把最后的 Attn 写回显存

伪代码执行流程(在 GPU 核心内部)

  1. 第一趟遍历(求分母)

    • 读取 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

  2. 第二趟遍历(写结果)

    • 再次读取 Scores 第 0 行 [3, 4],计算 e^3 / 20.085 = 1.0,有 Mask 的位置写 0.0,直接写入显存中的最终 Attn 矩阵。

    • 再次读取 Scores 第 1 行 [3, 7],计算 e^3/1116.715 ≈ 0.018e^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]

朴素写法(未融合):

python
# 步骤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 个内核中完成:

python
# 伪代码:读取 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 + residualrms_forward 内部的 powmean 放在同一个 @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 的权重自动拼接,实现第一种融合。

现在写的这两个模块(SiluAndMulRMSNorm),配上 @torch.compile,刚好踩中了算子融合的所有关键点,所以它们的运行效率会很接近底层 C++ 手写内核的水平!

 

 

 

posted on 2026-08-25 17:12  zhangkele  阅读(26)  评论(0)    收藏  举报

导航