多头注意力机制 - 原理图(上)
1 课程概览
本课讲解多头注意力机制原理图上半部分。多头注意力使用多个注意力机制获取多个关注点,将 512 维分成 8 个头,每个头 64 维。类比管理:一个人管 512 人不如分成 8 组每组 64 人。实际开发头数一般为 8、16、32,不会太多。结构:QKV → Linear → 分头 → Attention → Concat → Linear。
2 核心概念与定义
- Multi-Head Attention:多头注意力机制。
- 多头:多个注意力机制,获取多个关注点。
- 分头:将词向量维度切分成多个头。
- 头数(H):注意力机制的数量,如 8。
- 每头维度(d_k):每个头的维度,如 64。
- Linear:线性变换层。
- Concat:拼接多个头的结果。
3 模型与算法详解
多头注意力的作用
同时使用多个注意力机制获取多个不同的关注点。
- 将文本特征切分,分成多个人去观察
- 更有利于提取事物的特征
- 从多个角度分析问题,更专业
类比:管理班级
| 方式 | 说明 |
|---|---|
| 一个人管 512 人 | 不能更好知道每个人的特点 |
| 分成 8 组每组 64 人 | 组内互相了解,更清楚 |
头数的选择
| 头数 | 说明 |
|---|---|
| 8 | 常用 |
| 16 | 常用 |
| 32 | 较多 |
| 512 | 太碎,不推荐 |
实际开发头数一般为 8、16、32,不会太多。
多头注意力结构
Q ─→ Linear ─→ ┐
K ─→ Linear ─→ ┤→ 分头 → Attention → Concat → Linear → 输出
V ─→ Linear ─→ ┘
维度变化
| 步骤 | 形状 | 说明 |
|---|---|---|
| 原始 QKV | (2, 4, 512) | batch=2, seq_len=4, d_model=512 |
| Linear 后 | (2, 4, 512) | 维度不变,换个形式 |
| 分头后 | (2, 4, 8, 64) | 8 个头,每头 64 维 |
| Attention 后 | (2, 4, 8, 64) | 每个头单独计算 |
| Concat 后 | (2, 4, 512) | 合并 8 个头 |
| 最终输出 | (2, 4, 512) | 维度不变 |
512 维分成 8 个头
$$512 = 8 \times 64$$
- 每个头 64 维
- 8 个头并行计算
- 每个头单独计算自己的注意力
4 数学原理与推导
多头注意力公式
$$\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, ..., \text{head}_h)W^O$$
$$\text{head}_i = \text{Attention}(QW_i^Q, KW_i^K, VW_i^V)$$
参数说明
| 参数 | 含义 |
|---|---|
| h | 头数(如 8) |
| d_k | 每头维度(如 64) |
| d_model | 原始维度(如 512) |
| $W_i^Q, W_i^K, W_i^V$ | 线性变换矩阵 |
| $W^O$ | 输出线性变换矩阵 |
维度关系
$$d_{model} = h \times d_k$$ $$512 = 8 \times 64$$
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 # 每头维度
# 线性变换层(给 QKV 做投影)
self.W_q = nn.Linear(d_model, d_model)
self.W_k = nn.Linear(d_model, d_model)
self.W_v = nn.Linear(d_model, d_model)
# 输出线性层
self.W_o = nn.Linear(d_model, d_model)
# dropout
self.dropout = nn.Dropout(dropout)
def forward(self, Q, K, V, mask=None):
batch_size = Q.size(0)
# 1. 线性变换(换个形式,维度不变)
Q = self.W_q(Q) # (batch, seq_len, d_model)
K = self.W_k(K)
V = self.W_v(V)
# 2. 分头:(batch, seq_len, d_model) -> (batch, seq_len, n_heads, 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)
# 3. 计算注意力(每个头单独计算)
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)
# 4. Concat:合并多个头
output = output.transpose(1, 2).contiguous()
output = output.view(batch_size, -1, self.d_model)
# 5. 最终线性变换
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 重难点与易错提醒
- ❗重点:多头注意力使用多个注意力机制获取多个关注点。
- ❗重点:512 维分成 8 个头,每个头 64 维。
- ⚠️易错:头数不能太多,实际开发一般 8、16、32。
- 💡深入理解:类比管理班级,分组管理更高效。
- 💡深入理解:每个头单独计算自己的注意力。
7 课堂问答精选
Q1:多头注意力机制的作用是什么?
A:同时使用多个注意力机制获取多个不同的关注点。将文本特征切分,分成多个人去观察,更有利于提取事物的特征,从多个角度分析问题更专业。
Q2:512 维分成 8 个头,每个头多少维?
A:512 / 8 = 64,每个头 64 维。
Q3:头数一般选择多少?
A:实际开发头数一般为 8、16、32,不会太多。如果分成 512 个头,一个头管一个,太碎了。
Q4:多头注意力的结构是什么?
A:QKV → Linear(线性变换)→ 分头 → Attention(每个头单独计算)→ Concat(合并)→ Linear(输出线性变换)→ 输出。
8 本课小结
- 多头注意力:多个注意力机制获取多个关注点。
- 512 维分成 8 个头,每个头 64 维。
- 结构:QKV → Linear → 分头 → Attention → Concat → Linear。
- 头数一般 8、16、32。
- 类比管理班级,分组管理更高效。
9 延伸思考
- 多头注意力如何用代码实现?
- 分头和合并的具体操作是什么?