共计 3042 个字符,预计需要花费 8 分钟才能阅读完成。
背景:传统多头注意力的局限性
传统多头注意力机制(Multi-Head Attention, MHA)通过并行计算多个注意力头来捕捉不同子空间的语义信息,但其均匀分配计算资源的策略存在明显缺陷:

- 长序列低效性 :计算复杂度随序列长度呈平方增长($O(n^2d)$),难以处理长文本或高分辨率图像
- 粒度单一化 :所有位置采用相同的注意力粒度,无法自适应不同层次的特征交互需求(如段落级粗粒度与词级细粒度)
- 局部模式忽略 :全局注意力可能稀释局部重要模式,尤其在层次化数据结构中(如代码 / 数学公式)
C2F 多头注意力技术解析
分层注意力原理
C2F(Coarse-to-Fine)机制通过两级注意力实现层次化建模:
-
粗粒度阶段(Coarse):将输入序列划分为 $k$ 个块(chunks),计算块间注意力
$$
\text{Attention}_c(Q_c, K_c, V_c) = \text{softmax}(\frac{Q_cK_c^T}{\sqrt{d_k}})V_c
$$
其中 $Q_c, K_c, V_c \in \mathbb{R}^{n/k \times d}$,实现复杂度降为 $O((n/k)^2d)$ -
细粒度阶段(Fine):在每个块内部进行全连接注意力
$$
\text{Attention}_f(Q_f, K_f, V_f) = \text{softmax}(\frac{Q_fK_f^T}{\sqrt{d_k}})V_f
$$
保持原始维度 $Q_f, K_f, V_f \in \mathbb{R}^{k \times d}$,总复杂度 $O(nk d)$
权重共享机制
采用分层参数共享策略:
- 跨头共享 :所有注意力头共用相同的块划分策略
- 跨层共享 :不同 Transformer 层的块大小 $k$ 按比例递减(如深层网络使用更小的 $k$)
- 跨粒度共享 :Coarse 和 Fine 阶段共享投影矩阵 $W_Q, W_K, W_V$
复杂度对比
设序列长度 $n=1024$,头数 $h=8$,维度 $d=512$:
| 类型 | FLOPs | 内存占用 |
|---|---|---|
| 标准 MHA | 4.2G | 2.1GB |
| C2F-MHA ($k=32$) | 1.3G (-69%) | 0.7GB (-66%) |
PyTorch 实现
import torch
import torch.nn as nn
import math
class CoarseAttention(nn.Module):
def __init__(self, d_model, n_heads, chunk_size):
super().__init__()
assert d_model % n_heads == 0
self.d_k = d_model // n_heads
self.n_heads = n_heads
self.chunk_size = chunk_size
# 共享投影矩阵
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)
def forward(self, x):
# x: [batch, seq_len, d_model]
batch_size = x.size(0)
seq_len = x.size(1)
# 分块处理
x = x.view(batch_size, -1, self.chunk_size, x.size(-1)) # [batch, n_chunks, chunk_size, d_model]
# 投影到 Q /K/V
Q = self.W_q(x).view(batch_size, -1, self.n_heads, self.d_k) # [batch, n_chunks, heads, chunk_size, d_k]
K = self.W_k(x).view(batch_size, -1, self.n_heads, self.d_k)
V = self.W_v(x).view(batch_size, -1, self.n_heads, self.d_k)
# 块间注意力
attn_scores = torch.einsum('bqhd,bkhd->bhqk', [Q, K]) / math.sqrt(self.d_k)
attn_probs = torch.softmax(attn_scores, dim=-1)
out = torch.einsum('bhqk,bkhd->bqhd', [attn_probs, V])
return out.view(batch_size, seq_len, -1)
class FineAttention(nn.Module):
def __init__(self, d_model, n_heads):
super().__init__()
self.attn = nn.MultiheadAttention(d_model, n_heads)
def forward(self, x):
# x: [batch, seq_len, d_model]
return self.attn(x, x, x)[0]
class C2FAttention(nn.Module):
def __init__(self, d_model, n_heads, chunk_size):
super().__init__()
self.coarse = CoarseAttention(d_model, n_heads, chunk_size)
self.fine = FineAttention(d_model, n_heads)
def forward(self, x):
coarse_out = self.coarse(x)
fine_out = self.fine(x)
return coarse_out + fine_out # 残差连接
避坑指南
梯度消失问题
- 初始化策略 :对 Q / K 投影矩阵使用 Xavier 初始化,V 矩阵使用零均值高斯初始化(标准差 =0.02)
- 梯度裁剪 :限制注意力分数矩阵的梯度范数(推荐阈值 1.0)
- 混合精度训练 :使用 AMP(Automatic Mixed Precision)降低数值不稳定风险
粒度层级设置
根据任务特性选择 chunk_size:
| 任务类型 | 推荐 chunk_size | 理论依据 |
|---|---|---|
| 机器翻译 | 8-16 | 短语对齐需求 |
| 文本分类 | 32-64 | 段落级语义聚合 |
| 语音识别 | 4-8 | 局部声学特征主导 |
显存优化
- 动态 mask 生成 :避免预先分配全尺寸 attention mask
def get_mask(seq_len, chunk_size): return torch.block_diag(*[torch.ones(chunk_size, chunk_size) for _ in range(seq_len // chunk_size)]) - 梯度检查点 :在 backward 时重新计算中间结果
from torch.utils.checkpoint import checkpoint output = checkpoint(self.attn, x)
拓展思考:视觉 Transformer 适配
将 C2F 机制应用于 ViT 时需考虑:
1. 空间分块策略 :二维图像需改为网格划分(如 16×16→4×4 子网格)
2. 跨尺度特征融合 :浅层网络用大块(捕捉整体结构),深层网络用小块(捕获细节)
3. 位置编码适配 :需设计块内相对位置编码与块间绝对位置编码的混合方案
参考文献
- [Coarse-to-Fine Attention: arXiv:2107.02239]
- [Efficient Transformers Survey: arXiv:2009.06732]
- [PyTorch 官方实现: github.com/pytorch/pytorch]
