长上下文注意力机制论文
扩展语言模型的上下文长度是当前研究的热点方向,从 2K 到 128K 再到百万 Token,长上下文注意力机制的创新使得模型能够处理整本书、长代码库和完整对话历史。
核心挑战
标准自注意力的计算复杂度为 O(n²),长序列面临:
- 计算瓶颈:序列长度翻倍,注意力计算量翻四倍
- 内存瓶颈:KV 缓存随序列长度线性增长
- 位置编码外推:训练时未见过的位置导致性能下降
高效注意力机制
稀疏注意力
通过稀疏化注意力矩阵降低计算量:
python
class SlidingWindowAttention(nn.Module):
"""滑动窗口注意力"""
def __init__(self, d_model, num_heads, window_size=512):
super().__init__()
self.num_heads = num_heads
self.window_size = window_size
self.qkv = nn.Linear(d_model, 3 * d_model)
self.out = nn.Linear(d_model, d_model)
def forward(self, x):
B, T, D = x.shape
qkv = self.qkv(x).reshape(B, T, 3, self.num_heads, D // self.num_heads)
q, k, v = qkv.unbind(dim=2)
# 创建滑动窗口掩码
mask = torch.zeros(T, T, dtype=torch.bool, device=x.device)
for i in range(T):
start = max(0, i - self.window_size // 2)
end = min(T, i + self.window_size // 2 + 1)
mask[i, start:end] = True
# 仅在窗口内计算注意力
attn = torch.matmul(q, k.transpose(-2, -1)) / (D // self.num_heads) ** 0.5
attn = attn.masked_fill(~mask, float('-inf'))
attn = torch.softmax(attn, dim=-1)
out = torch.matmul(attn, v)
return self.out(out.reshape(B, T, D))FlashAttention
FlashAttention 通过优化 GPU 内存访问实现精确注意力的加速:
- 分块计算:将注意力矩阵分块,减少 HBM 访问
- 在线 Softmax:分块内累积 Softmax 统计量
- 内存节省:无需存储完整 N×N 注意力矩阵
分层注意力
- 全局+局部:少量 Token 使用全局注意力,其余使用局部窗口
- 聚类注意力:将相似 Token 聚类,类间计算全局注意力
位置编码扩展
RoPE 扩展
旋转位置编码(RoPE)是当前主流的位置编码方案:
| 方法 | 策略 | 特点 |
|---|---|---|
| 位置插值 | 缩放位置索引 | 简单但可能损失分辨率 |
| NTK-aware | 调整 RoPE 基频 | 更好的长度外推 |
| YaRN | 动态混合缩放 | 结合多种策略 |
| ALiBi | 线性偏置注意力 | 无需位置编码 |
NTK-aware 缩放
NTK-aware 缩放的核心洞察:位置编码的高频分量对局部关系敏感,低频分量对全局关系敏感。扩展上下文时,应调整 RoPE 的基频 θ,使得新的频率分布能覆盖更长的序列范围。
KV 缓存优化
长上下文的 KV 缓存是内存瓶颈:
- PagedAttention:虚拟内存式的 KV 缓存管理
- KV 压缩:注意力加权合并相似 KV 对
- StreamingLLM:保留开头和最近的 KV,丢弃中间
- 量化 KV:将 KV 缓存量化到更低精度