共计 1695 个字符,预计需要花费 5 分钟才能阅读完成。
Transformer 模型彻底改变了自然语言处理领域,其自注意力机制能够动态捕捉长距离依赖关系。相比 RNN 的序列计算,Transformer 的并行化架构大幅提升了训练效率。如今从 BERT 到 GPT 系列,Transformer 已成为现代 NLP 任务的基石架构。

自注意力机制的数学本质
自注意力机制的核心是三个权重矩阵:$W_Q\in\mathbb{R}^{d\times d_k}$(Query)、$W_K\in\mathbb{R}^{d\times d_k}$(Key)和 $W_V\in\mathbb{R}^{d\times d_v}$(Value)。在向量空间中:
- Query 向量代表当前需要关注的内容
- Key 向量相当于所有位置的索引标签
- Value 向量存储实际要提取的信息
计算过程可表示为:
$$
\text{Attention}(Q,K,V)=\text{softmax}(\frac{QK^T}{\sqrt{d_k}})V
$$
几何意义是:通过 QK^T 点积计算向量相似度,softmax 将其转化为概率分布,最后对 V 加权求和。
多头注意力的并行化实现
多头注意力(Multi-Head Attention)将输入拆分为 $h$ 个子空间:
- 输入矩阵 $X$ 分别乘以 $h$ 组不同的 $W_Q^i,W_K^i,W_V^i$
- 每个头计算独立的注意力结果 $\text{head}_i=\text{Attention}(XW_Q^i,XW_K^i,XW_V^i)$
- 所有头的结果拼接后通过 $W_O\in\mathbb{R}^{hd_v\times d}$ 输出
矩阵分块运算示意图:
# PyTorch 实现示例
assert x.dim() == 3 # [batch, seq_len, dim]
q = torch.chunk(x @ W_Q, num_heads, dim=-1) # 分块
k = torch.chunk(x @ W_K, num_heads, dim=-1)
v = torch.chunk(x @ W_V, num_heads, dim=-1)
# 各头独立计算
outputs = [attention(q[i], k[i], v[i]) for i in range(num_heads)]
位置编码的傅里叶分析
位置编码(Positional Encoding)采用正弦函数组合:
$$
PE(pos,2i)=\sin(pos/10000^{2i/d})
$$
$$
PE(pos,2i+1)=\cos(pos/10000^{2i/d})
$$
这实质上是傅里叶级数的离散采样:
- 不同 $i$ 对应不同频率分量
- $2i$ 和 $2i+1$ 构成正交基
- 位置变化体现为相位移动
Transformer 的电路建模
信息流动有向图
将 Transformer 视为计算电路时:
- 节点代表张量运算(矩阵乘 /softmax 等)
- 边表示数据流动方向
- 残差连接形成环路结构
残差连接的梯度分析
与传统 DNN 相比:
$$
\frac{\partial L}{\partial x}=\frac{\partial L}{\partial F(x)}+\frac{\partial L}{\partial (F(x)+x)}
$$
残差结构确保梯度至少有一条恒定通路,缓解梯度消失。
性能优化实践
稀疏注意力存储
对于稀疏注意力矩阵 $A$,采用 CSR 格式存储:
values = A[A != 0]
indices = (A != 0).nonzero()
FlashAttention 对比
| 算法 | 时间复杂度 | 空间复杂度 |
|---|---|---|
| 原始注意力 | O(N^2) | O(N^2) |
| FlashAttention | O(N^2/d) | O(N) |
数学验证挑战
问题 :推导 LayerNorm 如何缓解梯度消失
提示:考虑对输入 $X$ 的归一化:
$$
Y=\frac{X-\mu}{\sigma}\cdot\gamma+\beta
$$
参考答案:
$$
\frac{\partial Y}{\partial X} = \frac{\gamma}{\sigma}\left(I – \frac{1}{n}\mathbf{1}\mathbf{1}^T – \frac{(X-\mu)(X-\mu)^T}{n\sigma^2}\right)
$$
通过保持梯度幅度稳定,LayerNorm 能有效抑制梯度指数级衰减。
