格拉默矩阵和线性注意力:不稳定性是怎么被发现并逐步解决的

我最早跑 Performer 的时候,特别郁闷。

AI technology illustration

一行代码从标准 attention 换成 linear attention,然后训练 loss 就开始抽风——偶尔掉到 NaN,偶尔飞上十位数。我以为是学习率没调好,调了半天也没用。直到我翻了原论文才知道,问题不在优化器,而在注意力矩阵本身的「病态调性」。

后来我意识到:要理解线性注意力为什么这么脆弱,绕不开一个概念——格拉默矩阵(Gram matrix)。

线性注意力是怎么「偷懒」的?

标准 attention 要做一步「让每个 token 跟所有 token 打招呼」。

给定 query 矩阵 Q,key 矩阵 K,value 矩阵 V(都是 N 行 d 维),你算一个注意力矩阵 A = softmax(QKT/√d)。这个 A 是 N×N 的,N 越大,内存和计算耗费越高,硬生生把长度压到了几千 token。

线性注意力的核心想法是:能不能把 softmax 拆开?

如果存在一个核函数 φ,使得 softmax(QKT) ≈ φ(Q) φ(K)T,那么就可以利用乘法结合律:

Attention = (φ(Q) φ(K)^T) V = φ(Q) (φ(K)^T V)

这样复杂度从 O(N²) 降到 O(N d²),其中 N 是序列长度,d 是隐藏维度。对 d=64, N=8192 这种规模,计算量直接少两个数量级。

这个想法最早来自 Katharopoulos 等人 2020 年的论文,后来 Performer 把它升级成可逼近 softmax 的随机特征映射。总之,φ 要满足「内积≈softmax」。

但问题来了:这个替换真的卫生吗?

不稳定性,比预想中来得更早

Performer 论文里其实有提到,随机特征映射会有方差。但他们给出的理论误差下界看起来不错——只要特征维度 m 够大,误差就小。

然而在实际训练中,事情没这么简单。

先是很多人发现,Performer 在长序列上训练 loss 会震荡,有时直接 NaN。然后有人试图像普通 attention 那样加深网络,却发现梯度在第三四层就消失了。

更离谱的是,把线性注意力换成「直接除以一个常数」的简化版(比如只用 φ(Q) φ(K)^T,不归一化),模型几乎不收敛。

于是研究者们开始刨根问底,最后把目光落在了格拉默矩阵上。

格拉默矩阵,其实就藏在你的注意力里面

你可能不知道,标准 attention 里的矩阵 softmax(QKT) 本质上是一个「加权」过的格拉默矩阵。

格拉默矩阵的定义很简单:给定一组向量 v₁, v₂, …, v_N,它的 Gram 矩阵 G 满足 G_ij = ⟨vᵢ, vⱼ⟩,即每两个向量的内积。

快速回顾:什么是格拉默矩阵?

给定一组向量,它们的格拉默矩阵就是内积矩阵。它完全描述了这组向量之间的夹角和长度。如果格拉默矩阵的行列式为零,说明这些向量线性相关。

而在线性注意力里,φ(Q) φ(K)T 就是 φ(query) 和 φ(key) 两个向量集合的内积矩阵。所以整个注意力矩阵,就是一个格拉默矩阵。

格拉默矩阵有什么特别的?它的特征值直接反映了这组向量的「独立性」。

如果特征值从大到小衰减得特别快,只有少数几个主导特征值,说明这些向量高度相关——也就是条件数很大,矩阵几乎是「病态」的。

病态的格拉默矩阵,怎么把梯度搞坏的?

我们来算算梯度怎么穿过这个矩阵。

在标准的 softmax attention 中,注意力矩阵每行和恰好等于 1,因为 softmax 自己完成了归一化。这个归一化让每个 query 对 key 的注意力权重都落在一个可控区间,梯度也相对稳定。

而线性注意力没有软最大归一化,所以 φ(Q) φ(K)T 的值的大小完全由 φ 的尺度决定。

如果 φ 的范数在某些方向上特别大,注意力矩阵的某个特征值就可能极其膨胀。经过几层堆叠后,这个膨胀又会被后面的矩阵乘法放大,梯度因此爆炸。

举个例子。假设特征维度 d=16,随机初始化后,QKT 的格拉默矩阵的特征值可能从 10 到 0.001 不等。经过 softmax,这些值会被压缩到 0.3~0.4 之间;但线性注意力直接保住了原始差异,最大的特征值贡献了整个矩阵能量的 99%。你算一算条件数:10/0.001 = 10000。这就是梯度爆炸的来源。

我最早以为模型「没算对」,后来才意识到:它不是没算对,是根本没在算一个健康的矩阵。

你看,当格拉默矩阵的最大特征值比最小特征值高出上千倍时,这个矩阵在求导中就像一头疯牛——它的逆(或伪逆)会放大误差。

这给我们的启示是:稳定性的关键不在于核函数有多精确,而在于格拉默矩阵的谱分布是否「均匀」。

发现之旅:从随机到经验到理论

第一个正面应对这个问题的,是 Performer 的作者们。

他们在 2020 年的论文里提出了用「正交随机特征」(orthogonal random features)替代独立同分布随机特征,以减少特征映射的方差。正交化可以让 φ 的行向量更独立,从而降低格拉默矩阵的条件数。

随后,2022 年的一篇论文(Transformer Quality in Linear Time)观察到,即使不考虑随机性,线性注意力在堆叠时会出现「注意力输出去相关化」的趋势,导致信息瓶颈。他们提出的建议是在 Q、K、V 上做 LayerNorm,让输入特征有稳定的尺度。

再后来,研究者们从理论层面分析,发现线性注意力的误差下界不仅和特征维度有关,还和格拉默矩阵的平方范数有关。于是新的设计开始主动修正这个矩阵。

我记得有一篇叫 TransNormer 的论文(2022),它把 QK^T 中的内积做了一种「归一化」处理——不是标准化整个矩阵,而是减去格拉默矩阵的对角线,再乘一个可学习的温度。这个操作在实验中几乎消除了 NaN 现象。

直到最近 2024 年,像 Gated Linear Attention (GLA) 这样的模型,直接用门控机制动态调整每个 channel 的缩放,等于在隐式地控制格拉默矩阵的奇异值。

一张表告诉你,各家在稳定化上做了啥

方法 核心手段 针对的问题
Performer (2020) 正交随机特征 方差过大
QK LayerNorm (2022) 对 Q、K 做归一化 尺度失控
TransNormer (2022) 对角归一化 + 温度 条件数过大
GLA (2024) 门控 + 逐通道缩放 长序列衰减

你看,思路从「减少格拉默矩阵方差」逐步走向了「控制格拉默矩阵的谱」。

FAQ:你可能还想问

为什么不用 softmax 的归一化,直接手动标准化一下不就行了?

手动标准化确实能解决尺度问题,比如除以 √d。但线性注意力的目标是用矩阵乘法替换 softmax 的迭代计算。一旦你加入全局归一化,就需要求和所有列,这又回到了 O(N²)。目前的方案大多是「局部近似归一化」,比如用特征范数近似。

是不是所有线性注意力都不稳定?

不是。稳定性取决于特征映射和网络结构。只要控制好格拉默矩阵的条件数,线性注意力可以稳定训练。很多现代模型(如 RetNet、Mamba)都做得不错。

格拉默矩阵的浪漫,也是它的桎梏

现在你再回头看看线性注意力。它的高效本质来源于把 N×N 矩阵拆成两个小矩阵相乘。但代价是失去了 softmax 的天生归一化属性,必须自己照顾格拉默矩阵的卫生。

这让我想起一句话:免费的午餐总在别处收费。

好消息是,经过这几年的迭代,不稳定性已经基本被驯服。至少我现在用 GLA 训练 1 万 token 的序列,再也不用担心 NaN 了。

但坏消息是,目前所有的稳定性措施几乎都依赖于某种形式的信息瓶颈——要么牺牲精度,要么得加门控。没有哪个方法能在所有场景下同时做到稳定、高效、效果好。

线性注意力还有很长的路要走——不是提高运算速度的路,而是让格拉默矩阵「安分下来」的路。

原创文章,作者:guanweilu,如若转载,请注明出处:https://guanweilu.cn/article/597.html

(0)
上一篇 1天前
下一篇 1天前

相关推荐