09_attention_02

2499 字
12 分钟
09_attention_02
Attention(Q,K,V)=softmax(QKTdk)VAttention(Q,K,V)=softmax(\frac{QK^T}{\sqrt{d_k}})V

1. 单头自注意力的数学计算#

我们首先来看单头自注意力的数学计算过程。

以句子 “I am good” 为例,设词嵌入维度 dmodel=512d_{model}=512dk=64d_k=64.

dmodeld_{model}dkd_k

dmodel=512d_{model}=512 指的是模型维度 / 词嵌入维度,也就是每个 token(词)用一个多长的向量来表示,也是模型的“隐藏层宽度”。在我们举出的例子中,“I am good” 这三个词,每个词就被表示为一个 512 维的向量:

X=[Iamgood]3 个词    [x1,1x1,2x1,512x2,1x2,2x2,512x3,1x3,2x3,512]3×dmodel=3×512X = \underbrace{\begin{bmatrix} \text{I} \\ \text{am} \\ \text{good} \end{bmatrix}}_{3 \text{ 个词}} \;\longrightarrow\; \underbrace{\begin{bmatrix} x_{1,1} & x_{1,2} & \cdots & x_{1,512} \\ x_{2,1} & x_{2,2} & \cdots & x_{2,512} \\ x_{3,1} & x_{3,2} & \cdots & x_{3,512} \end{bmatrix}}_{3 \times d_{model} = 3 \times 512}

dk=64d_k=64 就不同了,它指的是每个注意力头 Key / Query 的维度,它主要决定的事每个注意力头“视野”的精细程度。在我们举出的例子中,通过权重矩阵 Wq(512×64)W_q(512\times 64),我们把 512 维的词向量投影到一个更小的 64 维空间来计算注意力:

XWq=[3×512]输入[512×64]Wq=[3×64]Q 矩阵X \cdot W_q = \underbrace{\begin{bmatrix} 3 \times 512 \end{bmatrix}}_{\text{输入}} \cdot \underbrace{\begin{bmatrix} 512 \times 64 \end{bmatrix}}_{W_q} = \underbrace{\begin{bmatrix} 3 \times 64 \end{bmatrix}}_{Q \text{ 矩阵}}

两者的关系为:

dk=dmodelhd_k=\frac{d_{model}}{h}

在 BERT 模型中,dmodel=768d_{model}=768,而 h=12h=12,所以 dk=768/12=64d_k=768/12=64.

也可以用下图总结:

d_model = 512 (整个向量的宽度)
┌──────────────────────────────────┐
│ ← head₁ (d_k=64) → │ │
│ │ ← head₂ → │ ...
│ ... │ │
└──────────────────────────────────┘
d_k = 每个头分到的一小段 = d_model / h

  1. 计算 QQKK 的转置点积 QKTQK^T
    • 先通过可学习的权重矩阵 Wq,Wk,WvW_q,W_k,W_v ,将源序列的词嵌入矩阵 XXn×dmodel=3×512n\times d_{model}=3\times 512nn 为序列长度)转化为 Q,K,VQ,K,V 矩阵(均为 n×dkn\times d_k).
    • 计算 QKTQK^T ,结果为 n×nn\times n 的相似度矩阵,每个值代表某一词的 QQ 与另一词的 KK 的匹配程度,即词间的语义关联度。

  1. 缩放,除以 dk\sqrt{d_k}
    • 将相似度矩阵的所有元素除以 dk\sqrt{d_k} .
    • 目的:防止 dkd_k 过大时,QKTQK^T 的值过大,导致 softmax 后梯度趋于 0 (梯度消失),保证训练时梯度的稳定性

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

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

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

2. 多头自注意力#

而多头自注意力的计算过程如下:

  1. Q,K,VQ,K,V 通过不同的权重矩阵拆分为 hh 个注意力头,每个头独立执行上述四步自注意力计算,得到 hh 个注意力输出矩阵 Z1,Z2,...,ZhZ_1,Z_2,...,Z_h .
  2. hhZiZ_i 按列拼接为一个大矩阵(n×h×dkn\times h\times d_k).
  3. 用一个可学习的权重矩阵 W0W_0 ​对拼接后的矩阵做线性变换,恢复为与原始词嵌入相同的维度 n×dmodeln×d_{model},得到多头自注意力的最终输出。

公式:

Multiheadattention=Concatenate(Z1,Z2,...,Zh)W0Multi-head\enspace attention=Concatenate(Z_1,Z_2,...,Z_h)W_0

BERT Self-layer 实际上是通过 MHA(Multi Head Attention) 分块矩阵实现的。

2.1 什么是 BERT layer?#

一个 BERT layer(BERT 层)是 BERT 模型的核心构建单元。BERT 模型由多个这样的层堆叠而成(BERT-base 有 12 层,BERT-large 有 24 层)。

每一层包含两个主要部分:

graph LR Input[输入向量] --> MHA[🔷 Multi-Head Attention<br/>多头注意力] MHA --> Add1[相加 + LayerNorm] Add1 --> FFN[🔶 Feed-Forward Network<br/>前馈网络] FFN --> Add2[相加 + LayerNorm] Add2 --> Output[输出向量]

可以通俗理解:Attention 负责”看上下文”(这个词和其他词有什么关系?),FFN 负责”思考加工”(把信息做非线性变换)。


2.2 什么是 MHA(Multi-Head Attention)?#

MHA 就是多头注意力。先理解”单头”:

  • 单头注意力:一个词只从一种角度去看其他词。比如”bank”这个词,可能只看金融角度的关联。

  • 多头注意力(MHA):同时从多个角度去看。比如:

    • 🟡 头 1:看语法关系(“bank”和”the”的关系)
    • 🟢 头 2:看语义关系(“bank”和”money”的关系)
    • 🔵 头 3:看指代关系(“bank”和”it”的关系)
    • … 共 hh 个头(BERT 中 h=12h=12

实现方式:每个头有自己独立的 Q,K,VQ,K,V 权重矩阵,并行计算,最后拼接起来:

headi=Attention(QWiQ,KWiK,VWiV)\text{head}_i = \text{Attention}(QW_i^Q, KW_i^K, VW_i^V)MHA=Concat(head1,head2,...,headh)WO\text{MHA} = \text{Concat}(\text{head}_1, \text{head}_2, ..., \text{head}_h) \cdot W^O

2.3 什么是分块矩阵?#

分块矩阵就是把一个大矩阵看成由若干小矩阵(“块”)拼接而成。

举个例子,一个 4×64 \times 6 的大矩阵可以分成 2 个 4×34 \times 3 的小块:

[a11a12a13a14a15a16a21a22a23a24a25a26a31a32a33a34a35a36a41a42a43a44a45a46]一个大矩阵 4×6=[A1A2]看成 2 个 4×3 的小块拼在一起\underbrace{\begin{bmatrix} a_{11} & a_{12} & a_{13} & | & a_{14} & a_{15} & a_{16} \\ a_{21} & a_{22} & a_{23} & | & a_{24} & a_{25} & a_{26} \\ a_{31} & a_{32} & a_{33} & | & a_{34} & a_{35} & a_{36} \\ a_{41} & a_{42} & a_{43} & | & a_{44} & a_{45} & a_{46} \end{bmatrix}}_{\text{一个大矩阵 } 4\times6} = \underbrace{\begin{bmatrix} \mathbf{A}_1 & | & \mathbf{A}_2 \end{bmatrix}}_{\text{看成 2 个 } 4\times3 \text{ 的小块拼在一起}}

2.4 为什么说 BERT layer 是由 MHA 分块矩阵实现的?#

这是关键的理解点!它说的是 MHA 的 Q,K,VQ,K,V 权重矩阵在实现上使用了分块矩阵的技巧

2.4.1 直观理解#

BERT 有 12 个注意力头,每个头都需要自己的 WiQ,WiK,WiVW_i^Q, W_i^K, W_i^V。与其分别做 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 个小的 WqW_q(每个 768×64768 \times 64)横向拼接成一个大的 WqW_q768×768768 \times 768),一次矩阵乘法就能得到所有 12 个头的 QQ,然后按块拆分给各个头。

2.4.3 总结一句话#

不是说 BERT layer 物理上被”切开”了,而是说 MHA 的实现利用分块矩阵的技巧,把多个注意力头的权重拼成一个大矩阵,用一次大矩阵乘法代替多次小矩阵乘法,在 GPU 上高效并行计算。

这正是 PyTorch 中 nn.MultiheadAttention 和 HuggingFace BERT 源码中实际采用的做法!我们可以在之前学过的 05_BERT+embedding_source_code 中看到源码验证这一点。

2.5 例子#

我们最后用一个例子来说明 BERT 中 Self-layer MHA 矩阵分块的实现过程。假设序列长度 n=10n=10,标准 BERT 模型 dmodel=768d_{model}=768, h=12h=12, dk=64d_k=64.

  1. 对于每个头,我们首先将词嵌入矩阵 X(10×768)X(10\times 768) 乘以权重矩阵 W(768×64)W(768\times 64) 来得到 Q,K,VQ,K,V
    • X@Wq=Q(10×64)X @ W_q=Q(10\times 64)
    • X@Wk=K(10×64)X @ W_k=K(10\times 64)
    • X@Wv=V(10×64)X @ W_v=V(10\times 64)
  2. 计算注意力得分矩阵:softmax(QKTdk)softmax(\frac{Q\cdot K^T}{\sqrt{d_k}}),得到 10×1010\times 10 的矩阵(每个词对每个词的关注度)。
  3. VV 相乘得到该头的输出:Zi=softmax(QKTdk)VZ_i = softmax(\frac{QK^T}{\sqrt{d_k}}) \cdot V,维度为 10×6410\times 64
  4. 对 12 个头分别执行步骤 1~3,得到 Z1,Z2,...,Z12Z_1, Z_2, ..., Z_{12}(每个都是 10×6410\times 64)。
  5. 按列拼接所有 ZiZ_iConcat(Z1,Z2,...,Z12)Concat(Z_1, Z_2, ..., Z_{12}),得到 10×76810\times 768 的大矩阵。
  6. 乘以输出权重矩阵 WO(768×768)W_O(768\times 768),得到最终 MHA 输出:10×76810\times 768

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 个 (768×64)(768\times 64) 的小 WqW_q 横向拼接成一个 (768×768)(768\times 768) 的大 WqW_q

Wqbig=[Wq1Wq2Wq12]768×768W_q^{\text{big}} = \underbrace{[W_q^1 \mid W_q^2 \mid \cdots \mid W_q^{12}]}_{768 \times 768}

然后一行代码算出所有头的 QQ

Q_big = X @ W_q_big # (10×768) @ (768×768) = (10×768)

得到的 QbigQ_{big}10×76810\times 768,但这个 768 列其实是 12 个 dk=64d_k=64 的块拼起来的:

Qbig=[Q1Q2Q12]10×768Q_{big} = \underbrace{[Q_1 \mid Q_2 \mid \cdots \mid Q_{12}]}_{10\times 768}

接下来用 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 个头上 并行广播

此时 Q,K,VQ,K,V 都是 (12,10,64)(12, 10, 64),即 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 对比总结#

概念上的做法(逐头循环)实际代码的做法(分块矩阵)
WqW_q 形状12 个 768×64768\times 641 个 768×768768\times 768
矩阵乘法次数12 次1 次
中间形状(10,64)×12(10,64)\times 12(10,768)(10,768),再 view 拆成 (12,10,64)(12,10,64)
GPU 效率低(串行)高(并行)

💡 一句话:分块矩阵就是把 12 个”小水管”焊成 1 根”大水管”,数据一次性流过,然后在出口处再分流给 12 个头。

文章分享

如果这篇文章对你有帮助,欢迎分享给更多人!

09_attention_02
https://github.com/chunhuizhang/bilibili_vlogs/
作者
HAC
发布于
2026-06-21
许可协议
CC BY-NC-SA 4.0

评论区

Profile Image of the Author
HAC
观之非易,行且克难
Greetings
欢迎来到我的博客!这里主要分享我的学习笔记与兴趣爱好。
音乐
封面

音乐

暂未播放

0:00 0:00
暂无歌词
分类
标签
站点统计
文章
32
分类
5
标签
13
总字数
79,889
运行时长
0
最后活动
0 天前

文章目录