迈向线性的关键:分析驱动的 Transformer 线性化

The Key to Going Linear: Analysis-Driven Transformer Linearization

arXiv: 2607.07706v1

论文信息

标题: The Key to Going Linear: Analysis-Driven Transformer Linearization

作者: Anna Kuzina, Paul N. Whatmough, Babak Ehteshami Bejnordi

发布日期: 2026-07-08

arXiv ID: 2607.07706v1

PDF 链接: 下载 PDF

3 分钟速览

  • 研究问题:如何在后训练线性化预训练 Transformer 时,确定哪种线性状态更新机制(如门控线性注意力 GLA、门控 Delta 网络 GDN)能最忠实地逼近原始 softmax 注意力,并补足由此产生的性能缺口。
  • 核心方法:从一阶泰勒展开与投影近似出发,证明 softmax 注意力天然蕴含键依赖的秩-1 正交投影,而 GDN 的 delta 更新恰好实现了这一几何修正;在严格冻结 backbone、仅训练新增参数的设定下,系统对比多种线性机制,进而引入滑动窗口、sink tokens 与短卷积等结构化补偿。
  • 关键结果:基于 GDN 的线性化 LLaMA3.1-8B 模型,在 64 token 缓存预算下 5-shot MMLU 达到 59.17,显著优于 LoLCATs 的 54.88 和 Liger-GLA 的 46.9;将预算扩至 128 token 后 MMLU 升至 63.22,与全注意力模型(65.25)的差距缩至仅约 3%(见表 4)。
  • 主要局限:一阶近似假设键向量单位化且高度集中,在序列前缀处可能不成立;固定预算的 sink token 策略对多条目检索任务仍显不足,且未结合自适应 token 选择机制。
  • 适合读者:关注大模型低成本长上下文推理、线性注意力机制设计,或寻求高效部署 Transformer 模型的工程师与研究人员。

论文背景和研究动机

因果自注意力的计算量与序列长度呈平方增长,并导致键值缓存(KV cache)急剧膨胀,这严重制约了预训练 Transformer 在长上下文场景下的推理效率。近年来兴起的后训练线性化(post hoc linearization)试图将现成的全注意力模型转换为线性时间架构,而无需从零预训练。然而,现有方案通常同时叠加低秩适配(LoRA)、滑动窗口注意力(SWA)、混合路由、蒸馏等多种干预,哪些组件真正对保持模型质量起决定性作用,始终缺乏清晰的图景。

本文由此提出一个核心问题:在严格冻结原有权重、只训练新增参数的 “纯净” 条件下,哪一种线性状态更新机制最能逼近预训练 softmax 注意力?论文把注意力替换问题拆解为独立的近似能力测试,并采用分析驱动的方式推导出 softmax 的动态特性,从而明确 Delta 式更新为何优于纯门控累积。在此基础上,论文逐步加回滑动窗口、sink tokens 等结构化补偿,使线性化模型的性能充分逼近全注意力基线。

核心方法和技术细节

线性注意力替换机制

论文考察两类线性注意力:基于归一化核的 Hedgehog,以及无显式归一化的门控线性注意力(GLA)和门控 Delta 网络(GDN)。GLA 通过输入相关的对角遗忘门进行状态累积:

St=GtSt−1+kt⊤vtS_t = G_t S_{t-1} + k_t^\top v_t

而 GDN 额外引入了键依赖的秩-1 更新:

St=(I−βtktkt⊤)GtSt−1+βtkt⊤vtS_t = (I - \beta_t k_t k_t^\top) G_t S_{t-1} + \beta_t k_t^\top v_t

所有实验均在冻结 backbone下进行:预训练的 Q、K、V 投影和 RMSNorm 全部固定,只训练机制专属的新增参数(如门投影、beta 标量等)。

从 softmax 到 delta 规则的一阶近似

论文通过三步引理揭示 softmax 与 delta 更新的内在联系。 首先,对 softmax 在零点做一阶泰勒展开,得到

Pifull≈1t+1tdqt⊤(ki−kˉ)P_i^{\text{full}} \approx \frac{1}{t} + \frac{1}{t\sqrt{d}} q_t^\top (k_i - \bar{k})

(引理 1)。 其次,将 “减去均值键” 替换为 “投影到均值键正交补”,在键归一化且方向集中的假设下引入可容忍的误差(引理 2)。 最后,利用秩-1 乘积展开,证明一系列 I−βjkjkj⊤I - \beta_j k_j k_j^\top 的乘积可近似为 I−c ΠkˉI - c\, \Pi_{\bar{k}}(引理 3),从而将投影形式改写为:

Pifull≈1tdq⊤∏j(I−βjkjkj⊤)ki+1tP_i^{\text{full}} \approx \frac{1}{t\sqrt{d}} q^\top \prod_j (I - \beta_j k_j k_j^\top) k_i + \frac{1}{t}

这与 GDN 的有效注意力权重形式(公式 9)高度吻合。相比之下,GLA 的隐式权重仅依赖查询无关的对角衰减矩阵,缺少这种键依赖的逐成分抑制。论文还提出一种键门控变体 kGLA,但仍受限于对角近似,无法完全匹配秩-1 更新的表达能力。

填补近似缺口的实用设计

理论分析揭示两种误差来源:

  1. 序列前缀偏差:泰勒展开包含 1/t1/t 的均匀偏置,开头若干 token 的近似误差较大。
  2. 键不集中:当键向量未能充分对齐到单一主方向时,投影残差增大,近似质量下降。

针对前缀偏差,论文采用sink tokens策略:将固定的 64 token 注意力预算分配为 8 个 sink token(始终保持全注意力)和 56 个滑动窗口 token。线性路径负责窗口外的其余 token,从而显著缓解开头难拟合的问题。进一步,可为 Q/K/V 投影引入短卷积或 LoRA,以弥补线性状态更新在表达能力上的固有不足。

创新点和贡献

  • 分析驱动的线性化视角:首次通过一阶近似严格证明,delta 规则实现的键依赖秩-1 修正正是 softmax 动态的核心,而纯门控线性注意力缺失这一特性。
  • 纯净实验设计:在冻结 backbone 的严格设定下对比多种线性机制,排除了其他干预的混杂效应,明确了 GDN 在近似能力上的优势。
  • 结构化补偿的轻量化方案:发现 sink tokens 能够以极小预算(仅 8 个 token)大幅提升 MMLU 得分,配合滑动窗口和短卷积即可使线性化模型接近全注意力性能(见表 2 和表 3),且超参数迁移性强。
  • 规模化验证:在 LLaMA3.1-8B、Qwen3-8B 乃至 32B 模型上展示了线性化性能随模型尺寸稳定增长(见图 4),并在 MMLU 上超越 LoLCATs、Liger-GLA 等已有后训练线性化基线(见表 4)。
  • 长上下文能力:在扩展的 RULER 基准上,GDN 线性化模型在 896 token 缓存预算下平均得分 47.52,与需要复杂自适应缓存策略的 STILL(47.9)相当,但实现更简洁(见表 9)。

实验结果分析

在纯净替换阶段,GDN 的层间归一化 MSE 远低于 GLA 和 Hedgehog,尤其在中间层优势明显(图 2b),且下游共同推理得分最高(表 1)。键集中度与投影残差的分析验证了理论假设:线性近似最差的区域恰好是键分散、残差大的序列开头(图 3),这直接启发了 sink tokens 的引入。

加入滑动窗口后,所有线性机制性能均大幅提升,但 GDN 在 MMLU 上仍以 51.57 领先(表 2)。当把缓存预算从纯滑动窗口变为 56 滑动 + 8 sink,GDN 的 MMLU 进一步飙升至 58.88,同时 Lambada 准确率从 41.92 升至 68.82,说明 sink tokens 对保持关键前缀信息至关重要。

在 Q/K/V 适应性方面,短卷积在 MMLU 上仅比全注意低约 6 个百分点,但 Lambada 差距缩小至约 4.5 个百分点(表 3);LoRA 虽然所需学习率更低,但也能达到类似效果。值得注意的是,短卷积使 GLA 与 GDN 之间的差距几乎消失,暗示额外表示能力可以部分弥补键依赖更新的缺失。

最终与先前工作相比,论文的 GDN 方案仅用 10M token 训练、64 token 缓存预算就在 MMLU 上获得 59.17,而 LoLCATs 和 Liger-GLA 分别只有 54.88 和 46.9(表 4)。当缓存预算扩至 128 token,MMLU 升至 63.22,与 Lizard(61.2)和 LoLA(57.6)相比同样具有竞争力。长上下文 S-NIAH 中,GDN 的检索准确率随缓存预算增加稳步上升,与需要自适应选择的 STILL 差距不大(表 8),表明固定预算策略已有一定实效。

实践建议

对于希望将预训练大模型部署为线性复杂度推理的团队,本文提供了一套经过验证的轻量级线性化方案:

  1. 选用 GDN 作为核心线性替换:基于论文的一阶近似理论和实验证据,GDN 的键依赖秩-1 更新与原 softmax 的几何特性最为契合,即使在冻结 backbone 下也能保持较低的 MSE 和较好的下游性能。
  2. 固定缓存预算分配:设置总注意力缓存(如 64 或 128),将其中一小部分(如 8 个)作为 sink tokens 保留全注意力,其余分配为滑动窗口。这种简单的划分无需任何内容感知选择算法,即可显著提升 MMLU 等知识基准(表 2)。
  3. 对 Q/K/V 投影施加短卷积:在 RoPE 之后、L2 归一化之前加入核大小为 8 的短卷积,可额外恢复约 4-6 个百分点的 MMLU 和 Lambada 准确率(表 3),且超参数通用性良好,成本增加极小。
  4. 训练与推理实现:所有线性状态更新均可基于 Flash Linear Attention(FLA)库的高效 Triton 内核实现。在单个 H100 GPU 上,10M token 训练约需 1 小时(GDN),推理时内存占用显著低于全注意力,可支持 32k 甚至更长上下文(图 2a)。
  5. 向更大模型迁移:论文已在 Qwen3 0.6B 至 32B 参数模型上验证线性化方案,各尺寸模型性能成比例下降(图 4),表明该方法具有可预测的扩展性,可直接应用于同类 Transformer 架构而无需大量调参。

上述配方无需知识蒸馏、不修改原有权重,仅引入数百万新增参数,既降低了工程复杂度,又可在常见 Benchmark 上获得接近原模型的性能,是当前后训练线性化领域一个兼具理论深度和实操价值的成熟选项。若应用场景对多条目并行检索有更高要求,可在固定预算基础上结合轻量级 token 选择策略,作者将此方向列为未来工作。