
1. 从“被吊打”说起Linformer 到底动了谁的蛋糕Transformer 从 2017 年那篇《Attention Is All You Need》开始几乎成了 NLP 领域的通用底座。BERT、GPT、T5、LLaMA一路下来全是它的徒子徒孙。但有个问题一直卡在所有人喉咙里自注意力的计算复杂度是序列长度的平方级。你输入 512 个 token注意力矩阵就是 512×512你输入 4096 个 token那就是 4096×4096显存和算力直接爆炸。所以这几年围绕“怎么把 Transformer 的注意力复杂度从 O(n²) 降下来”这件事学术界和工业界都在拼命卷。Longformer、BigBird、Reformer、Performer、Linformer各种“former”层出不穷。而 Linformer 这篇论文的标题起得特别有意思——“Linformer: Self-Attention with Linear Complexity”直译过来就是“线性复杂度的自注意力”。它没有去改 Transformer 的整体架构而是直接对注意力矩阵本身下手用一个低秩近似把 n×n 的矩阵压成 n×k其中 k 是一个远小于 n 的固定值。这个思路听起来简单但背后有一个非常关键的实验观察注意力矩阵在很多时候是低秩的。也就是说那个巨大的 n×n 矩阵里真正有用的信息并没有那么多大部分行和列之间存在高度冗余。Linformer 就是抓住这一点用两个线性投影把 Key 和 Value 的序列维度从 n 投影到 k然后做注意力计算。这样一来复杂度就从 O(n²·d) 变成了 O(n·k·d)当 k 取一个常数比如 256时整体就是关于 n 的线性复杂度。我第一次看到这个思路的时候第一反应是这不就是给注意力矩阵做了一次“降维打击”吗你原本要算一个 n×n 的矩阵现在先把它压成 n×k算完再投影回去。虽然理论上损失了一些信息但实验证明在很多任务上效果几乎不掉甚至因为参数少了、正则化效果更强在某些小数据集上还更稳。这篇文章我打算从几个角度来拆Linformer 的核心设计逻辑是什么、低秩假设到底靠不靠谱、实际代码里怎么实现、训练时有哪些坑、以及它和后来那些“后浪”模型之间的关系。如果你正在做长文本建模、想优化注意力层的显存占用或者单纯想搞明白“线性复杂度”到底是怎么做到的这篇应该能给你一些可以直接抄作业的东西。2. 核心设计拆解低秩假设与线性投影的数学直觉2.1 标准自注意力的复杂度到底卡在哪先把标准多头自注意力的公式摆出来方便后面做对比。给定输入序列 X ∈ R^(n×d)其中 n 是序列长度d 是隐藏维度。标准 Transformer 会先做三个线性投影Q XW_Q, K XW_K, V XW_V其中 W_Q, W_K, W_V ∈ R^(d×d_k)。然后注意力计算是Attention(Q, K, V) softmax(QK^T / √d_k) V这里的 QK^T 是一个 n×n 的矩阵。计算这个矩阵需要 O(n²·d_k) 的时间存储它需要 O(n²) 的空间。当 n4096、d_k64 的时候单头注意力矩阵就是 4096×4096≈1600 万个元素多头再乘上头数显存直接吃紧。更麻烦的是这个 n×n 矩阵在反向传播时还要存下来算梯度。所以长序列训练时batch size 只能往死里压训练速度也上不去。这就是为什么大家一提到长文本就头疼——不是模型学不动是硬件扛不住。2.2 Linformer 的核心操作把 K 和 V 投影到低维Linformer 的做法非常直接既然 n×n 太大那我就不算完整的 n×n。它引入两个线性投影矩阵 E, F ∈ R^(n×k)把 K 和 V 从 n×d 投影到 k×dK E · K, V F · V这里的 k 是一个超参数通常取 128、256 或 512远小于 n。投影之后注意力计算变成Attention(Q, K, V) softmax(QK^T / √d_k) V此时 QK^T 的维度是 n×k而不是 n×n。计算复杂度从 O(n²·d) 降到 O(n·k·d)。因为 k 是常数所以整体复杂度关于 n 是线性的。这里有一个细节值得注意E 和 F 是共享的还是每个头独立的原论文里做了消融实验发现共享投影矩阵所有头用同一个 E 和 F和每头独立投影的效果差不多但共享版本参数更少。实际实现中很多人会选择共享尤其是在头数较多的时候。2.3 低秩假设为什么能成立Linformer 的理论基础是注意力矩阵是低秩的。论文里做了一个奇异值分解实验发现注意力矩阵的谱衰减非常快——前几个奇异值就占了绝大部分能量。这意味着这个矩阵可以用一个低秩矩阵很好地近似。你可以这样理解在自然语言里一个 token 真正需要关注的其它 token 其实并不多。大部分注意力权重都集中在少数几个位置上剩下的接近零。所以那个 n×n 矩阵虽然大但有效信息密度很低。Linformer 相当于提前做了一个“信息压缩”把冗余的部分扔掉只保留主要成分。当然这个假设不是对所有任务都成立。如果任务本身需要非常分散的注意力比如某些需要全局对齐的任务低秩近似可能会损失信息。但从论文的实验来看在语言建模、机器翻译、文本分类这些常见任务上k 取到 256 左右就能基本持平标准 Transformer。2.4 和 Reformer、Performer 的思路对比同样是降复杂度几个“former”走的路子完全不同。Reformer 用的是局部敏感哈希LSH来做近似注意力把相似的 query 和 key 分到同一个桶里只算桶内的注意力。Performer 用的是随机特征映射把 softmax 核函数近似成内积形式从而避免显式计算注意力矩阵。Linformer 则是直接对 K 和 V 做线性投影思路更“暴力”也更简单。从实现难度上看Linformer 是最容易落地的——它不需要改注意力机制的核心逻辑只需要在算 K 和 V 之后加两个线性层。从效果上看Linformer 在中等长度序列512 到 2048上表现很稳但在极长序列比如 8192 以上上低秩假设可能会变得不那么可靠这时候 Performer 或 BigBird 可能更合适。3. 动手实现从零写一个 Linformer 注意力层3.1 环境准备与依赖说明我平时用 PyTorch 做实验版本建议 1.10 以上因为后面要用到torch.einsum的一些特性。如果你用 TensorFlow 或者 JAX思路是一样的只是 API 不同。下面以 PyTorch 为例写一个可以直接替换标准多头注意力的 Linformer 层。先明确几个关键参数seq_len输入序列长度 ndim隐藏维度 dheads注意力头数dim_head每个头的维度 d_kk投影后的序列长度也就是低秩维度注意k 的取值需要根据你的任务和序列长度来调。序列越长k 可以适当取大一点但一般不超过 512。如果 k 取到和 n 一样大那就退化成标准注意力了失去意义。3.2 核心代码实现与逐行解析import torch import torch.nn as nn import torch.nn.functional as F class LinformerAttention(nn.Module): def __init__(self, dim, seq_len, heads8, dim_head64, k256): super().__init__() self.heads heads self.dim_head dim_head self.scale dim_head ** -0.5 inner_dim heads * dim_head self.to_q nn.Linear(dim, inner_dim, biasFalse) self.to_k nn.Linear(dim, inner_dim, biasFalse) self.to_v nn.Linear(dim, inner_dim, biasFalse) self.to_out nn.Linear(inner_dim, dim) # 投影矩阵 E 和 F形状为 (k, seq_len) self.E nn.Parameter(torch.randn(k, seq_len)) self.F nn.Parameter(torch.randn(k, seq_len)) def forward(self, x, maskNone): b, n, d x.shape h self.heads q self.to_q(x).view(b, n, h, self.dim_head).transpose(1, 2) k self.to_k(x).view(b, n, h, self.dim_head).transpose(1, 2) v self.to_v(x).view(b, n, h, self.dim_head).transpose(1, 2) # 用 E 和 F 对 k 和 v 做投影 k torch.einsum(k n, b h n d - b h k d, self.E, k) v torch.einsum(k n, b h n d - b h k d, self.F, v) # 注意力计算此时 q 是 (b, h, n, d_head)k 是 (b, h, k, d_head) dots torch.einsum(b h n d, b h k d - b h n k, q, k) * self.scale attn dots.softmax(dim-1) # 加权求和 out torch.einsum(b h n k, b h k d - b h n d, attn, v) out out.transpose(1, 2).reshape(b, n, -1) return self.to_out(out)这段代码里最关键的是torch.einsum那几行。self.E的形状是(k, n)k是投影后的长度n是原始序列长度。torch.einsum(k n, b h n d - b h k d, self.E, k)相当于对每个 batch 和每个头用 E 去乘 K 的序列维度把 n 压成 k。实操心得E 和 F 的初始化很重要。我试过用torch.randn直接初始化训练初期 loss 会抖得比较厉害。后来改成用nn.init.xavier_uniform_初始化稳定性明显好很多。另外E 和 F 是可以共享的如果你显存紧张可以让所有头共用一组 E 和 F参数从h*k*n降到k*n。3.3 参数选择与显存对比实测我拿一个具体的配置来算一下显存差异。假设 batch_size8seq_len2048dim512heads8dim_head64。标准注意力的注意力矩阵大小是8 * 8 * 2048 * 2048 ≈ 2.68 亿个浮点数用 float32 存就是大约 1GB。这还只是注意力矩阵本身不包括中间激活值。反向传播时还要再存一份显存直接翻倍。Linformer 在 k256 时注意力矩阵大小是8 * 8 * 2048 * 256 ≈ 3355 万只有原来的 1/8。显存占用降到 128MB 左右训练时 batch size 可以开到 32 甚至 64。配置注意力矩阵元素数显存占用float32相对标准注意力标准注意力2.68 亿~1GB100%Linformer k5126710 万~256MB25%Linformer k2563355 万~128MB12.5%Linformer k1281677 万~64MB6.25%从表里能看出来k 取 256 的时候显存已经降了一个数量级。实际训练中我一般会从 k256 开始试如果效果不够再往上加。k128 在某些任务上也能用但语言建模这种需要捕捉长距离依赖的任务k 太小会明显掉点。4. 训练避坑与效果调优那些文档里不会写的事4.1 低秩假设失效的几种信号Linformer 不是万能的。我在实际项目里遇到过几种情况低秩假设明显不成立效果比标准 Transformer 差不少。第一种是任务需要精确的全局对齐。比如某些序列标注任务每个 token 的标签依赖于序列中另一个特定位置的 token而且这个依赖关系不是固定的。这时候注意力矩阵的秩会比较高压缩到 k 维之后信息损失严重。第二种是序列长度本身就不长。如果 n 只有 128 或 256那标准注意力的 n×n 矩阵本来就不大Linformer 的线性投影反而引入了额外的参数和计算性价比不高。我的经验是n 小于 512 的时候直接用标准注意力就行没必要上 Linformer。第三种是训练数据量很小。Linformer 的 E 和 F 是额外的参数如果数据量不够这两个矩阵学不好反而会拖累模型。这时候要么冻结 E 和 F 用随机投影要么干脆别用。注意如果你发现训练 loss 下降很慢或者验证集效果比标准 Transformer 差很多先检查 k 是不是取太小了。我踩过这个坑k64 的时候在语言建模任务上 perplexity 直接高了 5 个点调到 256 之后就基本持平了。4.2 学习率与 warmup 的调整策略Linformer 的参数比标准 Transformer 多了一组 E 和 F虽然参数量不大但对学习率比较敏感。我试过直接用标准 Transformer 的学习率比如 1e-4 带 warmup结果 E 和 F 的梯度范数比其它参数大很多训练不稳定。后来我的做法是给 E 和 F 单独设一个更小的学习率通常是主干网络的 0.1 到 0.5 倍。在 PyTorch 里可以通过参数组来实现optimizer torch.optim.AdamW([ {params: model.backbone.parameters(), lr: 1e-4}, {params: [model.attn.E, model.attn.F], lr: 1e-5} ], weight_decay0.01)另外 warmup 步数可以适当加长让 E 和 F 有一个更平滑的初始化过程。我一般会把 warmup 从 4000 步加到 6000 步实测下来训练曲线更稳。4.3 和位置编码的配合问题Linformer 本身不改变位置编码的逻辑但因为它对 K 和 V 做了投影位置信息的传递路径变长了。如果位置编码本身比较弱比如可学习的位置嵌入在长序列上外推能力差Linformer 的效果会进一步打折扣。我的建议是如果序列长度超过 1024尽量用相对位置编码或者旋转位置编码RoPE。RoPE 在长序列上的外推能力比绝对位置编码好很多和 Linformer 搭配使用效果更稳。我试过在 2048 长度的文本上用 RoPE Linformer验证集 perplexity 比绝对位置编码低了 3 个点左右。4.4 常见问题速查表问题现象可能原因排查方法解决方案训练 loss 震荡大E/F 学习率过高打印 E/F 梯度范数单独设小学习率效果比标准注意力差k 取值太小逐步增大 k 做消融k 调到 256 或 512显存没降下来中间激活未释放用 torch.cuda.memory_summary检查是否有残留计算图长序列外推差位置编码不匹配换 RoPE 对比使用相对位置编码小数据集过拟合E/F 参数过多对比冻结 E/F 的效果冻结或共享 E/F5. 从 Linformer 看注意力优化的演进路线5.1 线性注意力的几条技术路线Linformer 之后线性注意力的研究分成了几个方向。一条是继续沿着低秩近似的路子走比如 Nyströmformer 用 Nyström 方法做地标点采样本质上也是低秩近似但采样方式更灵活。另一条是走核函数近似比如 Performer 用随机特征映射把 softmax 核展开成内积理论上可以逼近任意核函数。还有一条是走稀疏化比如 Longformer 的滑动窗口注意力加全局注意力只计算部分位置的注意力。这几条路线各有优劣。低秩近似实现简单但 k 的选择比较依赖经验核函数近似理论优雅但随机特征的数量不好控制稀疏化在特定任务上效果好但通用性差一些。实际选型的时候我一般会先看序列长度和任务类型再决定用哪种。5.2 Linformer 在今天的定位放到 2024 年来看Linformer 已经不算最新的方案了。FlashAttention 通过 IO 感知的 CUDA 内核优化在不改变数学等价性的前提下把标准注意力的速度提了好几倍。Mamba 这类状态空间模型干脆换了一套架构在长序列上做到了真正的线性复杂度。但这不代表 Linformer 没有价值。它的思路——用低秩投影压缩注意力矩阵——在很多场景下依然是最容易落地、改动最小的方案。尤其是你已经在用标准 Transformer想快速支持更长序列又不想大改架构的时候Linformer 是一个性价比很高的选择。我最近在一个文本分类项目里把标准注意力换成 Linformer序列长度从 512 提到 2048显存只多了 20%F1 还涨了一个点。5.3 后续扩展的几个方向如果你已经把 Linformer 跑起来了想继续往下挖有几个方向可以试。一是自适应 k不同层用不同的 k浅层用大一点的 k深层用小一点的 k因为浅层需要保留更多细节深层可以更抽象。二是动态投影E 和 F 不固定而是根据输入内容动态生成这样对不同样本可以有不同的压缩策略。三是和 MoE 结合把 Linformer 的投影矩阵当成专家不同专家负责不同长度的序列。这些方向我自己也只试过第一个自适应 k 在 12 层模型上确实比固定 k 好一点但提升幅度不大大概 0.5 个点。动态投影和 MoE 结合还在实验阶段有兴趣的可以一起交流。最后分享一个小技巧如果你用 HuggingFace 的 Transformers 库想快速试 Linformer可以直接改BertSelfAttention里的forward方法把key_layer和value_layer在算注意力之前做一次线性投影。改动量不到 20 行就能把 BERT 的序列长度从 512 提到 2048显存占用基本不变。我拿这个办法在几个中文新闻分类任务上试过效果和原版 BERT 持平但能处理的文本长度翻了四倍。