多头注意力机制 - 原理图(下)
1 课程概览
本课讲解多头注意力机制原理图下半部分。QKV 三个 Linear 层给 QKV 做投影(换形式,维度不变),为分头做准备。多个头并行计算自己的注意力,得到自己视角的结果。Concat 将多个头的结果拼接(8×64=512)。最后 Linear 调整拼接好的结果,让不同头的信息更好融合,维度不变。
2 核心概念与定义
- Q(Query):查询张量,要查什么内容。
- K(Key):索引,被查的内容。
- V(Value):答案,实际要提取的信息。
- Linear 层:线性变换,给 QKV 做投影。
- 投影:换个形式,维度不变。
- 分头:将词向量维度分成多个头。
- 并行计算:每个头单独计算自己的注意力。
- Concat:拼接多个头的结果。
- 最终 Linear:调整拼接好的结果,让不同头的信息更好融合。
3 模型与算法详解
多头注意力流程(详细)
QKV → 三个 Linear 层(投影)→ 分头 → 并行 Attention → Concat → 最终 Linear → 输出
各步骤详解
步骤1:QKV
| 张量 | 含义 |
|---|---|
| Q | 查询张量,要查什么 |
| K | 索引,被查的内容 |
| V | 答案,实际要提取的信息 |
步骤2:三个 Linear 层(投影)
给 QKV 做投影,换个形式,维度不变,为分头做准备。
- 原始:(2, 4, 512)
- Linear 后:(2, 4, 512)
- 维度不变,只是换个形式
类比:原食材切片、切丝、切段,还是原食材,没有经过处理,只是改刀。
步骤3:分头
将 512 维分成 8 个头,每个头 64 维。
- 原始:(2, 4, 512)
- 分头后:(2, 8, 4, 64)
- 每个头单独计算自己的注意力
步骤4:并行 Attention
每个头单独计算自己的注意力,得到自己视角的结果。
- 每个头 64 维
- 8 个头并行计算
- 得到 8 个视角的注意力结果
类比:让各组长收作业,每个组知道自己组的情况,别的组不知道。
步骤5:Concat(拼接)
将多个头的结果拼接,8×64=512。
- 每个头结果:(2, 4, 64)
- 8 个头拼接:(2, 4, 512)
- 维度恢复为 512
类比:班长收齐各组的作业,再交给老师。
步骤6:最终 Linear
调整拼接好的结果,让不同头的信息更好融合,维度不变。
- 输入:(2, 4, 512)
- 输出:(2, 4, 512)
- 维度不变,但信息融合更好
类比:吃火锅时,青菜、羊肉、毛肚蘸麻酱,充分融合,感受丰富的层次感。
4 数学原理与推导
完整公式
$$\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, ..., \text{head}_h)W^O$$
各步骤公式
投影(Linear): $$Q' = QW^Q, \quad K' = KW^K, \quad V' = VW^V$$
分头: $$Q' \rightarrow [Q_1, Q_2, ..., Q_h]$$
并行 Attention: $$\text{head}_i = \text{Attention}(Q_i, K_i, V_i)$$
Concat: $$\text{Concat}(\text{head}_1, ..., \text{head}_h)$$
最终 Linear: $$\text{output} = \text{Concat}(\text{head}_1, ..., \text{head}_h)W^O$$
5 代码示例
import torch
import torch.nn as nn
import math
class MultiHeadAttention(nn.Module):
"""多头注意力机制"""
def __init__(self, d_model=512, n_heads=8, dropout=0.1):
super().__init__()
self.d_model = d_model
self.n_heads = n_heads
self.d_k = d_model // n_heads
# 步骤2:三个 Linear 层(投影)
self.W_q = nn.Linear(d_model, d_model) # 给 Q 做投影
self.W_k = nn.Linear(d_model, d_model) # 给 K 做投影
self.W_v = nn.Linear(d_model, d_model) # 给 V 做投影
# 步骤6:最终 Linear
self.W_o = nn.Linear(d_model, d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, Q, K, V, mask=None):
batch_size = Q.size(0)
# 步骤2:投影(换形式,维度不变)
Q = self.W_q(Q) # (batch, seq_len, d_model)
K = self.W_k(K)
V = self.W_v(V)
# 步骤3:分头 (batch, seq_len, d_model) -> (batch, n_heads, seq_len, d_k)
Q = Q.view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
K = K.view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
V = V.view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
# 步骤4:并行 Attention(每个头单独计算)
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
attn = torch.softmax(scores, dim=-1)
attn = self.dropout(attn)
output = torch.matmul(attn, V) # (batch, n_heads, seq_len, d_k)
# 步骤5:Concat(拼接多个头)
output = output.transpose(1, 2).contiguous()
output = output.view(batch_size, -1, self.d_model) # (batch, seq_len, d_model)
# 步骤6:最终 Linear(让不同头的信息更好融合)
return self.W_o(output)
# 测试
if __name__ == "__main__":
mha = MultiHeadAttention(d_model=512, n_heads=8)
Q = K = V = torch.randn(2, 4, 512)
output = mha(Q, K, V)
print(f"输入形状: {Q.shape}")
print(f"输出形状: {output.shape}")
6 重难点与易错提醒
- ❗重点:三个 Linear 层给 QKV 做投影,维度不变。
- ❗重点:每个头单独计算自己的注意力。
- ❗重点:Concat 将多个头的结果拼接,8×64=512。
- ❗重点:最终 Linear 让不同头的信息更好融合,维度不变。
- 💡深入理解:投影类比原食材改刀,还是原食材。
- 💡深入理解:最终 Linear 类比蘸麻酱,充分融合。
7 课堂问答精选
Q1:三个 Linear 层的作用是什么?
A:给 QKV 做投影,换个形式,维度不变,为分头做准备。类比原食材切片、切丝、切段,还是原食材,只是改刀。
Q2:每个头如何计算注意力?
A:每个头单独计算自己的注意力,得到自己视角的注意力结果。类比各组长收作业,每个组知道自己组的情况。
Q3:Concat 的作用是什么?
A:将多个头的结果拼接,8×64=512,维度恢复为 512。类比班长收齐各组的作业,再交给老师。
Q4:最终 Linear 的作用是什么?
A:调整拼接好的结果,让不同头的信息更好融合,维度不变。类比吃火锅时蘸麻酱,充分融合,感受丰富的层次感。
8 本课小结
- 三个 Linear 层:投影,维度不变。
- 分头:512 维分成 8 个头,每个头 64 维。
- 并行 Attention:每个头单独计算。
- Concat:拼接多个头,8×64=512。
- 最终 Linear:让不同头的信息更好融合,维度不变。
9 延伸思考
- 如何用代码实现多头注意力机制?
- 如何克隆多个相同的模块?