Transformer 的核心结构是什么?为什么它能取代 RNN?
简化版
Transformer 的核心是自注意力机制(Self-Attention):序列中每个位置都能一步直接和其他所有位置交互。相比 RNN 必须逐时间步串行、长距离依赖会随距离衰减,Transformer 能全序列并行计算,且任意两个位置的交互路径长度都是 O(1)——又快又能捕捉长依赖,这也是它成为现代大模型基座的根本原因。
详细版
原始论文里的 Transformer 是一个 Encoder-Decoder 结构,由若干相同的层堆叠而成,每层包含这些模块:
- 输入嵌入 + 位置编码:自注意力本身对顺序不敏感,必须用位置编码补充「谁在前谁在后」的信息。
- 多头自注意力(Multi-Head Attention):并行多组注意力,从不同子空间捕捉不同类型的关系。
- 前馈网络(FFN):对每个位置独立做一次非线性变换,增强表达能力。
- 残差连接 + LayerNorm:稳定深层网络训练、缓解梯度问题。
自注意力的核心公式:
Attention(Q, K, V) = softmax(Q · Kᵀ / √d_k) · V
它之所以能取代 RNN,关键是三点:并行(一次算完整个序列,不必按时间步串行)、长距离依赖(任意两位置路径长度 O(1),不像 RNN 随距离衰减)、可扩展(结构规整,容易堆叠加深、放大参数量)。GPT、BERT 等都建立在它之上。
完整版教学
记忆钩子:Attention 做 token 间路由,FFN 做 token 内变换,残差与归一化保证深层可训练。
一、先搞懂 RNN 的两个硬伤
RNN 按时间步一个一个处理:算第 t 个词必须先算完第 t-1 个。这带来两个问题:
- 无法并行:训练时序列内部严格串行,GPU 的并行能力用不上,长序列训练极慢。
- 长距离依赖衰减:信息要沿着时间步一路传递,早期信息经过多次变换后逐渐丢失(梯度消失),句首和句尾的关联很难学到。LSTM/GRU 用门控缓解,但没根治。
注意力机制的思路是:干脆让每个位置直接看到所有位置,不再逐步传递。
二、自注意力是怎么算的(QKV)
每个输入向量会被投影成三个向量:Query(查询)、Key(键)、Value(值)。可以类比一次「检索」:
- 用当前位置的 Q 去和所有位置的 K 做点积,得到相关性打分;
- 打分除以
√d_k后做 softmax,变成一组和为 1 的权重; - 用这组权重对所有位置的 V 加权求和,就是当前位置的新表示。
直观理解:每个词都在问「我该关注句子里的哪些词」,然后把被关注词的信息按权重汇聚过来。
三、为什么要除以 √d_k
若 Q、K 各维独立、均值为 0、方差约为 1,d_k 个乘积相加后的方差约为 d_k,标准差约为 √d_k。头维度越大,未经缩放的点积越容易把 softmax 推到近似 one-hot 的饱和区,使非最大项梯度接近 0。除以 √d_k 将 logits 的典型尺度拉回同一数量级,让不同头维度下的优化更稳定;它控制的是数值分布,不改变 O(n²) 复杂度。
四、多头注意力(Multi-Head)
只用一组注意力,模型只能学到一种「关注模式」。多头把 Q/K/V 拆成 h 组低维子空间,各自独立算注意力再拼接,最后过一次线性变换。不同的头可以分别捕捉语法关系、指代关系、局部搭配等,表达力更强。
五、位置编码:给顺序补课
自注意力有个特点——置换不变:打乱输入顺序,输出只是跟着换位置,模型本身分不清「猫追狗」和「狗追猫」。所以要显式注入位置信息:
- 原始论文用正弦/余弦位置编码(固定、可外推);
- 也可以用可学习位置嵌入;
- 现代大模型多用 RoPE(旋转位置编码)、ALiBi 等相对位置方案,长度外推更好。
六、Encoder、Decoder 与掩码
- Encoder 是双向的:每个位置能看到整句,适合做理解类任务(BERT 就是纯 Encoder)。
- Decoder 是自回归的:生成第 t 个词时只能看到前面的词,靠 因果掩码(causal mask) 挡住未来信息(GPT 就是纯 Decoder)。
- Encoder-Decoder 之间用交叉注意力连接(Decoder 的 Q 去查 Encoder 的 K/V),适合翻译这类序列到序列任务。
七、复杂度与工程瓶颈
自注意力要算所有位置两两之间的关系,复杂度是 O(n²·d)(n 为序列长度)。序列越长,显存和计算越吃紧,这是长文本场景的主要瓶颈。常见优化方向:
- FlashAttention:改进计算/访存方式,省显存、提速;
- 稀疏注意力 / 线性注意力:降低 O(n²) 的量级;
- KV Cache:推理时缓存历史的 K/V,避免重复计算,是大模型加速生成的关键手段。
八、把一个 Transformer 块走一遍
以 batch=2、序列长=128、隐藏维=768、12 个头为例,Q/K/V 投影后通常仍是 [2,128,768],拆头后变成 [2,12,128,64]。每个头产生 [128,128] 的注意力权重,再合并回 [2,128,768];随后 FFN 常把每个 token 从 768 升到约 3072 维再降回 768。残差连接要求子层输出形状回到原隐藏维。
[B,S,D] = [2,128,768]
-> Q,K,V: [2,12,128,64]
-> attention: [2,12,128,128]
-> merge heads: [2,128,768]
-> FFN: 768 -> 3072 -> 768
| 子模块 | token 之间是否交互 | 主要职责 |
|---|---|---|
| Self-Attention | 是 | 按内容聚合上下文 |
| FFN | 否,各位置独立 | 对每个 token 做非线性变换 |
| 残差与归一化 | 不新增连接 | 稳定深层优化与信息传递 |
这条形状链能同时解释多头、缩放点积、残差和 FFN 的位置。若面试官继续问性能,要区分 prefill 的大矩阵并行与 decode 的逐 token、带宽受限特征,而不能只背“注意力是 O(n²)”。
九、常见误区与追问
- 误区:Transformer 只有注意力,没有逐位置网络。 标准块还包含 FFN、残差连接与归一化;FFN 往往占据大量参数和计算。
- 追问:多头注意力为何不直接用一个大头? 多个投影子空间能并行学习不同关系,同时保持总隐藏维度;是否真的分工仍需分析而非假定。
- 误区:位置编码只是给 token 加序号。 它必须以模型可利用的几何方式注入顺序,并影响注意力对相对或绝对距离的表示。
- 追问:Encoder 与 Decoder 的核心差别是什么? Decoder 自注意力使用因果可见性,encoder-decoder 架构中的 Decoder 还多一个读取编码结果的交叉注意力。
- 追问:Transformer 的主要长序列瓶颈在哪里? 全局注意力形成 S×S 交互,计算与中间数据随长度平方增长;推理还要承担 KV Cache。
十、加强记忆
Transformer 块抓住两条主线:Self-Attention 负责 token 之间的信息路由,FFN 负责每个 token 内部的非线性变换;残差和归一化让深层网络可训练,位置编码补上顺序。形状上记住 [B,S,D] 拆头后成为 [B,H,S,D/H],注意力矩阵才是 [B,H,S,S]。回答再补上训练可并行但全局注意力随序列平方增长,就覆盖了结构、优势与代价。