共计 2813 个字符,预计需要花费 8 分钟才能阅读完成。
背景介绍
自注意力机制(Self-Attention)是 Transformer 架构的核心组件,也是 BERT 等预训练模型成功的关键。它通过计算输入序列中每个元素与其他元素的相关性,动态地学习上下文表示。相比传统的 RNN 和 CNN,自注意力机制能够更好地捕捉长距离依赖关系,并且具有天然的并行计算优势。

核心概念
Query、Key 和 Value 矩阵
- Query(Q):表示当前需要关注的元素,用于与其他元素的 Key 进行匹配。
- Key(K):表示其他元素的标识,用于与 Query 计算相似度。
- Value(V):包含实际的信息内容,根据相似度权重进行加权求和。
多头注意力
多头注意力(Multi-Head Attention)是将输入线性投影到多个子空间,每个子空间独立计算注意力,最后将结果拼接。这种方式可以让模型同时关注不同位置的不同特征。
代码实现
以下是一个完整的 BERT 多头自注意力机制的 PyTorch 实现,包含详细注释:
import torch
import torch.nn as nn
import math
class MultiHeadAttention(nn.Module):
def __init__(self, embed_dim, num_heads, dropout=0.1):
super(MultiHeadAttention, self).__init__()
assert embed_dim % num_heads == 0, "embed_dim must be divisible by num_heads"
self.embed_dim = embed_dim
self.num_heads = num_heads
self.head_dim = embed_dim // num_heads
# 线性变换层,用于生成 Q、K、V
self.q_linear = nn.Linear(embed_dim, embed_dim)
self.k_linear = nn.Linear(embed_dim, embed_dim)
self.v_linear = nn.Linear(embed_dim, embed_dim)
# 输出层和 dropout
self.out_linear = nn.Linear(embed_dim, embed_dim)
self.dropout = nn.Dropout(dropout)
def forward(self, x, mask=None):
batch_size = x.size(0)
# 生成 Q、K、V
q = self.q_linear(x)
k = self.k_linear(x)
v = self.v_linear(x)
# 分割多头
q = q.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
k = k.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
v = v.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
# 计算注意力分数
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.head_dim)
# 应用 mask(如有)if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
# 计算注意力权重
attention = torch.softmax(scores, dim=-1)
attention = self.dropout(attention)
# 加权求和
context = torch.matmul(attention, v)
# 合并多头
context = context.transpose(1, 2).contiguous().view(batch_size, -1, self.embed_dim)
# 输出
output = self.out_linear(context)
return output, attention
# 示例用法
if __name__ == "__main__":
embed_dim = 512
num_heads = 8
batch_size = 2
seq_len = 10
# 随机生成输入
x = torch.randn(batch_size, seq_len, embed_dim)
# 创建多头注意力层
mha = MultiHeadAttention(embed_dim, num_heads)
# 前向传播
output, attention_weights = mha(x)
print(f"输入形状: {x.shape}")
print(f"输出形状: {output.shape}")
print(f"注意力权重形状: {attention_weights.shape}")
常见问题
- 维度不匹配:
- 确保
embed_dim能被num_heads整除,否则会报错。 -
多头分割后,注意调整张量的形状和维度顺序。
-
权重初始化不当:
- 线性变换层的权重需要合理初始化(如 Xavier 初始化)。
-
偏差项(bias)通常初始化为零。
-
注意力分数溢出:
-
计算分数时,记得除以
sqrt(head_dim),避免 softmax 后梯度消失。 -
mask 应用错误:
- 确保 mask 的形状与注意力分数匹配。
- 将 mask 中需要忽略的位置设为 0,并替换为极小的负值(如
-1e9)。
优化建议
- 性能调优:
- 使用
torch.einsum替代matmul和transpose,减少显存占用。 -
开启 PyTorch 的自动混合精度(AMP),加速计算。
-
内存优化:
- 对于长序列,考虑使用稀疏注意力或分块计算。
-
梯度检查点(Gradient Checkpointing)可以减少训练时的显存占用。
-
数值稳定性:
- 在 softmax 前对分数做
masked_fill,避免无效位置影响计算结果。 - 使用
log_softmax替代softmax,提高数值稳定性。
实践建议
- 应用场景:
- 文本分类、命名实体识别等任务中,可以直接使用 BERT 的多头注意力。
-
对于生成任务(如机器翻译),需要结合编码器 - 解码器注意力。
-
注意事项:
- 训练时注意学习率设置,多头注意力对学习率比较敏感。
- 推理时可以通过缓存 Key 和 Value,减少重复计算。
延伸阅读
- Attention Is All You Need – Transformer 原论文
- BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding – BERT 原论文
- The Illustrated Transformer – 图解 Transformer
思考题
- 如何修改代码实现相对位置编码(Relative Positional Encoding)?
- 多头注意力中,为什么需要将
embed_dim分成num_heads份? - 如何实现跨语言的注意力机制(如翻译任务中的源语言和目标语言)?
正文完
