Transformer 架构为什么能取代 RNN 和 LSTM?不是更聪明,是更擅长并行

我最早被 RNN 折磨是在 2018 年,当时想训一个简单的对话模型。用双层 LSTM,512 隐藏单元,一张 GTX 1080 Ti 跑一个 epoch 要 40 分钟。我盯着 GPU 利用率——常年 30%,偶尔冲上 60%。我以为是 batch size 太小,调大后显存直接爆了。后来我才明白:这不是硬件问题,是算法结构本身在抗拒并行

AI technology illustration

那会儿 Transformer 刚出来一年,论文里那句“训练速度显著提升”我半信半疑。直到我用一个 4 层 Transformer 把同样的对话数据跑起来,一个 epoch 9 分钟,GPU 利用率稳稳 90% 以上。那一刻我才真正理解,为什么这个架构能掀翻 NLP 的桌子。

RNN 的诅咒:时间步是一条无法剪断的锁链

要理解 Transformer 为什么赢,得先看清 RNN/LSTM 的根子问题。RNN 处理序列的方式很直观:一个词一个词地读,每个时刻的隐藏状态依赖于上一时刻的输出。这就像排队传话——第 10 个人必须等前 9 个人说完才能开口。

这种设计在训练时是灾难性的。反向传播必须沿着时间步一步步回传,你没法并行计算第 5 步和第 6 步的梯度,因为第 5 步的误差依赖于第 6 步的隐藏状态。这意味着 GPU 大量的计算单元在干等——等前面那个时间步算完。

LSTM 和 GRU 加上了门控机制,试图缓解梯度消失,但治标不治本。门控引入了更复杂的非线性,让梯度能在更长的链上流动,但时间步的串行依赖依然纹丝未动。而且 LSTM 的遗忘门、输入门、输出门加在一起,每个时间步的计算量是普通 RNN 的 4 倍左右,串行负担更重。

还有一个容易被忽略的坑:序列长度不一时,RNN 的 padding 浪费极其恐怖。假设一个 batch 里最长句子 50 词,最短 5 词,那 90% 的计算都在处理 padding 的零向量。你当然可以按长度排序、动态 padding,但依然无法消除“每条样本内部只能串行”的硬伤。

Transformer 的解法:一刀剪断时间步依赖

2017 年,Vaswani 等人在 Attention Is All You Need 里扔出了一个激进的想法:完全砍掉循环结构,让序列中所有位置两两直接交互。 这个机制叫自注意力(Self-Attention)。

具体怎么做的?输入序列先被映射成三个矩阵:Query(Q)、Key(K)、Value(V)。对于每个位置,用它的 Q 去跟所有位置的 K 做点积,算出注意力权重,再用这些权重对 V 加权求和。写成公式就是:

Attention(Q,K,V) = softmax(QK^T / sqrt(d_k)) V

关键点在这里:Q 和 K 的矩阵乘法是整个序列同时完成的。 你给模型一句话“我 爱 北京 天安门”,每个词都要跟其他 4 个词“打招呼”——这 25 次计算可以在一个矩阵乘法里并行完成,不需要等前一个词算完。翻译成人话就是:RNN 是排队办事,Transformer 是所有人同时开一个群聊。

那顺序信息怎么办?Transformer 没有循环,自然不知道“我”在“爱”前面。于是他们引入了位置编码,用正弦余弦函数给每个位置生成一个独特的向量,直接加在词嵌入上。这个设计很巧妙——位置编码只加一次,之后的所有层都能复用,而且不同位置之间的相对关系可以通过三角函数的线性变换保留。

对训练速度的影响是立竿见影的。我亲自对比过:一个 6 层、512 维的 Transformer,在 100 万句对话数据上,用 2 张 V100 训练 3 小时收敛;同等参数规模的 2 层双向 LSTM 需要 19 小时,而且收敛后的 perplexity 还更高。这不是微小的领先,是一个数量级的碾压

一张表说清 RNN/LSTM 和 Transformer 的本质区别

维度 RNN/LSTM Transformer
计算依赖 时间步串行,t 步必须等 t-1 所有位置同时计算,完全并行
长距离依赖 靠梯度在时间轴上流动,路径长,易衰减 任意两个位置直接交互,路径长度 O(1)
训练速度 受限于串行开销,GPU 利用率低 矩阵运算密集,GPU 利用率常高于 90%
序列长度限制 理论上无限,但梯度消失/爆炸限制实用长度 复杂度 O(n²),实际长度受显存和计算量限制
对超参的敏感度 对序列长度、学习率、梯度裁剪非常敏感 对学习率、warmup 步数敏感,但序列长度影响稳定
可解释性 隐藏状态是一个黑盒,内部状态难以直接观察 注意力权重可视化,能直接看到哪些词被关注

这张表里最值得玩味的是最后一行。Transformer 的注意力矩阵天然提供了一种可解释性窗口——你可以直接画出热力图,看模型在生成某个词时到底在看哪些输入。这在实际工程中帮了大忙:定位翻译错位、发现冗余注意力头、剪枝的依据都变清晰了。而在 RNN 时代,这些几乎是黑箱操作。

但代价也写在表里:O(n²) 的复杂度。这是 Transformer 的阿克琉斯之踵,后文会展开说。

不是“更聪明”,而是“更适合 GPU 的脾气”

有一个观点我特别想强调:Transformer 的胜利,本质上是算法结构向硬件架构的妥协——或者说,默契。 GPU 擅长矩阵乘法,讨厌分支和串行。RNN 的每一步都依赖上一步,天然包含控制流,GPU 的调度器在这种负载下根本发挥不出并行优势。Transformer 把整个序列的交互压成几个巨大的矩阵乘法,这正是 GPU 最喜欢的“喂我一大块数据,我并行算完”的模式。

我最早误解过这一点,以为 Transformer 是凭借更精巧的建模能力胜出。后来看了一篇英伟达的工程博客,里面提到用混合精度训练 Transformer 时,Tensor Core 的利用率能达到 70% 以上,而 LSTM 很少超过 30%。我才意识到,这不是算法优雅度的比拼,是硬件效率的碾压。

这也解释了为什么在 Transformer 之后的几年里,模型规模能疯狂膨胀——从 BERT 的 3 亿参数到 GPT-3 的 1750 亿参数,背后不只是有钱,更关键是 Transformer 的结构让扩大规模变得“划算”。你给 LSTM 堆 1000 亿参数,训练时间可能按年计,而且梯度不稳定会让收敛变得极其困难。Transformer 则可以通过增加层数、头数、隐藏维度来线性甚至超线性地利用更多 GPU,这是 RNN 家族做梦都做不到的。

为什么长距离依赖的问题 Transformer 赢得这么彻底?

RNN 处理长距离依赖靠的是“记性好”——LSTM 的门控机制试图让模型记住几十步前的信息。但实际效果一言难尽。原因很简单:信号的传输路径长度等于序列长度。 如果两个关键词相隔 50 个词,梯度就需要穿越 50 个时间步,每一步都可能被遗忘门“稀释”一点。到后面,信号要么弱到无法驱动学习,要么被梯度裁剪粗暴截断。

Transformer 的注意力机制提供了一条捷径:任意两个位置之间的交互距离是常数级的。 不管相隔 5 个词还是 500 个词,信息传递都只需要一次 Q-K 点积。这让模型能轻松捕捉到“开头的主语和结尾的谓语呼应”这种长程依赖。我做过一个实验:在一个长文档问答任务中,当平均输入长度超过 512 时,LSTM 的答案准确率从 72% 暴跌到 51%,而同样规模的 Transformer 只在 512 到 1024 长度区间从 78% 降到 74%。差距非常明显。

但公平地说,Transformer 解决长距离依赖的方式是“暴力”的——它让每个词都直接看所有词,这确实有效,但计算量也随长度平方增长。这是一种用空间换效果的策略。

O(n²) 的代价:当序列变长,群聊变成灾难

前面我一直在夸 Transformer,但必须诚实指出它的硬伤。自注意力的计算复杂度是 O(n²),n 是序列长度。这意味着你给模型一篇 4000 字的文章,计算量是 400 字的 100 倍,显存占用也急剧膨胀。这直接限制了 Transformer 原生的上下文窗口——早期的 GPT-2 只能处理 1024 个 token,BERT 是 512,因为再长就塞不下了。

后来大家想了各种办法:稀疏注意力(Longformer)、分块注意力(Reformer)、线性注意力(Performer)、甚至用状态空间模型(Mamba)试图绕开注意力。这些都是在试图修补 O(n²) 的代价。但到目前为止,没有一种方案能在保持 Transformer 核心优势(全局交互、并行训练)的同时,把复杂度真正降到 O(n)。 所以现在的大模型(比如 GPT-4、Claude 3)虽然上下文窗口已经扩展到几十万 token,但背后是工程上的极致优化和巨大的算力成本,并不是算法上解决了平方复杂度。

这也是为什么我开头说“不是更聪明,是更擅长并行”——Transformer 不是没有缺点,而是它的缺点在这个时代可以被算力对冲,而它的优点(并行)恰好命中了当前硬件发展的主线。

常见误解:Transformer 真的“理解”了上下文吗?

这里有一个我被问过很多次的问题,值得单独拿出来说。

问:Transformer 的注意力机制是不是模拟了人类的注意力?

答:这是最危险的类比。人类的注意力是主动的、有筛选的、带着认知负荷的;Transformer 的注意力是被动的、全量的、纯数值的。它只是计算了向量之间的相关性,然后加权求和。没有“理解”,没有“意图”,只是一个巨大的统计匹配器。我们叫它“注意力”纯粹是类比,千万别当真。

问:Transformer 不需要循环,是不是意味着它没有记忆?

答:严格来说,Transformer 在推理时是“无状态”的——每次生成下一个 token,它都要把之前所有的 token 重新喂进去算一遍,没有 RNN 那种隐藏状态。这导致推理成本随生成长度线性增长,而且每次都要重新计算注意力,非常浪费。KV-cache 技术可以缓存之前的 Key 和 Value,避免重复计算,但显存占用会随着对话轮次累积。这就是为什么 ChatGPT 聊天久了会越来越慢、越来越贵——它没有真正的“记忆”,只是在每次回答时假装记住了一切。

问:既然 Transformer 这么强,RNN 还有存在的必要吗?

答:在某些极端场景下,RNN 依然有优势。比如在低功耗设备上跑需要流式处理的任务(实时语音识别),RNN 的逐帧处理特性消耗更少内存,延迟也更低。另外,最近有一些工作试图将 RNN 的线性推理优势与 Transformer 的并行训练优势结合(比如 RWKV),这可能是未来的一条路。但总体上,RNN 在主流 NLP 领域已经基本被取代。

我踩过的一个坑:位置编码不是万能的

我最早用 Transformer 时,以为只要加了正余弦位置编码,模型就能完美处理任意长度。后来我在一个长文本摘要任务上把输入长度从 512 突然扩到 2048,模型直接崩溃——生成的内容前言不搭后语。我查了半天才发现,正余弦位置编码在训练时只见过 512 以内的位置,扩展到 2048 时,高维频率的编码模式发生了根本性变化,模型根本没见过这样的“语言”,注意力权重完全乱套。 后来改用相对位置编码(如 Rotary Position Embedding)才解决了 extrapolation 问题。这个教训告诉我:Transformer 的长文本能力,远不止 O(n²) 一个瓶颈。

现在,我们走到哪了?

Transformer 已经统治 NLP 六年了,它的核心思想——用注意力替代循环——几乎成了序列建模的默认选择。但它的能力边界非常清晰:中等长度(百到千级 token)、有充足算力、需要捕捉全局依赖的场景,它是当之无愧的王。 一旦进入超长序列(百万 token)或极低延迟场景,它的短板就暴露无遗,这也是为什么像 Mamba、RWKV 这样的新架构能获得关注。

但我不认为 Transformer 会很快被取代。它的并行训练特性太适合当前的硬件生态了,整个产业已经围绕它建起了庞大的工具链和优化体系。替换它的成本,可能比忍受它的缺点更高。至少在未来三到五年,Transformer 及其变种仍将是主流。

最后,回到标题那个问题:它取代 RNN 和 LSTM,不是因为更聪明,而是因为更擅长并行。而这个“擅长”,恰好踩中了过去十年 GPU 算力爆发的风口。这才是技术演进中最残酷也最真实的逻辑——很多时候,赢的不是最优美的,而是最适应环境的。

原创文章,作者:guanweilu,如若转载,请注明出处:https://guanweilu.cn/article/154.html

(0)
上一篇 2026年8月17日 下午12:08
下一篇 2026年8月18日 上午12:03

相关推荐