Nyströmformer:通过 Nyström 方法在线性时间和内存中近似自注意力
TL;DR
Nyströmformer 是一种高效的 Transformer 变体,通过近似标准自注意力机制实现线性时间和内存复杂度 $O(n)$,而非传统的二次复杂度 $O(n^2)$。这使得模型能够处理显著更长的输入序列,同时在自然语言处理和计算机视觉任务中保持竞争力的性能。
Nyström 方法用于矩阵近似
Nyström 方法通过采样 $m$ 行和 $m$ 列来近似一个大的矩阵 $P^{n \times n}$。将这些样本排列成块矩阵后,原矩阵被划分为四个子矩阵:$A_{P}$($m \times m$)、$B_{P}$($m \times (n - m)$)、$F_{P}$($(n - m) \times m$)以及 $C_{P}$($(n - m) \times (n - m)$)。
由于 $A_{P}$、$B_{P}$ 和 $F_{P}$ 的条目通过采样已知,未知的子矩阵 $C_{P}$ 可以估计为:
$$ C_{P} = F_{P} A_{P}^{+} B_{P} $$
其中 $+$ 表示 Moore‑Penrose 伪逆。得到的 Nyström 近似 $\hat{P}$ 表示为三个矩阵的乘积,从而降低了计算完整矩阵的计算负担。
将 Nyström 方法应用于自注意力
标准自注意力依赖于软最大矩阵 $S = \text{softmax}(\frac{QK^T}{\sqrt{d}})$。直接将 Nyström 方法应用于 $S$ 是不可行的,因为计算软最大矩阵的单列需要了解所有其他列以计算行软最大分母。
为了解决此问题,Nyströmformer 从查询 ($Q$) 和键 ($K$) 中采样“地标”(Nyström 点),而不是从得到的软最大矩阵中采样。模型基于查询地标 $\tilde{Q}$ 和键地标 $\tilde{K}$ 定义了三个矩阵:
- $\tilde{F} = \text{softmax}(\frac{Q\tilde{K}^T}{\sqrt{d}})$
- $\tilde{A} = \text{softmax}(\frac{\tilde{Q}\tilde{K}^T}{\sqrt{d}})^{+}$
- $\tilde{B} = \text{softmax}(\frac{\tilde{Q}K^T}{\sqrt{d}})$
近似的软最大矩阵 $\hat{S}$ 随后计算为 $\hat{S} = \tilde{F} \tilde{A} \tilde{B}$。通过将该近似与值 ($V$) 相乘,模型实现了自注意力的线性近似,而无需计算完整的 $QK^T$ 矩阵乘积。
地标选择与实现
Nyströmformer 使用段均值来构建 $\tilde{Q}$ 和 $\tilde{K}$。将 $n$ 个 token 划分为 $m$ 段,并计算每段的均值作为地标。研究表明,即使在 $n = 4096$ 或 $8192$ 的序列长度下,使用仅 32 或 64 个地标也能提供竞争力的性能。
实现细节
- Depthwise Convolution(深度可分离卷积): 该架构在值上加入了 1D 深度可分离卷积(DConv)作为跳跃连接。
- Complexity(复杂度): 该方法成功避免了 $O(n^2)$ 的复杂度,使得在长序列上也能高效运行。
在 Hugging Face 上的可用性与使用
Nyströmformer 可通过 Hugging Face Transformers 库用于掩码语言模型(MLM)。提供了四个基于序列长度的检查点:nystromformer-512、nystromformer-1024、nystromformer-2048 和 nystromformer-4096。
用户可以通过 NystromformerConfig 中的 num_landmarks 参数来控制地标数量 $m$。模型可使用 NystromformerForMaskedLM 类或高级的 pipeline API(如 fill-mask 任务)进行部署。