共计 2140 个字符,预计需要花费 6 分钟才能阅读完成。
1. Transformer 与自注意力机制简介
2017 年诞生的 Transformer 架构彻底改变了 NLP 领域,其核心创新就是 自注意力机制。与传统 RNN 不同,自注意力能直接捕捉序列中任意两个元素的关系,解决了长距离依赖问题。举个简单例子:

- 在句子 ”The animal didn’t cross the street because it was too tired” 中,自注意力能自动发现 ”it” 与 ”animal” 的高关联度,而无需像 RNN 那样逐步传递信息
2. 2D 多头注意力分步拆解
2.1 整体计算流程图
graph TD
A[输入序列 X] --> B[线性投影得到 Q,K,V]
B --> C[拆分为多头 Q,K,V]
C --> D[缩放点积注意力计算]
D --> E[多头结果拼接]
E --> F[最终输出]
2.2 关键公式与解释
-
QKV 投影:
$$\begin{aligned}
Q = XW_Q, \quad K = XW_K, \quad V = XW_V
\end{aligned}$$ -
这三个矩阵的维度通常为
(seq_len, d_model) -
投影权重 $W_Q/W_K/W_V$ 是可训练参数
-
多头拆分(以头数 h = 8 为例):
$$\text{MultiHead}(Q,K,V) = \text{Concat}(head_1,…,head_h)W_O$$ -
使用
einops.rearrange实现优雅的维度变换:q = rearrange(q, 'b s (h d) -> b h s d', h=self.num_heads) -
缩放点积注意力:
$$\text{Attention}(Q,K,V) = \text{softmax}(\frac{QK^T}{\sqrt{d_k}})V$$ -
缩放因子 $\sqrt{d_k}$ 防止梯度消失
- 计算复杂度为 $O(n^2)$,n 为序列长度
3. PyTorch 完整实现
import torch
import torch.nn as nn
from einops import rearrange
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, num_heads=8):
super().__init__()
self.d_model = d_model
self.num_heads = num_heads
assert d_model % num_heads == 0, "d_model 必须能被 num_heads 整除"
# 定义 QKV 投影矩阵
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):
b, s, _ = x.shape # batch_size, seq_len, d_model
# 1. 线性投影
q = self.w_q(x) # (b,s,d)
k = self.w_k(x)
v = self.w_v(x)
# 2. 拆分为多头
q = rearrange(q, 'b s (h d) -> b h s d', h=self.num_heads)
k = rearrange(k, 'b s (h d) -> b h s d', h=self.num_heads)
v = rearrange(v, 'b s (h d) -> b h s d', h=self.num_heads)
# 3. 缩放点积注意力
scores = torch.matmul(q, k.transpose(-2,-1)) / (self.d_model ** 0.5)
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
attn = torch.softmax(scores, dim=-1)
# 4. 多头拼接
output = torch.matmul(attn, v) # (b,h,s,d)
output = rearrange(output, 'b h s d -> b s (h d)')
return self.w_o(output)
4. 性能优化实战
4.1 计算复杂度分析
- 空间复杂度:存储 $QK^T$ 矩阵需要 $O(n^2)$ 内存
- 处理 1024 长度序列时,显存占用已达 1GB(float32)
4.2 头数选择经验
| 头数 | 计算速度 | 显存占用 | 效果 |
|---|---|---|---|
| 4 | 最快 | 最低 | 一般 |
| 8 | 适中 | 中等 | 推荐 |
| 16 | 较慢 | 较高 | 提升有限 |
5. 避坑指南
5.1 梯度爆炸处理
- 出现 NaN 值时,尝试调大缩放因子:
scale = (d_model / num_heads) ** 0.5 # 可调整为 1.0~2.0
5.2 变长序列技巧
# 创建 padding mask 示例
mask = (x != 0).unsqueeze(1).unsqueeze(2) # (b,1,1,s)
6. 拓展思考:Flash Attention
最新提出的 Flash Attention 通过以下方式优化:
1. 分块计算避免存储完整 $QK^T$ 矩阵
2. 融合 kernel 减少内存访问
3. 支持半精度计算
实现示例:
from flash_attn import flash_attention
output = flash_attention(q, k, v)
建议读者尝试将本文实现迁移到 Flash Attention,比较两者的速度和内存差异。
正文完
发表至: 未分类
近两天内
