共计 2018 个字符,预计需要花费 6 分钟才能阅读完成。
1. 自注意力计算复杂度基础
自注意力机制的核心计算可以用公式表示为:

$$
\text{Attention}(Q,K,V) = \text{softmax}(\frac{QK^T}{\sqrt{d}})V
$$
其中关键复杂度来自三个部分:
- QK^T 矩阵乘法 :形状为(l,d) 和(d,l)的矩阵相乘,产生 O(l²d)计算量
- softmax 计算:对 l×l 矩阵的逐行归一化,复杂度 O(l²)
- 权重与 V 的乘积 :l×l 与 l×d 矩阵相乘,再次产生 O(l²d) 复杂度
综合可得整体复杂度为 O(l²d)。这意味着:
- 当序列长度翻倍时,计算量变为 4 倍
- 特征维度增加主要影响与 d 成正比的项
2. 注意力变体复杂度对比
| 注意力类型 | FLOPs | 内存占用 | 适用场景 |
|---|---|---|---|
| 标准注意力 | O(l²d) | O(l²) | 短序列任务 |
| 稀疏注意力 | O(l√l·d) | O(l√l) | 长文档建模 |
| 局部注意力 | O(lk·d) | O(lk) | 固定窗口模式(k= 窗口大小) |
3. PyTorch 标准实现
import torch
import torch.nn as nn
import torch.nn.functional as F
class SelfAttention(nn.Module):
def __init__(self, d_model, heads=8):
super().__init__()
self.d_head = d_model // heads
self.heads = heads
# 线性变换层
self.to_qkv = nn.Linear(d_model, 3*d_model)
self.scale = self.d_head ** -0.5
def forward(self, x):
"""
输入: [batch, seq_len, d_model]
输出: [batch, seq_len, d_model]
"""
b, l, d = x.shape
# 生成 QKV [3, b, heads, l, d_head]
qkv = self.to_qkv(x).chunk(3, dim=-1)
q, k, v = map(lambda t: t.view(b, l, self.heads, -1).transpose(1, 2),
qkv
)
# 注意力得分 [b, heads, l, l]
scores = (q @ k.transpose(-2, -1)) * self.scale
attn = scores.softmax(dim=-1)
# 上下文向量 [b, heads, l, d_head]
out = attn @ v
# 合并多头 [b, l, d_model]
out = out.transpose(1, 2).reshape(b, l, -1)
return out
4. 关键优化技术
4.1 Flash Attention
通过避免中间注意力矩阵的显式存储,将内存占用从 O(l²)降到 O(l):
# 需要安装 flash-attn 包
torch.backends.cuda.enable_flash_sdp(True)
def flash_attention(q, k, v):
return F.scaled_dot_product_attention(q, k, v)
4.2 序列分块处理
将长序列切分为可管理的块(chunk_size=512 典型值):
def chunked_attention(x, chunk_size=512):
chunks = x.split(chunk_size, dim=1)
return torch.cat([self_attention(chunk) for chunk in chunks], dim=1)
4.3 混合精度训练
实测在 A100 上可获得 1.8-2.3 倍加速:
with torch.autocast(device_type='cuda', dtype=torch.float16):
output = model(input)
5. 生产环境实践
5.1 内存管理技巧
-
使用梯度检查点:牺牲 30% 计算时间换取 50% 内存节省
from torch.utils.checkpoint import checkpoint output = checkpoint(self_attention, x) -
及时释放中间变量:
with torch.no_grad(): # 中间计算代码 torch.cuda.empty_cache()
5.2 硬件适配建议
| 硬件 | 优化重点 | 典型配置建议 |
|---|---|---|
| GPU | 增大 batch_size 利用并行度 | 开启 Tensor Core |
| TPU | 避免小矩阵运算 | 使用 XLA 优化器 |
| CPU | 限制线程数避免争抢资源 | 设置 OMP_NUM_THREADS |
6. 深度思考与实践
当特征维度 d 远大于序列长度 l 时(例如 d =1024, l=32):
- 计算量分布:l²d=1M vs ld²=33M → 此时优化 d 更关键
- 可尝试策略:
- 降维投影
- 特征分组注意力
- 低秩近似
建议在 Colab 上运行以下对比实验:
for d in [256, 512, 1024]:
for l in [64, 128, 256]:
x = torch.randn(32, l, d).cuda()
%timeit self_attention(x)
通过实际测量理解不同参数对性能的影响规律。
正文完
发表至: 未分类
近三天内
