共计 3662 个字符,预计需要花费 10 分钟才能阅读完成。
传统序列模型的局限性
在自然语言处理(NLP)任务中,传统的循环神经网络(RNN)和卷积神经网络(CNN)存在一些明显的局限性。让我们通过两个例子来说明:

-
情感分析:当处理句子 ” 这部电影虽然特效很棒,但剧情太拖沓 ” 时,RNN 需要逐步处理每个词,且后面的词对前面词的影响有限。这导致模型难以捕捉 ” 虽然 … 但 …” 这种转折关系。
-
命名实体识别:在句子 ” 苹果公司宣布新款 iPhone” 中,要识别 ” 苹果 ” 是公司名而非水果,需要同时考虑前后文信息。传统单向模型很难做到这点。
这些局限性促使了自注意力机制(Self-Attention)的发展,特别是 BERT 中采用的双向自注意力机制。
自注意力机制基础
自注意力机制的核心思想是:每个词都可以直接关注输入序列中的所有其他词,计算它们之间的相关性。这与传统 RNN 的顺序处理形成鲜明对比。
单向 vs 双向注意力
-
单向注意力:只能关注当前位置及之前的信息(类似 RNN),公式表示为:
$$Attention(Q,K,V) = softmax(\frac{QK^T}{\sqrt{d_k}})V$$
其中 mask 会屏蔽未来信息。 -
双向注意力(BERT 采用):可以同时关注前后所有位置的信息,计算时不需要 mask 未来位置。
QKV 矩阵运算
自注意力机制涉及三个关键矩阵:
1. Query(查询)矩阵 Q
2. Key(键)矩阵 K
3. Value(值)矩阵 V
计算过程如下:
-
首先计算注意力分数:
$$Attention_Scores = \frac{QK^T}{\sqrt{d_k}}$$ -
应用 softmax 归一化:
$$Attention_Weights = softmax(Attention_Scores)$$ -
最后加权求和:
$$Output = Attention_Weights \times V$$
其中 $d_k$ 是 Key 向量的维度,用于缩放点积结果,防止 softmax 的梯度太小。
PyTorch 实现简化版 BERT 自注意力层
import torch
import torch.nn as nn
import torch.nn.functional as F
import matplotlib.pyplot as plt
class SelfAttention(nn.Module):
def __init__(self, embed_size, heads):
super(SelfAttention, self).__init__()
self.embed_size = embed_size
self.heads = heads
self.head_dim = embed_size // heads
assert (self.head_dim * heads == embed_size), "Embedding size needs to be divisible by heads"
# 线性变换得到 Q,K,V
self.values = nn.Linear(self.head_dim, self.head_dim, bias=False)
self.keys = nn.Linear(self.head_dim, self.head_dim, bias=False)
self.queries = nn.Linear(self.head_dim, self.head_dim, bias=False)
# 最终输出线性变换
self.fc_out = nn.Linear(heads * self.head_dim, embed_size)
def forward(self, values, keys, query, mask):
# 获取 batch size
N = query.shape[0]
# 获取序列长度
value_len, key_len, query_len = values.shape[1], keys.shape[1], query.shape[1]
# 分割 embedding 到多个头
values = values.reshape(N, value_len, self.heads, self.head_dim)
keys = keys.reshape(N, key_len, self.heads, self.head_dim)
queries = query.reshape(N, query_len, self.heads, self.head_dim)
# 计算注意力分数
energy = torch.einsum("nqhd,nkhd->nhqk", [queries, keys])
# 应用 mask(如果有)if mask is not None:
energy = energy.masked_fill(mask == 0, float("-1e20"))
# 计算注意力权重
attention = torch.softmax(energy / (self.embed_size ** (1/2)), dim=3)
# 可视化注意力权重(取第一个样本的第一个头)if query_len == 1: # 避免可视化太长的序列
self.visualize_attention(attention[0,0].detach().cpu().numpy())
# 计算输出
out = torch.einsum("nhql,nlhd->nqhd", [attention, values])
out = out.reshape(N, query_len, self.heads * self.head_dim)
# 通过最后的线性层
out = self.fc_out(out)
return out
def visualize_attention(self, attention_weights):
plt.matshow(attention_weights)
plt.title("Attention Weights")
plt.colorbar()
plt.show()
# 示例用法
embed_size = 256
heads = 8
seq_len = 32
batch_size = 4
# 创建输入(模拟 token 嵌入)inputs = torch.rand((batch_size, seq_len, embed_size))
# 创建 mask(可选)mask = torch.ones((batch_size, heads, seq_len, seq_len))
# 初始化自注意力层
attention = SelfAttention(embed_size, heads)
# 前向传播
output = attention(inputs, inputs, inputs, mask)
print("Output shape:", output.shape) # 应该是 [4, 32, 256]
性能优化策略
多头注意力并行计算
BERT 使用多头注意力(Multi-Head Attention)来并行计算多个注意力头,每个头学习不同的注意力模式。在实现上,可以通过矩阵运算的并行性来加速:
- 将 Q、K、V 矩阵分割为多个头
- 每个头独立计算注意力
- 拼接所有头的输出
显存占用估算
自注意力层的显存占用主要来自注意力矩阵,其大小为 $[batch_size, heads, seq_len, seq_len]$。可以通过以下公式估算:
$$Memory\ (bytes) = batch_size \times heads \times seq_len^2 \times 4$$
(假设使用 float32,每个元素占 4 字节)
可以使用 torch.cuda.memory_allocated() 验证实际显存使用。
长序列处理策略
当序列长度超过 512 时:
- 分块计算:将长序列分割为多个 512 的块
- 稀疏注意力:只计算局部注意力或关键位置间的注意力
- 内存高效注意力:使用重组计算技术减少中间存储
避坑指南
初始化技巧
自注意力层容易出现梯度消失问题,推荐使用:
- Xavier/Glorot 初始化:
nn.init.xavier_uniform_(self.queries.weight) nn.init.xavier_uniform_(self.keys.weight) nn.init.xavier_uniform_(self.values.weight)
混合精度训练
使用混合精度训练时:
- 设置适当的 scale 因子防止梯度下溢
- 使用
torch.cuda.amp.GradScaler()自动管理
数值稳定性
处理注意力矩阵时:
- 添加小的 epsilon 防止除零错误
- 对注意力分数进行缩放(除以 $\sqrt{d_k}$)
- 对 softmax 输入进行截断
思考题
- 如何设计实验验证双向注意力相比单向注意力的优势?
-
可以比较相同模型在掩码语言模型任务中单向和双向版本的性能差异
-
自注意力机制与人类阅读习惯有哪些认知差异?
-
人类阅读通常是顺序的、局部的,而自注意力是全局的、并行的
-
在移动端部署 BERT 时,有哪些量化方案可以选择?
- 动态量化、静态量化、量化感知训练等不同方案各有优劣
总结
双向自注意力机制是 BERT 等现代 Transformer 模型的核心组件。通过本文的讲解和代码实现,希望读者能够理解其工作原理,并在实际项目中正确应用。虽然自注意力机制计算复杂度较高,但通过多头并行、优化实现等技术,可以在合理资源消耗下获得强大的上下文建模能力。
