为什么需要Transformer
RNN逐token处理,Transformer并行处理所有token,这一架构选择改变了2017年后深度学习的所有扩展曲线
为什么需要Transformer — RNN的问题
RNN逐个token处理序列,Transformer一次性处理所有token。这一架构选择改变了2017年后深度学习的所有扩展曲线。
类型: 学习 语言: Python 前置知识: 阶段3(深度学习核心), 阶段5 · 09(序列到序列), 阶段5 · 10(注意力机制) 预计时间: ~45分钟
问题所在
2017年之前,全球每一个最先进的序列模型——语言、翻译、语音——都是循环神经网络。LSTM和GRU在相当于ImageNet的翻译基准测试上称霸了五年。它们是所有人唯一的工具。
它们有三个致命弱点。顺序计算意味着你无法沿时间轴并行化:token t+1需要token t的隐藏状态。一个1,024个token的序列意味着在每周期可执行1,000,000次浮点运算的GPU上需要1,024个串行步骤。训练挂钟时间在为并行性设计的硬件上随序列长度线性增长。
梯度消失意味着50个token之前的信息已经被压缩通过了50个非线性。门控循环单元(LSTM, GRU)缓解了这个问题但从未消除它。长程依赖——“我去年夏天在飞往京都的飞机上读的那本书是……”——经常失败。
固定宽度的隐藏状态意味着编码器在解码器看到任何东西之前将整个源序列压缩成单个向量。无论源序列是5个token还是500个都无所谓;瓶颈的形状是一样的。
2017年的论文”Attention Is All You Need”提出了一些激进的东西:完全放弃循环。让每个位置并行地关注每个其他位置。用一次大型矩阵乘法代替1,024次顺序乘法。
到2026年,这个结果主导了每一个模态。语言(GPT-5, Claude 4, Llama 4),视觉(ViT, DINOv2, SAM 3),音频(Whisper),生物学(AlphaFold 3),机器人(RT-2)。相同的块,不同的输入。
核心概念
循环作为瓶颈。 RNN计算 h_t = f(h_{t-1}, x_t)。每一步依赖前一步。你无法在 h_4 之前计算 h_5。在拥有10,000+并行核心的现代GPU上,这在长序列上浪费了99%的硅片。
注意力作为广播。 自注意力为每一对 (i, j) 同时计算 output_i = sum_j(a_ij * v_j)。整个N x N注意力矩阵在一次批量矩阵乘法中填充。没有步骤依赖另一个。GPU喜欢它。
加速不是常数。 它是 O(N) 串行深度和 O(1) 串行深度之间的差异。在实践中,在匹配硬件上N=512时,transformer每个epoch训练快5-10倍,差距随序列长度增大,直到触及注意力的 O(N^2) 内存墙(Flash Attention后来修复了这个问题——见第12课)。
Transformer的代价。 注意力内存按 O(N^2) 扩展。对于2K上下文,没问题。对于128K上下文,你需要滑动窗口、RoPE外推、Flash Attention分块或线性注意力变体。循环在时间和内存上都是 O(N);transformer用内存换时间,然后通过并行性赢回时间。
归纳偏移的转变。 RNN假设局部性和近因性。Transformer不假设任何东西——每一对都是注意力的候选。这就是为什么transformer需要更多数据才能训练好,但一旦有了数据就能扩展更远。Chinchilla(2022)形式化了这一点:给定足够的token,transformer总是击败等参数量的RNN。
动手构建
这里没有神经网络——我们用数值模拟核心瓶颈,让你在笔记本上感受差距。
步骤1:测量串行深度
参见 code/main.py。我们构建两个函数。一个将序列编码为加法链(串行,类似RNN)。一个将其编码为并行归约(广播,类似注意力)。相同的数学,不同的依赖图。
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是独立的
我们在长达100,000个元素的序列上对两者计时。RNN版本是O(N)且是单CPU流水线。即使在纯Python中,注意力风格的归约在长度>= 1,000时也能胜出,因为Python的 sum() 是用C实现的,无需每步解释器开销即可迭代。
步骤2:计算理论运算量
两种算法都做N次加法。区别在于依赖深度:在下一个操作开始之前必须顺序执行多少操作。RNN深度 = N。注意力深度 = 使用树归约时为log(N),或使用并行扫描时为1。深度,而非操作数,决定GPU时间。
步骤3:长序列上的经验扩展
我们打印一个使O(N)差距可见的计时表。在2026年的Mac笔记本上,1,000个元素以下的序列太快而无法测量。100,000个元素的序列显示出清晰的线性扫描。将其扩展到具有12层LSTM等价物的16,384-token transformer,你就会理解为什么训练挂钟时间在2016年是个障碍。
实际应用
2026年何时仍然选择RNN:
| 场景 | 选择 |
|---|---|
| 流式推理,一次一个token,恒定内存 | RNN或状态空间模型(Mamba, RWKV) |
| 超长序列(>1M token),注意力内存爆炸 | 线性注意力, Mamba 2, Hyena |
| 没有矩阵乘法加速器的边缘设备 | 深度可分离RNN在FLOPs/瓦特上仍然胜出 |
| 其他任何场景(训练,批量推理,上下文高达128K) | Transformer |
状态空间模型(SSM)如Mamba本质上是具有结构化参数化的RNN,赋予它们两者的优点:O(N)扫描内存,通过选择性扫描实现并行训练。它们恢复了90%的transformer质量,同时具有更好的长上下文扩展性。2026年大多数前沿实验室训练混合SSM+transformer模型(如Jamba, Samba)——循环没有死,它是一个组件。
交付成果
参见 outputs/skill-architecture-picker.md。该技能根据长度、吞吐量和训练预算约束为新序列问题选择架构。它应该始终拒绝在没有说明权衡的情况下为超过10亿token的训练运行推荐纯RNN。
练习
- 简单。 从
code/main.py中取rnn_style,将标量隐藏状态替换为长度为64的隐藏状态向量。重新测量。串行开销随隐藏状态维度增长多少? - 中等。 用纯Python实现并行前缀和(Hillis-Steele扫描)。验证它在长度1024上产生与串行扫描相同的数值输出。计算深度。
- 困难。 将注意力风格的归约移植到GPU上的PyTorch。在序列长度从64到65,536的范围内对两者计时。绘制并解释曲线形状。
关键术语
| 术语 | 人们怎么说 | 实际含义 |
|---|---|---|
| 循环 | ”RNN是顺序的” | 步骤t依赖步骤t-1的计算,迫使沿时间轴串行执行。 |
| 串行深度 | ”图有多深” | 依赖操作的最长链;即使在无限硬件上也限制挂钟时间。 |
| 注意力 | ”让token互相看” | 加权和 sum_j a_ij v_j,其中 a_ij 来自位置i和j之间的相似度分数。 |
| 上下文窗口 | ”模型能看到多少” | 注意力层可作为输入的位置数;二次内存成本在此扩展。 |
| 归纳偏置 | ”架构中嵌入的假设” | 关于数据外观的先验;CNN假设平移不变性,RNN假设近因性。 |
| 状态空间模型 | ”有代数基础的RNN” | 通过结构化状态空间矩阵参数化以实现并行训练的循环。 |
| 二次瓶颈 | ”为什么上下文这么贵” | 注意力内存 = 序列长度的 O(N^2);Flash Attention隐藏了常数,而非扩展。 |
延伸阅读
- Vaswani et al. (2017). Attention Is All You Need — 终结主流NLP中循环的论文。
- Bahdanau, Cho, Bengio (2014). Neural MT by Jointly Learning to Align and Translate — 注意力诞生的地方,附加在RNN上。
- Hochreiter, Schmidhuber (1997). Long Short-Term Memory — 原始LSTM论文,作为记录。
- Gu, Dao (2023). Mamba: Linear-Time Sequence Modeling with Selective State Spaces — 现代循环对transformer的回答。