🎯 课程主题
多头注意力机制的代码实现——支持三种注意力类型的通用模块。
📝 核心知识点
1. 通用设计
- 概念说明:一套代码实现三种注意力机制。
- 关键细节:
- 自注意力:Q=K=V 来自同一输入
- 因果掩码注意力:Q=K=V + 因果 mask
- 交叉注意力:Q 来自解码器,K/V 来自编码器
- 通过输入参数 X1, X2, X3 区分
2. 参数共享优化
- 概念说明:用一个大全连接层代替多个小全连接层。
- 关键细节:
- 原理:每个头有独立的 $W_Q, W_K, W_V$($d \times d_k$)
- 代码:用一个 $d \times (h \cdot d_k)$ 的全连接层代替
- $h$ 个头的参数合并,效率更高
- 共定义 4 个全连接层:$W_Q, W_K, W_V, W_O$
3. 前向传播流程
- 概念说明:多头注意力的完整计算流程。
- 关键细节:
- 输入 X1, X2, X3(可能来自不同层)
- 处理 mask(增加 batch 维度)
- 获取 batch_size 用于维度转换
- X1, X2, X3 分别通过 $W_Q, W_K, W_V$ 映射
- 维度转换:将结果切分为 $h$ 份(列表)
- 每份送入 attention 函数
- 得到 $h$ 个输出 $Z_0, Z_1, ..., Z_{h-1}$
- concat 拼接所有输出
- 通过 $W_O$ 映射得到最终输出
4. 维度变化
- 概念说明:多头切分和拼接的维度操作。
- 关键细节:
- 输入:$[batch, L, d_{model}]$
- 映射后:$[batch, L, h \cdot d_k]$
- 切分为 $h$ 份:$h \times [batch, L, d_k]$
- 注意力计算后:$h \times [batch, L, d_k]$
- concat 拼接:$[batch, L, h \cdot d_k]$
- $W_O$ 映射:$[batch, L, d_{model}]$
🧮 核心公式与推导
多头注意力公式:
$$\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, ..., \text{head}_h) W^O$$
$$\text{head}_i = \text{Attention}(Q W_i^Q, K W_i^K, V W_i^V)$$
参数维度:
- $W_i^Q \in \mathbb{R}^{d_{model} \times d_k}$
- $W_i^K \in \mathbb{R}^{d_{model} \times d_k}$
- $W_i^V \in \mathbb{R}^{d_{model} \times d_v}$
- $W^O \in \mathbb{R}^{hd_v \times d_{model}}$
代码优化:合并 $h$ 个头的参数
- $W^Q \in \mathbb{R}^{d_{model} \times hd_k}$(代替 $h$ 个 $W_i^Q$)
🏗️ 模型架构与数据流向
💻 代码实战
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
import copy
def attention(Q, K, V, mask=None, dropout=None):
"""缩放点积注意力函数"""
d_k = Q.size(-1)
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k)
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
attn_weights = F.softmax(scores, dim=-1)
if dropout is not None:
attn_weights = dropout(attn_weights)
output = torch.matmul(attn_weights, V)
return output, attn_weights
def clones(module, N):
"""克隆模块N次"""
return nn.ModuleList([copy.deepcopy(module) for _ in range(N)])
class MultiHeadedAttention(nn.Module):
"""多头注意力机制 - 支持自注意力/因果注意力/交叉注意力"""
def __init__(self, h, d_model, dropout=0.1):
"""
h: 注意力头数 (如8)
d_model: 模型维度 (如512)
"""
super().__init__()
assert d_model % h == 0, "d_model必须能被头数h整除"
self.d_k = d_model // h # 每个头的维度
self.h = h # 头数
self.d_model = d_model
# 4个全连接层: W_Q, W_K, W_V, W_O
# 用一个大的全连接层代替h个小全连接层
self.linears = clones(nn.Linear(d_model, d_model), 4)
self.attn = None # 保存注意力得分, 用于可视化
self.dropout = nn.Dropout(p=dropout)
def forward(self, query, key, value, mask=None):
"""
query, key, value: [batch, seq_len, d_model]
mask: [batch, 1, seq_len] 或 None
支持三种模式:
- 自注意力: query=key=value
- 因果注意力: query=key=value + causal mask
- 交叉注意力: query来自解码器, key=value来自编码器
"""
if mask is not None:
# mask增加头维度: [batch, 1, 1, seq_len]
mask = mask.unsqueeze(1)
batch_size = query.size(0)
# 1. 线性映射并切分多头
# [batch, seq_len, d_model] → [batch, seq_len, h, d_k] → [batch, h, seq_len, d_k]
query, key, value = [
l(x).view(batch_size, -1, self.h, self.d_k).transpose(1, 2)
for l, x in zip(self.linears, (query, key, value))
]
# 2. 注意力计算
x, self.attn = attention(
query, key, value, mask=mask, dropout=self.dropout
)
# 3. 拼接多头
# [batch, h, seq_len, d_k] → [batch, seq_len, h, d_k] → [batch, seq_len, d_model]
x = x.transpose(1, 2).contiguous().view(
batch_size, -1, self.h * self.d_k
)
# 4. 最终线性映射
return self.linears[-1](x)
# 使用示例
if __name__ == "__main__":
d_model = 512
h = 8
batch_size = 2
seq_len = 10
mha = MultiHeadedAttention(h, d_model)
# 1. 自注意力 (编码器中)
x = torch.randn(batch_size, seq_len, d_model)
output = mha(x, x, x, mask=None)
print(f"自注意力输出: {output.shape}") # [2, 10, 512]
# 2. 交叉注意力 (解码器中)
dec_input = torch.randn(batch_size, 5, d_model) # 解码器输入
enc_output = torch.randn(batch_size, 10, d_model) # 编码器输出
output = mha(dec_input, enc_output, enc_output)
print(f"交叉注意力输出: {output.shape}") # [2, 5, 512]
⚠️ 常见问题与避坑指南
d_model必须能被头数 $h$ 整除(如 512/8=64)- 用一个大全连接层代替 $h$ 个小全连接层,参数量相同但效率更高
- 维度转换时注意
transpose(1, 2)将头维度提前,便于并行计算 - concat 后需
contiguous()保证内存连续
💡 个人总结与延伸
多头注意力机制通过并行计算多个注意力头,让模型从不同子空间关注不同信息。代码实现巧妙地用一个大全连接层代替多个小全连接层,简化了实现。现代大模型中,MQA(Multi-Query Attention)和 GQA(Grouped-Query Attention)进一步优化了注意力计算,减少 KV 缓存内存占用。