共计 2478 个字符,预计需要花费 7 分钟才能阅读完成。
1. 核心概念:Q/K/ V 矩阵与缩放点积注意力
自注意力机制的核心是通过三个矩阵(Query、Key、Value)建立序列元素间的关联。给定输入序列 $X \in \mathbb{R}^{n \times d}$(n 为序列长度,d 为特征维度),计算过程如下:

-
线性投影:
$$Q = XW_Q, \quad K = XW_K, \quad V = XW_V$$
其中 $W_Q, W_K, W_V \in \mathbb{R}^{d \times d_k}$ 为可训练参数矩阵 -
注意力分数:
$$\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V$$
缩放因子 $\sqrt{d_k}$ 用于防止点积结果过大导致 Softmax 梯度消失
2. PyTorch 完整实现
import torch
import torch.nn as nn
import math
class SelfAttention(nn.Module):
def __init__(self, embed_size, heads):
super().__init__()
self.embed_size = embed_size
self.heads = heads
self.head_dim = embed_size // heads
assert self.head_dim * heads == embed_size, "Embed size needs division by heads"
# 线性投影层
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, x, mask=None):
N = x.shape[0]
seq_len = x.shape[1]
# 分割多头 (batch, seq_len, heads, head_dim)
x = x.view(N, seq_len, self.heads, self.head_dim)
queries = self.queries(x)
keys = self.keys(x)
values = self.values(x)
# 计算缩放点积注意力
energy = torch.einsum("nqhd,nkhd->nhqk", [queries, keys])
if mask is not None:
energy = energy.masked_fill(mask == 0, float("-1e20"))
attention = torch.softmax(energy / (self.embed_size ** (1/2)), dim=3)
out = torch.einsum("nhql,nlhd->nqhd", [attention, values])
# 合并多头输出
out = out.reshape(N, seq_len, -1)
return self.fc_out(out)
# 位置编码示例
class PositionalEncoding(nn.Module):
def __init__(self, d_model, max_len=100):
super().__init__()
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
self.register_buffer('pe', pe.unsqueeze(0))
def forward(self, x):
return x + self.pe[:, :x.size(1)]
3. 性能优化策略
- 计算复杂度分析:
- QK^T 乘法:$O(n^2d)$
- Softmax 计算:$O(n^2)$
-
与 V 相乘:$O(n^2d)$
-
优化技巧:
- 使用
torch.einsum替代逐元素计算 - 混合精度训练(AMP)减少显存占用
- 当序列长度 >512 时考虑使用 Flash Attention
4. 常见错误与解决方案
- 错误 1:未缩放 Attention 分数
- 现象:训练初期出现 NaN 损失
-
修复:确保除以 $\sqrt{d_k}$
-
错误 2:忽略 padding 影响
- 现象:模型对填充位置过度关注
-
修复:添加 mask 矩阵
masked_fill(-1e20) -
错误 3:多头维度分配不当
- 现象:head_dim 非整数导致维度错误
- 修复:添加
assert embed_size % heads == 0
5. 扩展思考与实践
-
Attention 可视化:
# 获取 attention 权重 attention = model.get_attention(inputs) plt.imshow(attention[0].detach().numpy()) # 可视化第一个头 -
长序列优化方案:
- 使用稀疏 Attention(如 Longformer)
- 采用分块计算(Reformer 的 LSH Attention)
- 尝试线性 Attention 变体(Performer)
实践建议
从简单的序列分类任务开始(如 IMDB 影评分类),逐步尝试以下改进:
1. 对比单头与多头注意力的效果差异
2. 添加 / 移除位置编码观察性能变化
3. 在自定义数据集上可视化 Attention 热力图
代码仓库推荐参考 HuggingFace 的 transformers 库实现,其中包含了工业级的 Attention 优化方案。
正文完
