The Illusion of State in State-Space Models
William Merrill, Jackson Petty, Ashish Sabharwal
ICML 2024
cs.LG, cs.CC, cs.CL, cs.FL
2024-04-13
NYU与Allen AI证明S4、Mamba等线性SSM和Transformer一样只能表达TC0;实验里单层RNN与输入依赖的IDS4能学任意长A5置换,S4和Mamba深度必须随序列增长。
Transformer被证明只能表达复杂度类TC0里的问题:常数深度、多项式规模的阈值电路能算的那些。状态跟踪是「必须按顺序把更新叠上去」的计算,很多落在NC1完全问题里。最干净的例子是五个物体偶置换群A5上的字问题:给你一串置换,问乘完之后每个位置落到哪。RNN一层就能表达,Transformer不能。
S4和Mamba这类状态空间模型(SSM)被写成RNN的并行替代。Gu等人还论证过线性SSM能模拟一般循环模型。它们是不是补上了Transformer缺的状态跟踪能力?这篇的答案是否定的。
分析对象是广义线性SSM层:隐状态按 hi = Āi h{i-1} + B̄i xi 更新,卷积形式把这段展开成矩阵连乘再求和。S4对应Ā不随输入变;Mamba用的S6里Ā是对角阵,可以随输入变(选择性),对角性让连乘退化成逐维标量连乘。
关键引理:只要任意区间上的Ā连乘能在L-uniform TC0里算出来,整层卷积形式就能被TC0电路模拟。两条路径把常见SSM送进这个框。
所以S4、S6/Mamba和Transformer一样,只能表达TC0。在公认猜想TC0 ≠ NC1下,它们表达不了S5/A5字问题,也就表达不了能编码S5的任务:UCI记法下跟踪棋步、按特定格式评Python、长叙事里的实体跟踪。
Gu等人那条「SSM能模拟RNN」依赖无限层数。有界层数下,这条路不通。
补丁有两条。一条是每步循环更新加非线性,变成RNN-SSM,一层就能认任意正则语言,SCAN那种前缀和并行化没了。另一条是让Āi成为完整的输入依赖矩阵,叫IDS4,接近Liquid S4:一般矩阵的迭代乘积不在TC0,一层就能模拟DFA,SCAN还能用。
实验把字问题做成逐步打标,每一步的标签是前缀乘积。对照三个都是60个元素的群:交换群Z60、可解非交换群A4×Z5、非可解群A5。模型包括Transformer、RNN、S4、Mamba、IDS4。看达到90%验证准确率所需的最少层数如何随序列长度变。
图3的结论很硬。单层RNN和单层IDS4在三个群上都能处理任意长序列。Transformer、S4、Mamba在A5上要深度随长度单调增加。
它们在理论上属于TC0的A4×Z5上同样要加深。两种解释都说得通:这些架构实际能表达的是TC0的真子集;或者常数深度解存在,但学不出来。S4和Mamba在近似状态跟踪上比Transformer省层,省的是常数,不是渐近。
棋步跟踪的NC1完全性只对UCI的(源格, 目标格)记法成立,标准SAN可能更简单。实体跟踪的难度也随题面格式变。
「SSM更循环所以更能跟状态」这句话,在表达力上不成立。Mamba的选择性让Ā随输入变,对角约束把连乘留在TC0里。想要真状态跟踪,至少得放弃对角,做成IDS4那种输入依赖的满矩阵,或者把非线性塞进循环步。前者还能并行。大规模语言模型里能不能训、梯度会不会炸,论文自己标成开放问题。
做代码执行、长程实体跟踪、棋类状态的人,不该默认换Mamba就能过Transformer过不去的那道坎。
整条证明绑在卷积形式与循环形式算同一函数上,浮点下两者并不严格相等。精度模型是c log n比特浮点,有限精度会更弱。扩展构造要求层输入维k大于字母表大小。实验是合成群乘法,没有在真实语言模型上验证「跟不住实体」。IDS4只在小任务上证明能学,没有语言建模实验。H3因为上下文不是单向量,不在定理范围内。