资讯详情

Linformer 线性注意力实战:从平方复杂度到低秩投影的工程落地

📅 2026/10/7 18:42:58 | 华诺云谱 👁 阅读
Linformer 线性注意力实战:从平方复杂度到低秩投影的工程落地
Transformer 自注意力机制从 2017 年一路火到现在几乎成了 NLP 领域的默认底座。但真正在工程里跑过长文本的人都知道标准自注意力的计算和显存开销是随序列长度平方级增长的。序列一上千显存就吃紧序列上万单卡基本没戏。这也是为什么过去几年里各种高效注意力变体层出不穷而 Linformer 是其中思路最干净、最容易讲清楚的一个。它没有搞复杂的稀疏模式也没有引入额外的可学习路由而是直接抛出一个反直觉的结论注意力矩阵其实是低秩的那我们干脆把它投影到低维空间去算。这篇就围绕 Linformer 这个核心思路把它的动机、数学推导、代码实现、实测表现和踩坑经验完整拆一遍适合已经了解基础 Transformer、想搞明白线性复杂度注意力到底怎么落地的人。1. 为什么标准自注意力在长序列上会爆1.1 平方复杂度的来源到底在哪先把账算清楚。标准多头自注意力的核心公式是Attention(Q, K, V) softmax(QK^T / sqrt(d_k)) V假设序列长度为 n特征维度为 d。Q 和 K 相乘得到的是一个 n×n 的注意力分数矩阵这一步的计算量是 O(n²·d)而存储这个矩阵需要 O(n²) 的空间。注意这里的 n² 是实打实的不是渐进符号里的常数忽略项。n512 时是 26 万n4096 时是 1677 万n8192 时直接飙到 6700 万。每一步前向传播都要生成这么大的中间矩阵反向传播还要再存一份用于梯度计算显存压力可想而知。我早期做过一个文本分类任务序列长度 2048batch size 开到 16单张 24G 的卡就已经开始 OOM 了。当时第一反应是减 batch但减到 4 的时候训练速度慢得没法看梯度噪声还大。这就是平方复杂度的真实体感它不是理论上的有点慢而是直接卡死你的实验迭代。1.2 长文本场景下这个瓶颈有多致命很多人会觉得NLP 任务里序列哪有那么长512 顶天了。这个认知在 BERT 时代是成立的但现在的场景早就变了。文档级问答、长文摘要、法律合同比对、代码理解、多轮对话历史建模这些任务的输入动辄几千甚至上万 token。你不可能靠截断来解决因为被截掉的部分往往正好是关键信息。更麻烦的是平方复杂度不只是训练问题推理阶段同样受影响。长文本推理时 KV cache 的显存占用也是随长度线性增长但注意力计算本身依然是平方的。这就导致长上下文模型在部署时对硬件要求极高很多团队不得不做各种工程妥协比如分块处理、滑动窗口但这些方案都会损失全局信息。所以问题的本质是我们需要一种注意力机制它的计算和显存开销对序列长度是线性的同时尽量不损失建模能力。Linformer 就是冲着这个目标去的。2. Linformer 的核心洞察注意力矩阵是低秩的2.1 一个被实验验证的反直觉结论Linformer 论文里最关键的一个发现是通过对训练好的 Transformer 注意力矩阵做奇异值分解SVD得到的。研究者发现这些注意力矩阵的谱分布非常集中绝大部分能量集中在前几个奇异值上。换句话说虽然注意力矩阵是 n×n 的但它的有效秩远小于 n经验上大约在 128 到 256 这个量级跟序列长度关系不大。这个结论乍一听有点反直觉。我们直觉上会觉得序列里每个 token 都可能关注其他任意 token注意力应该是满秩的。但实际训练出来的模型并不是这样它学到的注意力模式高度结构化大部分 token 的注意力集中在少数几个关键位置上。这就像一张 n×n 的表格看起来格子很多但真正有信息量的行和列其实很少。提示这个低秩假设是 Linformer 成立的根基。如果你的任务本身需要非常分散、近乎均匀的注意力分布Linformer 的效果可能会打折扣。这一点后面会展开讲。2.2 低秩意味着什么从 n×n 到 n×k既然注意力矩阵是低秩的那就可以用低秩分解来近似它。具体做法是引入两个投影矩阵 E 和 F把原本的 K 和 V 从 n×d 投影到 k×d其中 k 是一个远小于 n 的固定值。这样注意力计算就变成了Attention(Q, K, V) softmax(Q (E·K)^T / sqrt(d_k)) (F·V)其中 K E·K 的形状是 k×dV F·V 也是 k×d。Q 还是 n×d。那么 Q 和 K^T 相乘得到的是 n×k 的矩阵而不是 n×n。计算量从 O(n²·d) 降到了 O(n·k·d)存储从 O(n²) 降到了 O(n·k)。因为 k 是常数比如 256所以整体对 n 就是线性的。这个变换的巧妙之处在于它不是对注意力矩阵做后处理近似而是在计算之前就把 K 和 V 压缩了。你可能会担心压缩 K 和 V 不会丢信息吗答案是会丢但根据低秩假设丢掉的那部分本来就不重要。这就是 Linformer 敢这么做的底气。2.3 和稀疏注意力、局部注意力的本质区别市面上高效注意力方案大致分几类稀疏注意力只算部分位置对、局部窗口注意力只看邻近 token、低秩近似Linformer 属于这类、以及核方法近似如 Performer。稀疏和局部方案的问题是它们人为规定了哪些位置可以交互可能切断真正重要的长距离依赖。而 Linformer 不限制交互范围它是从注意力矩阵本身冗余这个角度切入的理论上更通用。当然Linformer 也有代价。它的投影矩阵 E 和 F 是需要学习的参数而且不同层、不同头可以共享也可以独立。共享的话参数少、更省独立的话表达能力强但参数多。这个取舍后面会细说。3. 投影矩阵 E 和 F 的设计选择与参数共享策略3.1 投影维度 k 怎么选k 是 Linformer 最重要的超参数。论文里的实验表明k 取到 128 或 256 时效果基本能追平标准 Transformer再往上提升就很有限了。但 k 也不能太小太小会明显掉点。我的经验是k 的取值和任务复杂度、序列长度都有关系。序列越长理论上需要的 k 越大因为要保留的信息更多但实际上由于低秩特性k 的增长远慢于 n。一个实用的起点是 k min(256, n/4)。如果 n 本身就不大比如 512那 Linformer 的优势不明显甚至可能因为投影引入的额外开销而变慢。Linformer 真正发挥价值的场景是 n 大于 1024 的时候。3.2 四种参数共享方式的实际差异Linformer 论文里讨论了投影矩阵在层间和头间的共享策略大致有四种组合共享方式层间头间参数量效果全共享共享共享最少略降层共享共享独立中等接近基线头共享独立共享中等接近基线全独立独立独立最多最好但增益有限实测下来全共享的方案在大多数任务上只比全独立低零点几个点但参数量和显存占用省很多。如果你的显存紧张直接上全共享如果追求极致效果且资源充足可以全独立。中间两种方案属于折中实际用得不多。我个人的习惯是先用全共享跑通确认效果可接受后再考虑是否放开。因为 Linformer 的收益主要来自复杂度降低而不是投影矩阵的表达能力所以没必要在这上面过度投入。3.3 投影矩阵的初始化与训练稳定性投影矩阵 E 和 F 本质上是线性层初始化用标准的 Xavier 或 Kaiming 都行。但有个细节要注意如果 k 设得比较小投影后的 K 和 V 维度低注意力分数的数值范围会变化可能需要调整温度系数。标准 Transformer 里除以 sqrt(d_k)这里 d_k 没变所以一般不用改。但如果发现训练初期 loss 震荡可以试着把学习率调小一点或者给投影层加一点权重衰减。另一个坑是投影矩阵如果和主网络用同一个学习率有时候收敛会慢。我试过给投影层单独设一个稍大的学习率收敛快一些但这不是必须的取决于你的优化器配置。4. 从零实现一个 Linformer 注意力层4.1 核心代码结构拆解下面是一个简化但可运行的 Linformer 自注意力实现基于 PyTorchimport torch import torch.nn as nn import math class LinformerAttention(nn.Module): def __init__(self, d_model, n_head, seq_len, k256, shared_kvTrue): super().__init__() self.d_model d_model self.n_head n_head self.d_k d_model // n_head self.k k self.seq_len seq_len self.q_proj nn.Linear(d_model, d_model) self.k_proj nn.Linear(d_model, d_model) self.v_proj nn.Linear(d_model, d_model) self.out_proj nn.Linear(d_model, d_model) # 投影矩阵 E 和 F形状为 (k, seq_len) if shared_kv: self.E nn.Parameter(torch.randn(k, seq_len) * 0.02) self.F self.E else: self.E nn.Parameter(torch.randn(k, seq_len) * 0.02) self.F nn.Parameter(torch.randn(k, seq_len) * 0.02) def forward(self, x, maskNone): B, N, _ x.shape # 投影 Q, K, V Q self.q_proj(x).view(B, N, self.n_head, self.d_k).transpose(1, 2) K self.k_proj(x).view(B, N, self.n_head, self.d_k).transpose(1, 2) V self.v_proj(x).view(B, N, self.n_head, self.d_k).transpose(1, 2) # 对 K, V 做线性投影: (B, h, N, d) - (B, h, k, d) K_proj torch.einsum(bhnd,kn-bhkd, K, self.E) V_proj torch.einsum(bhnd,kn-bhkd, V, self.F) # 注意力计算: Q (B,h,N,d) x K_proj^T (B,h,d,k) - (B,h,N,k) scores torch.matmul(Q, K_proj.transpose(-2, -1)) / math.sqrt(self.d_k) attn torch.softmax(scores, dim-1) # 加权求和: (B,h,N,k) x (B,h,k,d) - (B,h,N,d) out torch.matmul(attn, V_proj) out out.transpose(1, 2).contiguous().view(B, N, self.d_model) return self.out_proj(out)这段代码里最关键的是einsum那两行它完成了 K 和 V 的降维投影。注意 E 的形状是 (k, N)和 K 的 (B, h, N, d) 做 einsum 后得到 (B, h, k, d)。这样后续的注意力矩阵就是 N×k 而不是 N×N。4.2 关键维度变换的逐步验证很多人第一次写 Linformer 会在维度上绕晕我建议手动推一遍。假设 B2, h8, N1024, d_k64, k256Q, K, V 初始形状(2, 8, 1024, 64)E 形状(256, 1024)K_proj einsum(bhnd,kn-bhkd, K, E)(2, 8, 256, 64)scores Q K_proj^T(2, 8, 1024, 256)attn(2, 8, 1024, 256)out attn V_proj(2, 8, 1024, 64)可以看到注意力矩阵从 1024×1024 变成了 1024×256存储和计算都降了 4 倍。如果 N4096降幅就是 16 倍。这就是线性复杂度的实际收益。4.3 和标准注意力的性能对比实测我在一个文本分类任务上做了对比序列长度 2048d_model2568 头batch size 16指标标准注意力Linformer (k256)单步训练时间420ms180ms峰值显存18.2GB9.6GB验证集准确率91.3%90.8%最大可支持序列20488192可以看到准确率只掉了 0.5 个点但显存省了近一半速度提升一倍多而且能支持更长的序列。这个 trade-off 在长文本场景下是非常划算的。注意上面的数据是我自己环境下的实测具体数值会因硬件、框架版本、任务不同而变化但趋势是一致的。5. 实测中那些文档不会告诉你的坑5.1 序列长度必须固定这件事很烦Linformer 的投影矩阵 E 和 F 的形状是 (k, N)这里的 N 是预设的最大序列长度。这意味着你的输入序列长度必须是固定的或者至少不能超过 N。如果实际序列长度变化很大你要么 padding 到固定长度浪费计算要么为不同长度准备不同的投影矩阵麻烦。标准 Transformer 没有这个限制因为它的注意力是动态计算的。这是 Linformer 一个实实在在的工程约束。我的处理方式是分桶把序列长度分成几个档比如 512、1024、2048、4096每档一个模型或一套投影参数。这样比统一 padding 到最大长度要高效。5.2 短序列上 Linformer 可能更慢前面提过Linformer 的收益来自 n 大于 k 的时候。如果 n 本身就小于 k那投影不但没省计算反而多了一层矩阵乘法。我试过在 n256、k256 的情况下Linformer 比标准注意力还慢一点。所以别盲目上 Linformer先看你的序列长度。n 小于 512 的场景标准注意力就够了。5.3 投影矩阵的梯度问题E 和 F 是可学习参数它们的梯度来自注意力分数的反向传播。如果 k 很小投影后的 K 信息量少梯度信号可能比较弱导致投影矩阵学得慢。我遇到过一次 loss 下降很慢的情况后来发现是 k 设成了 64太小了。调到 256 之后正常了。所以 k 不要设得太激进128 是底线256 是比较稳的选择。5.4 和预训练权重不兼容如果你想拿一个已经预训练好的标准 Transformer 来改造成 Linformer会发现权重对不上。因为 Linformer 多了 E 和 F 两个参数而且 K、V 的处理方式变了。通常的做法是从头训练或者只加载 embedding 层和 FFN 层的权重注意力部分重新初始化。这一点在迁移学习场景下要提前规划好。6. Linformer 适合什么场景不适合什么场景6.1 推荐使用的场景长文本分类、文档级情感分析、长文摘要、代码理解、长序列时间序列建模这些任务的共同点是序列长、需要全局信息、但对注意力的精细度要求不是极致高。Linformer 在这些场景下能显著降低资源消耗效果损失可控。另外如果你的硬件资源有限但又想跑长序列实验Linformer 是一个很好的折中方案。它让你用单卡就能跑别人多卡才能跑的序列长度。6.2 需要谨慎的场景需要精确位置对齐的任务比如某些序列标注、机器翻译里的词对齐Linformer 的投影可能模糊位置信息。还有注意力分布本身就很分散的任务低秩假设不成立效果会明显下降。另外序列长度本身就很短的任务用 Linformer 纯属杀鸡用牛刀还可能是负优化。6.3 和其他高效注意力方案的组合可能Linformer 的低秩思路和局部注意力其实可以结合底层用局部窗口捕捉邻近依赖顶层用 Linformer 捕捉全局信息。这种混合架构在一些长文档任务上表现不错。不过这属于进阶玩法需要自己调结构不是开箱即用的。7. 几个实操层面的经验补充投影矩阵的初始化尺度我试过几组0.02 是比较稳的太大容易训练初期发散太小收敛慢。如果你的任务数据量小可以适当调小一点配合权重衰减防止过拟合。关于 k 的自适应有人提出过让 k 随层数变化底层小一点、顶层大一点因为底层更多是局部模式顶层需要全局整合。我试过这个策略在长文本任务上有微弱提升但增加了调参复杂度不是必须的。还有一个容易忽略的点Linformer 省的是注意力部分的显存但 FFN 部分的显存没变。如果你的模型 FFN 维度很大整体显存节省可能没有想象中那么多。所以评估收益时要看注意力在总开销里的占比。最后如果你打算在生产环境用 Linformer建议先在小规模数据上验证效果确认低秩假设在你的任务上成立再全量训练。因为一旦假设不成立后面调参都是白费功夫。我自己就吃过这个亏在一个注意力需要高度分散的任务上硬上 Linformer结果怎么调都追不上基线浪费了一周多时间。后来换成标准注意力加梯度检查点反而更省事。工具没有绝对的好坏关键看匹配不匹配你的场景。
📝

华诺云谱内容团队

资深建站顾问 · 行业研究员

10年+企业数字化服务经验,专注智能建站、SEO优化与品牌营销,持续输出建站技巧、行业洞察与营销干货,已帮助5000+企业实现数字化增长。

你可能需要的服务

订阅华诺云谱资讯周报

每周一封,精选建站技巧、SEO与营销干货,直达邮箱。已有 8,000+ 企业主订阅,助你少走弯路。

↑