为什么需要 Transformer:RNN 的问题


文档摘要

为什么需要 Transformer:RNN 的问题 本节摘要:2017 年之前,地球上每一个登顶的序列模型——语言、翻译、语音——都是循环神经网络(RNN)。它们快、它们能跑,但它们有三个致命弱点:串行计算让 GPU 上 99% 的算力在长序列上空转、梯度消失让 50 个 token 之前的信息被压成渣、固定宽度的隐状态把整条源序列硬塞进一个向量。2017 年的《Attention Is All You Need》做了一个激进的决定:彻底丢掉循环,让每个位置同时与所有其他位置「互相看见」。这一个架构层面的赌注,改写了 2017 年之后深度学习的每一条缩放曲线。

为什么需要 Transformer:RNN 的问题

本节摘要:2017 年之前,地球上每一个登顶的序列模型——语言、翻译、语音——都是循环神经网络(RNN)。它们快、它们能跑,但它们有三个致命弱点:串行计算让 GPU 上 99% 的算力在长序列上空转、梯度消失让 50 个 token 之前的信息被压成渣、固定宽度的隐状态把整条源序列硬塞进一个向量。2017 年的《Attention Is All You Need》做了一个激进的决定:彻底丢掉循环,让每个位置同时与所有其他位置「互相看见」。这一个架构层面的赌注,改写了 2017 年之后深度学习的每一条缩放曲线。本节将带你用数值实验亲身体会 RNN 与 Transformer 在**依赖深度(Dependency Depth)**上的鸿沟,理解为什么注意力是「广播」而非「接力」,以及代价是什么——O(N²) 的显存墙。读完本节,你能说清为什么 2026 年 Transformer 统治了所有模态,以及什么场景下仍然该选 RNN 或状态空间模型(SSM)。

学习目标

阅读完本节,你应当能够:

  1. 说清 RNN 的三个致命弱点:串行计算、梯度消失、固定宽度隐状态瓶颈,以及它们各自在工程上的后果。
  2. 解释 Transformer 用并行广播式注意力替代循环后,为什么训练速度从 O(N) 串行深度降到 O(1),以及这为什么不是常数级加速。
  3. 算出注意力的 O(N²) 显存代价,并知道在什么上下文长度下会撞墙、撞墙后该用什么手段(滑动窗口、RoPE 外推、Flash Attention、线性注意力)。
  4. 识别 RNN/Transformer/SSM 各自的归纳偏置(Inductive Bias),据此为新任务选择合适的架构。

一、问题与直觉

2017 年之前,LSTM 和 GRU 统治了相当于 ImageNet 的翻译基准长达五年,它们是当时唯一的工具。但它们有三个致命弱点。

弱点一:串行计算。 RNN 计算 h_t = f(h_{t-1}, x_t),每一步都依赖前一步——你想算 h_5,就必须先算完 h_4。在每秒能做一百万次浮点运算的 GPU 上,一条 1024 token 的序列意味着 1020 步串行等待。训练墙钟时间随序列长度线性增长,而这块硬件是为并行设计的。

弱点二:梯度消失。 50 个 token 之前的信息,要被压过 50 个非线性变换才能传到现在。门控单元(LSTM、GRU)缓解了挤压,但从未消除它。于是长程依赖——「去年夏天飞往京都的飞机上我读的那本书是……」——经常失效。

弱点三:固定宽度隐状态。 编码器把整条源序列塞进一个向量,解码器才开始工作。源序列是 5 个 token 还是 500 个,瓶颈的形状都一样。

2017 年的论文提出了一个激进方案:彻底丢掉循环,让每个位置同时关注所有其他位置,用一次大矩阵乘法代替 1024 次串行运算。到 2026 年,同一个 Transformer 块统治了所有模态:语言(GPT-5、Claude 4、Llama 4)、视觉(ViT、DINOv2、SAM 3)、音频(Whisper)、生物(AlphaFold 3)、机器人(RT-2)。同一个块,不同的输入。

💡 关键直觉:RNN 是接力赛,棒子必须一棒一棒传;Transformer 是广播网,所有人同时向所有人喊话。前者受限于人的速度,后者受限于广播网的带宽(显存)。

二、为什么注意力比循环快:依赖深度

这一节没有神经网络,我们用纯数值实验让你在笔记本上摸到这道鸿沟。

循环风格 vs 注意力风格。 同样的数学(把序列归约成一个值),依赖图却截然不同:

def rnn_style(xs): h = 0.0 for x in xs: h = 0.9 * h + x # 无法并行:h 依赖前一个 h return h def attention_style(xs): return sum(xs) / len(xs) # 每个 x 相互独立

两种算法都做 N 次加法,差别在依赖深度(Dependency Depth)——下一拍开始前必须串行完成多少步。RNN 的深度是 N,注意力用树形归约是 log(N),用并行扫描是 1。决定 GPU 时间的不是操作数,而是深度。

实测长序列的缩放

完整代码见原课程 phases/07-transformers-deep-dive/01-why-transformers/code/main.py。在 2026 年的笔记本上,1000 元素以下的序列快到测不出差异;10 万元素的序列,RNN 风格显示一条干净的线性扫描,注意力风格因为是纯 C 实现的 sum() 而几乎瞬间完成。把这个差距放大到一个 16384 token 的 Transformer 对比 12 层 LSTM 等价物,你就明白为什么 2016 年训练墙钟时间是个拦路虎。

设计要点:加速不是常数。它是 O(N) 串行深度与 O(1) 串行深度的差别。在 N=512、相同硬件下,Transformer 每轮训练快 5~10 倍,而且差距随序列长度拉大——直到你撞上注意力 O(N²) 的显存墙(Flash Attention 后来修掉了常数项,见第 12 节)。

三、代价:注意力的二次方瓶颈

Transformer 不是免费午餐。注意力显存随序列长度按 O(N²) 增长。

上下文长度 注意力矩阵大小 是否需要工程手段
2K 约 1600 万 完全够用
32K 约 10 亿 开始吃紧
128K 约 160 亿 必须用滑动窗口 / Flash Attention 分块
1M+ 千亿级 改用线性注意力 / Mamba 2 / Hyena

循环在时间和显存上都是 O(N);Transformer 用显存换时间,再靠并行把时间赢回来。这套权衡的工程结论是:短中上下文 Transformer 完胜,超长上下文需要专门的稀疏化或线性化手段

归纳偏置的转变

RNN 假设数据是「局部 + 近因」的;Transformer 不假设任何东西——任意两个位置都可能产生注意力。这就是为什么 Transformer 需要更多数据才能训好,但一旦有了数据就能走得更远。Chinchilla(2022)把这个直觉形式化了:给定足够多的 token,Transformer 总是击败同等参数量的 RNN。

四、可复用产物

原课程 outputs/ 目录提供一个架构选择器 Skill(skill-architecture-picker.md):给定新任务的序列长度、吞吐需求、训练预算,推荐合适的架构(RNN / SSM / Transformer / 线性注意力)。它有一条硬规则——任何超过 10 亿 token 的训练任务,绝不推荐纯 RNN,除非同时说明代价

五、练习

  1. (Easy)rnn_style 的标量隐状态换成 64 维向量,重新计时。串行开销随隐状态维度增长多少?
  2. (Medium) 用纯 Python 实现并行前缀和(Hillis-Steele 扫描),验证它在长度 1024 上与串行扫描给出相同的数值结果,并数出依赖深度。
  3. (Hard) 把注意力风格的归约移植到 PyTorch GPU 上,把序列长度从 64 扫到 65536,绘图并解释曲线形状(尤其是显存随长度的二次增长拐点)。

本节要点回顾

  1. RNN 的三个致命伤:串行计算(GPU 空转)、梯度消失(长程依赖失效)、固定宽度隐状态(编码瓶颈),2017 年前唯一能跑但已到极限。
  2. 核心区别是依赖深度:RNN 深度 N,注意力深度 log(N) 或 1;决定 GPU 时间的是深度,不是操作数。
  3. 加速不是常数:N=512 时 Transformer 每轮快 5~10 倍,差距随长度拉大。
  4. 代价是 O(N²) 显存:2K 够用,128K 必须上 Flash Attention / 滑动窗口 / RoPE 外推,百万级改线性注意力。
  5. 归纳偏置转变:RNN 假设近因局部,Transformer 不假设;后者更吃数据但更会 scale。
  6. Chinchilla 结论:给定足够 token,Transformer 总是击败等参 RNN。
  7. RNN 没死,变成了组件:SSM(Mamba、RWKV)是带结构化参数的 RNN,2026 年前沿实验室普遍训混合 SSM+Transformer(Jamba、Samba)。

下一节,我们将亲手实现自注意力(Self-Attention)的核心公式 softmax(QKᵀ/√dk)V,把「每个 token 同时向所有 token 提问」落到可运行的 NumPy 代码里。


发布者: 作者: Rohit Gupta 转发
评论区 (0)
U