从零手推 Self-Attention 与 Multi-Head Attention
极简数值算例推导、多头拆解、FLOPs算力实测与位置编码全解析
全网 10w+ 爆款深度复刻与完整重构:从 RNN 串行瓶颈到 Transformer 并行架构,结合 12 张高清全流程图解、纯手工矩阵数值推导演练、多头机制工程切分等价性、FLOPs 算力守恒实测、位置编码实验与 2026 现代大模型注意力架构跃迁。
从 2017 年 Google 提出《Attention Is All You Need》,到如今大语言模型(LLM)、多模态大模型与 Vision Transformer(ViT)席卷全球,Self-Attention(自注意力)与 Multi-Head Attention(多头注意力) 已成为整个人工智能时代的底层动力心脏。
如果你之前在网上找过 Self-Attention 或 Transformer 的相关资料,大概率看到的都是论文里那几张抽象的矩阵方块图与高深公式,正如李宏毅老师在课程里所调侃:“不懂的人再怎么看原图也不会懂”。
本文完整搬运并系统整理了国内著名技术博主“太阳花的小绿豆”的万字经典图解教程,包含全部 12 张高清全彩手绘流程图,结合李宏毅老师的教学脉络与底层 PyTorch 源码实现,带你从最原始的两节点极简数值一步一步手工算清每一个矩阵!
📖 一、前言:为什么 Transformer 能终结 RNN 时代?
在 Transformer 诞生之前,自然语言处理(NLP)领域几乎被循环神经网络(RNN、LSTM、GRU)彻底统治。然而,传统 RNN 结构存在两个无法逾越的致命缺陷:
- 无法并行化(Sequential Bottleneck):RNN 必须按时间步 $t_1, t_2, \dots, t_n$ 依次向前串行计算,只有计算完 $t_i$ 时刻的隐状态数据,才能开始计算 $t_{i+1}$ 时刻的数据,计算效率极低,无法发挥现代 GPU 大规模矩阵并行的硬件优势。
- 长距离遗忘与信息瓶颈:虽然 LSTM/GRU 引入了门控机制,但在面对长文本时,早期输入的语义信息依然会在反复的循环迭代中被剧烈稀释或遗忘。
2017 年,Google 团队在论文《Attention Is All You Need》中抛弃了所有的循环与卷积结构,提出了基于纯注意力机制的 Transformer。它不仅能够在一瞬间捕捉序列中任意两个词之间的全局依赖关系,更实现了序列内部所有 Token 的全并行化计算!

图 1:原论文《Attention Is All You Need》中给出的 Scaled Dot-Product Attention(左)与 Multi-Head Attention(右)架构图
🧮 二、Self-Attention(自注意力机制):极简数值手推演练
为了让所有人都能彻底理解,我们构造一个序列长度为 2 的极简示例:输入就两个节点 $x_1, x_2$。
整个计算流程分为以下核心步骤:
- 特征嵌入(Input Embedding):通过映射函数 $f(x)$(通常为 Embedding 查找层或全连接层)将离散输入节点映射为连续的词向量 $a_1, a_2$;
- 线性变换生成 Q、K、V:分别将 $a_1, a_2$ 与三个变换矩阵 $W^q, W^k, W^v$ 相乘(这三个参数在整个序列中是共享且可训练的,工程源码中通常通过全连接层
nn.Linear实现,此处为方便手推理解忽略偏置 bias),得到对应的 $q_i, k_i, v_i$:- $q$(Query,查询向量):后续会主动去和每一个词的 $k$ 进行点乘匹配;
- $k$(Key,键向量):后续会被每一个词的 $q$ 匹配;
- $v$(Value,值向量):代表从原始输入 $a$ 中提取得到的实际特征信息载荷;
- $q$ 与 $k$ 匹配的过程可以直观理解为计算两个词之间的相关性(相似度)——相关性越大,对应 $v$ 所分配到的注意力权重就越大。

图 2:Self-Attention 输入节点映射与 Q、K、V 线性变换全流程手绘拆解图
1. 手工演练:Q、K、V 的并行矩阵计算
为了计算简便,我们设定特征维度为 2,具体数值如下:
$$a_1 = (1, 1), \quad a_2 = (1, 0)$$ $$W^q = \begin{pmatrix} 1 & 1 \\ 0 & 1 \end{pmatrix}, \quad W^k = \begin{pmatrix} 1 & 0 \\ 0 & 1 \end{pmatrix}, \quad W^v = \begin{pmatrix} 1 & 0 \\ 0 & 1 \end{pmatrix}$$首先单独计算各个节点的 Query 向量:
$$q^1 = a_1 W^q = (1, 1) \begin{pmatrix} 1 & 1 \\ 0 & 1 \end{pmatrix} = (1 \times 1 + 1 \times 0, \; 1 \times 1 + 1 \times 1) = (1, 2)$$ $$q^2 = a_2 W^q = (1, 0) \begin{pmatrix} 1 & 1 \\ 0 & 1 \end{pmatrix} = (1 \times 1 + 0 \times 0, \; 1 \times 1 + 0 \times 1) = (1, 1)$$由于 Transformer 是全并行架构,在实际硬件中我们无需串行逐个计算,直接将输入向量按行堆叠拼成矩阵 $A$ 一次性完成矩阵乘法:
$$Q = \begin{pmatrix} q^1 \\ q^2 \end{pmatrix} = \begin{pmatrix} a_1 \\ a_2 \end{pmatrix} W^q = \begin{pmatrix} 1 & 1 \\ 1 & 0 \end{pmatrix} \begin{pmatrix} 1 & 1 \\ 0 & 1 \end{pmatrix} = \begin{pmatrix} 1 & 2 \\ 1 & 1 \end{pmatrix}$$同理,可一次性并行计算出 Key 矩阵 $K$ 与 Value 矩阵 $V$:
$$K = \begin{pmatrix} k^1 \\ k^2 \end{pmatrix} = \begin{pmatrix} a_1 \\ a_2 \end{pmatrix} W^k = \begin{pmatrix} 1 & 1 \\ 1 & 0 \end{pmatrix} \begin{pmatrix} 1 & 0 \\ 0 & 1 \end{pmatrix} = \begin{pmatrix} 1 & 0 \\ 0 & 1 \end{pmatrix}$$ $$V = \begin{pmatrix} v^1 \\ v^2 \end{pmatrix} = \begin{pmatrix} a_1 \\ a_2 \end{pmatrix} W^v = \begin{pmatrix} 1 & 1 \\ 1 & 0 \end{pmatrix} \begin{pmatrix} 1 & 0 \\ 0 & 1 \end{pmatrix} = \begin{pmatrix} 1 & 0 \\ 0 & 1 \end{pmatrix}$$(注:在作者图解算例中,为直观演示取基底向量 $k^1=(1,0), k^2=(0,1)$ 及 $v^1=(1,0), v^2=(0,1)$,以下严格按照该数值进行点积匹配演算)
2. 匹配得分计算、缩放因子 $\sqrt{d}$ 与 Softmax 归一化
接着,拿查询向量 $q^1$ 与每一个 $k$ 分别进行点乘匹配(Dot-Product),并除以缩放因子 $\sqrt{d}$ 得到未归一化的相似度得分 $\alpha$(其中 $d$ 为 Key 向量的维度,在本例中 $d=2$,故 $\sqrt{d} = \sqrt{2} \approx 1.414$):
$$\alpha_{1, 1} = \frac{q^1 \cdot k^1}{\sqrt{d}} = \frac{1 \times 1 + 2 \times 0}{\sqrt{2}} = \frac{1}{\sqrt{2}} \approx 0.71$$ $$\alpha_{1, 2} = \frac{q^1 \cdot k^2}{\sqrt{d}} = \frac{1 \times 0 + 2 \times 1}{\sqrt{2}} = \frac{2}{\sqrt{2}} \approx 1.41$$
同理,拿查询向量 $q^2$ 去匹配所有的 $k$:
$$\alpha_{2, 1} = \frac{q^2 \cdot k^1}{\sqrt{d}} = \frac{1 \times 1 + 1 \times 0}{\sqrt{2}} = \frac{1}{\sqrt{2}} \approx 0.71$$ $$\alpha_{2, 2} = \frac{q^2 \cdot k^2}{\sqrt{d}} = \frac{1 \times 0 + 1 \times 1}{\sqrt{2}} = \frac{1}{\sqrt{2}} \approx 0.71$$将其统一写成全局矩阵运算形式:
$$\begin{pmatrix} \alpha_{1, 1} & \alpha_{1, 2} \\ \alpha_{2, 1} & \alpha_{2, 2} \end{pmatrix} = \frac{\begin{pmatrix} q^1 \\ q^2 \end{pmatrix} \begin{pmatrix} k^1 \\ k^2 \end{pmatrix}^T}{\sqrt{d}} = \frac{Q K^T}{\sqrt{d}} = \begin{pmatrix} 0.71 & 1.41 \\ 0.71 & 0.71 \end{pmatrix}$$紧接着,对矩阵的每一行分别进行 Softmax 归一化处理,得到最终归一化后的注意力权重分布 $\hat{\alpha}$:
$$\hat{\alpha}_{1, 1} = \frac{e^{0.71}}{e^{0.71} + e^{1.41}} \approx 0.33, \quad \hat{\alpha}_{1, 2} = \frac{e^{1.41}}{e^{0.71} + e^{1.41}} \approx 0.67$$ $$\hat{\alpha}_{2, 1} = \frac{e^{0.71}}{e^{0.71} + e^{0.71}} = 0.50, \quad \hat{\alpha}_{2, 2} = \frac{e^{0.71}}{e^{0.71} + e^{0.71}} = 0.50$$到这里,我们就完整推导完了公式中 $\text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)$ 的全部中间结果。

图 3:Q 与 K 矩阵转置点乘、缩放因子 $\sqrt{d}$ 以及逐行 Softmax 归一化手绘推导图
原论文的解释非常精辟:假设向量 $q$ 与 $k$ 的各个分量是均值为 0、方差为 1 的独立同分布随机变量,那么它们的点积 $q \cdot k = \sum_{i=1}^{d_k} q_i k_i$ 的均值为 0,但方差会随着维度线性膨胀为 $d_k$(标准差为 $\sqrt{d_k}$)。
当维度 $d_k$ 很大时(例如 64 或 128),点积数值的绝对值会非常大,导致输入到 Softmax 函数后落入导数极小的饱和区(Saturation Region),引发致命的梯度消失(Gradient Vanishing)!因此除以 $\sqrt{d_k}$ 能将方差稳定重置为 1,确保反向传播梯度通畅。
3. 与 Value 向量加权求和得到最终输出 $b$
计算出针对每个 $v$ 的注意力权重 $\hat{\alpha}$ 后,最后一步就是与 Value 向量进行加权求和,得到融合了全局上下文信息的最终输出特征向量 $b_1, b_2$:
$$b_1 = \hat{\alpha}_{1, 1} v^1 + \hat{\alpha}_{1, 2} v^2 = 0.33 \times (1, 0) + 0.67 \times (0, 1) = (0.33, 0.67)$$ $$b_2 = \hat{\alpha}_{2, 1} v^1 + \hat{\alpha}_{2, 2} v^2 = 0.50 \times (1, 0) + 0.50 \times (0, 1) = (0.50, 0.50)$$统一写成标准矩阵乘法形式:
$$\begin{pmatrix} b_1 \\ b_2 \end{pmatrix} = \begin{pmatrix} \hat{\alpha}_{1, 1} & \hat{\alpha}_{1, 2} \\ \hat{\alpha}_{2, 1} & \hat{\alpha}_{2, 2} \end{pmatrix} \begin{pmatrix} v^1 \\ v^2 \end{pmatrix} = \begin{pmatrix} 0.33 & 0.67 \\ 0.50 & 0.50 \end{pmatrix} \begin{pmatrix} 1 & 0 \\ 0 & 1 \end{pmatrix} = \begin{pmatrix} 0.33 & 0.67 \\ 0.50 & 0.50 \end{pmatrix}$$到这里,Self-Attention 的全部计算过程就推导完毕了。总结下来正是论文中名垂青史的标准数学公式:
$$\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V$$
图 4:注意力权重矩阵与 Value 矩阵相乘得到最终上下文融合向量 $b_1, b_2$
👥 三、Multi-Head Attention(多头注意力机制):原理与拆解
在实际工业级模型与原论文中,几乎从不使用单头 Self-Attention,而是全部采用 Multi-Head Attention(多头注意力机制)。
“Multi-head attention allows the model to jointly attend to information from different representation subspaces at different positions.”
—— 多头注意力机制允许模型在不同的空间位置,同时联合关注来自不同表征子空间的信息。
其实只要搞懂了 Self-Attention 模块,Multi-Head Attention 就变得极其简单直观:
- 首先依然和 Self-Attention 一样,将输入节点特征通过 $W^q, W^k, W^v$ 映射得到总的 $q_i, k_i, v_i$(总特征维度为 $D$);
- 然后再根据设定的 head 数量 $h$,直接在通道维度上将得到的 $q_i, k_i, v_i$ 均分成 $h$ 份(每个头的子维度 $d_k = D / h$);
- 例如下图中假设 $q^1 \in \mathbb{R}^4$,拆分成 $q^{1, 1} \in \mathbb{R}^2$ 与 $q^{1, 2} \in \mathbb{R}^2$,那么 $q^{1, 1}$ 分配给 Head 1,而 $q^{1, 2}$ 分配给 Head 2。

图 5:根据 Head 数量在通道维度上将 $q, k, v$ 均分为不同子空间的拆解图
1. 论文公式 vs 工程源码:线性映射与切分的数学等价性
看到这里,读过原论文的读者可能会产生疑问:论文中写的是通过 $h$ 组独立的投影矩阵 $W_i^Q, W_i^K, W_i^V$ 分别映射得到各个 head 的 $Q_i, K_i, V_i$:
$$\text{head}_i = \text{Attention}(Q W_i^Q, K W_i^K, V W_i^V)$$为什么在 GitHub 官方与各大开源库(如 HuggingFace、PyTorch 官方 nn.MultiheadAttention)中,都是先算一个大矩阵再直接切分(Split / Reshape)?
原因在于两者的数学本质是完全等价的! 我们可以将 $W_i^Q$ 设定为对角掩码切片矩阵,如下图所示,大矩阵 $Q$ 乘以特定的块投影矩阵,得到的结果与直接在张量通道上切片完全一致,但后者的计算能够以单个超大 GEMM 矩阵乘法充分打满 GPU Tensor Core 的并行吞吐!

图 6:论文中各 Head 独立投影矩阵 $W_i^Q$ 与工程中直接切分均分的数学等价性证明
2. 各 Head 独立计算、通道 Concat 拼接与 $W^O$ 融合
通过上述切分得到每个 Head 对应的 $Q_i, K_i, V_i$ 后,接下来针对每个 Head 独立执行标准的 Self-Attention 计算:
$$\text{head}_i = \text{Attention}(Q_i, K_i, V_i) = \text{softmax}\left(\frac{Q_i K_i^T}{\sqrt{d_k}}\right)V_i$$
图 7:各个 Head 在各自独立的子空间中并行执行 Self-Attention 运算
接着将各个 head 独立计算得到的结果在通道维度上进行 Concat 拼接:例如下图中将 Head 1 得到的 $b_1^1$ 与 Head 2 得到的 $b_1^2$ 拼在一起得到第一个节点的拼接特征;将 Head 1 的 $b_2^1$ 与 Head 2 的 $b_2^2$ 拼在一起得到第二个节点的拼接特征:

图 8:各 Head 输出的子特征向量在通道维度重新 Concat 拼接回原始维度
最后,将拼接后的结果通过一个可学习的输出融合矩阵 $W^O \in \mathbb{R}^{D \times D}$ 进行线性变换融合,得到最终的多头聚合输出向量 $b_1, b_2$:

图 9:Concat 拼接结果通过可学习权重矩阵 $W^O$ 线性融合得到最终的多头输出
总结下来,Multi-Head Attention 正是由原论文中的两个核心公式完全定义:
$$\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, \dots, \text{head}_h)W^O$$ $$\text{where} \quad \text{head}_i = \text{Attention}(Q W_i^Q, K W_i^K, V W_i^V)$$⚡ 四、FLOPs 算力之谜:单头与多头的计算量完全一致?!
在原论文第 3.2.2 节最后,作者写下了一句非常反直觉的结论:
“Due to the reduced dimension of each head, the total computational cost is similar to that of single-head attention with full dimensionality.”
—— 由于每个头的维度按比例缩小,多头注意力的总计算成本与全维度的单头注意力几乎完全一致!
为了严谨验证这一结论,我们使用 Facebook 开源的 fvcore 库进行真实的 FLOPs(浮点运算量)对比测试:
import torch
from fvcore.nn import FlopCountAnalysis
from model import Attention # 标准多头注意力模块实现
def main():
# 1. 构造输入张量:[Batch_Size=32, Num_Tokens=1024, Total_Embed_Dim=512]
t = (torch.rand(32, 1024, 512),)
# 2. 单头 Self-Attention (num_heads=1),并将输出投影矩阵 Wo 替换为 Identity (不作任何操作)
a1 = Attention(dim=512, num_heads=1)
a1.proj = torch.nn.Identity() # 移除单头中原本不需要的 Wo
flops1 = FlopCountAnalysis(a1, t)
print("Self-Attention (无 Wo) FLOPs:", flops1.total())
# 3. 完整的 8头 Multi-Head Attention (num_heads=8, 包含 Wo 线性映射层)
a2 = Attention(dim=512, num_heads=8)
flops2 = FlopCountAnalysis(a2, t)
print("Multi-Head Attention (含 Wo) FLOPs:", flops2.total())
# 4. 将 8头 Multi-Head Attention 的 Wo 也替换为 Identity
a2.proj = torch.nn.Identity()
flops3 = FlopCountAnalysis(a2, t)
print("Multi-Head Attention (无 Wo) FLOPs:", flops3.total())
if __name__ == '__main__':
main()
数学证明与本质揭秘:
设序列长度为 $L$,特征总维度为 $D$,头数为 $h$,每个头分配到的子维度 $d_k = D / h$。
在核心注意力矩阵乘法阶段($Q K^T$ 以及 $\text{Attn} \times V$):
- 单头注意力($h=1, d_k=D$):$Q K^T$ 的乘法次数为 $L \times D \times L = L^2 D$;$\text{Attn} \times V$ 的乘法次数为 $L \times L \times D = L^2 D$;
- 多头注意力($h$ 个头,每个头 $d_k = D/h$):每个头计算 $Q_i K_i^T$ 的乘法次数为 $L \times (D/h) \times L = L^2 (D/h)$;全部 $h$ 个头的乘法总次数为 $h \times L^2 (D/h) = L^2 D$!同理,$\text{Attn}_i \times V_i$ 的多头总乘法次数也为 $h \times L^2 (D/h) = L^2 D$!
📍 五、置换等变性与位置编码(Positional Encoding)
如果仔细观察 Self-Attention 的计算公式,会发现一个重大的内在特性:注意力计算对输入 Token 的先后顺序是完全无感的(Permutation Equivariance 置换等变性)!
假设输入序列为 $a_1, a_2, a_3$,对于 $a_1$ 而言,$a_2$ 和 $a_3$ 离它的距离计算规则完全相同,没有任何先后顺序概念。如果将输入顺序调换为 $a_1, a_3, a_2$,输出向量 $b_1$ 的数值不会发生哪怕任何一点点改变!
我们使用 PyTorch 官方的 nn.MultiheadAttention 编写一段极简测试来实证这一点:
import torch
import torch.nn as nn
# 创建一个单头注意力模块
m = nn.MultiheadAttention(embed_dim=2, num_heads=1)
# 输入两组序列:t1 顺序为 (1, 2, 3);t2 将第 2 与第 3 个词的顺序对调为 (1, 3, 2)
t1 = [[[1., 2.], # q1, k1, v1
[2., 3.], # q2, k2, v2
[3., 4.]]] # q3, k3, v3
t2 = [[[1., 2.], # q1, k1, v1
[3., 4.], # q3, k3, v3
[2., 3.]]] # q2, k2, v2
# 将 t1 输入模块前向传播
q, k, v = torch.as_tensor(t1), torch.as_tensor(t1), torch.as_tensor(t1)
print("result1 (t1 前向输出): \n", m(q, k, v)[0])
# 将 t2 输入模块前向传播
q, k, v = torch.as_tensor(t2), torch.as_tensor(t2), torch.as_tensor(t2)
print("result2 (t2 前向输出): \n", m(q, k, v)[0])

图 10:PyTorch 实测:调换输入 Token 顺序后,第一个 Token 的输出向量数值完全一致
实验结果表明,$b_1$ 的输出毫无变化!然而在人类自然语言中,“张三打了李四”和“李四打了张三”语义截然相反,语序至关重要。
为了让模型感知语序,原论文在输入端的 Embedding 向量上直接逐元素相加(Element-wise Add)了位置编码(Positional Encodings):
$$x_{\text{input}} = x_{\text{embedding}} + PE$$即 $\{pe_1, pe_2, \dots, pe_n\}$ 与 $\{a_1, a_2, \dots, a_n\}$ 拥有完全相同的维度大小。

图 11:位置编码向量与词嵌入向量逐元素相加示意图
关于位置编码,原论文提出了两种主流方案:
- 固定三角函数绝对编码(Sinusoidal Functions):利用不同频率的正弦和余弦函数直接计算生成位置向量,无需额外训练参数,且具备推断比训练长度更长序列的理论潜力: $$PE_{(pos, 2i)} = \sin\left(\frac{pos}{10000^{2i / d_{\text{model}}}}\right)$$ $$PE_{(pos, 2i+1)} = \cos\left(\frac{pos}{10000^{2i / d_{\text{model}}}}\right)$$
- 可学习位置编码(Learnable Positional Embedding):将位置看作可学习的嵌入查找表参数矩阵(如 BERT 与 Vision Transformer ViT 中所采用)。论文实验表明两者的最终下游性能基本持平。
📊 六、Transformer 核心超参数精解(论文 Table 3)
原论文《Attention Is All You Need》在第 3 节与表 3 中给出了不同模型规模下的标准超参数配置:
- $N$:重复堆叠 Transformer Encoder / Decoder Block 的层数;
- $d_{\text{model}}$:Multi-Head Self-Attention 输入与输出的 Token 隐藏层总维度;
- $d_{\text{ff}}$:Feed-Forward 前馈网络(MLP)中间隐藏层的节点维度;
- $h$:Multi-Head Self-Attention 中 Head 的个数;
- $d_k, d_v$:每个单头分配到的 Key / Query / Value 维度($d_k = d_v = d_{\text{model}} / h$);
- $P_{\text{drop}}$:Dropout 层的丢弃率(防止过拟合)。

图 12:原论文 Table 3 给出的一系列 Transformer 超参数变体与性能指标对照表
| 模型版本 | 堆叠层数 $N$ | 隐藏总维度 $d_{\text{model}}$ | MLP 中间维度 $d_{\text{ff}}$ | 头数 $h$ | 单头维度 $d_k, d_v$ | Dropout $P_{\text{drop}}$ | 参数总量 (Params) |
|---|---|---|---|---|---|---|---|
| Base 模型 | 6 | 512 | 2048 | 8 | 64 | 0.1 | 65M |
| Big 模型 | 6 | 1024 | 4096 | 16 | 64 | 0.3 | 213M |
💻 七、工业级 PyTorch 源码实现(手写纯净版)
以下是包含 QKV 大矩阵融合投影、通道重整(Reshape/Transpose)、Scaled Dot-Product 注意力计算与最终 $W^O$ 映射的标准 PyTorch 模块实现:
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
class ScaledDotProductAttention(nn.Module):
"""
Scaled Dot-Product Attention (缩放点积注意力核心算子)
"""
def __init__(self, dropout: float = 0.0):
super().__init__()
self.dropout = nn.Dropout(dropout) if dropout > 0.0 else nn.Identity()
def forward(self, q, k, v, mask=None):
# q, k, v 形状: [Batch_Size, Num_Heads, Seq_Len, Head_Dim]
d_k = q.size(-1)
# 1. Q 与 K^T 点积匹配并除以 sqrt(d_k)
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d_k)
# 2. 如果存在注意力掩码 (如因果掩码 / padding mask),将掩码区域置为 -inf
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
# 3. 逐行 Softmax 归一化得到注意力权重
attn_weights = F.softmax(scores, dim=-1)
attn_weights = self.dropout(attn_weights)
# 4. 与 V 矩阵相乘加权求和
output = torch.matmul(attn_weights, v)
return output, attn_weights
class MultiHeadAttention(nn.Module):
"""
Multi-Head Attention (标准多头注意力模块)
"""
def __init__(self, d_model: int = 512, num_heads: int = 8, dropout: float = 0.1):
super().__init__()
assert d_model % num_heads == 0, "d_model 必须能被 num_heads 整除!"
self.d_model = d_model
self.num_heads = num_heads
self.d_k = d_model // num_heads
# 融合投影层:一次性计算 Q, K, V,大幅提升 GPU 吞吐
self.qkv_proj = nn.Linear(d_model, 3 * d_model, bias=False)
self.attention = ScaledDotProductAttention(dropout=dropout)
self.out_proj = nn.Linear(d_model, d_model, bias=False)
def forward(self, x, mask=None):
B, L, D = x.shape
# 1. 计算 QKV 并在通道维度拆分
qkv = self.qkv_proj(x) # [B, L, 3 * D]
qkv = qkv.reshape(B, L, 3, self.num_heads, self.d_k) # [B, L, 3, H, d_k]
qkv = qkv.permute(2, 0, 3, 1, 4) # [3, B, H, L, d_k]
q, k, v = qkv[0], qkv[1], qkv[2]
# 2. 并行执行多头 Scaled Dot-Product Attention
out, weights = self.attention(q, k, v, mask=mask) # [B, H, L, d_k]
# 3. 将各 Head 结果重新 Concat 拼接回 [B, L, D]
out = out.permute(0, 2, 1, 3).contiguous().reshape(B, L, D)
# 4. 经过输出矩阵 Wo 线性融合
return self.out_proj(out)
# 单元测试验证
if __name__ == '__main__':
x = torch.randn(2, 16, 512) # Batch=2, Seq_Len=16, Dim=512
mha = MultiHeadAttention(d_model=512, num_heads=8)
out = mha(x)
assert out.shape == (2, 16, 512), "输出形状异常!"
print("MultiHeadAttention 前向测试通过!输出张量形状:", out.shape)
🚀 八、2026 现代大模型注意力演进前沿
在当代超长上下文大模型(如 DeepSeek-V3、LLaMA-3、Qwen-2.5)中,传统的 Multi-Head Attention 与绝对位置编码已经演进出了四大颠覆性技术标准:
1. RoPE (旋转位置编码)
通过正交复数二维旋转矩阵将相对位置信息直接编码进 $Q$ 和 $K$ 的内积中($\langle R_m q, R_n k \rangle = q^T R_{n-m} k$),具备极其出色的外推性(Context Length Extrapolation),已成为现代开源大模型的绝对通用标配。
<div class="p-4 rounded-lg bg-white/70 dark:bg-slate-900/70 border border-cyan-200 dark:border-cyan-800">
<h4 class="font-bold text-slate-900 dark:text-white text-sm mb-1">2. GQA (分组查询注意力) 与 MQA</h4>
<p class="text-base text-slate-600 dark:text-slate-300">
介于 MHA 与 MQA 之间,让多个 Query 头共享同一组 Key/Value 头(如 8 个 Query 共享 1 个 KV 头),将生成阶段的 <strong>KV Cache 显存占用直接降低 80% 以上</strong>,显著突破大模型高并发吞吐瓶颈。
</p>
</div>
<div class="p-4 rounded-lg bg-white/70 dark:bg-slate-900/70 border border-cyan-200 dark:border-cyan-800">
<h4 class="font-bold text-slate-900 dark:text-white text-sm mb-1">3. FlashAttention 3 (IO 感知硬件加速)</h4>
<p class="text-base text-slate-600 dark:text-slate-300">
利用 GPU 超高速 SRAM 片上缓存进行分块计算(Tiling)与在线 Softmax 动态更新,彻底消除了将 $L \times L$ 庞大注意力矩阵反复读写显存(HBM)的带宽墙瓶颈,实现 3~5 倍端到端加速与线性显存开销。
</p>
</div>
<div class="p-4 rounded-lg bg-white/70 dark:bg-slate-900/70 border border-cyan-200 dark:border-cyan-800">
<h4 class="font-bold text-slate-900 dark:text-white text-sm mb-1">4. PagedAttention 与 MLA (多头潜在注意力)</h4>
<p class="text-base text-slate-600 dark:text-slate-300">
借用虚拟内存分页机制解决 KV Cache 显存碎片问题(vLLM 核心);同时 DeepSeek 提出的 MLA 进一步将 KV 压缩为极低维度的潜在隐向量(Latent Vector),在大模型规模扩展中实现超高效推理。
</p>
</div>
📚 资料出处与致谢
- 原创核心参考出处:CSDN 博客 · 《详解Transformer中Self-Attention以及Multi-Head Attention》(作者:太阳花的小绿豆)
- Bilibili 配套精讲视频:太阳花的小绿豆 · Transformer 核心架构精讲视频教程
- 学术原著论文:Vaswani, Ashish, et al. “Attention Is All You Need.” Advances in Neural Information Processing Systems (NeurIPS 2017) · arXiv:1706.03762
- 全量搬运、图例重构与 2026 现代架构增补:Jenny Zhang · Jenny’s Space(人工智能与视觉专栏)