从RNN到Transformer:序列建模的演进逻辑与工程启示
先说明一下我自己在自然语言处理这个方向摸爬滚打了快十年从当年用RNN做机器翻译、文本分类到后来看着Attention机制一步步从“辅助模块”变成“主力架构”再到Transformer横扫一切。这一路走过来最大的感受是如果只学一个“RNN好用”或者“Attention好用”的结论那很快就会过时。真正有价值的是搞清楚——为什么RNN会走向瓶颈为什么Attention能补上它的短板以及Transformer又是凭什么把Attention推向了极致。这篇文章就想把这套演进逻辑完整地拆一遍适合刚接触深度学习序列建模的开发者也适合那些已经在用大模型、但想回头补一补基础知识的工程师。这篇文章会从RNN的核心设计动机讲起接着分析它在长程依赖上的天然缺陷再拆解Attention到底在做什么、为什么它有效最后看Transformer如何把Attention变成唯一的主角。文中的公式我会控制在必要的范围尽量用直白的方式解释清楚大家跟着走一遍就能建立完整的认知框架。1. 先把RNN的“循环”讲透它到底在循环什么很多人一上来就背RNN的结构图输入x、隐藏状态h、输出y然后一堆箭头……但从来没想明白过一个更基础的问题——为什么要“循环”1.1 序列建模的核心困境一句话里每个词都不是孤立的做自然语言处理或者说做任何序列数据的建模最绕不开的一个事实是数据点之间是有顺序依赖的。比如“我在北京吃烤鸭”这句话如果打乱顺序变成“吃烤鸭我在北京”意思就变了。“北京”这个词之所以重要不只是因为它是一个地名而是因为它在“在”的后面、在“吃烤鸭”的前面这个位置关系本身承载了语义。用全连接网络做这件事会非常尴尬。如果你把一句话的所有词一次性塞进一个固定维度的向量里等于假设“第几个位置”是固定的但句子长短不一、词序千变万化这个假设根本不成立。更重要的是全连接网络学不到“第3个词受第1个词影响”这种时间维度上的依赖关系——因为它处理所有输入的方式是一样的没有先后之分。RNN的设计动机就是为了解决这个问题它希望模型能够像一个正在阅读的人那样读到第5个词的时候还记得第1个词和第3个词带来的上下文信息并且用这个上下文来辅助理解当前的词。1.2 隐藏状态RNN把“记忆”变成了一个可计算的向量RNN最核心的概念是隐藏状态h你可以把它理解成一个“记忆容器”。在每一个时间步t模型做两件事读取当前输入x_t然后结合上一个时间步传下来的隐藏状态h_{t-1}更新出一个新的隐藏状态h_t。公式非常简单h_t tanh(W_h · h_{t-1} W_x · x_t b)这个公式里W_h是处理“旧的记忆”的权重W_x是处理“当前输入”的权重b是偏置。tanh是激活函数把结果压缩到-1到1之间保证数值稳定。如果你把这个公式展开来看就会得到一个很有意思的链条h_1只由x_1决定h_2由x_2和h_1决定所以h_2里面其实包含了x_1的信息h_3由x_3和h_2决定所以h_3里面包含x_1和x_2的信息……以此类推。这就是RNN中被反复提到的“展开结构”——从结构图上看它是一个循环但展开后就是一条从左到右的链每一个时间步都在“吞”新的输入同时把历史信息往下传。我在很多教程里看到他们把RNN的展开图画得特别复杂各种箭头绕来绕去。其实你只需要抓住一点展开后的RNN本质上是一个“按时间顺序共享参数的深层网络”每一层的参数完全一样只是输入不同。你不需要管它那些花哨的变体理解这个主干就够了。1.3 为什么参数共享对RNN如此重要这里有个很容易被忽略的细节就是RNN在所有时间步用的是同一组参数W_h、W_x、b而不是每个时间步一套独立参数。这个设计是刻意的原因也很实际。第一是参数量可控。如果序列长度是50每个时间步都用不同的参数那就是50套独立的权重模型体积直接爆炸而且训练数据根本不够支撑这么多参数收敛。参数共享意味着RNN学到的不是一个“针对某个位置的规则”而是一个“通用的更新规则”——不管输入出现在句子的第几个位置模型都用同一套逻辑去更新记忆这天然符合语音、文本、时间序列这类数据的规律。第二是模型更不容易过拟合。参数越少模型的自由度和表达能力就越受限但反过来看这种限制本身就是一种正则化。在处理长度可变的序列时参数共享保证模型不会“记位置”而只能“学规律”泛化能力会好很多。不过参数共享也有它的问题既然每个时间步都在用同一组参数不断更新同一个状态向量h_t那信息在传递过程中就必须反复穿过同一个非线性变换。这里就埋下了一个巨大的雷——梯度消失。下面这段是整条演进路线里最关键的一环。2. 长程依赖这道坎为什么LSTM依然没能彻底解决它2.1 梯度沿着时间反向传播为什么连乘会出事先说清楚梯度消失到底怎么来的。RNN训练用的是BPTTBackpropagation Through Time时间反向传播核心思想是把展开后的网络当成一个普通深层网络然后用标准反向传播计算梯度。假设我们在时间步t算出了损失L_t要把这个损失传递回时间步t-k去更新参数。中间经过k步每一步都要对h求一次偏导。在最简化的情形下这部分梯度约等于k个“W_h^T乘以一个对角矩阵”的连乘。也就是说梯度在时间维度上要连续乘以同一个矩阵W_h^T一共k次。如果W_h的特征值小于1连乘k次之后梯度会指数级趋近于0如果特征值大于1梯度又会指数级爆炸。梯度爆炸可以靠梯度截断来勉强压住但梯度消失是真正的噩梦——因为梯度消失意味着远距离的词对当前词的参数更新几乎贡献不了任何信号。模型想学“我在北京住了十年所以我很喜欢这座城市”里的“北京”与“这座城市”之间的指代关系但梯度传到“北京”那儿的时候已经衰减成零了模型根本没法学。这就是RNN“记不住长距离信息”的数学根源。你可能会想把W_h初始化大一点行不行不行因为你把W_h放大到大于1之后短距离的梯度就爆炸了。这是一个两难的困境本质原因是RNN只有一条“高危通道”来传递信息既不安全也不可靠。2.2 LSTM的应对之道给梯度开一条“高速公路”LSTMLong Short-Term Memory长短期记忆网络的出发点非常直接既然信息经过太多层非线性变换会丢失那我干脆专门开一条“传送带”让信息可以原封不动地往前传。这条“传送带”就是记忆单元c_t。LSTM的更新方式不再是把所有信息都压进一个状态里反复非线性变换而是用三个门控来控制信息的写入、遗忘和输出遗忘门f_t决定上一个时刻的记忆c_{t-1}有多少要保留下来输入门i_t决定当前候选的新信息c̃_t有多少要写入记忆输出门o_t决定当前记忆c_t有多少要输出给隐藏状态h_t其中记忆单元的更新公式是c_t f_t ⊙ c_{t-1} i_t ⊙ c̃_t注意看这个公式c_t到c_{t-1}的路径是一条逐元素乘法的捷径。如果遗忘门恰好接近1那么梯度就可以几乎无损地沿着这条路径传回很远之前的时间步。这就像是在高速公路上开了一条应急车道堵车梯度消失的问题一下子缓解了很多。我在实际项目中用过LSTM做序列标注直观感受是对于20到50个token左右的中短序列LSTM的记忆能力比原生RNN好一个量级。具体到任务上比如根据前文预测后文原生RNN经常只能记住最近的3到5个词LSTM则可以稳定记住10到20个词。这个差距在真实任务里非常明显。2.3 GRU的精简和RNN家族的共同天花板GRUGated Recurrent Unit门控循环单元是LSTM的一个精简变体把三个门缩减成了两个——重置门和更新门参数更少计算更快在很多中小规模任务上性能和LSTM基本持平。如果你做时间序列预测、行为识别这类任务GRU往往是更快落地的选择。但这里我想说的是LSTM和GRU只是把长程依赖问题的“症状”缓解了并没有根治。门控机制让信息有了更顺畅的传递路径但序列终究还是串行处理的h_t必须等h_{t-1}算完才能算GPU再强也帮不上忙因为时间步之间有严格的数据依赖。我举一个直观的数字对比训练一个标准的Seq2Seq翻译模型如果用LSTM做编码器一个batch的数据要按时间步顺序跑很多次。当时我们在一张V100上用LSTM训练一个中等规模的英德翻译模型一个step要跑大概0.5秒换成Transformer之后同样的batch size和训练步数速度能提升好几倍。这个差距直接决定了你在实际工程里能不能快速迭代实验。还有一个更致命的问题——信息瓶颈。这句话我先放这儿下一章细讲。3. Attention的出现它没有替代RNN而是补上了最要命的短板3.1 神经机器翻译里的那个“固定向量”瓶颈2014年Sutskever等人提出了经典的Seq2Seq模型用两个RNN分别作为编码器和解码器。编码器负责把整个源句子压缩成一个固定维度的向量通常是最后一步的隐藏状态解码器再从这个向量出发一个词一个词地生成译文。这个方法在当时已经比统计机器翻译强很多了但它有一个非常让人难受的限制不管源句子是5个词还是50个词编码器都只能把它压成一个长度固定的向量。这个向量就像一个“万能口袋”你必须把所有信息都塞进去但口袋的容量是有限的。句子越长信息丢失得越严重。我当时做过实验用Seq2Seq翻译句子当源句子长度超过20个词时翻译质量就会明显下滑长句子经常出现漏译、重复翻译的情况。这个现象后来被很多文章都验证过本质上就是因为编码器输出的那个固定向量是信息的“瓶颈”——解码器每一步都只能从这个固定向量里取信息它不知道哪部分信息对应源句的哪个位置。3.2 Attention的数学形式query、key、value到底在做什么Bahdanau在2015年提出Attention机制的思路非常优雅与其把所有信息都塞进一个固定向量不如让解码器在生成每一步时自己决定“去源句的哪些位置找信息”。Attention机制里你会反复听到三个词query、key、value。这套定义最早其实是借鉴了数据库和检索系统的概念后来成为Transformer的基本语言。这里我用最简单的方式解释query当前解码器“想要查找什么信息”。在翻译任务里它就是当前要生成目标词时解码器所处的状态。key源句里每个位置“能够提供什么信息”的标签。每个源词对应一个key。value源句里每个位置“实际拥有的信息内容”。Attention的计算分三步用query和每个key计算相似度得分常见的打分函数有加性注意力Bahdanau和点积注意力Luong。点积形式最简单score(q, k) q^T k对这些得分做softmax归一化得到一组和为1的权重系数α_i表示“当前查询应该以多大比例关注这个源位置”用这些权重对所有value做加权求和得到最终的context向量数学形式是context Σ_i α_i · v_i其中α_i softmax(score(q, k_i))。3.3 从对齐到软寻址Attention的本质是信息检索如果你还在想“Attention到底解决了什么问题”我给你一个本质性的答案Attention把“编码”和“检索”解耦了。在纯RNN的Seq2Seq里编码和信息提取是杂糅的——编码器需要把所有信息都装进一个隐藏状态里解码时只能被动地从这个状态里拿。而Attention让解码器可以“主动地”去访问源句的每一个位置就像你在图书馆查资料一样你有明确的问题query根据图书馆的索引目录key找到可能相关的书然后翻开书找到你需要的内容value。这个视角下Attention机制可以被理解为一种软寻址——它不是直接从某个位置取信息而是对所有位置做加权平均权重取决于query和key的相关程度。这样既保留了“精确定位”的能力又因为有加权求和这个平滑操作梯度可以顺畅地传到所有源位置。我在理解Attention的时候心里一直有一个类比纯RNN的编码器就像一个人听了一段话之后被人要求把整段话背下来再转述而Attention机制则是这个人在转述时可以随时回头“看笔记”需要用到哪段就翻到哪段去快速查阅。这个差距在面对长文本时是决定性的。还要特别提一句Attention最初不是用来替代RNN的而是作为RNN decoder的辅助模块出现的。当时所有的主流模型还是在用RNN做序列建模Attention只是负责帮助解码器“看得更远、看得更准”。但谁也没想到这个“辅助模块”最后会喧宾夺主成为整个深度学习架构的核心原语。4. Transformer的赌注彻底抛弃循环依赖之后发生了什么4.1 Self-Attention让每个位置都能“看见”全部2017年Google的Vaswani等人发表了《Attention is All You Need》这篇论文的标题本身就是一场豪赌如果我们把Attention作为唯一的计算原语彻底不要RNN结果会怎样Transformer的核心设计是Self-Attention自注意力。所谓“自注意力”就是在一个序列内部做注意力计算——序列里的每一个token都和其他所有token计算相关性再根据相关性更新自己的表示。注意这里不再是解码器在查询源句而是序列中的每个token都在查询其他所有token。在普通的RNN中一个token要“看到”它前面的所有token必须通过隐藏状态一个接一个地把信息传下来距离越远传递损耗越大。而在Self-Attention里不管你离得远还是近都直接通过注意力权重建立联系——“我”要关注“谁”完全由query和key的匹配决定没有中间商赚差价。举个例子“The animal didnt cross the street because it was too tired”这里的“it”指代的是什么RNN需要沿着序列慢慢推理而Self-Attention可以让“it”这个位置的query直接和“animal”的key做匹配学到“it → animal”的高注意力权重。这种“处处可直达”的特性正是大模型语境下反复提到的“全局感知能力”。它让模型不再受限于循环结构带来的“距离衰减”为学习长距离依赖提供了全新的路径。4.2 并行化的巨大红利与位置编码的必要性RNN在训练时让人最头疼的问题就是串行依赖——第t个时间步必须等前t-1个时间步全部算完才能计算。这意味着无论你GPU有多强有效计算利用率上不去。而Self-Attention的计算本质是一系列矩阵乘法所有位置之间的相关性可以一次性并行算出GPU的并行算力能够被完全利用起来。举一个我实际经历过的实验数据同样训练一个Transformer-base模型和LSTM Seq2Seq模型在相同的翻译任务上Transformer的训练速度更快因为并行度高并且在BLEU分数上大幅领先。这也是Transformer能快速取代RNN成为主流架构的直接原因——它可能不是第一个想到“全局建模”的模型但它是第一个把并行效率做到极致的模型。不过抛弃RNN换来一个副作用Self-Attention本身对“顺序”是盲的。你把序列里的token调换顺序Attention计算的权重完全不变因为矩阵乘法是对称的。所以Transformer必须显式地给每个token注入位置信息也就是Positional Encoding位置编码让模型知道“这个词在第几个位置”。位置编码的设计也是有讲究的。论文里用的是正弦余弦函数pos是位置索引i是维度索引PE(pos, 2i) sin(pos / 10000^(2i / d_model)) PE(pos, 2i1) cos(pos / 10000^(2i / d_model))为什么用三角函数而不用简单的“第几个位置就编码成几”因为sin/cos的编码可以支持模型外推到更长的序列而且相对的“位置差”比如pos_a与pos_b的差可以表示为两个位置编码向量的线性组合这给了模型感知“相对位置”的能力。后来很多工作发现绝对位置编码和相对位置编码各有优劣但最初的正弦位置编码确实解决了一个很基础的需求让无状态的自注意力知道“先后顺序”。4.3 Multi-Head、残差、LayerNormAttention落地成模块的完整拼图Transformer里还有一个容易被初学者跳过但实际至关重要的设计——Multi-Head Attention多头注意力。这个设计的动机很直观基础的Self-Attention只有一组query、key、value相当于用一把“尺子”去度量所有位置的关系。但实际上句子里的关系是多元的有的位置关系体现语法依赖有的体现指代关系有的体现语义相关性。用一组q、k、v可能只能捕捉其中一种关系。Multi-Head Attention的做法是把query、key、value分别投影到h个子空间比如h8在每个子空间独立做Attention计算再把h个结果拼接起来做一次线性变换。这样模型就能并行地学习多种不同的“关系模式”互不干扰。这有点像CNN里的多个卷积核每个核关注不同的特征模式。另外两个配套设计也极为重要残差连接。Transformer每一层的输出都会和输入相加即output x Sublayer(x)。没有残差连接深层网络在反向传播时梯度经过很多层后仍然容易消失或变得不稳定训练深层Transformer几乎不可能收敛。残差连接让梯度可以“跳过”子层直接回传这是一个非常基础的工程细节但很多人刚开始自己搭Transformer时不接残差训练一上来就NaN或者loss震荡十有八九就是这里出了问题。LayerNorm。它把每一层的输出归一到均值为0、方差为1的分布上。相比BatchNormLayerNorm不依赖batch大小在变长序列和在线推理时更稳定。到这里一个完整的Transformer Encoder block由三件事构成Multi-Head Self-Attention、残差连接和LayerNorm、前馈网络FFN。前馈网络对每个位置的token独立做两轮线性变换加ReLU激活给模型补充非线性表达能力。还需要提一句Attention效率的后续故事。Self-Attention的时间和空间复杂度都是O(n²)输入序列一长计算量和显存占用就急剧上升。这也是后来FlashAttention、稀疏注意力、线性注意力等一系列工作要解决的问题。很多人现在用大模型处理长文本时会遇到“上下文窗口不够用”的困境本质上就是Attention的O(n²)复杂度在物理资源上的限制在起作用。工程上想要继续扩展序列长度必须在注意力计算方式上做文章这就是题外话了但理解这一点对于使用现代大模型非常重要。5. 回头看这条演进路线几个值得深思的工程启示5.1 RNN并没有“过时”它依然在某些场景里更合适每次讲RNN到Attention的演进总会有人直接得出“RNN已经被淘汰了”的结论。作为一个实际做过多个NLP项目的工程师我想说这个结论太粗暴了。RNN尤其是LSTM、GRU在以下场景里依然有不可替代的价值小规模数据。Attention模型尤其是Transformer是大规模参数的“饕餮”动辄几千万上亿参数数据量不够的时候很容易过拟合。RNN参数量相对小在小数据集上反而能学得更稳。时间序列预测。像股票走势、电力负荷、传感器读数这类数据序列长度通常不会特别长且没有明显的“长距多跳依赖”LSTM/GRU的归纳偏置天然按时间顺序建模往往比Transformer更契合训练成本还低很多。低延迟推理场景。RNN推理时不需要缓存整句的q、k、v只需要维护一个隐藏状态。在流式语音识别、在线推荐等低延迟场景里RNN类模型的内存访问模式和计算量都更可控。强顺序依赖的生成任务。有些生成任务要求输出一个严格从左到右、依赖关系非常紧密的序列这时RNN天生的“自回归状态推进”方式反而比Transformer的“全体并行”更自然。我个人的经验是先估摸你的数据规模和任务性质再决定上不上Transformer不要为了追新而追新。5.2 O(n²)复杂度是Attention的软肋Attention机制用“全量两两交互”换来了全局建模能力但代价就是平方级的计算复杂度。在Transformer处理长文档、长视频、高分辨率图像时这个问题尤其明显。工程师做长文本任务时经常要在序列长度上做截断本质就是因为显存装不下完整的注意力矩阵。工程上有几条路线可以缓解局部注意力Local Attention把序列切成固定大小的块每个token只和块内的token做注意力计算把复杂度降到O(n)。稀疏注意力Sparse Attention预先定义一些“重要”的注意力模式比如局部窗口加全局token只计算这些模式的位置代表性工作是Longformer和BigBird。线性注意力Linear Attention把softmax拆解成核函数的形式让Attention的计算从O(n²)降低到O(n)但通常会损失一部分表达能力。FlashAttention通过分块计算和IO优化在不改变Attention数学定义的前提下大幅降低显存占用和计算时间。现在大模型预训练几乎都离不开这套优化。如果你以后要做大模型推理优化或者在长序列场景里调模型这些方向几乎绕不开。理解Attention的复杂度瓶颈是你判断优化方向的第一块基石。5.3 我自己踩过的几个坑给后来者的一些实操建议最后分享几点我在学习和工程实践中踩过坑之后沉淀下来的经验希望对大家有用。第一学习Attention时不要一上来就啃Multi-Head的实现代码。先把单头的Attention公式手写一遍用一个小矩阵验证输入的shape变化再去看多头实现。很多人刚开始做Transformer的代码实现q、k、v拆成多个头之后维度对不上各种报错就是因为单头的逻辑没吃透。第二调试Transformer时如果loss不降先检查残差连接和LayerNorm的位置。Post-Norm和Pre-Norm两种结构在训练稳定性上有明显差异很多新手从论文里抄来一个block结果训练超过一定深度就崩。如果你想快速稳定训练用Pre-LN先LayerNorm再做Attention通常更稳。第三做长文本任务前先算一下显存预算。Self-Attention的显存占用是batch_size × num_heads × seq_len² × head_dim序列长度从512涨到1024显存需求直接翻四倍不是两倍。上线前先把这条算清楚能省去很多痛苦。第四RNN和Attention不是水火不容的。在实际系统里你可以用Attention来增强RNN的Decoder正是Bahdanau的原始思路也可以用CNN或RNN来压缩序列长度再喂给Transformer很多混合架构在实际业务里表现相当好。理解两者的优劣和使用边界比单纯追求“最新架构”要重要得多。回头再看“从RNN到Attention”这整条演进路线它其实讲了一个很朴素但深刻的道理一个结构能不能成为主流取决于它在当前的硬件、数据规模和任务需求下是否给出了最合适的“信息流动方式”。RNN用“串行记忆”传递信息Attention用“全量直连加权检索”传递信息Transformer则把后者做到极致并抛弃了循环结构。理解了这个道理你看到未来出现更新的结构时就不会觉得是在追一个又一个新名词而是能看懂它到底在优化信息流动的哪个环节。