09_attention_02
1. 单头自注意力的数学计算
我们首先来看单头自注意力的数学计算过程。
以句子 “I am good” 为例,设词嵌入维度 ,.
指的是模型维度 / 词嵌入维度,也就是每个 token(词)用一个多长的向量来表示,也是模型的“隐藏层宽度”。在我们举出的例子中,“I am good” 这三个词,每个词就被表示为一个 512 维的向量:
就不同了,它指的是每个注意力头 Key / Query 的维度,它主要决定的事每个注意力头“视野”的精细程度。在我们举出的例子中,通过权重矩阵 ,我们把 512 维的词向量投影到一个更小的 64 维空间来计算注意力:
两者的关系为:
在 BERT 模型中,,而 ,所以 .
也可以用下图总结:
d_model = 512 (整个向量的宽度) ┌──────────────────────────────────┐ │ ← head₁ (d_k=64) → │ │ │ │ ← head₂ → │ ... │ ... │ │ └──────────────────────────────────┘
d_k = 每个头分到的一小段 = d_model / h

- 计算 与 的转置点积
- 先通过可学习的权重矩阵 ,将源序列的词嵌入矩阵 (, 为序列长度)转化为 矩阵(均为 ).
- 计算 ,结果为 的相似度矩阵,每个值代表某一词的 与另一词的 的匹配程度,即词间的语义关联度。

- 缩放,除以
- 将相似度矩阵的所有元素除以 .
- 目的:防止 过大时, 的值过大,导致 softmax 后梯度趋于 0 (梯度消失),保证训练时梯度的稳定性。

- Softmax 归一化
- 对缩放后的相似度矩阵做 Softmax 运算,将每个行的数值归一化到 0~1 的范围,且每行求和为 1 ;
- 归一化后的数值为注意力得分,代表某一词对序列中其他词的 “关注程度”。

- 与 矩阵相乘生成注意力输出
- 将 Softmax 后的注意力得分矩阵()与 矩阵()相乘,得到最终的自注意力输出矩阵 ();
- 中每个词的特征向量,是原 矩阵中所有词的特征向量按注意力得分加权求和的结果,即融合了序列内所有相关词的语义信息。


- 通过上述自注意力机制,我们就能明白句子中每个词与其他词的关联程度。

2. 多头自注意力
而多头自注意力的计算过程如下:
- 将 通过不同的权重矩阵拆分为 个注意力头,每个头独立执行上述四步自注意力计算,得到 个注意力输出矩阵 .
- 将 个 按列拼接为一个大矩阵().
- 用一个可学习的权重矩阵 对拼接后的矩阵做线性变换,恢复为与原始词嵌入相同的维度 ,得到多头自注意力的最终输出。
公式:
BERT Self-layer 实际上是通过 MHA(Multi Head Attention) 分块矩阵实现的。
2.1 什么是 BERT layer?
一个 BERT layer(BERT 层)是 BERT 模型的核心构建单元。BERT 模型由多个这样的层堆叠而成(BERT-base 有 12 层,BERT-large 有 24 层)。
每一层包含两个主要部分:
可以通俗理解:Attention 负责”看上下文”(这个词和其他词有什么关系?),FFN 负责”思考加工”(把信息做非线性变换)。
2.2 什么是 MHA(Multi-Head Attention)?
MHA 就是多头注意力。先理解”单头”:
-
单头注意力:一个词只从一种角度去看其他词。比如”bank”这个词,可能只看金融角度的关联。
-
多头注意力(MHA):同时从多个角度去看。比如:
- 🟡 头 1:看语法关系(“bank”和”the”的关系)
- 🟢 头 2:看语义关系(“bank”和”money”的关系)
- 🔵 头 3:看指代关系(“bank”和”it”的关系)
- … 共 个头(BERT 中 )
实现方式:每个头有自己独立的 权重矩阵,并行计算,最后拼接起来:
2.3 什么是分块矩阵?
分块矩阵就是把一个大矩阵看成由若干小矩阵(“块”)拼接而成。
举个例子,一个 的大矩阵可以分成 2 个 的小块:
2.4 为什么说 BERT layer 是由 MHA 分块矩阵实现的?
这是关键的理解点!它说的是 MHA 的 权重矩阵在实现上使用了分块矩阵的技巧。
2.4.1 直观理解
BERT 有 12 个注意力头,每个头都需要自己的 。与其分别做 12 次小矩阵乘法:
# ❌ 低效做法:逐个计算每个头for i in range(12): Qi = Q @ Wq_i # 每个 Wq_i 是 [768, 64] Ki = K @ Wk_i Vi = V @ Vk_i不如把它们拼成一个大矩阵,一次性计算:
# ✅ 高效做法:分块矩阵Wq_big = torch.cat([Wq_1, Wq_2, ..., Wq_12], dim=-1) # [768, 768]Q_big = Q @ Wq_big # 一次性计算所有头的 Q!2.4.2 用图来理解
每个头独立的小 W_q 拼成大矩阵 W_q ┌──────┐ ┌──────┐ ┌──────┐ ┌─────────────────────┐ │ W_q¹ │ │ W_q² │ ... │ W_q¹²│ │ W_q¹ │ W_q² │...│W_q¹²│ │768×64│ │768×64│ │768×64│ │ 768×768 │ └──────┘ └──────┘ └──────┘ └─────────────────────┘ ↑ ↑ 12次矩阵乘法 1次矩阵乘法搞定!这就是分块矩阵的思想:12 个小的 (每个 )横向拼接成一个大的 (),一次矩阵乘法就能得到所有 12 个头的 ,然后按块拆分给各个头。
2.4.3 总结一句话
不是说 BERT layer 物理上被”切开”了,而是说 MHA 的实现利用分块矩阵的技巧,把多个注意力头的权重拼成一个大矩阵,用一次大矩阵乘法代替多次小矩阵乘法,在 GPU 上高效并行计算。
这正是 PyTorch 中 nn.MultiheadAttention 和 HuggingFace BERT 源码中实际采用的做法!我们可以在之前学过的 05_BERT+embedding_source_code 中看到源码验证这一点。
2.5 例子
我们最后用一个例子来说明 BERT 中 Self-layer MHA 矩阵分块的实现过程。假设序列长度 ,标准 BERT 模型 , , .
- 对于每个头,我们首先将词嵌入矩阵 乘以权重矩阵 来得到 :
- 计算注意力得分矩阵:,得到 的矩阵(每个词对每个词的关注度)。
- 与 相乘得到该头的输出:,维度为 。
- 对 12 个头分别执行步骤 1~3,得到 (每个都是 )。
- 按列拼接所有 :,得到 的大矩阵。
- 乘以输出权重矩阵 ,得到最终 MHA 输出:。
2.5.1 维度变化总览
X (10×768) │ ├─[头1]──→ Q₁(10×64), K₁(10×64), V₁(10×64) ──→ Z₁(10×64) ├─[头2]──→ Q₂(10×64), K₂(10×64), V₂(10×64) ──→ Z₂(10×64) ├─ ... ... └─[头12]─→ Q₁₂(10×64), K₁₂(10×64), V₁₂(10×64) ──→ Z₁₂(10×64) │ 按列拼接 Concat ↓ (10×768) │ × W_O(768×768) ↓ MHA输出 (10×768)2.5.2 分块矩阵实现(实际代码的做法)
上述”逐个计算 12 个头”只是为了理解,实际代码中不会用 for 循环,而是利用分块矩阵一次性完成。
核心技巧:把 12 个 的小 横向拼接成一个 的大 :
然后一行代码算出所有头的 :
Q_big = X @ W_q_big # (10×768) @ (768×768) = (10×768)得到的 是 ,但这个 768 列其实是 12 个 的块拼起来的:
接下来用 reshape + transpose 把它拆成多头格式:
# (10, 768) → (10, 12, 64) → (12, 10, 64)Q = Q_big.view(10, 12, 64).transpose(0, 1) # 把 batch 和 head 互换K = K_big.view(10, 12, 64).transpose(0, 1)V = V_big.view(10, 12, 64).transpose(0, 1).transpose().transpose(a,b) 就是交换张量的第 a 维和第 b 维。
这里互换 batch 和 head 的原因是我们后续要批量计算注意力,所以需要把“头数”放在第 0 维。这样,当我们调用 Pytorch 的矩阵乘法时,就会自动在 12 个头上 并行广播。
此时 都是 ,即 12 个头 × 10 个词 × 64 维,可以直接批量计算注意力:
# 批量计算所有头的注意力:(12, 10, 64) @ (12, 64, 10) = (12, 10, 10)attn_scores = softmax(Q @ K.transpose(-2, -1) / sqrt(64)) # 转置每个头的 K 矩阵,即 K^T
# (12, 10, 10) @ (12, 10, 64) = (12, 10, 64)Z = attn_scores @ V最后把多头拼回去:
# (12, 10, 64) → (10, 12, 64) → (10, 768)Z_concat = Z.transpose(0, 1).reshape(10, 768)output = Z_concat @ W_O # (10, 768) @ (768, 768) = (10, 768)2.5.3 对比总结
| 概念上的做法(逐头循环) | 实际代码的做法(分块矩阵) | |
|---|---|---|
| 形状 | 12 个 | 1 个 |
| 矩阵乘法次数 | 12 次 | 1 次 |
| 中间形状 | 先 ,再 view 拆成 | |
| GPU 效率 | 低(串行) | 高(并行) |
💡 一句话:分块矩阵就是把 12 个”小水管”焊成 1 根”大水管”,数据一次性流过,然后在出口处再分流给 12 个头。
文章分享
如果这篇文章对你有帮助,欢迎分享给更多人!