【Transformer 与注意力机制】56|状态空间模型:Mamba、S4 的线性复杂度路径

💡 原文中文,约16200字,阅读约需39分钟。
📝

内容提要

本文介绍状态空间模型(SSM)如何通过固定大小状态压缩历史,解决Transformer长序列瓶颈。S4用HiPPO矩阵实现长程记忆,Mamba引入选择性机制按内容更新状态,并用并行扫描保持训练效率。推理时SSM显存O(1)优于KV Cache,但在精确检索任务上受容量限制。SSD揭示SSM与注意力对偶,混合架构如Jamba平衡两者优势。

🔎

延伸解读

状态压缩与显式查表:两种信息取舍的代价

本文的核心对比是状态压缩与显式查表。注意力机制不压缩上下文,显式保存所有历史K/V,因此检索精确但显存和计算随序列长度增长;SSM用固定大小状态折叠历史,显存恒定但必须决定保留什么、丢弃什么。理解这一取舍有助于判断不同任务适合哪种架构:需要精确检索的任务可能更适合注意力,而长序列、高吞吐场景SSM更有优势。

选择性机制:从LTI到内容感知的关键一步

S4的LTI特性使其无法根据输入内容调整状态更新,导致在选择性复制等任务上表现不佳。Mamba通过让Δ、B、C依赖输入,引入选择性机制,使模型能按内容决定记忆或遗忘,从而在语言建模和合成任务上显著提升性能。这一改变打破了卷积等价性,但通过并行扫描算法保持了训练效率,是SSM走向实用的关键。

推理优势与检索短板:SSM的适用边界

SSM推理时状态显存O(1),相比KV Cache的O(n)在长上下文和高并发场景下吞吐优势明显,但固定状态容量限制了精确复制和检索能力。Jelassi等人的理论证明,状态大小固定的模型在复制任务上存在数学上界,即使增大状态也只能线性提升。因此,SSM更适合非检索密集型任务,而混合架构如Jamba通过保留少量注意力层来弥补这一短板。

Q&A

状态空间模型(SSM)是如何解决Transformer长序列瓶颈的?

SSM通过固定大小的状态向量来压缩历史信息,而不是像Transformer那样显式保存所有历史token的K/V。这样推理时状态大小与序列长度无关,显存占用为O(1),避免了KV Cache的线性增长和二次复杂度。

S4模型是如何实现长程记忆的?

S4使用HiPPO矩阵作为状态转移矩阵,该矩阵将输入历史投影到正交多项式基上,使得状态更新在数学上等价于对历史的有原则压缩,从而解决了朴素SSM的梯度消失/爆炸问题,实现了对上万步之前信息的记忆。

Mamba的选择性机制具体做了什么?

Mamba让状态更新中的参数Δ、B、C都成为当前输入的函数,使得模型能够根据输入内容决定是重置状态、聚焦当前输入,还是保持旧状态、忽略当前输入。这打破了S4的线性时不变(LTI)限制,实现了按内容取舍的状态更新。

parallel scan是如何让递归既线性又可并行训练的?

parallel scan利用仿射变换组合的结合律,将递归计算组织成平衡二叉树,自底向上合并,将串行深度从O(L)降到O(log L),同时总工作量保持O(L)。这样既保持了线性复杂度,又能在GPU上并行执行。

SSM在推理时相比Transformer的KV Cache有什么优势?

SSM的推理状态大小固定为O(N),与序列长度无关,因此显存占用不随序列增长,而Transformer的KV Cache随序列长度线性增长。在长上下文和高并发场景下,SSM可以节省显存,支持更大的batch size,从而提升吞吐量。

SSM在哪些任务上会输给Transformer?

SSM在需要精确复制或检索上下文信息的任务上会输给Transformer。Jelassi等人的理论和实验表明,固定大小状态的模型在复制任务上存在容量上界,无法像Transformer那样精确取回任意历史token,即使语言建模困惑度更低。

SSD和混合架构(如Jamba)是如何结合SSM和注意力机制的?

SSD证明了SSM的线性递归和结构化掩码注意力是同一类半可分矩阵的两种分解,为混合架构提供了理论基础。Jamba按约每8层放1层注意力、其余用Mamba层的比例交替堆叠,并混入MoE,利用注意力层负责精确检索,Mamba层负责长程压缩,从而平衡两者优势。

🏷️

标签

➡️

继续阅读