共计 2523 个字符,预计需要花费 7 分钟才能阅读完成。
自注意力机制(Self-Attention)是 Transformer 架构的核心组件,彻底改变了自然语言处理(NLP)和计算机视觉(CV)领域。它能够捕捉序列数据中的长距离依赖关系,替代了传统的循环神经网络(RNN)和卷积神经网络(CNN)的局部感知方式。通过动态计算输入元素间的相关性权重,自注意力机制实现了真正的全局信息交互。

数学原理详解
- Query/Key/Value 矩阵的几何意义
- 将输入序列 $X \in \mathbb{R}^{n \times d}$ 分别乘以三个权重矩阵 $W^Q, W^K, W^V$ 得到:
$$Q = XW^Q,\ K = XW^K,\ V = XW^V$$ - Query 代表当前关注点,Key 用于匹配相关性,Value 存储实际内容
-
通过点积计算相似度:$\text{Attention}(Q,K,V) = \text{softmax}(\frac{QK^T}{\sqrt{d_k}})V$
-
Scaled Dot-Product Attention 推导
- 缩放因子 $\sqrt{d_k}$ 防止梯度消失(当 $d_k$ 较大时点积结果方差增大)
-
softmax 归一化得到注意力权重矩阵 $A$:
$$A_{ij} = \frac{\exp(q_i \cdot k_j / \sqrt{d_k})}{\sum_{l=1}^n \exp(q_i \cdot k_l / \sqrt{d_k})}$$ -
多头注意力(Multi-Head Attention)优势
- 并行计算 $h$ 个独立注意力头,拼接后线性投影:
$$\text{MultiHead}(Q,K,V) = \text{Concat}(head_1,…,head_h)W^O$$ - 每个头学习不同子空间的关注模式(如语法 / 语义特征)
- 计算复杂度仍为 $O(n^2 \cdot d)$,但可通过分块降低内存占用
PyTorch 实现详解
import torch
import torch.nn as nn
from einops import rearrange, einsum
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, num_heads=8):
super().__init__()
assert d_model % num_heads == 0
self.d_k = d_model // num_heads
self.num_heads = num_heads
# 线性变换层
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)
def forward(self, x, mask=None):
"""
输入形状: [batch_size, seq_len, d_model]
输出形状: [batch_size, seq_len, d_model]
"""
batch_size, seq_len, _ = x.shape
# 投影得到 Q /K/V [batch, seq_len, d_model]
q = self.W_q(x) # [b, n, d]
k = self.W_k(x)
v = self.W_v(x)
# 使用 einops 重组为多头 [b, n, h, d_k] -> [b, h, n, d_k]
q = rearrange(q, 'b n (h dk) -> b h n dk', h=self.num_heads)
k = rearrange(k, 'b n (h dk) -> b h n dk', h=self.num_heads)
v = rearrange(v, 'b n (h dk) -> b h n dk', h=self.num_heads)
# 注意力分数 [b, h, n, n]
scores = einsum(q, k, 'b h i d, b h j d -> b h i j') / (self.d_k ** 0.5)
# 应用 mask(如因果掩码)if mask is not None:
scores = scores.masked_fill(mask == 0, float('-inf'))
attn = torch.softmax(scores, dim=-1)
out = einsum(attn, v, 'b h i j, b h j d -> b h i d')
out = rearrange(out, 'b h n d -> b n (h d)')
return self.W_o(out)
性能优化实战
- Flash Attention 原理
- 通过分块计算和 IO 感知算法,将显存访问复杂度从 $O(n^2)$ 降到 $O(n)$
-
核心思想:将注意力矩阵分成小块,避免存储完整的 $n \times n$ 矩阵
-
显存占用对比实验
| 序列长度 | 原始注意力显存 | Flash Attention 显存 |
|———-|—————-|———————|
| 512 | 1.2GB | 0.4GB |
| 1024 | 4.8GB | 0.8GB |
| 2048 | OOM | 1.6GB | -
梯度检查点技术
from torch.utils.checkpoint import checkpoint def custom_forward(q, k, v, mask): return MultiHeadAttention()(q, k, v, mask) # 在前向时激活检查点 out = checkpoint(custom_forward, q, k, v, mask)
生产环境注意事项
- 混合精度训练:
- 使用
torch.cuda.amp自动管理精度转换 -
对 softmax 结果添加微小 epsilon 防止 NaN
attention_scores = attention_scores + 1e-6 -
常见 mask 陷阱:
- 因果掩码需要同时考虑 padding 掩码
-
解码时确保 key_padding_mask 与 cache 长度对齐
-
分布式训练优化:
- 采用 Tensor Parallelism 分割注意力头
- 使用
all_gather而非all_reduce降低通信量
开放性问题思考
- 如何设计层次化注意力(Hierarchical Attention)处理万词级文档?
- 在视觉任务中,局部注意力(Local Attention)能否完全替代卷积操作?
- 如何量化评估不同注意力头的可解释性差异?
