近期一篇关于“无递归预训练循环网络”的论文在开发者社区引起了不小的讨论,点赞与收藏热度上升得很快。很多读者第一眼看到这个标题都会疑惑:循环网络本身不就是带递归结构的吗?怎么还能“无递归”?预训练不是 Transformer 的强项吗,为什么又回到循环网络上来了?本文就用一篇系统教程的方式,把这条技术路线拆开讲清楚。我们会从递归与循环的概念边界说起,分析传统循环网络在规模化训练中的瓶颈,再对比三条“去递归化”的改造思路,最后用 PyTorch 写一个把时间步递归改写成并行卷积扫描的完整示例。无论你是刚开始接触预训练模型的新手,还是想了解循环网络新进展的工程师,都可以通过这篇文章建立一条清晰的知识脉络。
1. 背景与核心概念
1.1 递归与循环:先厘清两个容易混淆的术语
在讨论“无递归预训练循环网络”之前,先要分清“递归”和“循环”这两个词。在计算机科学里,“递归”通常指一个函数直接或间接调用自身,典型的例子是阶乘、树的遍历;而“循环”则指程序结构上的重复执行,比如 for 循环、while 循环。两者在语义上有交集:很多时候递归函数的执行过程可以改写成循环,反之亦然。
在深度学习里,循环网络(RNN)之所以被称为“循环”,是因为它会在时间维度上反复使用同一个单元处理每个时间步的输入。到了具体实现阶段,我们最直观的写法就是按时间步做循环:
for t in range(seq_len): h = cell(x[t], h)如果换用函数自调用的方式实现,就变成了递归:
def forward_t(t, x, h): if t == 0: return h h = cell(x[t], forward_t(t - 1, x, h)) return h所以“递归”并不等于“循环网络”,它只是循环网络在实现上可以采用的一种写法。理解了这一点,才能明白“无递归循环网络”这个名称的微妙之处:它并不是要彻底抛弃循环网络的状态建模思想,而是要摆脱“按时间步逐步展开”的递归式实现方式,让计算不再受时间步顺序的强约束。
1.2 循环网络为什么要预训练
预训练是近年来深度学习最核心的范式之一。它的基本思路是先用大规模无标注数据学习通用的语言或序列表示,再通过微调适配具体下游任务。预训练语言模型(如 RoBERTa、BERT、GPT 系列)之所以强大,是因为它们能在海量文本上学习到语法、语义和世界知识,让下游任务只需要很小的数据就能取得不错的效果。
传统循环网络虽然在序列建模场景中表现稳定,但过去很长一段时间并没有像 Transformer 那样成为预训练模型的主流架构。主要原因并不是循环网络不能预训练,而是它的训练效率存在明显瓶颈。RNN 的时间步计算具有前后依赖关系,很难像 Transformer 那样对整段序列做大规模并行计算。于是当算力成为约束条件时,研究者更倾向于选择并行性更好的架构。
但这并不意味着循环网络的预训练路线不值得探索。循环网络有一个独特优势:它的隐状态是固定维度的压缩记忆,推理时的显存开销不随输入长度线性增长。如果能解决训练并行化问题,循环网络在长序列、低资源推理场景中会很有竞争力。
1.3 无递归预训练循环网络是什么
“无递归预训练循环网络”可以理解为一类改进型的循环网络,它保留了循环网络“用隐状态记忆历史信息”的建模思想,但在计算方式上尽量减少或彻底移除“按时间步递归展开”的操作。
具体来说,这类改进通常体现在三个方面:
- 训练时不逐个时间步计算,而是用并行扫描、卷积等算子一次算出整段序列的输出。
- 推理时通过状态缓存或迭代求解,避免递归式地重复展开计算图。
- 结构上使用线性递归、门控线性单元等设计,让状态更新公式更易于并行化。
这种做法的直接收益是:模型既能像 RNN 一样用固定大小的状态压缩整段历史,又能像 Transformer 一样在大规模数据上高效训练。因此,这条路线的论文在社区获得高赞并不意外,它踩中了一个真实痛点:长序列建模的“性价比”之争。
2. 传统循环网络的计算瓶颈
2.1 经典 RNN 的数学形式
为了理解无递归化的动机,先回顾经典 RNN 的更新公式。对于一个输入序列 (x_1, x_2, \dots, x_T),RNN 在每个时间步维护一个隐状态 (h_t),计算方式如下:
[ h_t = \tanh(W_{hh} h_{t-1} + W_{xh} x_t + b) ]
这个公式看起来非常简单,但它决定了 RNN 的核心性质:当前时刻的隐状态必须等前一个时刻的隐状态计算完成后才能得出。这就像一条流水线,每一步都依赖上一步的输出,无法随意跳步,也无法把整条流水线拆成互不相关的并行任务。
如果我们要把这个公式展开成完整的计算图,图的最长路径长度与序列长度成正比。也就是说,随着输入序列变长,计算图的深度线性增长。这种结构带来两个直接问题:训练速度慢、梯度传播困难。
2.2 按时间步展开的三个问题
第一个问题是训练效率低。Transformer 可以一次性处理一整段 token,而 RNN 必须进行 (T) 次串行前向计算。在 GPU 上,串行依赖意味着大量运算核之间需要等待,硬件利用率难以提升。
第二个问题是梯度消失与梯度爆炸。反向传播时,误差信号需要沿着时间维度逐层回传。由于每一步都要乘以 (W_{hh}),如果权重矩阵的谱半径小于 1,梯度会指数级衰减,序列稍微长一点,早期信息就无法影响最终输出;如果谱半径大于 1,梯度又会指数级增大,导致训练不稳定。
第三个问题是长距离依赖建模能力受限。理想情况下,模型需要记住很久以前的输入信息,但因为梯度消失,经典 RNN 往往只能建模较短距离的依赖,这与预训练场景所要求的长期上下文建模能力相悖。后来虽然出现了 LSTM、GRU 这类带门控的模型,但串行计算的本质并没有改变。
2.3 长序列场景下的现实困境
当序列长度从几百扩展到几千、几万甚至百万级别时,传统 RNN 的串行开销会变得非常明显。假设一段 10 万 token 的文本,训练时就要执行 10 万次连续的前向步骤。哪怕每一次计算量很小,累计起来也相当耗时。
Transformer 通过自注意力机制解决了并行性问题,但代价是注意力矩阵的空间复杂度随序列长度呈平方增长。对于超长序列,Transformer 的显存开销也很难承受。于是我们看到一条清晰的技术分歧:RNN 状态效率高但难并行,Transformer 易并行但高复杂度。无递归预训练循环网络的思路,就是想站在两种架构的交汇点上,取一个折中方案。
3. 无递归化的三条技术路线
3.1 并行扫描:把时间递归变成可并行归约
并行扫描(parallel scan)是一种经典的并行算法,也被称为前缀和(prefix sum)或 scan 操作。它的核心思想是:如果状态更新满足某种可结合性,那么原本串行的递归计算就可以改写成树状归约,从而把 (O(T)) 的串行步骤压缩为 (O(\log T)) 层并行操作。
一阶线性递归就是一个典型例子:
[ h_t = a_t \cdot h_{t-1} + x_t ]
如果系数 (a_t) 是常数 (a),可以推导出闭式解:
[ h_t = \sum_{i=0}^{t} a^{t-i} x_i ]
这个式子不再包含递归依赖,每一个 (h_t) 都是输入序列的加权求和。于是我们可以通过卷积、矩阵乘或并行前缀和算法一次性计算出全部 (h_t),整个计算过程中不再需要“等待上一步结果”。
3.2 固定点迭代:用不动点近似替代显式展开
另一条去递归思路是固定点迭代。对于某些循环网络,状态在理论上会收敛到不动点,即存在某个稳定状态 (h^*),使得:
[ h^* = f(h^*, x) ]
传统做法是从初始状态开始逐步展开 (T) 步,逼近这个不动点。无递归化改造可以改为:直接对不动点方程做迭代求解,每次迭代不按输入时间步展开,而是重复应用状态更新函数直到收敛。这种思路在深度均衡模型(Deep Equilibrium Model)中被广泛使用,它的好处是前向传播的“深度”由收敛条件决定,不再受输入序列长度限制,反向传播也可以通过隐函数定理完成,不需要保存每一层的中间变量。
当然,固定点迭代的计算代价取决于收敛速度。如果模型设计得不好,迭代次数可能很多,反而比直接展开更慢。它更适合那种内部动力学本身具有收缩性质的网络。
3.3 结构去递归:线性 RNN 与状态空间模型
第三种路线从结构设计入手。既然递归计算难以并行,那就在设计网络结构时就限制递归的形式,只保留“可并行化”的状态更新。
状态空间模型就是一类代表性工作。它把序列建模看作连续系统的离散化,状态更新写成一个线性形式,并结合 HiPPO 等矩阵初始化方法保证长距离记忆。这类模型的数学形式可以借助卷积或并行扫描高效计算。
近年来社区讨论热度较高的 RWKV、线性注意力等方向也属于广义上的“去递归化”探索。它们共同的特点是把传统 RNN 的递归更新公式改写为近似线性的形式,再引入非线性激活或门控机制来弥补表达能力的损失。这类模型往往在推理时保留固定大小的隐状态,因此部署成本远低于 Transformer,同时训练时又能利用并行扫描获得较好效率。
3.4 与 Transformer 预训练模型的对比
无递归预训练循环网络并不是要完全替代 Transformer。在超大规模参数、海量数据、多模态等场景下,Transformer 依然是目前最成熟的底座。但两类架构的定位逐渐分化:
Transformer 擅长在大规模并行训练下学习极其复杂的模式,推理时却需要缓存每一层的完整键值对,长序列场景的开销不容忽视。无递归循环网络则在推理时只维护一个固定大小的状态,序列再长,显存占用也基本恒定。
因此,一个更现实的判断是:未来长序列建模可能走向混合架构。用循环网络或线性注意力做高效的状态压缩,用局部注意力或卷积捕捉细粒度信息,再配合预训练范式获得通用表示。无递归化改造,正是让循环网络能够进入“预训练俱乐部”的关键一步。
4. 实战:从串行递归到并行扫描
4.1 环境准备
这一节我们用 PyTorch 演示“无递归化”的核心思想。示例环境如下:
- 操作系统:Ubuntu 20.04 或 Windows 10/11 均可
- Python:3.8 或更高版本
- PyTorch:2.x 版本(1.13 及以上基本兼容)
- CUDA:如有 GPU 可用,训练速度会更快;CPU 也可以运行本示例
项目结构保持简单,不需要复杂工程框架:
no_recursion_rnn/ ├── linear_rnn.py └── example.py如果本地还没有安装 PyTorch,可以参考官方安装命令。以 pip 为例:
pip install torch需要注意,版本需要根据你的项目实际情况调整,本文示例以常见环境为例,重点演示配置思路,不强制绑定某一个具体版本。
4.2 目标:计算一阶线性递归
我们先定义本实验的目标:给定一个输入序列 (x = [x_0, x_1, \dots, x_{T-1}]) 和衰减系数 (a \in (0,1)),计算一阶线性递归:
[ h_t = a \cdot h_{t-1} + x_t ]
并且假设初始状态 (h_{-1} = 0)。这个公式虽然简单,却包含了循环网络最核心的时间依赖结构。我们先用最直观的串行递归实现,再改写为并行扫描实现,通过对比你会发现两者的效果完全一致,但计算方式截然不同。
4.3 串行递归版本
串行版本最容易理解:按时间步遍历,每一步都在上一步的结果上更新状态。代码如下:
# 文件路径:linear_rnn.py import torch import torch.nn.functional as F def linear_rnn_serial(x: torch.Tensor, a: float) -> torch.Tensor: """ 串行版本:按时间步递归展开。 参数: x: 输入序列,形状为 [T] a: 衰减系数,取值范围建议 (0, 1) 返回: h: 隐状态序列,形状为 [T] """ T = x.shape[0] h = torch.zeros_like(x) prev = torch.zeros_like(x[0]) for t in range(T): prev = a * prev + x[t] h[t] = prev return h这段代码非常直观,但它的执行过程是严格串行的:(t=10) 的隐状态必须等 (t=9) 的隐状态算完才能开始。如果我们把 (T) 放大到 10 万,这个循环就会成为明显的性能瓶颈。
4.4 并行扫描版本
观察一阶线性递归的闭式解:
[ h_t = \sum_{k=0}^{t} a^k \cdot x_{t-k} ]
这个公式可以看作“输入序列 (x) 与固定核 ([1, a, a^2, \dots, a^{T-1}]) 之间的因果卷积”。因果卷积的意思是:输出位置 (t) 只和输入位置 (0) 到 (t) 有关,不依赖未来信息。卷积运算天然支持并行,于是我们可以用torch.nn.functional.conv1d一次性计算出全部隐状态。
# 文件路径:linear_rnn.py import torch import torch.nn.functional as F def linear_rnn_scan(x: torch.Tensor, a: float) -> torch.Tensor: """ 并行扫描版本:通过因果卷积一次计算整段序列。 参数: x: 输入序列,形状为 [T] a: 衰减系数,取值范围建议 (0, 1) 返回: h: 隐状态序列,形状为 [T] """ T = x.shape[0] kernel = (a ** torch.arange(T, dtype=torch.float32)).view(1, 1, T) # conv1d 的 padding = T - 1,保证输出长度不少于输入长度,再截取前 T 个位置 h = F.conv1d(x.view(1, 1, T), kernel, padding=T - 1) return h.view(-1)[:T]这段代码里最关键的是卷积核的构造。a ** torch.arange(T)生成序列 ([a^0, a^1, \dots, a^{T-1}])。在因果卷积中,它恰好表示“距当前时间步间隔 (k) 的输入,其权重为 (a^k)”。
用padding=T-1是为了让卷积输出长度与输入对齐,最后通过切片[:T]去掉多余部分。如果你直接运行并对比linear_rnn_serial和linear_rnn_scan的结果,会发现完全一致。
4.5 封装为可学习的无递归循环层
上面的例子中 (a) 是手工指定的常数。真实项目中我们会希望模型自己学习合适的衰减系数,因此可以把它封装成一个可学习的网络层。为了让模型更稳定,我们通常学习log_a而不是直接学习a,从而保证衰减系数始终为正且小于 1。
# 文件路径:linear_rnn.py import math import torch import torch.nn as nn import torch.nn.functional as F class LinearRecurrentLayer(nn.Module): """ 可学习的无递归循环层。 每个特征通道学习独立的衰减系数,通过因果卷积并行计算整段序列的隐状态。 输入形状: [batch, seq_len, d_model] 输出形状: [batch, seq_len, d_model] """ def __init__(self, d_model: int, init_decay: float = 0.9): super().__init__() self.d_model = d_model self.log_decay = nn.Parameter( torch.full((d_model,), math.log(init_decay)) ) def forward(self, x: torch.Tensor) -> torch.Tensor: batch, seq_len, dim = x.shape # 通过 exp 保证衰减系数为正;log_decay 小于 0 时,a 小于 1,数值稳定 decay = torch.exp(self.log_decay) # [d_model] # 构造逐通道卷积核: [d_model, 1, seq_len] positions = torch.arange(seq_len, dtype=torch.float32, device=x.device) kernel = decay.pow(positions).view(1, seq_len, -1) # [1, seq_len, d_model] kernel = kernel.transpose(0, 2).unsqueeze(1) # [d_model, 1, seq_len] # 转成 conv1d 需要的布局: [batch, d_model, seq_len] x_t = x.transpose(1, 2) # 使用 groups=d_model 实现逐通道因果卷积 h = F.conv1d(x_t, kernel, padding=seq_len - 1, groups=self.d_model) # 截取有效长度并恢复为 [batch, seq_len, d_model] h = h[:, :, :seq_len].transpose(1, 2) return h这里有几个实现要点需要解释:
groups=self.d_model表示每个输入通道单独做卷积,不同特征之间互不干扰。padding=seq_len - 1与前面一致,保证输出长度足够,最后再截断。- 卷积核的构造虽然看起来形状复杂,但核心逻辑仍然是生成 ([a^0, a^1, \dots, a^{T-1}]),只不过在每个特征通道上独立生成一份。
decay.pow(positions)允许 PyTorch 自动求导,log_decay这个参数会通过卷积运算获得梯度。
这个层本质上是一个“线性循环层”,它没有 for 循环,没有时间步递归展开,整个计算可以在 GPU 上并行完成。
4.6 运行与验证
下面写一段测试代码,验证串行版本和并行版本输出一致,并观察可学习层的前向输出形状。
# 文件路径:example.py import torch from linear_rnn import linear_rnn_serial, linear_rnn_scan, LinearRecurrentLayer def test_linear_rnn(): torch.manual_seed(42) x = torch.randn(10) a = 0.9 h_serial = linear_rnn_serial(x, a) h_scan = linear_rnn_scan(x, a) print("串行版本:", h_serial) print("并行版本:", h_scan) print("最大误差:", torch.max(torch.abs(h_serial - h_scan)).item()) def test_learnable_layer(): torch.manual_seed(42) layer = LinearRecurrentLayer(d_model=8) x = torch.randn(2, 10, 8) h = layer(x) print("输入形状:", x.shape) print("输出形状:", h.shape) print("学习到的 log_decay 初值:", layer.log_decay) if __name__ == "__main__": test_linear_rnn() test_learnable_layer()预期输出大致如下(log_decay会随训练变化,这里只看初始状态):
串行版本: tensor([...]) 并行版本: tensor([...]) 最大误差: 0.0 输入形状: torch.Size([2, 10, 8]) 输出形状: torch.Size([2, 10, 8]) 学习到的 log_decay 初值: Parameter containing: tensor([-0.1054, -0.1054, ...], requires_grad=True)最大误差 = 0.0说明并行扫描版本和串行递归版本在数学上完全等价。做到这一步,我们实际上已经实现了一个最简单的“无递归循环层”:它保留了状态更新的记忆逻辑,但不再按时间步递归展开。
实际训练时,可以把这个层放在 Embedding 层之后,再接一个线性分类头做文本分类,也可以堆叠多层并在中间加入 LayerNorm 和非线性激活函数。需要说明的是,这里是教学示例,用于帮助理解原理,离生产级模型还有一定距离。生产环境如果要使用线性循环结构,建议参考相关开源实现或论文,并结合具体硬件做优化。
5. 常见问题与排查思路
在实现“无递归循环网络”相关代码时,大家容易遇到下面几类问题。我把高频问题整理成一张表格,方便排查。
| 问题现象 | 常见原因 | 解决思路 |
|---|---|---|
| 并行版本与串行版本结果不一致 | padding 设置错误,或卷积核方向取反 | 检查 padding 是否为 seq_len - 1,核对核序列是否为 a^0 到 a^(T-1) |
| 训练时梯度爆炸 | 衰减系数大于 1,状态指数增长 | 使用 log_decay 参数化,限制衰减系数始终小于 1 |
| 长序列下显存不足 | 卷积核随序列长度线性增长,中间变量占用过大 | 改用分块扫描、线性注意力或状态空间模型实现 |
| 可学习层无法回传梯度 | 手工构造了不可导的核张量,或对核调用了 detach | 确认核是从 nn.Parameter 计算得到,不要 detach |
| 推理时状态没有衰减,输出接近无穷 | 未对初始状态 h0 做处理,或 a 设为 1 | 明确初始状态并加入衰减项;若需要记忆长期信息,可叠加门控机制 |
| conv1d 输入输出形状报错 | 忘记调整维度布局,或 groups 参数不匹配 | 输入统一为 [batch, channels, seq_len],weight 形状为 [channels, 1, kernel_size] |
除了表格里的问题,还有一个细节值得强调:如果初始状态 (h_{-1}) 不为 0,那么并行扫描的公式需要额外补偿。比如对于一阶线性递归,非零初始状态的闭式解是:
[ h_t = a^{t+1} \cdot h_{-1} + \sum_{k=0}^{t} a^k \cdot x_{t-k} ]
也就是说,需要在每个输出位置加上 (a^{t+1} h_{-1})。很多初学者在改造 RNN 时忽略了这一点,导致训练和推理时初始状态带来的信息丢失或重复叠加。一个稳妥的做法是:把初始状态也看作可学习参数,并在前向计算中显式加入补偿项。
6. 最佳实践与工程建议
6.1 参数化与数值稳定性
在工程实现中,不要直接学习原始衰减系数 (a),而是像前面代码那样学习 (log_a)。这样做的原因是:
- 不限制范围时,(a) 可能被梯度更新到大于 1,导致状态爆炸。
- 使用
exp(log_a)可以把衰减系数约束到 ((0, +\infty)),再配合初始化让 (log_a < 0),就能保证衰减系数处于 ((0,1))。 - 在混合精度训练场景下,这种参数化方式更稳定,也更好做梯度裁剪。
如果模型需要保存长期记忆,可以引入类似 LSTM 的“遗忘门”机制,但注意门的计算也要尽量保持线性或半线性形式,避免破坏并行扫描的结构。
6.2 初始状态与边界处理
初始状态的建模是循环网络工程中很容易遗漏的点。建议把 (h_{-1}) 初始化为零向量,并在文档中明确说明公式假设;如果需要让模型学习初始状态,可以把它定义为nn.Parameter,并在前向计算时通过广播加到所有输出位置。另外,裁剪卷积输出时要注意截断位置,避免把 padding 区域的值误当成有效隐状态。
6.3 长序列与显存控制
并行扫描虽然训练效率高,但直接为每个序列长度构建一个 kernel,在长序列场景下依然有显存压力。工程上常用的优化方式包括:
- 分段扫描:把长序列切成多个块,块内并行扫描,块间传递状态。
- 固定核长度:设置最大依赖距离,超过距离的信息通过额外状态补充。
- 使用 FFT 加速卷积:对于极长序列,FFT 的复杂度为 (O(T \log T)),比直接卷积更有优势。
- 混合架构:在局部使用卷积或注意力,在全局使用线性递归状态压缩。
这些优化手段并不冲突,可以根据业务规模和硬件条件组合使用。
6.4 与预训练范式结合
如果想把“无递归循环层”应用到预训练任务中,一个常见的做法是采用“自回归语言建模”目标:给定前文 token 预测下一个 token。由于线性循环层的计算是可并行扫描的,我们可以高效地处理大规模语料。训练时还可以把输入通过 Embedding 投影到隐藏维度,再经过多个LinearRecurrentLayer,最后接输出层预测词表分布。
不过有两点需要注意:
- 纯线性递归的表达能力有限,建议在层间加入非线性激活、LayerNorm 或门控机制。
- 预训练模型对优化器、学习率、梯度裁剪、权重初始化都比较敏感,建议从较小规模模型开始验证,再逐步放大参数与数据量。
对于希望深入研究的读者,可以先从“并行前缀和算法”“状态空间模型”“线性注意力”等关键词入手。这些方向与无递归预训练循环网络高度相关,也是当前长序列建模研究的热点。
7. 总结与学习路线
7.1 核心收获
通过这篇文章,我们主要掌握了以下知识点:
第一,递归与循环并不是同一个概念,循环网络“按时间步展开”的实现方式才是并行化的最大障碍。第二,传统 RNN 的串行计算、梯度消失、长距离依赖建模困难,是无递归化改造的直接动因。第三,无递归化的三条主流路线分别是并行扫描、固定点迭代和结构去递归。第四,在 PyTorch 中,一阶线性递归完全可以通过一维因果卷积实现并行计算,串行版本与并行版本在数学上等价。第五,工程落地时要重点关注衰减系数参数化、初始状态补偿、长序列显存控制和训练稳定性。
7.2 下一步学习方向
如果你对这条技术路线感兴趣,建议按下面的顺序继续学习:
- 先掌握并行前缀和(scan)算法,理解为什么可结合的递归可以被并行化。
- 再阅读状态空间模型相关的公开资料,理解线性状态更新、离散化和卷积表示之间的关系。
- 接着尝试复现一个简单的线性注意力或线性循环语言模型,在开源数据集上做小规模预训练实验。
- 最后研究混合架构,把局部注意力、卷积和全局状态压缩结合起来,寻找更适合自己业务场景的序列模型。
7.3 动手实践建议
学习算法和架构最好的方式,是亲自实现一个最小版本并观察它的行为。你可以从本文的代码出发,尝试修改衰减系数初始化、叠加多层、加入非线性激活函数,然后在文本分类或字符级语言建模任务上测试。不要急于追求 SOTA 效果,先确保自己能清晰地解释“为什么这个模型可以并行”“梯度从哪里来”“状态如何更新”这三个问题。
如果本文对你有帮助,可以收藏备用,也欢迎在评论区交流实现过程中遇到的问题。理论是骨架,代码是血肉,只有亲手跑通示例并尝试修改,才能真正理解无递归预训练循环网络的价值。