高斯混合注意力:通过概率潜在路由实现线性时间序列混合
摘要
本文介绍了高斯混合注意力(Gaussian Mixture Attention,GMA),这是一种概率性注意力机制,它用通过学习得到的高斯混合组件进行路由,取代了显式的成对查询-键比较,从而在序列长度上实现了线性时间复杂度。实验表明,在长上下文任务中,它凭借固定K的线性内存扩展展现出了有竞争力的性能。
查看缓存全文
缓存时间: 2026/06/18 05:39
# 高斯混合注意力:通过概率潜在路由实现线性时间序列混合 来源:https://arxiv.org/html/2606.18283 Yongchao Huang¹¹[email protected] Raza²²[email protected] (16/05/2026) ###### 摘要 标准点积注意力的密集 token-to-token 交互模式仍是缩放 Transformer 架构以适应长上下文的核心瓶颈。我们提出 **高斯混合注意力(GMA)** ,这是一种概率性的注意力式序列混合器,它用通过 \(K\) 个学习到的高斯混合组件进行路由,取代了显式的逐对查询-键比较。查询和键被映射到共享潜在路由空间上的后验 **责任向量**;它们的重叠定义了隐式的责任空间亲和性,而值则被写入并读取自一个 \(K\) 槽位的潜在记忆。通过利用矩阵乘法的结合性,GMA 避免了具体化生成的 \(N \times N\) 亲和矩阵,而是使用两个责任矩阵,对于固定的 \(K\),其主导激活存储规模为 \(\mathcal{O}(NK)\) 而非 \(\mathcal{O}(N^2)\)。我们制定了双向和因果变体的 GMA,提供了高斯混合组件的端到端可微参数化,并分析了其责任调制梯度结构、约束非负低秩亲和性解释以及局部路由稳定性。实验上,GMA 展现了预期的固定 \(K\) 线性内存缩放,并在长上下文分类任务中与注意力式基线模型竞争。同时,因果 GMA 在 WikiText-103 上优于测试的线性/随机特征注意力变体,但在当前实现中仍落后于优化的因果 SDPA 和 Mamba。对学习到的责任的分析进一步显示了广泛的组件使用以及与表层形式 token 类别的一定程度对齐,支持 GMA 作为一种概率性、可解释、固定 \(K\) 的线性时间注意力式替代方案,而非对优化 softmax 注意力或状态空间模型的通用替代。 ## 1 引言 **Transformer** 架构在语言、视觉和多模态学习领域取得了卓越性能,这主要归功于点积多头注意力(MHA)的表示能力和并行性(Vaswani et al., 2017;Dosovitskiy et al., 2021;Radford et al., 2021)。然而,标准自注意力计算序列中所有 \(N\) 个 token 之间的成对交互。对于一个注意力头,设 \(Q \in \mathbb{R}^{N \times d_k}\) 为查询矩阵,\(K_{\mathrm{att}} \in \mathbb{R}^{N \times d_k}\) 为键矩阵,\(V \in \mathbb{R}^{N \times d_v}\) 为值矩阵。这里 \(d_k\) 是用于计算点积分数的查询/键通道维度,而 \(d_v\) 是被聚合向量的值维度。缩放点积注意力计算: \[ O = \operatorname{softmax}\left( \frac{Q K_{\mathrm{att}}^\top}{\sqrt{d_k}} \right) V, \qquad O \in \mathbb{R}^{N \times d_v}. \tag{1} \] 中间分数矩阵 \(Q K_{\mathrm{att}}^\top \in \mathbb{R}^{N \times N}\) 包含所有查询和键位置之间的 token-to-token 分数。尽管最终输出的维度为 \(O \in \mathbb{R}^{N \times d_v}\),但计算密集注意力需要形成或隐式表示所有位置对之间的交互。在标准显式形式中,这导致 \(\mathcal{O}(N^2)\) 的注意力分数存储,以及 \(\mathcal{O}(N^2 d_k + N^2 d_v)\) 的分数和值乘法的算术运算。当 \(d_k\) 和 \(d_v\) 被视为固定值时,这给出了熟悉的序列长度 **二次缩放**。这种二次依赖性使得长上下文建模成本高昂,并促使在长文档处理、字节级分类、高分辨率视觉、基因组学以及自回归语言建模等场景中寻求高效替代方案(Beltagy et al., 2020;Zaheer et al., 2020;Tay et al., 2020;Wang et al., 2020;Choromanski et al., 2021;Katharopoulos et al., 2020;Avsec et al., 2021)。 大量工作试图通过稀疏注意力模式、低秩投影、核近似、优化的精确注意力核以及循环或状态空间序列模型来减少二次注意力瓶颈(Tay et al., 2022;Beltagy et al., 2020;Zaheer et al., 2020;Wang et al., 2020;Katharopoulos et al., 2020;Choromanski et al., 2021;Dao et al., 2022;Gu et al., 2022;Gu and Dao, 2024)。稀疏注意力方法,如 **Longformer** 和 **BigBird**,通过限制每个 token 只关注本地、全局或结构化位置的子集来减少计算量(Beltagy et al., 2020;Zaheer et al., 2020)。低秩方法如 **Linformer** 通过沿序列维度学习投影来压缩注意力矩阵(Wang et al., 2020)。基于核的方法用特征映射构造替换 softmax 注意力核,包括 **Linear Transformer** 中的确定性正特征映射以及 **Performer** 中的随机特征近似(Katharopoulos et al., 2020;Choromanski et al., 2021)。IO 感知实现如 **FlashAttention** 通过硬件感知分块降低了内存流量,同时保留了精确的 softmax 注意力(Dao et al., 2022)。最近,结构化状态空间和选择性状态空间模型,包括 **S4** 和 **Mamba**,通过用循环或状态空间动力学替代显式注意力,实现了线性或近线性序列建模(Gu et al., 2022;Gu and Dao, 2024)。这些方法展现了不同的权衡:它们可能非常高效,但其内部路由结构通常不如原始的 softmax 注意力权重那样直接具有概率性或易于解释为 token-to-token 分布。 在这项工作中,我们提出 **高斯混合注意力(GMA)** ,这是一种点积注意力的概率性替代方案,它将序列混合重新概念化为 **潜在责任路由**。GMA 并非显式计算所有成对 token-to-token 相似度,而是在一个投影的 **路由表示空间** 中引入 \(K\) 个学习到的高斯混合组件。查询和键表示被映射到这些组件上的后验责任向量。键责任将值写入一个潜在记忆 \(\tilde{V} \in \mathbb{R}^{K \times d_v}\) 以及一个组件级的归一化器 \(Z \in \mathbb{R}^K\),而查询责任则从这个归一化的潜在记忆中读取以产生 token 级输出。因此,GMA 用两个非负责任矩阵 \(\Gamma^Q, \Gamma^K \in \mathbb{R}^{N \times K}\) 替代了显式的 \(N \times N\) token-to-token 注意力矩阵。尽管代数上诱导的亲和性 \(\Gamma^Q (\Gamma^K)^\top\) 仍然是一个 \(N \times N\) 矩阵,但 GMA 并未具体化它。相反,它利用矩阵乘法的结合性,先计算键-值潜在记忆 \((\Gamma^K)^\top V_X\),然后乘以 \(\Gamma^Q\)。对于固定的 \(K\),这给出了主导激活存储的线性于 \(N\) 的缩放,同时保留了归一化的注意力式路由解释。 GMA 的一个关键动机是,责任矩阵不仅仅是计算中间产物:它们是可分析的概率对象。混合组件的边际使用可以诊断潜在路由空间是被广泛使用还是坍缩到一小部分组件,而由 \(z_i = \arg\max_k \gamma_{i,k}\) 导出的硬分配可以与 token 类别或其他注释进行比较。这提供了一种可解释性手段,这在随机特征注意力近似或隐式循环/状态空间隐藏动力学中不那么直接。我们随后的分析表明,学习到的 GMA 责任使用了大多数可用组件,并表现出与表层形式 token 类别的一定程度对齐,尽管这些组件不应被解释为清晰的语义类别(第 5.4 节)。 我们的贡献和发现如下: 1. 我们提出了 **高斯混合注意力(GMA)** ,一种基于归一化责任的序列混合器,它用通过学习到的高斯混合组件进行路由,取代了显式的 token-to-token 注意力。 2. 我们推导了双向和因果 GMA。因果变体使用前缀潜在记忆和前缀归一化器,使得位置 \(i\) 的自回归预测仅依赖于位置 \(j \le i\),同时保持固定的 \(K\) 线性于序列长度的缩放。 3. 我们分析了 GMA 的优化和表示结构,包括责任调制梯度、约束非负低秩亲和性解释,以及在有限输入和方差下限下的局部 Lipschitz 连续性。 4. 我们在 4 个实验设置中评估 GMA:受控系统分析(表 1)、Long Range Arena (LRA) 长上下文分类(Tay et al., 2020)(表 2)、WikiText-103 自回归语言建模(Merity et al., 2016)(表 3),以及潜在责任可解释性分析(表 4;图 2–3)。结果表明,GMA 展现了预期的线性内存缩放,并在我们流程中评估的注意力式基线模型中取得了有竞争力的性能。在 LRA 上,它在这些注意力式基线模型中取得了最强的平均性能(表 2)。在 WikiText-103 上,因果 GMA 优于 Linear Transformer(Katharopoulos et al., 2020)和 Performer(Choromanski et al., 2021),尽管优化的因果 SDPA(Vaswani et al., 2017)和 Mamba(Gu and Dao, 2024)仍然更强(表 3)。最后,学习到的 GMA 责任使用了大多数可用组件,并表现出与表层形式 token 类别的一定程度对齐(表 4;图 2–3)。 5. 我们讨论了 GMA 框架的未来扩展,包括优化的和混合的 GMA 实现、交叉注意力和多模态路由、用于自适应组件加权的贝叶斯和狄利克雷过程变体,以及概率性混合专家路由。 ## 2 相关工作 #### 注意力作为学习到的兼容性。 注意力机制可以广泛地视为计算 **查询表示** 与 **上下文表示** 之间 **兼容性** 的方法,然后利用得到的权重来聚合值。早期的神经编码器-解码器模型使用学习到的对齐机制来聚焦解码到相关的源位置(Bahdanau et al., 2016;Luong et al., 2015)。**Transformer** 通过 \(Q K_{\mathrm{att}}^\top\) 计算 token-to-token 分数并应用行式 softmax 归一化,将 **缩放点积注意力** 确立为主导形式(Vaswani et al., 2017),如公式 (1) 所示。这种设计极具表达力且并行性高:每个 token 可以使用批量矩阵乘法与所有其他 token 形成查询-键分数,值聚合也可以通过现代加速器上的密集线性代数来计算。这种并行结构是 Transformer 架构成功的主要原因,但它也需要形成或隐式表示一个 \(N \times N\) 的交互模式,从而产生了熟悉的序列长度二次缩放。从更广泛的设计视角来看,点积只是可能的兼容性函数之一。注意力分数也可以基于学习到的加法分数、核函数、距离、散度、稀疏模式或潜在路由结构。**GMA** 遵循这种更广泛的视角,用概率责任空间中的兼容性取代了直接的成对查询-键比较。关于相似度、距离、散度和潜在空间注意力设计的更明确分类法见附录 G。 #### 稀疏和结构化注意力。 高效 Transformer 研究的一个重要方向是减少需要比较的查询-键对数量。**Sparse Transformer** 使用结构化的稀疏模式,比密集注意力更高效地生成长序列(Child et al., 2019)。**Reformer** 用局部敏感哈希注意力取代了密集点积注意力,并使用可逆残差层来减少激活存储(Kitaev et al., 2020)。**Longformer** 结合了局部滑动窗口模式与任务驱动的全局注意力,实现了序列长度的线性缩放,并在长文档任务上表现强劲(Beltagy et al., 2020)。**BigBird** 结合了局部、随机和全局注意力模式,获得了线性稀疏注意力,同时保留了全注意力的重要理论性质,包括在其稀疏模式下的通用逼近和图灵完备性结果(Zaheer et al., 2020)。这些方法保留了 token-to-token 注意力,但限制了注意力图。因此,它们的效率依赖于设计或采样的稀疏模式。**GMA** 的不同之处在于,它通过潜在混合组件保持密集全局信息的可用性,而不是选择一组稀疏的 token 对。 #### 低秩和地标近似。 第二个方向是用低维表示来近似注意力矩阵。**Linformer** 认为自注意力可以用低秩矩阵近似,并通过沿序列维度投影来降低注意力的复杂度至线性。
相似文章
GQLA: 面向硬件自适应大语言模型解码的分组查询潜在注意力
GQLA 提出了对多头潜在注意力(MLA)的极小修改,在相同训练权重上同时暴露 MQA 吸收路径和 GQA 路径,从而无需重新训练即可实现硬件自适应解码。该方法压缩 KV 缓存并支持张量并行性,通过将 LLaMA-3-8B 从 GQA 转换为 GQLA 得到验证。
Hierarchical Global Attention (HGA)
Hierarchical Global Attention (HGA) 是一种可直接替换预训练长上下文Transformer中密集因果注意力的方法。它采用分层两级路由机制,使得能够对一个小规模路由工作集进行精确注意力计算,从而允许像 Qwen3-30B 这样的模型在单个 RTX 5090 上以64K上下文运行,且质量损失极小。
动态线性注意力
本文提出DLA,一种用于多状态线性注意力的动态内存建模框架,它能根据令牌信息变化自适应地合并状态,并维护固定大小的状态缓存,从而在无需标准注意力二次复杂度的前提下实现更好的长上下文表示。
MISA:用于长上下文大语言模型推理的索引器混合稀疏注意力机制
本文介绍了 MISA,这是一种将混合专家(MoE)方法应用于稀疏注意力机制中索引器头部的技术,在保持性能的同时显著降低了长上下文大语言模型推理的计算成本。
从格劳伯轨迹中学习不依赖混合的高斯图模型
本文提出了一种多项式时间算法,用于从单条格劳伯动力学轨迹中学习高斯图模型的结构,其轨迹长度保证不依赖于混合时间。