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 等都建立在它之上。
完整版教学
记忆钩子:架构题不要只背名字,要沿着「张量怎么变、复杂度怎么变、训练/推理代价怎么变」三步讲。
一、先搞懂 RNN 的两个硬伤
RNN 按时间步一个一个处理:算第 t 个词必须先算完第 t-1 个。这带来两个问题:
- 无法并行:训练时序列内部严格串行,GPU 的并行能力用不上,长序列训练极慢。
- 长距离依赖衰减:信息要沿着时间步一路传递,早期信息经过多次变换后逐渐丢失(梯度消失),句首和句尾的关联很难学到。LSTM/GRU 用门控缓解,但没根治。
注意力机制的思路是:干脆让每个位置直接看到所有位置,不再逐步传递。
二、自注意力是怎么算的(QKV)
每个输入向量会被投影成三个向量:Query(查询)、Key(键)、Value(值)。可以类比一次「检索」:
- 用当前位置的 Q 去和所有位置的 K 做点积,得到相关性打分;
- 打分除以
√d_k后做 softmax,变成一组和为 1 的权重; - 用这组权重对所有位置的 V 加权求和,就是当前位置的新表示。
直观理解:每个词都在问「我该关注句子里的哪些词」,然后把被关注词的信息按权重汇聚过来。
三、为什么要除以 √d_k
当维度 d_k 较大时,Q·K 点积的数值方差会变大,容易把 softmax 推到「非常尖锐」的区域,导致梯度极小、训练不稳定。除以 √d_k 做缩放,把点积拉回合适范围,这也是它叫 Scaled Dot-Product Attention 的原因。
四、多头注意力(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,避免重复计算,是大模型加速生成的关键手段。
八、面试拆解算例
这类题最怕只讲术语,最好把它落成一次资源账。假设模型有 32 层、hidden size 为 4096、序列长度从 2K 增到 16K,标准注意力的相关度矩阵规模会从 2K × 2K 变成 16K × 16K,理论元素数量放大 64 倍。即使具体算子不会真的把所有中间矩阵都落到 HBM,复杂度曲线仍然决定了 prefill 延迟和显存压力会快速上升。
attention_scores_per_head = seq_len * seq_len
2K -> 2,048 * 2,048 ≈ 4.19M
16K -> 16,384 * 16,384 ≈ 268.44M
放大倍数 ≈ 64
| 观察维度 | 面试要说清的问题 | 工程判断 |
|---|---|---|
| 张量形状 | Q/K/V、hidden state 或路由权重怎样变化 | 能否解释实现差异 |
| 复杂度 | 随层数、序列长度、head 数怎样增长 | 谁先成为瓶颈 |
| 质量风险 | 是否改变训练分布或表达能力 | 会不会掉点 |
| 部署代价 | 算子、框架、缓存是否支持 | 能不能稳定上线 |
输入 token -> embedding -> 注意力/FFN/归一化模块 -> hidden state -> logits
| | |
位置/掩码 显存与吞吐 质量与稳定性
所以回答「Transformer 的核心结构是什么?为什么它能取代 RNN?」时,推荐先讲结构变化,再讲这条变化怎样影响资源曲线,最后补一句质量和部署的边界。这样比单纯背「某某结构更快、更省」更像工程答案。
九、常见误区与追问
- 误区:把「Transformer 的核心结构是什么?为什么它能取代 RNN?」理解成单个模块的孤立优化。 架构题通常要同时联系训练稳定性、推理显存、吞吐和长上下文表现,孤立背结论不够。
- 追问:这个设计改变了哪一类张量或计算? 回答时要能指出 Q/K/V、隐藏状态、归一化、FFN 或路由中的具体变化,否则容易停留在概念层。
- 误区:新结构一定全方位优于旧结构。 很多优化是在显存、并发、质量、实现复杂度之间做交换,不存在无代价替代。
- 追问:batch、序列长度或并发变大时会发生什么? 架构设计最终会落到复杂度和资源曲线上,要能解释哪一项先成为瓶颈。
- 误区:论文里的指标可以直接迁移到业务线上。 线上还要看框架支持、算子融合、缓存命中、模型规模和请求分布,实验结论需要重新压测。
十、加强记忆
RNN 是「排队传话」,越传越失真且不能并行;Transformer 是「开会」,所有人一次性互相看到彼此,靠注意力决定听谁的——快、记得远、还容易堆大。