3.11 从Transformer到Mamba 状态空间模型

1.Transformer的算力账:序列一长就吃力

说起来,前面第36章我们花了一整章讲Transformer,把它讲成了这几年深度学习的顶梁柱。它的核心叫自注意力,每一个位置都要去和序列里所有别的位置算一遍相似度,再决定要关注谁。这么做的好处是灵活,不管两个词隔多远,它都能直接建立联系。可问题也藏在这份灵活里头。

我们把这串输入的长度记作 NNNN 就是序列里一共有多少个位置(一个位置可以是一个词,也可以是一个图像块、一帧视频、一段基因序列)。自注意力要算的是每一对位置之间的关系,那么两两组合的总数大约是 N×NN \times N,也就是 N2N^2。我们把这种随着 NN 增大按平方增长的代价写作 O(N2)O(N^2),这里的 OO 读作大O,是一种描述计算量增长快慢的记号,括号里的式子告诉你输入变大时计算量大概按什么形状往上走。

O(N2)O(N^2) 这件事其实很要命。NN10001000 涨到 1000010000,计算量就是原来的 100100 倍。NN 再涨到 100000100000,又是原来的 100100 倍。说白了,序列一旦拉长,算力和显存都会以平方的速度往上冲。普通一句话两句话的翻译还好,可一旦拿它去处理整本书、长视频、基因组(人的一个基因组大约三十亿个碱基对),显卡就要顶不住了。

我打个比方。Transformer的自注意力有点像一场大型座谈会,到场的每一位客人都要和在场的所有别的客人逐一寒暄一遍,人一多,整个会场就忙得喘不过气来。来十个人还其乐融融,来一千个人就谁也顾不上谁了。小张之前手头有一份几百页的财报要分析,他试着把全文一次性塞进一个长上下文的Transformer,结果显卡显存当场吃满,机器嗡嗡响了大半天才吐出结果。这就是 O(N2)O(N^2) 在现实里的代价。

那么有没有办法绕开这份平方级的负担呢。有,状态空间模型就是其中一条路。

2.状态空间模型:另一种记东西的方式

状态空间模型(英文State Space Model,大家习惯简称SSM)这一类模型其实来历不小,它最早并不在深度学习里,而是在控制论和信号处理那边。研究自动控制的人很早就用一套方程来描述一个随时间变化的系统,比如一个电路里的电压怎么变、一架无人机的姿态怎么演化。这套工具到了2020年前后被人重新捡起来,做成了能和Transformer并驾齐驱的序列模型,代表作叫Structured State Space Sequence Model,大家常写作S4。

SSM的核心想法和注意力很不一样。注意力是让每个位置都和别的位置两两比较,相当于把整段历史摊在桌上随时翻看。SSM呢,它更像一个人拿着一本笔记本顺着时间一路读下去,每读到一个新的输入,就在笔记本上更新一下当前的摘要,下一个时刻只看这本摘要就够了,已经读过的原文不必再回头翻。

我们把这个摘要叫作状态,记作 h(t)h(t)hh 就是状态向量,括号里的 tt 表示连续时间(这是从控制论里沿袭下来的写法,tt 是一个连续的实数,而不是一个一个离散的步子)。SSM用两条方程来描述它怎么动:

h(t)=Ah(t)+Bx(t)h'(t) = Ah(t) + Bx(t) y(t)=Ch(t)y(t) = Ch(t)

这两条公式读起来其实不吓人,我们一条一条说。先看第一条,h(t)h'(t) 表示状态 h(t)h(t) 关于时间 tt 的导数,也就是状态随时间变化的快慢。等号右边有两项,Ah(t)Ah(t) 是说当前状态 h(t)h(t) 经过一个参数矩阵 AA 做一次线性变换,代表状态自己的演化趋势。Bx(t)Bx(t) 是说当前的输入 x(t)x(t) 经过另一个参数矩阵 BB 进入状态,代表外界新来的信号往状态里塞了多少。这里 AABB 都是模型要学习的参数矩阵,x(t)x(t)tt 时刻的输入。整条方程的意思是,状态的变化,一部分来自它自身的惯性,一部分来自当前输入的推动。

第二条就更简单了。y(t)y(t)tt 时刻的输出,CC 是又一个参数矩阵,输出就是状态 h(t)h(t) 经过 CC 变换之后的结果。说白了,状态里攒下来的东西,经过 CC 一过滤,就成了这一刻的输出。

真正用起来的时候,我们处理的是一段一段离散的数据(比如一句话里的第1个词、第2个词),所以要把上面这两条连续时间的方程离散化,改成一步一步递推的形式。我们记离散时间步的编号为 nnn=1,2,3,n = 1, 2, 3, \dots),离散后的状态记作 HnH_n,第 nn 步的输入记作 xnx_n,再引入一个叫步长的量 Δ\DeltaΔ\Delta 表示每一步对应的连续时间间隔,也是模型学出来的),递推就大致写成:

Hn=AˉHn1+BˉxnH_n = \bar{A}H_{n-1} + \bar{B}x_n

这里 Aˉ\bar{A}Bˉ\bar{B} 是把连续的 AABB 经过离散化之后得到的新的参数矩阵(头上的横线表示这是离散化后的版本,和原来的 AABB 区分开)。这条式子读起来特别顺,这一步的状态 HnH_n,等于上一步的状态 Hn1H_{n-1} 推一下,再加上这一步输入 xnx_n 的新贡献。每来一个新输入就更新一次状态,而状态本身的大小和序列长度 NN 没关系,永远是那么大一个向量。

这就是SSM最讨喜的地方。它处理一个长度为 NN 的序列,只需要顺着走 NN 步,每一步做一次常数大小的矩阵乘法,总计算量是 O(N)O(N),也就是随着序列长度线性增长,而不是平方增长。

我记得我之前翻过一本讲信号与系统的旧书,里头讲这种递推式的状态更新,用的例子是一个老式收音机里的滤波电路,电压随着时间一点点演化,每一步都只依赖上一步的状态,不用回头管历史。这种思路搬到深度学习里,就给了我们一条绕开 O(N2)O(N^2) 的路子。

可光线性还不够。早期SSM有一个让人头疼的毛病,AABBCC 这三组参数在整段序列上是固定不变的,不管输入是什么,状态更新的方式都一样。这就好比那位记笔记的朋友只会机械地往下抄,什么内容进来都按同一个节奏记,重要的和不重要的一视同仁。Transformer之所以灵活,恰恰在于它能根据内容动态地决定关注谁,而SSM少了这份眼力见儿。

3.Mamba的突破:让SSM学会挑挑拣拣

Mamba是2023年底提出来的,一作叫Albert Gu。它做的事,说穿了就一句话,让SSM学会根据内容动态地决定记什么、忘什么。作者管这个叫选择性状态空间模型(Selective State Space Model)。

我们先看选择机制到底要解决什么问题。回头看上一节那条递推式 Hn=AˉHn1+BˉxnH_n = \bar{A}H_{n-1} + \bar{B}x_n,里面的 Aˉ\bar{A}Bˉ\bar{B} 都是固定的参数,不管输入 xnx_n 是什么,它们都是一个样。Mamba做的事情,是让其中一部分参数(特别是控制输入怎么进入状态的那个 Bˉ\bar{B},还有前面提到的步长 Δ\Delta)变成输入 xnx_n 的函数,也就是说,每来一个新的输入,模型都会根据这个输入本身的内容,临时算出一份属于自己的 Bˉ\bar{B}Δ\Delta

这么一改,效果就完全不一样了。当输入是一段重要信息的时候,模型可以让 Δ\Delta 变大、Bˉ\bar{B} 变得显著,把当前输入充分写进状态里。当输入是一段无关紧要的填充内容时,模型又可以让这些参数变小,几乎把这次输入忽略掉,状态原地不动。这就好比一个老到的编辑看长稿,遇到关键的论点和数据会停下笔细细标注,遇到水分大的段落就直接翻过去,时间和精力都花在了刀刃上。

把这件事讲得更细一点。步长 Δ\Delta 在原来的SSM里是一个全局固定的标量,控制着离散化的粒度。Mamba让它由输入 xnx_n 决定,于是模型就拥有了调节时间分辨率的自由,要看仔细的地方把 Δ\Delta 调大,相当于把镜头拉近。可以略过的地方把 Δ\Delta 调小,相当于快进。这种内容相关的取舍,正是注意力机制最拿手的事,而Mamba把它做到了SSM里头。

不过天下没有白捡的午餐。原来的SSM之所以能算得飞快,是因为参数固定,整个序列可以预先用一种叫卷积的运算一次性并行算完。参数一旦依赖输入,这种并行就做不下去了,只能老老实实顺着时间一步一步递推。Mamba为了把速度找回来,在工程上做了非常细致的优化,写了一套专门针对GPU的内核,让这种逐时间步的递推在硬件上跑得也不算太慢。这套巧思,是Mamba能真正落地的关键。

总结一下。Mamba继承了SSM的线性复杂度 O(N)O(N),又借由选择机制获得了类似注意力的内容感知能力,在长序列上既算得动,又能挑出关键信息。它刚出来那阵子确实让不少人兴奋,社区里一时间各种Mamba变体层出不穷。

4.一路往后:Mamba-2、Vision Mamba和Jamba

Mamba火了之后,后续工作很快跟了上来,我们挑几个有代表性的说一说。

先是Mamba-2,作者还是同一批人。这一版最值得讲的是它在数学上把SSM和注意力统一了起来。作者发现,SSM里那种状态递推,换个角度看,和一种特殊的注意力在数学上是同构的,作者把这种关系叫作结构化状态空间对偶(Structured State Space Duality,简称SSD)。这个统一的视角非常漂亮,说起来有点像物理学里电和磁原本各算各的,后来被麦克斯韦方程组统一成一个东西,看上去是两回事,底子里是同一个数学结构。有了这个统一,SSM和注意力就成了一家人,你可以把它们看作同一个谱上的两个端点,在这个谱上自由地选位置。Mamba-2还顺便把状态维度做得更大,整体性能又有提升。

再往下是视觉那头。原本大家觉得注意力在视觉任务上已经做得很好了(图像块之间的两两关系,ViT那一套),换SSM图什么呢。其实图的还是那份线性复杂度,尤其是把图像切得特别细,或者要处理医学影像、遥感图这种大图的时候,O(N2)O(N^2) 同样压人。Vision Mamba就是把Mamba那一套搬进视觉,先把图像切成小块再排成序列,然后让SSM一路扫过去做分类、检测、分割。在分辨率特别高的任务上,它的优势会比较明显。

第三个想提的是Jamba。它走的是混合路线,把Mamba和Transformer混在了一起,层与层之间交替堆叠。这么做的原因很直接,Transformer在短序列和精细推理上仍然稳,Mamba在长序列上算得轻省,两边各取所长。Jamba还顺手在中间掺了混合专家(Mixture of Experts,简称MoE)模块,把每一层里活跃的参数量也压了下来。这样一来,它既能撑很长的上下文,又能在常见推理任务上保持竞争力。类似的混合思路在DeepSeek、Mixtral这些模型上也能见到端倪,只是各家配比不同。

5.说几句现状:别神化,也别小看

走到这一步,大概能看清Mamba这类模型的位置了。

在超长序列上头,Mamba确有它独到的优势。序列长度 NN 涨到几十万、上百万这种量级的时候,O(N2)O(N^2) 的Transformer基本跑不动,O(N)O(N) 的Mamba还能稳稳出结果。这点在做长文档问答、长视频理解、基因组分析、长代码阅读的人那里感受最深。小明之前跟着一个生信课题,要把整段基因组丢进去找调控元件,纯Transformer压根喂不进去,换成Mamba一系的模型,至少能跑起来,出来的结果也说得过去。

但在中短序列上,Transformer仍然是大家更熟、调得更顺的那一个。生态、工具、预训练经验,Transformer这边都厚实得多。很多常见的任务,序列本来就不长,O(N2)O(N^2) 的代价并不明显,这时候用Transformer反而更省心。Mamba要彻底替代它,目前看还差些火候。

说到底,现在越来越多的工作倾向把两者结合起来,按任务需要搭配使用,未必非要分个高下。我个人觉得这种态度比较实在。一项技术从论文走到产业,往往要好几年的打磨,Transformer当年也是这么过来的。Mamba这一类模型还在快速演进,到底能走到哪一步,我们不妨再看看。

练习

Q1. Transformer 的自注意力为什么在长序列上吃不消?SSM 又是怎么绕开它的?

自注意力要算每一对位置之间的关系,两两组合总数约 N×NN\times N,复杂度是 O(N2)O(N^2),序列一长算力和显存就以平方速度暴涨。SSM(状态空间模型)则像拿本子顺时间记摘要,每来一个新输入就更新一次状态,状态大小跟序列长度无关,顺着走 NN 步、每步一次常数大小的矩阵乘法,总复杂度是 O(N)O(N),线性增长,绕开了平方级负担。

Q2. 序列长度 NN 从 1000 涨到 10000,再涨到 100000,Transformer 注意力的计算量大约各涨到原来的多少倍?

注意力计算量约 O(N2)O(N^2)NN 从 1000 涨到 10000(10 倍),计算量是原来的 102=10010^2=100 倍;再涨到 100000(相对 10000 又是 10 倍),又是原来的 100 倍。所以序列一旦拉长,显卡很快顶不住,这正是要找替代方案的原因。

Q3. 早期的 SSM(比如 S4)参数固定,这个毛病具体体现在哪?Mamba 又是怎么补上这块的?

早期 SSM 里 AABBCC 参数在整段序列上固定不变,不管输入是什么,状态更新的方式都一样,重要和不重要的内容一视同仁,缺少根据内容动态决定记什么、忘什么的"眼力见儿"。Mamba 让控制输入进入状态的 Bˉ\bar{B} 和步长 Δ\Delta 变成输入 xnx_n 的函数,遇到重要信息就把 Δ\Delta 调大、Bˉ\bar{B} 调显著充分写入状态,遇到无关内容就调小几乎忽略,从而获得了类似注意力的内容感知能力。

Q4.(面试题) Mamba 在超长序列上有优势,在中短序列上却往往不如 Transformer,请分析两者的取舍。Mamba-2 又做了什么统一工作?

超长序列(几十万、上百万)下 O(N2)O(N^2) 的 Transformer 基本跑不动,O(N)O(N) 的 Mamba 还能稳出结果,这是它的主场。但中短序列里 O(N2)O(N^2) 代价并不明显,而 Transformer 生态、工具、预训练经验都厚实得多,调起来更顺,这时候用 Transformer 反而更省心;Mamba 选择机制让参数依赖输入后,原本的并行卷积做不下去了,只能逐时间步递推,得靠专门的 GPU 内核优化才能不慢。所以现在多按任务搭配使用。Mamba-2 在数学上把 SSM 的状态递推和一种特殊注意力统一起来,提出了结构化状态空间对偶(SSD),说明两者是同一谱上的两个端点,还顺手把状态维度做得更大、性能提升。

相关标签
深度学习Mamba状态空间模型Transformer