Recurrent Neural Networks.

循环神经网络(Recurrent Neural Network,RNN)用一个随时间演化的隐状态来概括已经看过的全部历史,从而能够建模序列数据内部的相关性,并处理长度不固定的文本、语音、时间序列等数据。它的核心假设可以写成一个递推式:

\[h_t = f(h_{t-1},x_t), \quad y_t = g(h_t)\]

这个递推把无界的历史压缩进有限的状态,因此推理成本与序列长度无关(这是它相对注意力机制的根本优势);但同一个压缩过程既会丢失信息,又会让梯度沿时间轴反复相乘,因此长程依赖与并行化成为贯穿整个领域的两条主线。本文按这两条主线组织:先讨论建模动机与理论性质,再讨论训练与长程依赖,然后系统梳理门控循环网络,最后给出深层结构和正则化建议。

  1. 循环神经网络的建模动机
    • (1) 序列建模与权值共享
    • (2) vanilla RNN
    • (3) RNN的理论性质
  2. RNN的训练:梯度计算与长程依赖
    • (1) 随时间反向传播 BPTT
    • (2) 实时循环学习 RTRL
    • (3) 长程依赖问题:梯度消失与爆炸
    • (4) 缓解长程依赖的非门控手段
  3. 门控循环网络
    • 3.1 经典门控单元
    • 3.2 面向并行化的简化循环单元
    • 3.3 结构化与连续时间的记忆
    • 3.4 隐状态即模型:测试时训练
    • 3.5 门控循环网络的复兴
  4. 深层、双向与其他结构设计
    • (1) 堆叠循环神经网络 Stacked RNN
    • (2) 双向循环神经网络 Bidirectional RNN
    • (3) 残差与高速连接
    • (4) 循环网络专用的正则化
    • (5) 编码器-解码器结构与注意力的由来

符号约定:全文用\(x_t \in \Bbb{R}^M\)表示$t$时刻的输入,\(h_t \in \Bbb{R}^D\)表示隐状态(hidden state),$c_t$表示记忆状态(cell state),$y_t$表示输出,$T$表示序列长度;用$\sigma(\cdot)$专门表示Sigmoid函数,$\odot$表示逐元素乘积。vanilla RNN中用$W_{xh},W_{hh},W_{hy}$分别表示输入到隐状态、隐状态到隐状态、隐状态到输出的权重;在各类门控单元中简写为$W_{\bullet}$(作用于输入$x_t$)与$U_{\bullet}$(作用于$h_{t-1}$),$b_{\bullet}$为偏置。

1. 循环神经网络的建模动机

(1) 序列建模与权值共享

用前馈网络处理序列数据有两条朴素的路线,它们各自都不令人满意:

循环神经网络给出的是第三条路线:只用同一组参数$f$,把它沿时间反复复合。这带来三个直接后果:

  1. 权值共享。参数量与序列长度完全无关,且某个时间模式一旦学会,在任意时刻都能复用。
  2. 无界的理论感受野。$h_t$依赖于$x_1,\dots,x_t$的全部信息,不存在卷积那样的窗口上界。
  3. 常数的推理开销。每步只需一次$f$的计算与$O(D)$的状态存储,与已处理的长度无关。这正是RNN在长上下文推理、流式/在线场景中始终具有吸引力的原因。

代价同样来自反复复合:前向计算天然是串行的($h_t$必须等$h_{t-1}$算完),反向传播要穿过$T$层复合而极易出现梯度消失/爆炸。本文其余部分几乎都在处理这两个代价。

值得区分的是递归神经网络:它沿语法树等递归结构而非时间轴共享权值,RNN可以看作递归结构退化为一条链的特例。

(2) vanilla RNN

一个简单的循环神经网络包括输入层、一层隐藏层和输出层。

令向量\(x_t \in \Bbb{R}^M\)表示$t$时刻网络的输入,\(h_t \in \Bbb{R}^D\)表示隐状态,则:

\[\begin{aligned} h_t &= f(W_{hh}h_{t-1}+W_{xh}x_t+b) \\ y_t &= W_{hy}h_t \end{aligned}\]

其中$f(\cdot)$是激活函数,常用Sigmoid或Tanh函数;参数$W_{hh}$、$W_{xh}$、$W_{hy}$、$b$在时间维度上权值共享。

上式的反馈来自隐状态$h_{t-1}$,这种结构称为Elman网络,是当今RNN一词的默认含义。与之相对的Jordan网络把反馈接在输出上,即\(h_t = f(W_{hh}y_{t-1}+W_{xh}x_t+b)\)。两者的差别不只是接线位置:Jordan网络的循环状态被输出维度\(\dim(y)\)卡住(分类任务中往往只有几十维),记忆容量远小于Elman网络的$D$维隐状态;此外Jordan网络在推理时反馈的是自己的预测,训练时若反馈真实标签就会产生teacher forcing式的暴露偏差。因此现代序列模型几乎一律采用Elman式的隐状态反馈,Jordan式的输出反馈只保留在自回归解码(把上一步生成的token重新喂回输入)这一场景中。

⚪ vanilla RNN的Pytorch实现

可以通过torch.nn.RNN构建RNN:

点击展开代码
rnn = torch.nn.RNN(
    input_size=10, # 输入序列的特征维度
    hidden_size=20, # 隐藏层状态的特征维度
    num_layers=2, # 循环层的数量,默认为1,用于实现Stacked RNN
    nonlinearity='tanh', # 激活函数,可选'relu'
    bias=True, # 是否使用偏差项b
    batch_first=False, # 若设置为True,则输入尺寸应为[Batch Size, Sequence Length, Input Size]
    dropout=0, # 设置神经元的dropout
    bidirectional=False, # 设置双向RNN
    )

input = torch.randn(5, 3, 10) # [Sequence Length, Batch Size, Input Size]
h0 = torch.randn(2, 3, 20) # [Bidirectional, Batch Size, Hidden Size]
output, hn = rnn(input, h0)

(3) RNN的理论性质

⚪ 性质1:通用近似定理(Universal Approximation Theory)

如果一个完全连接的循环神经网络有足够数量的sigmoid神经元,它可以以任意的准确率去近似任何一个非线性动力系统。

一个非线性动力系统可以用常微分方程(Ordinary Differential Equation, ODE)描述:

\[\dot{x}(t) = f(x(t),t)\]

ODE通常比较难以求出解析解,可以采用欧拉解法,即用$\frac{x(t+h)-x(t)}{h}$近似导数项$\dot{x}(t)$,则迭代公式为:

\[x(t+h)=x(t)+hf(x(t),t)\]

因此ODE的欧拉解法是RNN的一个特例,这也说明了RNN对于时间序列数据具有很强的拟合能力。这条联系并非只有理论意义:把上式改写成\(h_t = h_{t-1}+\Delta t\cdot f(h_{t-1},x_t)\)就得到一个带残差连接的循环单元,其中的步长$\Delta t$若改为可学习、甚至改为输入相关的量,就分别得到$3.2$节的IndRNN/LRU与$3.3$节的连续时间记忆单元。换个角度看,门控机制中的遗忘门$f_t$正是一个逐通道可学习的$1-\Delta t$。

⚪ 性质2:图灵完备性(Turing Completeness)

图灵完备是指一种数据操作规则,比如一种计算机编程语言,可以实现图灵机(Turing Machine)的所有功能,解决所有的可计算问题。目前主流的编程语言(比如C++、Java、Python等)都是图灵完备的。

RNN的图灵完备性是指所有的图灵机都可以被一个由使用Sigmoid神经元构成的全连接循环网络来进行模拟。

⭐ 讨论:理论表达力与实际可学习性的鸿沟

上面两条性质只是存在性结论:存在一组参数使RNN逼近任意动力系统/模拟任意图灵机,却完全没有说这组参数能否用梯度下降找到。图灵完备性的构造还依赖于无限精度的实数权重与状态(用有限精度的$D$维状态模拟无限长的图灵机纸带是不可能的)。真正决定RNN实际能力的是第$2$节讨论的优化性质:由于梯度沿时间指数衰减,vanilla RNN实际能利用的上下文往往只有十几步。这也解释了一个反复出现的现象:RNN领域的进展几乎从不来自“提升表达力”,而是来自“让已有的表达力变得可优化、可并行”。

2. RNN的训练:梯度计算与长程依赖

RNN的参数可以通过梯度下降方法来进行学习,在RNN中主要有两种计算梯度的方式:

  1. 随时间反向传播(BPTT)算法
  2. 实时循环学习(RTRL)算法

为了更好地说明RNN的参数更新过程,引入中间变量$z_t$表示应用激活函数前的隐状态,则RNN的前向传播过程可写作:

\[\begin{aligned} z_t &= W_{hh}h_{t-1}+W_{xh}x_t+b \\ h_t &= f(z_t) \\ y_t &= W_{hy}h_t \end{aligned}\]

(1) 随时间反向传播 BPTT

随时间反向传播(BackPropagation Through Time,BPTT)算法将循环神经网络看作一个展开的多层前馈网络,其中每一层对应循环网络中的每个时刻。

定义误差项\(\delta_{t,k} = \frac{\partial L_t}{\partial z_k}\),则误差的反向传播:

\[\begin{aligned} \delta_{t,k} &= \frac{\partial L_t}{\partial z_k} = \frac{\partial L_t}{\partial z_{k+1}} \frac{\partial z_{k+1}}{\partial z_k} \\&= \frac{\partial L_t}{\partial z_{k+1}} \frac{\partial z_{k+1}}{\partial h_k} \frac{\partial h_k}{\partial z_k} \\&= \delta_{t,k+1}W_{hh}f'(z_k) \end{aligned}\]

RNN所有层的参数是共享的,因此参数的真实梯度是所有展开层的参数梯度之和:

\[\begin{aligned} \frac{\partial L}{\partial W_{hh}} &= \sum_{t=1}^{T} {\sum_{k=1}^{t} {\delta_{t,k}h_{k-1}^\top}} \\ \frac{\partial L}{\partial W_{xh}} &= \sum_{t=1}^{T} {\sum_{k=1}^{t} {\delta_{t,k}x_{k}^\top}} \\ \frac{\partial L}{\partial b} &= \sum_{t=1}^{T} {\sum_{k=1}^{t} {\delta_{t,k}}} \end{aligned}\]

(2) 实时循环学习 RTRL

实时循环学习(Real-Time Recurrent Learning,RTRL)是通过前向传播的方式来计算梯度,以$W_{hh}$为例:

\[\begin{aligned} \frac{\partial h_{t+1}}{\partial W_{hh}} &= \frac{\partial h_{t+1}}{\partial z_{t+1}} \frac{\partial z_{t+1}}{\partial W_{hh}} = \frac{\partial h_{t+1}}{\partial z_{t+1}} \left(\frac{\partial W_{hh}h_t}{\partial W_{hh}} + W_{hh}\frac{\partial h_{t}}{\partial W_{hh}}\right) \\ \frac{\partial L_{t}}{\partial W_{hh}} &= \frac{\partial h_{t}}{\partial W_{hh}} \frac{\partial L_{t}}{\partial h_{t}} \end{aligned}\]

即RTRL沿时间递推地维护状态对参数的雅可比\(\partial h_t/\partial W\),每读入一个时间步就把它更新一次,因此可以在不回看历史的情况下立即得到梯度。

⭐ 讨论:BPTT与RTRL的取舍

RTRL算法和BPTT算法都是基于梯度下降的算法,分别通过前向模式和反向模式应用链式法则来计算梯度,两者得到的梯度在数学上完全相同,差别只在于复杂度与适用场景。设隐状态维度为$D$、参数量为$P=O(D^2)$、序列长度为$T$:

  时间复杂度 空间复杂度 是否需要回看历史
BPTT \(O(TD^2)\) \(O(TD)\) 需要(保存所有时刻的中间量)
RTRL \(O(TD^4)\) \(O(D^3)\) 不需要(可在线更新)

(3) 长程依赖问题:梯度消失与爆炸

RNN反向传播中的误差项\(\delta_{t,k}\)满足:

\[\delta_{t,k} = \delta_{t,k+1}W_{hh}f'(z_k)\]

若记\(\gamma \approx \Vert W_{hh}f'(z_k) \Vert\),则:

\[\delta_{t,k} \sim \gamma^{t-k}\delta_{t,t}\]

由于RNN经常使用Sigmoid函数或Tanh函数作为非线性激活函数,其导数值都小于$1$(Sigmoid的导数上界甚至只有$1/4$),因而经常会出现梯度消失问题。更精确地说,$\gamma$由$W_{hh}$的谱半径(最大奇异值)与激活函数导数共同决定:只有当谱半径恰好落在$1$附近时梯度才既不消失也不爆炸,而这是一个测度为零的临界条件;这也解释了为什么长程依赖不是一个可以靠调参绕过的工程问题,而是循环结构的内在性质。

值得注意的是,梯度消失并不是参数$W$的梯度\(\frac{\partial L_{t}}{\partial W}\)消失了,而是隐藏层状态$h_{k}$的梯度\(\frac{\partial L_{t}}{\partial h_{k}}\)消失了。也就是说参数$W$的更新主要靠最近时刻的几个相邻状态贡献,而长距离的状态则无法产生影响。

虽然RNN理论上可以建立长时间间隔的状态之间的依赖关系,但是由于梯度消失问题,实际上只能学习到短期的依赖关系。这个问题称作长程依赖问题(Long-Term Dependencies Problem)。

两类问题的性质并不对称,缓解手段也不同:

(4) 缓解长程依赖的非门控手段

在引入门控之前,有若干与门控正交、可以叠加使用的手段。

⚪ 梯度裁剪

按梯度的范数或元素值对更新量做截断,是防止梯度爆炸最直接有效的做法(循环网络与语言模型训练中几乎是标配,常用全局范数阈值$1.0$)。详细公式与自适应变体见梯度裁剪。

⚪ 恒等初始化与IRNN

IRNN指出:如果把循环矩阵初始化为单位矩阵、偏置初始化为$0$,并把激活函数换成ReLU,那么在初始时刻网络的行为恰好是“把历史状态原样累加”:

\[h_t = \text{ReLU}\left(W_{xh}x_t+W_{hh}h_{t-1}+b\right), \quad W_{hh} \leftarrow I,\ b\leftarrow 0\]

此时\(\partial h_t/\partial h_{t-1}=I\)(在正半轴上),梯度既不衰减也不放大,长程信息在训练初期就能自由流动。IRNN在长序列的加法任务与置换序列MNIST上可以达到接近LSTM的水平,代价是必须使用很小的学习率;因为ReLU不再有界,状态存在发散风险。

一般地,把循环边权重初始化为正交矩阵是循环网络最重要的初始化技巧:循环矩阵会被复合$T$次,任何偏离$1$的奇异值都会被放大到$T$次幂。详见正交、恒等与等距初始化。

⚪ 酉循环网络 uRNN

正交初始化只保证初始时刻的范数保持性,训练过程中$W_{hh}$会漂移。uRNN更激进:把循环矩阵约束为复数酉矩阵($W^{*}W=I$),从而在整个训练过程中严格保证\(\Vert Wh \Vert = \Vert h \Vert\),梯度范数不衰减也不爆炸。

直接在酉矩阵流形上做优化代价很高,作者给出了一个只需$O(D)$参数、$O(D\log D)$计算的参数化,它由若干种本身即为酉矩阵的结构算子复合而成:

\[W = D_3 R_2 \mathcal{F}^{-1} D_2 \Pi R_1 \mathcal{F} D_1\]

其中$D$为对角相位矩阵,$R$为Householder反射,$\mathcal{F}$为离散Fourier变换,$\Pi$为固定的置换。配套的激活函数为只改变复数模长、保留相位的modReLU:

\[\sigma_{\text{modReLU}}(z) = \begin{cases} \left(\lvert z \rvert+b\right)\frac{z}{\lvert z \rvert}, & \lvert z \rvert+b \geq 0 \\ 0, & \lvert z \rvert+b < 0 \end{cases}\]

uRNN这条路线(以及后续的oRNN、scoRNN、expRNN等)在纯记忆类合成任务上表现极佳,但在真实语言任务上并未超过LSTM:严格保范意味着模型不能主动遗忘,而遗忘恰恰是有用的。它更重要的价值在于把“用复数对角/酉矩阵参数化循环”这一思路留给了后来的状态空间模型与$3.2$节的LRU。

⚪ 跳跃连接与多时间尺度

另一类思路是直接缩短梯度需要穿过的路径长度:

⚪ 截断BPTT

完整的BPTT要求把整条序列展开,空间开销$O(TD)$,对语言模型这种动辄上万步的序列不可行。实践中的标准做法是截断随时间反向传播(Truncated BPTT, TBPTT):把长序列切成长度为$k_2$的连续片段,前向传播时跨片段传递隐状态(不清零),但反向传播只在片段内进行(把进入片段的隐状态视为常量,即h = h.detach())。

⚪ 归一化

在循环网络中直接套用BatchNorm会遇到“每个时间步的统计量都不同、且测试时序列可能更长”的困难,实践中一般改用LayerNorm(作用在门控的预激活上),或使用逐时间步维护统计量的Recurrent BatchNorm。详见深度学习中的归一化方法。

3. 门控循环网络

为了改善循环神经网络的长程依赖问题,引入了门控机制(Gated Mechanism):用一个取值在$(0,1)$的、输入相关的向量来控制信息在时间轴上的通断。它的作用可以从两个互补的角度理解:

本节按四条演化线索组织:经典门控单元($3.1$)、为并行化而简化门控($3.2$)、给记忆施加结构或连续时间约束($3.3$)、把隐状态本身升级为一个模型($3.4$),以及$2024$年前后“回归门控RNN”的复兴($3.5$)。

3.1 经典门控单元

⚪ LSTM:长短期记忆网络

长短期记忆网络(Long Short-Term Memory Network,LSTM)引入了门控机制来控制信息传递的路径,可以有效地缓解RNN的梯度消失问题。

LSTM网络引入了输入门$i$ (input gate)、遗忘门$f$ (forget gate),和输出门$o$ (output gate);并把输入和隐状态整合为记忆状态$c$(cell state);根据遗忘门和输入门更新记忆状态后,根据输出门更新隐状态。记忆状态通过线性的循环信息控制缓解了梯度消失问题。

\[\begin{aligned} i_t &= \sigma(W_{i}x_t+U_{i}h_{t-1}+b_i) \\ f_t &= \sigma(W_{f}x_t+U_{f}h_{t-1}+b_f) \\ o_t &= \sigma(W_{o}x_t+U_{o}h_{t-1}+b_o) \\ \tilde{c}_t &= \text{tanh}(W_{c}x_t+U_{c}h_{t-1}+b_c) \\ c_t &= c_{t-1} \odot f_t + i_t \odot \tilde{c}_t \\ h_{t} &= o_t \odot \text{tanh}(c_t) \end{aligned}\]

LSTM这个名字对应三种时间尺度的记忆:

结构上最关键的一点是:$c_t$的更新路径上没有矩阵乘法与饱和激活,只有逐元素的乘加。这条“恒等通路”与残差连接的作用完全同源,也是LSTM沿用至今的根本原因。

LSTM有若干经过长期检验的变体与训练细节:

1)遗忘门偏置初始化。遗忘门参数的初始化如果比较小,会在训练初期丢弃前一时刻的大部分信息,很难捕捉到长距离的依赖。因此遗忘门的偏置初始值一般设为一个正数(通常$b_f=1$,也有工作建议\(1\sim 3\)),使\(\sigma(b_f)\)接近$1$、初始时状态近似恒等传递。

2)peephole连接。输入门$i$、遗忘门$f$和输出门$o$不但依赖于输入$x_t$和上一时刻的隐状态$h_{t-1}$,也直接依赖于记忆单元$c$(通常用逐元素的对角权重):

\[\begin{aligned} i_t &= \sigma(W_{i}x_t+U_{i}h_{t-1}+ p_i \odot c_{t-1} +b_i) \\ f_t &= \sigma(W_{f}x_t+U_{f}h_{t-1}+ p_f \odot c_{t-1} +b_f) \\ o_t &= \sigma(W_{o}x_t+U_{o}h_{t-1}+ p_o \odot c_{t} +b_o) \end{aligned}\]

这让门控能“看到”被输出门遮挡住的记忆内容,在需要精确计时的任务上有帮助。

3)耦合输入门与遗忘门(CIFG)。输入门和遗忘门有互补关系,可以把两者耦合为一个门,参数量减少约$1/4$而效果基本不变:

\[\begin{aligned} f_t &= 1-i_t \\ c_t &= c_{t-1} \odot (1-i_t) + i_t \odot \tilde{c}_t \end{aligned}\]

这一形式使$c_t$成为$c_{t-1}$与\(\tilde c_t\)的凸组合,因此记忆状态天然有界,是GRU、QRNN、SRU、minGRU共同采用的写法。

下面给出单个LSTM模块的定义和序列的循环处理过程:

点击展开代码
import torch
import torch.nn as nn

class LSTMCell(nn.Module):
    def __init__(self, input_size, hidden_size, num_layers = 1, dropout = 0.1):
        super(LSTMCell, self).__init__()

        self.num_layers = num_layers
        self.dropout = nn.Dropout(p=dropout)

        ih, hh = [], []
        for i in range(num_layers):
            if i==0:
                ih.append(nn.Linear(input_size, 4 * hidden_size))
                hh.append(nn.Linear(hidden_size, 4 * hidden_size))
            else:
                ih.append(nn.Linear(hidden_size, 4 * hidden_size))
                hh.append(nn.Linear(hidden_size, 4 * hidden_size))
        self.w_ih = nn.ModuleList(ih)
        self.w_hh = nn.ModuleList(hh)

    def forward(self, input, hidden):
        if hidden[0].shape[0] != self.num_layers:
            hidden = (
                torch.tile(hidden[0], [self.num_layers,1,1]),
                torch.tile(hidden[1], [self.num_layers,1,1]))

        hy, cy = [], []
        for i in range(self.num_layers):
            hx, cx = hidden[0][i], hidden[1][i]
            gates = self.w_ih[i](input) + self.w_hh[i](hx)
            i_gate, f_gate, c_gate, o_gate = gates.chunk(4, 1)
            i_gate = torch.sigmoid(i_gate)
            f_gate = torch.sigmoid(f_gate)
            c_gate = torch.tanh(c_gate)
            o_gate = torch.sigmoid(o_gate)
            ncx = (f_gate * cx) + (i_gate * c_gate)
            nhx = o_gate * torch.tanh(ncx)
            cy.append(ncx)
            hy.append(nhx)
            input = self.dropout(nhx)

        hy, cy = torch.stack(hy, 0), torch.stack(cy, 0)  # number of layer * batch * hidden
        return hy, cy

lstm = LSTMCell(10, 20, 2)
input = torch.randn(5, 3, 10) # [Sequence Length, Batch Size, Input Size]
hx = torch.randn(3, 20) # [Batch Size, Hidden Size]
cx = torch.randn(3, 20) # [Batch Size, Cell Size]
output = []
for i in range(input.size()[0]):
    hx, cx = lstm(input[i], (hx, cx))
    output.append(hx)
output = torch.stack(output, dim=0)

也可以通过torch.nn.LSTM构建LSTM:

点击展开代码
lstm = torch.nn.LSTM(
    input_size=10, # 输入序列的特征维度
    hidden_size=20, # 隐藏层状态的特征维度
    num_layers=2, # 循环层的数量,默认为1,用于实现Stacked RNN
    bias=True, # 是否使用偏差项b
    batch_first=False, # 若设置为True,则输入尺寸应为[Batch Size, Sequence Length, Input Size]
    dropout=0, # 设置神经元的dropout
    bidirectional=False, # 设置双向RNN
    proj_size=0, # 记忆状态的特征维度,默认等于hidden_size
    )

input = torch.randn(5, 3, 10) # [Sequence Length, Batch Size, Input Size]
h0 = torch.randn(2, 3, 20) # [Bidirectional, Batch Size, Hidden Size]
c0 = torch.randn(2, 3, 20) # [Bidirectional, Batch Size, Cell Size]
output, (hn, cn) = lstm(input, (h0, c0))

⚪ GRU:门控循环单元

门控循环单元(Gated Recurrent Unit,GRU)比LSTM结构更加简单:它取消了独立的记忆状态$c$(只保留$h$),并把三个门缩减为两个:更新门$z$ (update gate) 和重置门$r$ (reset gate):

\[\begin{aligned} z_t &= \sigma(W_{z}x_t+U_{z}h_{t-1}+b_z) \\ r_t &= \sigma(W_{r}x_t+U_{r}h_{t-1}+b_r) \\ \tilde{h}_t &= \text{tanh}(W_{h}x_t+U_{h}(h_{t-1}\odot r_t)+b_h) \\h_{t} &= z_t \odot h_{t-1} + (1-z_t) \odot \tilde{h}_{t} \end{aligned}\]

两个门的分工是不同的:更新门$z$作用在时间轴上,决定“保留多少历史”(对应LSTM耦合后的遗忘/输入门);重置门$r$作用在候选状态的计算上,决定“生成新内容时参考多少历史”,$r_t\to 0$时相当于把序列在此处切断、重新开始。注意上式的$z$按“保留历史的比例”定义(与PyTorch实现一致),文献中也常见把$z$定义为“写入新内容的比例”的相反写法。

GRU的参数量约为LSTM的$3/4$,在中小规模数据上通常与LSTM相当甚至更好;但GRU的状态有界(凸组合)且没有独立的记忆通道,在需要精确计数/计时的任务上表达力严格弱于LSTM。

可以通过torch.nn.GRU构建GRU:

点击展开代码
gru = torch.nn.GRU(
    input_size=10, # 输入序列的特征维度
    hidden_size=20, # 隐藏层状态的特征维度
    num_layers=2, # 循环层的数量,默认为1,用于实现Stacked RNN
    bias=True, # 是否使用偏差项b
    batch_first=False, # 若设置为True,则输入尺寸应为[Batch Size, Sequence Length, Input Size]
    dropout=0, # 设置神经元的dropout
    bidirectional=False, # 设置双向RNN
    )

input = torch.randn(5, 3, 10) # [Sequence Length, Batch Size, Input Size]
h0 = torch.randn(2, 3, 20) # [Bidirectional, Batch Size, Hidden Size]
output, hn = gru(input, h0)

⚪ MGU:最小门控单元

既然GRU已经把LSTM的三个门减到两个,还能不能只留一个门?最小门控单元(Minimal Gated Unit, MGU)给出的答案是可以:它把GRU的更新门与重置门合并为同一个门$f_t$(称为遗忘门),同时承担“保留多少历史”与“生成候选时参考多少历史”两项职责:

\[\begin{aligned} f_t &= \sigma(W_{f}x_t+U_{f}h_{t-1}+b_f) \\ \tilde{h}_t &= \text{tanh}(W_{h}x_t+U_{h}(f_t \odot h_{t-1})+b_h) \\ h_{t} &= (1-f_t) \odot h_{t-1} + f_t \odot \tilde{h}_{t} \end{aligned}\]

MGU只有两组权重矩阵(LSTM四组、GRU三组),参数量约为GRU的$2/3$,在多个基准上与GRU基本持平。它的价值主要在于结构层面的结论:门控RNN中真正不可缺少的只有一个“控制历史保留比例”的门;其余的门都是可选的表达力增益。这一结论在$3.5$节的minGRU中被再次确认。

⚪ JANET与chrono初始化

JANET (Just Another NETwork)从另一个方向做减法:它保留LSTM的记忆状态$c$,但只保留遗忘门,去掉输入门(与遗忘门耦合)与输出门(直接令$h_t=c_t$),并去掉输出端的$\tanh$:

\[\begin{aligned} s_t &= W_{f}x_t+U_{f}h_{t-1}+b_f \\ \tilde{c}_t &= \text{tanh}(W_{c}x_t+U_{c}h_{t-1}+b_c) \\ c_t &= \sigma(s_t) \odot c_{t-1} + \left(1-\sigma(s_t-\beta)\right) \odot \tilde{c}_t \\ h_{t} &= c_t \end{aligned}\]

这里的$\beta$是一个值得注意的细节:如果直接令输入门为\(1-\sigma(s_t)\)(即标准的耦合门CIFG),那么遗忘门与输入门被强制满足\(f_t+i_t=1\),“多记住历史”与“多写入新信息”完全是零和的。从输入门的预激活里减去一个正的偏移$\beta$后,\(1-\sigma(s_t-\beta) > 1-\sigma(s_t)\),两者之和略大于$1$:在初始化处(\(s_t\approx 0\))遗忘门为$0.5$而输入门约为$0.525$,新信息的权重略高于历史。原文取$\beta=0.1$,并报告这一偏移使模型在训练早期更容易遗忘,从而更快摆脱随机初始化时无意义的历史状态。

JANET真正的关键在于遗忘门偏置的初始化,即chrono初始化(chrono initialization)。把遗忘门写成\(f=\sigma(b_f)\),则该通道的记忆时间常数约为\(\tau \approx 1/(1-f)\);若希望模型的各个通道覆盖从$1$到\(T_{\max}\)的各种时间尺度,就应当让$\tau$在\([1,T_{\max}]\)上大致均匀,对应的偏置初始化为:

\[b_f \sim \log\left(\mathcal{U}\left(\left[1, T_{\max}-1\right]\right)\right), \quad b_i = -b_f\]

其中\(T_{\max}\)取任务中需要建模的最长依赖长度。这是“遗忘门偏置初始化为$1$”这一经验做法的原理化推广:后者只是把所有通道都设成同一个(且相当短的)时间尺度。chrono初始化可以直接套用在标准LSTM上,是长序列任务上一个几乎零成本的改进。

⭐ 讨论:到底哪些门是必要的

从LSTM($4$组权重、$3$个门)到GRU($3$组、$2$个门)到MGU/JANET($2$组、$1$个门),门控单元的演化史本质上是一场持续二十年的消融实验。综合大规模结构搜索(LSTM: A Search Space Odyssey、An Empirical Exploration of Recurrent Network Architectures)与上述工作,可以归纳出几条相当稳定的结论:

3.2 面向并行化的简化循环单元

LSTM/GRU的门控依赖于$h_{t-1}$,因此每个时间步都必须等上一步算完才能做矩阵乘法,训练时GPU利用率很低。本小节的方法共享同一个思路:把依赖$h_{t-1}$的矩阵乘法移出循环体,只在循环体内保留逐元素的乘加。这样一来,占据绝大部分计算量的矩阵乘法可以对整条序列并行完成,剩下的逐元素递推还可以用并行扫描进一步加速。

⚪ QRNN:准循环神经网络

标准的循环神经网络需要循环地实现,即每次处理输入序列的一个token,无法并行化处理输入序列。这是因为在参数化处理输入序列时依赖于上一时刻的隐状态。准循环神经网络 (Quasi-Recurrent Neural Network, QRNN)通过把卷积层引入RNN,实现了输入序列的并行处理,同时输出结果依赖于序列顺序。

QRNN使用一维(因果)卷积处理输入序列,设置卷积核大小为$k$,则根据最近$k$时刻的输入$x_{t-k+1:t}$生成当前时刻的输入门$z_t$、遗忘门$f_t$和输出门$o_t$。与LSTM不同,这一步不依赖于隐状态$h_{t-1}$,因此可以以矩阵方式并行地运算:

\[\begin{aligned} z_t &= \text{tanh}(W_{z}^{1}x_{t-k+1}+W_{z}^{2}x_{t-k+2}+\cdots + W_{z}^{k}x_{t}) \\ f_t &= \sigma(W_{f}^{1}x_{t-k+1}+W_{f}^{2}x_{t-k+2}+\cdots + W_{f}^{k}x_{t}) \\ o_t &= \sigma(W_{o}^{1}x_{t-k+1}+W_{o}^{2}x_{t-k+2}+\cdots + W_{o}^{k}x_{t}) \end{aligned}\]

然后通过动态平均池化构造序列的输出(隐状态),这一步是循环实现的,但是一个无参数函数:

\[\begin{aligned} c_t &= c_{t-1} \odot f_t + z_t \odot (1-f_t) \\ h_{t} &= o_t \odot c_t \end{aligned}\]

QRNN因此是“卷积负责混合局部信息、循环只负责传递长程状态”的分工:卷积核大小$k$决定了每一步能看到的局部上下文,而循环部分保证了理论上无界的感受野。

使用torchqrnn库构造QRNN:

点击展开代码
import torch
from torchqrnn import QRNN

seq_len, batch_size, hidden_size = 7, 20, 256
qrnn = QRNN(hidden_size, hidden_size, num_layers=2, dropout=0.4)
input = torch.randn(seq_len, batch_size, hidden_size)
output, hidden = qrnn(input)

⚪ SRU:简单循环单元

简单循环单元 (Simple Recurrent Unit, SRU)的设计思路与QRNN类似,通过把矩阵乘法放在串行循环之外,能够并行地处理输入序列,提升了运算速度。

SRU中每个时间步的门控计算只依赖于当前时间步的输入(连卷积窗口都不需要),并在输出(隐状态)中添加了高速连接(highway connection):

\[\begin{aligned} \tilde{x}_t &= Wx_t \\ f_t &= \sigma(W_{f}x_{t}+b_f) \\ r_t &= \sigma(W_{r}x_{t}+b_r) \\ c_t &= c_{t-1} \odot f_t + \tilde{x}_t \odot (1-f_t) \\ h_{t} &= r_t \odot g(c_t) + (1-r_t) \odot x_t \end{aligned}\]

其中$r_t$称为重置门,它把“循环通路”\(g(c_t)\)与“直连通路”$x_t$做插值,使得堆叠很多层SRU时梯度仍能沿直连通路传播——这与Transformer中残差连接的作用相同,也是SRU能堆到很深的原因。由于\(\tilde x_t, f_t, r_t\)都只依赖$x_t$,三次矩阵乘法可以合并成一次批量GEMM对整条序列完成,作者报告的训练速度比cuDNN LSTM快$5\sim 9$倍。

⚪ IndRNN:独立循环神经元

IndRNN采取了一个更彻底的简化:把循环矩阵$W_{hh}$退化为一个向量$u$,循环连接改为逐元素乘法:

\[h_t = \sigma\left(W x_t + u \odot h_{t-1} + b\right)\]

于是同一层内的神经元在时间上彼此独立:第$n$个神经元只接收自己的历史\(h_{t-1,n}\),神经元之间的信息交互全部交给层间连接完成。这个改动带来三个好处:

  1. 梯度可以被解析地控制。第$n$个神经元从时刻$t$回传到时刻$k$的梯度因子是\(u_n^{t-k}\prod \sigma'\),是一个标量幂而非矩阵幂。要求梯度落在\([\epsilon,\gamma]\)内,只需约束\(\lvert u_n \rvert \in \left[\epsilon^{1/(t-k)},\ \gamma^{1/(t-k)}\right]\);实践中直接对$u$做裁剪即可。这使得IndRNN可以安全地使用ReLU而不发散,并处理数千步的序列。
  2. 可以堆得很深。由于逐层之间是普通的全连接/卷积,可以直接套用BatchNorm与残差连接,作者堆叠了$21$层IndRNN。
  3. 循环部分只有逐元素运算,与本小节其他方法一样便于并行与加速。

代价是单层的表达力下降(层内不再有跨神经元的时间混合),需要靠深度来补偿。这一“对角循环”的思路后来被状态空间模型全面继承。

⚪ LRU:线性循环单元

线性循环单元(Linear Recurrent Unit, LRU)把上述简化推到极致:去掉循环体内的所有非线性,只保留一个复数对角线性递推,非线性全部交给层间的MLP:

\[h_t = \Lambda h_{t-1} + \gamma \odot \left(B x_t\right), \quad y_t = \Re\left[C h_t\right] + D x_t\]

其中\(\Lambda = \text{diag}(\lambda_1,\dots,\lambda_D)\)为复数对角矩阵。为了保证稳定性(要求\(\lvert \lambda_j\rvert<1\))并让特征值能够贴近单位圆,采用指数参数化与归一化:

\[\lambda_j = \exp\left(-\exp(\nu_j^{\log}) + i \exp(\theta_j^{\log})\right), \quad \gamma_j = \sqrt{1-\lvert \lambda_j \rvert^2}\]

这篇工作的意义在于给出了一条“消融路径”:从vanilla RNN出发,依次做线性化、对角化(复数)、稳定的指数参数化、归一化四步,就能在长序列基准上达到与S4相当的水平;也就是说,状态空间模型的强大表现主要来自这些参数化技巧,而不是HiPPO那样的连续时间理论。由于LRU的核心卖点已经是线性递推与并行扫描,它在谱系上更接近状态空间模型。

⭐ 讨论:并行扫描——把串行递推压缩到对数深度

本小节所有方法(以及$3.5$节的minGRU/mLSTM和整个SSM家族)之所以能并行训练,靠的是同一个数学事实。它们的递推都能写成输入相关系数的一阶线性形式:

\[h_t = a_t \odot h_{t-1} + b_t\]

其中\(a_t,b_t\)只依赖于$x_t$,与$h_{t-1}$无关。把它展开:

\[h_t = \sum_{k=1}^{t} \left(\prod_{j=k+1}^{t} a_j \right)\odot b_k\]

括号里的累积乘积可以用前缀积一次性算出,因此整个序列的\(\{h_t\}\)可以由关联扫描(associative scan)在\(O(\log T)\)的并行深度、$O(T)$的总工作量内完成(Blelloch扫描)。相比之下,LSTM/GRU的$a_t$依赖$h_{t-1}$,扫描的结合律被破坏,只能$O(T)$串行。

两个实践要点:

3.3 结构化与连续时间的记忆

上一小节为了并行而削弱记忆结构,本小节相反:主动给记忆状态施加结构约束,让它承载超出“一堆数”的语义。

⚪ ON-LSTM:有序神经元LSTM

LSTM的神经元是无序的:如果把运算过程中涉及到的所有向量的位置按同一方式重新打乱,并相应地打乱权重的顺序,则输出结果只是原向量的重新排序,信息量完全不变。也就是说,标准LSTM的记忆状态各维度之间没有任何结构关系。

而自然语言句子通常具有层级结构:字、词、词组、短语、从句逐级嵌套。层级越高、颗粒度越粗,它在句子中的跨度就越大。ON-LSTM通过有序神经元(Ordered Neuron)把这种层级结构(树结构)整合到LSTM中,从而允许LSTM无监督地学习到层级结构信息(如句子的句法结构)。

ON-LSTM假设记忆神经元$c_t$已经排好序:$c_t$中索引值越小的元素表示越低层级的信息,索引值越大的元素表示越高层级的信息。每次更新$c_t$之前,先预测两个整数$d_f$与$d_i$,分别表示历史信息$h_{t-1}$的层级与当前输入$x_t$的层级,然后分区间更新:

\[c_t = \begin{pmatrix} \hat{c}_{t,<d_f} \\ f_{t,[d_f,d_i]} \odot c_{t-1,[d_f,d_i]} + i_{t,[d_f,d_i]} \odot \hat{c}_{t,[d_f,d_i]} \\ c_{t-1,>d_i} \end{pmatrix}\] \[c_t = \begin{pmatrix} \hat{c}_{t,\leq d_i} \\ 0_{(d_i,d_f)} \\ c_{t-1,\geq d_f} \end{pmatrix}\]

上图给出了分区间更新的全貌:上方$c_{t-1}$为历史信息(黄色部分为历史信息层级),下方\(\hat c_t\)为当前输入(绿色部分为输入信息层级),中间$c_t$为整合结果;黄色是直接复制的历史信息,绿色是直接复制的输入信息,紫色是按LSTM方式融合的交集,白色是互不相关的全零状态。

基于这种更新方式,高层信息可能保留相当长的距离(高层区间直接复制历史,历史因此可能被不断复制而不改变),而低层信息在每一步都可能被重写(低层区间直接复制输入,输入不断改变)。不同分组的信息传播跨度不同,这些跨度就构成了输入序列的层级结构。

为了把两种情况合并成一个可微的公式,引入记号\(1_k\)表示第$k$位为$1$的one-hot向量,并定义cumsum操作:

\[\text{cumsum}([x_1,x_2,\cdots ,x_n]) =[x_1,x_1+x_2,\cdots, x_1+x_2+\cdots +x_n]\]

则主遗忘门(master forget gate)\(\tilde f_t\)用于标记$c_{t-1}$中的历史信息层级、主输入门(master input gate)\(\tilde i_t\)用于标记\(\hat c_t\)中的输入信息层级,两者的交集为\(w_t\):

\[\tilde{f}_t = \text{cumsum}(1_{d_f}), \quad \tilde{i}_t = 1- \text{cumsum}(1_{d_i}), \quad w_t = \tilde{f}_t \odot \tilde{i}_t\] \[c_t = w_t\odot (f_t \odot c_{t-1} + i_t \odot \hat{c}_t) + (\tilde{f}_t-w_t) \cdot c_{t-1} + (\tilde{i}_t-w_t) \cdot \hat{c}_t\]

one-hot向量\(1_{d_f}\)和\(1_{d_i}\)的构造过程是不可导的,根据函数的光滑化结论,one-hot向量的光滑近似为softmax函数。至此ON-LSTM的完整更新过程为:

\[\begin{aligned} i_t &= \sigma(W_{i}x_t+U_{i}h_{t-1}+b_i) \\ f_t &= \sigma(W_{f}x_t+U_{f}h_{t-1}+b_f) \\ o_t &= \sigma(W_{o}x_t+U_{o}h_{t-1}+b_o) \\ \hat{c}_t &= \text{tanh}(W_{c}x_t+U_{c}h_{t-1}+b_c) \\ \tilde{f}_t &= \text{cumsum}(\text{softmax}(W_{\tilde{f}}x_t+U_{\tilde{f}}h_{t-1}+b_{\tilde{f}})) \\ \tilde{i}_t &= 1- \text{cumsum}(\text{softmax}(W_{\tilde{i}}x_t+U_{\tilde{i}}h_{t-1}+b_{\tilde{i}})) \\ w_t &= \tilde{f}_t \odot \tilde{i}_t \\ c_t &= w_t\odot (f_t \odot c_{t-1} + i_t \odot \hat{c}_t) + (\tilde{f}_t-w_t) \cdot c_{t-1} + (\tilde{i}_t-w_t) \cdot \hat{c}_t \\ h_{t} &= o_t \odot \text{tanh}(c_t) \end{aligned}\]

实现细节:\(\tilde f_t,\tilde i_t\)表示的是层级,而序列的总层级数一般不大,因此这两个向量不应该太长;但它们需要与$f_t$逐元素相乘,而LSTM的隐状态维度通常较大。解决办法是把隐层维度分解为\(n=pq\),只构造$p$维的\(\tilde f_t,\tilde i_t\),再把其中每个元素重复$q$次。这既减少了层级总数,也减少了参数量。

层级结构的无监督提取:由于\(1_{d_f}\approx \text{softmax}(\cdot)\),层级\(d_f\)可由\(\arg\max\)给出;用\(\arg\max\)的光滑近似可以推出一个只依赖主遗忘门的闭式表达:

\[d_{f,t} \approx n+1-\sum_{k=1}^n \tilde{f}_t[k]\]

给定输入序列与训练好的ON-LSTM(例如一个语言模型),按上式得到层级序列\(\{d_{f,t}\}\),再用贪心算法递归切分:找出层级序列中最大值所在下标$k$,把序列分区为\([x_{t<k},[x_k,x_{t>k}]]\),然后对两个子序列重复该操作直到长度为$1$。其直觉是:层级最高处包含的历史信息最少、与前文联系最薄弱,最有可能是一个新子结构的开始。这样就能在没有句法标注的情况下从语言模型中析出成分句法树。

⚪ LMU:勒让德记忆单元

LMU给记忆施加的是连续时间的结构:它要求记忆向量$m_t$在任意时刻都是输入信号在最近$\theta$长度窗口内历史的勒让德多项式正交投影。也就是说,$m_t$不是“随便学出来的一组数”,而是“滑动窗口内历史函数的一组正交基系数”。

推导的起点是一个连续时间延迟系统\(F(s)=e^{-\theta s}\),它可以由如下ODE表示($\theta$为时间窗口长度):

\[\theta \dot{\mathbf{m}}(t) = A\mathbf{m}(t) + B\mathbf{u}(t)\]

用勒让德多项式构造的标准正交基\(g_n(t)=\sqrt{\frac{2n+1}{2}}L_n(t)\)去最优逼近窗口内的输入历史,即最小化

\[\mathop{\arg \min}_{c_0,\dots,c_N} \int_a^b \left[ \mathbf{u}(t) - \sum_{n=0}^N c_n g_n(t) \right]^2 dt\]

其最优系数为\(c_n = \int_a^b \mathbf{u}(t)g_n(t)dt\)。把积分区间改为随时间滑动的窗口\([T-\theta, T]\)并对$T$求导,再利用勒让德多项式的导数递推关系\(L_{n+1}'(s)=\sum_{k=0}^{n}(2k+1)\left[1-(n-k)\%2\right]L_k(s)\),就能把\(\{c_n(T)\}\)的演化整理成一个线性ODE,其系数矩阵(经适当缩放后)即为LMU使用的矩阵$A,B$。

LMU层的设计包括一个$n$维状态向量$h_t$和一个$d$维记忆向量$m_t$,它们通过线性和非线性变换动态耦合。记忆向量$m_t$表示输入$u_t$在勒让德多项式上投影的滑动窗口,从而实现历史信息的压缩和高效存储:

\[\begin{aligned} h_t &= f(W_xx_t+W_hh_{t-1}+W_mm_t) \\ m_t &= Am_{t-1}+Bu_t \\ u_t &= e_x^\top x_t+e_h^\top h_{t-1}+e_m^\top m_{t-1} \\ \mathbf{A}&=[A_{n,k}]\in\mathbb{R}^{d\times d}, \quad A_{n,k} = (2n+1)\begin{cases} -1, & n<k \\ (-1)^{n-k+1}, & n\geq k\end{cases} \\ \mathbf{B}&=[B_{n}]\in\mathbb{R}^{d\times 1},\quad B_n = (2n+1)(-1)^n \end{aligned}\]

其中\(W_x,W_h,W_m\)是可学习的权重矩阵,\(e_x,e_h,e_m\)是可学习的编码向量,而$A,B$是由上述推导解析给定的常数矩阵(可以固定不训练)。LMU是第一个能够处理跨越$10^5$个时间步依赖的循环架构,其原因正是记忆的更新是线性且理论上正交的:既不存在门控的乘性衰减,也不会因为学习而破坏基函数的正交性。

LMU在谱系上是连接门控RNN与状态空间模型的关键一环:把上式中“用勒让德基压缩历史”的思想一般化为对任意测度的最优投影,就得到了HiPPO框架,进而导出S4与Mamba。

⚪ 连续时间RNN与Neural ODE

沿着$1$节$(3)$中“RNN是ODE欧拉解”的联系继续走,可以把隐状态的演化直接写成一个由神经网络参数化的ODE:\(\dot h(t)=f_\theta(h(t),t)\),用数值积分器在任意时间点求解(Neural ODE);配合在观测时刻用GRU式门控做一次状态跳变,就得到能处理不规则采样时间序列的ODE-RNN / GRU-ODE / Latent ODE族。这类模型在医疗时序、物理系统辨识中很有用,代价是每步都要调用数值积分器,训练开销远高于离散RNN。

3.4 隐状态即模型:测试时训练

前面所有方法的隐状态都是一个向量或矩阵,更新规则由人工设计。本小节的思路是把这两者都换掉:隐状态是一个模型的权重,更新规则是一步梯度下降。

⚪ TTT:测试时训练层

这项工作的出发点是对RNN瓶颈的一个重新表述:RNN必须把无界的历史压缩进固定大小的隐状态,这个压缩过程会丢信息;而“如何把大量数据高效压缩进固定大小的参数、并保留深层语义关系”,正是机器学习训练本身在做的事情。既然如此,不如让隐状态就是一个正在被训练的模型。

测试时训练(Test-Time Training, TTT)层是一个改进的RNN层,其隐藏状态是一个可训练的机器学习模型\(f(\cdot;W_t)\)。作者提出了两个具体的TTT层实例:TTT-Linear(隐藏状态是一个线性模型)和TTT-MLP(隐藏状态是一个两层MLP)。

更新规则定义为对这个隐藏状态模型进行一步自监督学习的梯度下降。定义一个自监督损失$\ell$(比如去噪重构),计算它关于当前模型权重\(W_{t-1}\)的梯度并据此更新:

\[W_t = W_{t-1} - \eta \nabla \ell(W_{t-1}; x_t)\]

自监督任务的具体形式是:把$x_t$投影到一个低维空间造成“破坏”,再要求模型从破坏后的版本重构出$x_t$的另一个投影:

\[\ell(W;x_t) = \left\Vert f\left(\theta_K x_t;W\right) - \theta_V x_t \right\Vert^2\]

TTT层在时间步$t$的输出$z_t$,就是用刚刚更新过的隐藏状态模型对当前输入的查询投影做一次前向推理:

\[z_t = f\left(\theta_Q x_t; W_t\right)\]

其中\(\theta_K,\theta_V,\theta_Q\)是可学习的投影矩阵,它们定义了具体的自监督任务。这构成一个双层优化结构:内循环在TTT层内部用自监督损失优化隐状态权重$W_t$(推理时也在进行,故名“测试时训练”);外循环用整体目标(如next-token prediction)优化TTT层之外的所有参数以及\(\theta_K,\theta_V,\theta_Q\)。

朴素实现的效率很低,因为每个时间步都依赖上一步的权重。作者给出两项关键优化:

⭐ 讨论:隐状态在线学习——一条正在快速演化的路线

TTT最有价值的部分是它给出的统一视角:任何序列模型都可以被描述为“隐状态(一个学习器)+ 更新规则(一个优化器)”这一对选择。作者证明了两个等价性:当隐状态是线性模型且内循环采用批梯度下降时,TTT层在数学上等价于线性注意力;当“学习器”是非参数的Nadaraya-Watson估计器时,TTT层等价于标准的自注意力。于是RNN与注意力不再是两个模型,而是同一框架下学习器容量的两端。

这个视角立刻提示了大量设计空间:把内循环优化器从SGD换成带动量的、带权重衰减的、或Muon式的更新,就分别得到Titans(带“惊奇度”驱动的动量与遗忘的长期记忆模块)、DeltaNet / Gated DeltaNet(用delta rule做记忆的精确覆写而非累加)、TTT-Muon等一系列工作;把隐状态模型从线性换成MLP乃至更深的网络,就得到容量更大但并行更困难的变体。

3.5 门控循环网络的复兴

$2024$年前后出现了一批“回归门控RNN形式”的工作。它们的共同背景是:注意力的\(O(T^2)\)训练成本与\(O(T)\)推理缓存在长上下文下越来越难以承受,而并行扫描的成熟让“训练也能并行”的循环模型重新具备竞争力。与$3.2$节的动机相同,但这一次的目标是直接替换大语言模型中的注意力层。

⚪ xLSTM:扩展长短期记忆

xLSTM明确诊断了LSTM在大模型时代的三个缺陷,并逐一给出修改:

  1. 无法修正存储决策:LSTM的门是Sigmoid,一旦某个信息被写入就很难在后面被修正。$\to$ 用指数门控放开门的取值范围。
  2. 存储容量有限:所有信息被压进一个标量对\((c,h)\)的逐通道结构里,做检索类任务时容量不够。$\to$ 用矩阵记忆替换标量记忆。
  3. 无法并行:门依赖\(h_{t-1}\)。$\to$ 在矩阵记忆版本中去掉门对\(h_{t-1}\)的依赖。

据此得到两种单元。

sLSTM

sLSTM保留标量记忆与循环连接(因此仍然串行,但表达力最强),把输入门改为指数函数,并引入一个归一化状态$n_t$来抵消指数门带来的尺度膨胀:

\[\begin{aligned} c_t &= f_t \odot c_{t-1} + i_t \odot z_t \\ n_t &= f_t \odot n_{t-1} + i_t \\ h_t &= o_t \odot \tilde{h}_t, \quad \tilde{h}_t = c_t / n_t \\ z_t &= \varphi\left(w_z^\top x_t + r_z h_{t-1} + b_z\right) \\ i_t &= \exp\left(\tilde{i}_t\right), \quad \tilde{i}_t = w_i^\top x_t + r_i h_{t-1} + b_i \\ f_t &= \sigma\left(\tilde{f}_t\right) \text{ 或 } \exp\left(\tilde{f}_t\right), \quad \tilde{f}_t = w_f^\top x_t + r_f h_{t-1} + b_f \\ o_t &= \sigma\left(w_o^\top x_t + r_o h_{t-1} + b_o\right) \end{aligned}\]

注意\(h_t=o_t\odot (c_t/n_t)\)中已经没有$\tanh$:有界性由除以$n_t$的归一化保证。由于\(i_t\)可以取到\(e^{\tilde i_t}\)这样的大值,新写入的内容能够压过历史的累积,“修正存储决策”的能力就来自这里。

指数运算极易溢出,xLSTM为此引入稳定器状态(stabilizer state) \(m_t\),在对数域上追踪门控累积乘积的最大指数:

\[\begin{aligned} m_t &= \max\left(\log f_t + m_{t-1},\ \log i_t\right) \\ i_t' &= \exp\left(\log i_t - m_t\right) = \exp\left(\tilde{i}_t - m_t\right) \\ f_t' &= \exp\left(\log f_t + m_{t-1} - m_t\right) \end{aligned}\]

把\(i_t,f_t\)替换为\(i_t',f_t'\)后,\(c_t\)与\(n_t\)被同一个因子\(e^{-m_t}\)缩放,因此比值\(c_t/n_t\)完全不变,数值上却被限制在安全范围内。这是“\(\log\)-域softmax减最大值”这一经典技巧在循环形式下的推广,也与$3.2$节讨论的并行扫描对数域实现同源。

mLSTM

mLSTM把记忆从向量升级为矩阵\(C_t \in \mathbb{R}^{d\times d}\),用协方差式(外积)更新规则写入键值对:

\[\begin{aligned} q_t &= W_q x_t + b_q, \quad k_t = \frac{1}{\sqrt{d}}W_k x_t + b_k, \quad v_t = W_v x_t + b_v \\ C_t &= f_t C_{t-1} + i_t\, v_t k_t^\top \\ n_t &= f_t n_{t-1} + i_t\, k_t \\ h_t &= o_t \odot \frac{C_t q_t}{\max\left(\lvert n_t^\top q_t \rvert,\ 1\right)} \end{aligned}\]

读出操作\(C_t q_t\)恰好是“用查询\(q_t\)检索之前写入的所有值”,因此mLSTM在形式上与线性注意力同构;分母\(\max(\lvert n_t^\top q_t\rvert,1)\)的下界$1$保证了当查询与所有历史键都不匹配时输出不会被放大。关键的一点是:mLSTM的门\(i_t,f_t,o_t\)只依赖\(x_t\),因此这个递推是\(C_t=f_tC_{t-1}+i_tv_tk_t^\top\)的形式,可以用并行扫描(或分块的“块内并行、块间循环”实现)在训练时完全并行化。记忆容量从$O(d)$提升到\(O(d^2)\),这正是它在联想检索任务上强于LSTM的原因。

把这两种单元套上残差块(sLSTM用post-up-projection的LSTM式块,mLSTM用pre-up-projection的Transformer式块)并堆叠起来,就得到xLSTM架构。实践上通常大部分层用mLSTM(并行、容量大),少量层用sLSTM(提供状态跟踪能力)。

⚪ minLSTM与minGRU:RNN本来就够了吗

这项工作提出了一个反向的问题:如果只是想让门控RNN能并行训练,最少需要删掉什么?答案出乎意料地简单:只要把门对\(h_{t-1}\)的依赖去掉,剩下的部分就自动落进$3.2$节讨论的并行扫描形式。

minGRU在GRU基础上做两步删减:$1)$ 去掉重置门\(r_t\)(它的唯一作用就是把\(h_{t-1}\)引入候选状态);$2)$ 让更新门与候选状态都只依赖\(x_t\):

\[\begin{aligned} z_t &= \sigma\left(\text{Linear}(x_t)\right) \\ \tilde{h}_t &= \text{Linear}(x_t) \\ h_t &= \left(1-z_t\right) \odot h_{t-1} + z_t \odot \tilde{h}_t \end{aligned}\]

minLSTM同样删掉门对\(h_{t-1}\)的依赖,去掉输出门与记忆状态,并把遗忘门与输入门归一化为一组和为$1$的系数(替代LSTM中隐含的凸组合约束,保证状态尺度与时间无关):

\[\begin{aligned} f_t &= \sigma\left(\text{Linear}(x_t)\right), \quad i_t = \sigma\left(\text{Linear}(x_t)\right) \\ f_t' &= \frac{f_t}{f_t+i_t}, \quad i_t' = \frac{i_t}{f_t+i_t} \\ \tilde{h}_t &= g\left(\text{Linear}(x_t)\right) \\ h_t &= f_t' \odot h_{t-1} + i_t' \odot \tilde{h}_t \end{aligned}\]

两者都是\(h_t=a_t\odot h_{t-1}+b_t\)的形式(其中\(a_t,b_t\)只依赖\(x_t\)),因此可以用并行扫描在\(O(\log T)\)深度内训练;作者在$T=512$时报告了约\(175\)倍的训练加速。为了在扫描中安全地取对数,候选状态使用了值域为正的\(g(\cdot)\)(Softplus式的连续正函数)而不是$\tanh$,并整个在对数域实现。

需要注意这两个模型的取舍:门不再依赖\(h_{t-1}\)意味着模型失去了“根据当前状态决定如何更新状态”的能力,理论上无法完成需要状态跟踪(如奇偶校验、有限状态机模拟)的任务。这篇论文的价值更多在于基线校准:在语言建模等任务上,这样极简的门控RNN已经能与Mamba等复杂设计打成平手,说明近年的大量收益可能来自并行化与训练配方,而非结构本身。

⭐ 讨论:循环、卷积与注意力的收敛

把$3.2$、$3.4$、$3.5$三小节放在一起看,会发现三条原本独立的技术路线正在收敛到同一个数学对象上:

因此当前活跃的设计空间可以用三个正交的选择来描述:(a) 状态的形状(向量 / 矩阵 / 一个小网络的权重);(b) 状态的更新规则(乘性遗忘 / 外积累加 / delta rule / 带动量的梯度步);(c) 门控信号的依赖(只依赖\(x_t\)则可并行,依赖\(h_{t-1}\)则表达力更强但必须串行)。“RNN vs Transformer”这个二分法已经基本失效,取而代之的是在这三个维度上的权衡。

4. 深层、双向与其他结构设计

上一节讨论的是单个时间步内部的结构,本节讨论把循环单元组织成完整网络的方式。

(1) 堆叠循环神经网络 Stacked RNN

深层RNN通过增加循环神经网络的深度(即堆叠循环层的数量)增强循环神经网络的特征提取能力,即增加同一时刻网络输入到输出之间的路径。堆叠循环神经网络(Stacked RNN)是将多个循环网络堆叠起来,第$l$层在时刻$t$的输入是第$l-1$层在时刻$t$的输出。

需要区分RNN中的两种“深度”:时间深度(沿$t$展开,可达数千步)与空间深度(沿层堆叠)。前者由权值共享决定、天然很深;后者才是通常意义上的网络深度。实践中循环网络的空间深度远小于CNN,一般只用\(2\sim 4\)层(语音识别中可到\(5\sim 8\)层),因为层数增加同时会拉长梯度路径,收益迅速饱和。

(2) 双向循环神经网络 Bidirectional RNN

双向循环神经网络(Bidirectional RNN)由两层循环神经网络组成,它们的输入相同,只是信息传递的方向不同:一层从$t=1$向前,一层从$t=T$向后,最后把两个方向的隐状态拼接作为输出。

双向结构让每个时刻的表示都同时包含左右上下文,在语音识别、序列标注、机器阅读理解等整条序列已知的任务上几乎总是优于单向。但它有两个硬性限制:不能用于自回归生成(会泄露未来信息),不能用于流式/在线场景(必须等到序列结束才能计算反向通路)。因此语言模型只用单向RNN,而BERT式的双向表示需要靠掩码语言建模而非双向循环来实现。

(3) 残差与高速连接

沿空间深度方向堆叠时同样会出现梯度问题,解决办法与CNN中相同:在层与层之间添加残差连接\(h^{(l)}_t = h^{(l-1)}_t + \text{RNN}^{(l)}(h^{(l-1)}_t)\)或高速连接(highway connection)\(h^{(l)}_t = g\odot h^{(l-1)}_t + (1-g)\odot \text{RNN}^{(l)}(h^{(l-1)}_t)\)。深层RNN(如Google NMT的$8$层LSTM编码器)离不开这类跨层连接;$3.2$节的SRU把高速连接直接内置进了单元定义。

(4) 循环网络专用的正则化

标准Dropout不能直接用在循环连接上:如果每个时间步都独立采样掩码,噪声会沿时间轴累积$T$次,把长程信息彻底破坏。循环网络因此发展出两种专门的方案(更完整的正则化谱系见深度学习中的正则化方法)。

⚪ Variational / Recurrent Dropout:跨时间步共享掩码

从变分贝叶斯的视角看,Dropout近似的是对权重的后验采样;既然循环网络在所有时间步共享同一份权重,那么掩码也应当在整条序列上保持不变:

\[h_t = f\left(W_{xh}\left(x_t \odot z_x\right)+W_{hh}\left(h_{t-1} \odot z_h\right)+b\right), \quad z_x, z_h \sim \text{Bernoulli}(1-p)\]

关键在于\(z_x,z_h\)每条序列采样一次、在\(t=1,\dots,T\)中复用。这样噪声等价于对权重矩阵的一次扰动,不会随时间累积。实践中通常还对输入/输出的嵌入层使用同一掩码(embedding dropout)。这就是AWD-LSTM等强基线中“weight-dropped LSTM + variational dropout”配方的来源。

另一种同类思路是recurrent dropout:只对候选状态(LSTM的\(\tilde c_t\))施加Dropout,而不触碰\(c_{t-1}\)的恒等通路,这样也不会破坏长程梯度。

⚪ Zoneout:随机保持上一时刻的状态

Zoneout把Dropout的“随机置零”换成“随机保持不变”:以概率\(p\)让某个单元直接沿用上一时刻的值,否则按正常规则更新。若记正常的更新结果为\(\tilde h_t\)、掩码为\(d_t\sim\text{Bernoulli}(p)\),则

\[h_t = d_t \odot h_{t-1} + \left(1-d_t\right)\odot \tilde{h}_t\]

对LSTM则分别对记忆状态与隐状态使用独立掩码:

\[\begin{aligned} c_t &= d_t^c \odot c_{t-1} + \left(1-d_t^c\right)\odot \left(f_t \odot c_{t-1}+i_t\odot \tilde{c}_t\right) \\ h_t &= d_t^h \odot h_{t-1} + \left(1-d_t^h\right)\odot \left(o_t \odot \text{tanh}(c_t)\right) \end{aligned}\]

Zoneout相比Dropout的关键优势在于:被“zone out”的单元其梯度是恒等地传给上一时刻的(\(\partial h_t/\partial h_{t-1}=1\)),因此它不但不破坏、反而缩短了梯度沿时间的传播路径;可以看成随机版本的跳跃连接。它也可以视为一个随机的遗忘门(把\(f_t\)随机地置为$1$)。常用配置是对记忆状态用较大的概率(\(\approx 0.5\))、对隐状态用较小的概率(\(0.05\sim 0.2\));Zoneout与Dropout可以叠加使用。

(5) 编码器-解码器结构与注意力的由来

RNN最有影响力的应用形态是序列到序列(Sequence to Sequence, Seq2Seq)的编码器-解码器结构:用一个RNN(编码器)把输入序列读成一个固定长度的向量\(c=h_T^{\text{enc}}\),再用另一个RNN(解码器)以$c$为初始状态自回归地生成输出序列。它把“变长到变长”的映射统一了起来,直接催生了神经机器翻译、文本摘要、图像描述、语音识别等任务的端到端方案。按输入输出的对齐方式,RNN的任务形态可以分为:一对多(图像描述)、多对一(文本分类、情感分析)、同步多对多(序列标注、逐帧分类)、异步多对多(Seq2Seq,即机器翻译)。详见序列到序列模型。

正是这个结构暴露出了一个致命瓶颈:无论输入多长,都必须压缩进同一个固定维度的向量$c$。句子越长,翻译质量下降越明显。解决办法是让解码器在每一步都回头去看编码器的全部隐状态\(\{h_1^{\text{enc}},\dots,h_T^{\text{enc}}\}\),并按当前解码状态计算一组权重做加权求和——这就是注意力机制的起源(Bahdanau 等人的神经机器翻译工作)。此后的发展路径是:注意力先作为RNN的补丁出现,随后被发现“注意力本身就足够了”,RNN主干被彻底移除,于是有了自注意力与Transformer。有趣的是,第$3.5$节的工作又从反方向回到了这里:把注意力写成循环形式以换取线性复杂度。注意力机制的公式细节见序列到序列模型中的注意力机制与自注意力机制。