多头注意力机制 - 代码实现
1 课程概览
本课实现多头注意力机制代码。包括两个关键函数:①clones(module, n) 克隆 N 个相同模块(深拷贝);②MultiHeadAttention 类实现多头注意力。使用 assert 确保维度能整除。步骤:初始化 → 投影 → 分头 → Attention → Concat → 输出 Linear。
2 核心概念与定义
- clones():克隆 N 个相同模块,返回列表,深拷贝。
- copy.deepcopy:深拷贝,每个模块拥有独立参数。
- assert:断言,条件成立才继续执行,否则报错。
- nn.ModuleList:模块列表。
- MultiHeadAttention:多头注意力机制类。
3 模型与算法详解
clones 函数
创建 N 个相同模块,返回列表,深拷贝。
def clones(module, n):
return nn.ModuleList([copy.deepcopy(module) for _ in range(n)])
作用:
- 写一次模块,复制 N 次
- 深拷贝,每个模块独立参数
- 用于编码器(6 个编码器层)等
assert 断言
条件成立才继续执行,否则报错。
assert d_model % n_heads == 0 # 确保能整除
示例:
10 % 2 == 0:成立,继续执行10 % 3 == 0:不成立,报错,不继续10 % 0:报错(除以 0)
MultiHeadAttention 类
MultiHeadAttention(nn.Module)
├── __init__(d_model, n_heads, dropout)
│ ├── assert d_model % n_heads == 0
│ ├── d_k = d_model // n_heads
│ ├── W_q, W_k, W_v (Linear)
│ ├── W_o (Linear)
│ └── dropout
└── forward(Q, K, V, mask)
├── 投影:Q' = W_q(Q), K' = W_k(K), V' = W_v(V)
├── 分头:view + transpose
├── Attention:scores = Q·K^T / √d_k → softmax → ·V
├── Concat:transpose + view
└── 输出:W_o(output)
维度变化
| 步骤 | 形状 | 说明 |
|---|---|---|
| 输入 QKV | (2, 4, 512) | batch=2, seq_len=4, d_model=512 |
| 投影后 | (2, 4, 512) | 维度不变 |
| 分头后 | (2, 8, 4, 64) | 8 个头,每头 64 维 |
| Attention | (2, 8, 4, 64) | 每个头单独计算 |
| Concat | (2, 4, 512) | 合并 8 个头 |
| 输出 | (2, 4, 512) | 维度不变 |
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)$$
维度关系
$$d_{model} = h \times d_k$$ $$512 = 8 \times 64$$
assert 条件
$$d_{model} \mod h = 0$$
5 代码示例
import torch
import torch.nn as nn
import math
import copy
def clones(module, n):
"""
克隆 N 个相同模块
:param module: 要被克隆的模块
:param n: 克隆次数
:return: 包含 N 个相同模块的列表(深拷贝)
"""
return nn.ModuleList([copy.deepcopy(module) for _ in range(n)])
def attention(query, key, value, mask=None, dropout=None):
"""注意力计算"""
d_k = query.size(-1)
scores = torch.matmul(query, key.transpose(-2, -1)) / math.sqrt(d_k)
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
p_attn = torch.softmax(scores, dim=-1)
if dropout is not None:
p_attn = dropout(p_attn)
return torch.matmul(p_attn, value), p_attn
class MultiHeadAttention(nn.Module):
"""多头注意力机制"""
def __init__(self, d_model=512, n_heads=8, dropout=0.1):
super().__init__()
# 确保维度能整除
assert d_model % n_heads == 0, f"d_model ({d_model}) 必须能被 n_heads ({n_heads}) 整除"
self.d_model = d_model
self.n_heads = n_heads
self.d_k = d_model // n_heads # 每头维度
# 4 个线性变换层
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)
self.attn = None
self.dropout = nn.Dropout(dropout)
def forward(self, Q, K, V, mask=None):
"""
前向传播
:param Q: 查询张量 (batch, seq_len, d_model)
:param K: 键张量 (batch, seq_len, d_model)
:param V: 值张量 (batch, seq_len, d_model)
:param mask: 掩码张量
:return: 多头注意力输出
"""
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, 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)
# 3. 计算注意力(每个头并行计算)
output, self.attn = attention(Q, K, V, mask=mask, dropout=self.dropout)
# 4. Concat:(batch, n_heads, seq_len, d_k) -> (batch, seq_len, d_model)
output = output.transpose(1, 2).contiguous()
output = output.view(batch_size, -1, self.d_model)
# 5. 最终线性变换
return self.W_o(output)
# 测试代码
if __name__ == "__main__":
# 测试 clones 函数
linear = nn.Linear(512, 512)
cloned = clones(linear, 6)
print(f"克隆数量: {len(cloned)}")
# 测试 MultiHeadAttention
mha = MultiHeadAttention(d_model=512, n_heads=8, dropout=0.1)
Q = K = V = torch.randn(2, 4, 512)
output = mha(Q, K, V)
print(f"输入形状: {Q.shape}")
print(f"输出形状: {output.shape}")
print(f"d_model: {mha.d_model}")
print(f"n_heads: {mha.n_heads}")
print(f"d_k: {mha.d_k}")
代码说明
| 代码 | 说明 |
|---|---|
copy.deepcopy(module) | 深拷贝模块 |
nn.ModuleList([...]) | 模块列表 |
assert d_model % n_heads == 0 | 确保能整除 |
d_model // n_heads | 每头维度 |
Q.view(batch, -1, n_heads, d_k) | 分头 |
.transpose(1, 2) | 交换维度 |
.contiguous() | 内存连续 |
6 重难点与易错提醒
- ❗重点:
clones()用深拷贝,每个模块独立参数。 - ❗重点:
assert确保维度能整除。 - ❗重点:分头用
view+transpose。 - ⚠️易错:
view前要确保内存连续(用contiguous())。 - 💡深入理解:深拷贝保证每个模块参数独立。
7 课堂问答精选
Q1:clones 函数的作用是什么?
A:克隆 N 个相同模块,返回列表。用深拷贝,每个模块拥有独立参数。用于编码器(6 个编码器层)等。
Q2:assert 的作用是什么?
A:断言,条件成立才继续执行,否则报错。例如 assert d_model % n_heads == 0 确保维度能整除。
Q3:为什么要用深拷贝?
A:深拷贝保证每个模块拥有独立参数。如果不用深拷贝,所有模块共享同一参数,修改一个会影响其他。
Q4:分头如何实现?
A:用 view 改变形状,再用 transpose 交换维度。例如 (2, 4, 512) → view(2, 4, 8, 64) → transpose(1, 2) → (2, 8, 4, 64)。
8 本课小结
clones(module, n):克隆 N 个模块,深拷贝。assert:条件成立才继续。- MultiHeadAttention:投影 → 分头 → Attention → Concat → 输出。
- 分头:
view+transpose。 - 深拷贝保证参数独立。
9 延伸思考
- 如何测试多头注意力机制?
- 多头注意力机制的效果如何验证?