Transformer中⼀直强调的self-attention是什么
1️⃣ 考察意图
面试官想确认你不仅背过公式,还能讲清 self-attention 为什么是 Transformer 的“灵魂”,以及它如何解决 RNN/CNN 的硬伤。这是典型的“背概念 + 工程取舍”混合题,刁钻点在于:很多人只会说“QKV 和 softmax”,但讲不清为什么除 √d、为什么多头、为什么不用 RNN。答好了能展示你对序列建模本质的理解,以及从论文到落地的工程直觉。
2️⃣ 标准答
Self-attention 的核心思想是:让序列中每个 token 通过加权聚合所有 token 的信息来更新自己的表示,权重由 token 间的相似度动态决定。公式是:
Attention(Q, K, V) = softmax(QK^T / √d_k) V
其中 Q、K、V 来自同一输入 X 乘以三个不同的权重矩阵 W_Q、W_K、W_V。
为什么叫“self”? 因为 Q、K、V 都来自同一个序列,不像传统 attention(如 seq2seq 中 decoder query 对 encoder key-value 做 cross-attention)。Self-attention 让序列内部做“自我对齐”。
关键设计取舍:
- 除以 √d_k:防止 QK^T 内积值随维度 d_k 增大而过大,导致 softmax 梯度消失。这是论文中的核心 trick,不是随意加的。
- 多头机制:单头 attention 只能学习一种关系模式(如语法或语义)。多头(如 8 头)让模型并行学习不同子空间表示,比如一个头关注主谓关系,另一个头关注修饰关系。最后拼接 + 线性投影。
- 并行计算 vs RNN 的串行:RNN 必须按时间步递推,无法并行。Self-attention 一次计算所有位置对,训练时可 batch 化,GPU 利用率极高。代价是 O(n²) 复杂度,长序列(如 8k+ tokens)需用 FlashAttention 或稀疏 attention 优化。
实际落地的坑 + 解法:
- 坑 1:长序列 OOM。标准 self-attention 显存随序列长度平方增长。解法:用 FlashAttention(通过 tiling 和 kernel fusion 减少显存读写)或 sliding window attention(如 Mistral 的 4k 窗口)。
- 坑 2:位置信息缺失。Self-attention 是排列不变的(permutation invariant),即“我打你”和“你打我”的 QK^T 矩阵相同。解法:加位置编码,如 RoPE(旋转位置编码)或 ALiBi(线性偏置)。RoPE 通过旋转矩阵注入相对位置,被 LLaMA 和 GPT-4 采用。
- 坑 3:训练不稳定。深层 Transformer 中 attention 输出方差大。解法:Pre-LN(层归一化放在 attention 前)比 Post-LN 更稳定,被几乎所有现代模型采用。
在 Transformer 中的角色:
- 编码器:每层 self-attention 让 token 互相“看”对方,捕捉双向上下文。
- 解码器:带掩码的 self-attention(masked self-attention)只允许看过去 token,保证自回归生成。然后接 cross-attention 看编码器输出。
一句话总结:Self-attention 通过动态加权聚合,用 O(n²) 计算换来了长距离依赖捕捉和并行训练,是 Transformer 取代 RNN 的根本原因。
3️⃣ 答题模板(30 秒电梯版)
“这个问题我从原理、公式、工程取舍三个层面回答。原理上,self-attention 让序列每个 token 通过注意力权重聚合所有 token 信息,权重由 QK 相似度决定。公式是 softmax(QK^T/√d)V,其中除以 √d 防止梯度消失。工程上,多头机制并行学习不同关系,但 O(n²) 复杂度需用 FlashAttention 优化,位置缺失用 RoPE 解决。总结一句:self-attention 是 Transformer 并行化、长程建模的核心,也是其计算瓶颈所在。”
4️⃣ 高频追问 & 应对
追问 1:为什么 self-attention 中 Q 和 K 要分开,不能直接用 X 自己算权重?
如果直接用 X 算权重,相当于 Q=K=X,那么权重矩阵是对称的(X·X^T),只能捕捉“自身与自身的相似度”,无法区分“谁在查询、谁在被查询”。分开后,Q 和 K 通过不同线性变换学习不同子空间,Q 负责“主动关注”,K 负责“被关注”,这样模型可以学到非对称关系(如“主语”关注“谓语”而非反过来)。这是 attention 机制“查询-键值”范式的核心。
追问 2:多头注意力中,头数怎么选?头越多越好吗?
不是。头数增加会提升模型容量,但每个头的维度 d_k = d_model / h 会变小,过小的头维度(如 <32)会导致表达力不足。Google 论文发现 8 头在 d_model=512 时效果最好。实际工程中,头数受 GPU 并行度限制:A100 上 16 头比 32 头更快,因为矩阵乘法更规整。建议:d_k 保持在 64-128 之间,头数选 8 或 16。
追问 3:Self-attention 的 O(n²) 复杂度怎么优化?具体方法有哪些?
三类主流方法:1)稀疏 attention:如 Longformer 的 sliding window + global tokens,复杂度降到 O(n·w)。2)线性 attention:用核函数近似 softmax(如 Performer 的 FAVOR+),复杂度 O(n)。3)FlashAttention:不改变算法,通过 tiling 和 kernel fusion 减少显存读写,实际加速 2-4 倍,被 GPT-4 和 LLaMA 采用。注意:稀疏 attention 会丢失全局信息,线性 attention 精度有损,FlashAttention 是当前最实用的 trade-off。
5️⃣ 避坑 · 常见错误答法
- ❌ “Self-attention 就是让每个词看所有词,权重由点积决定。” → ✅ 必须讲清 QKV 的由来(三个线性变换)、除以 √d 的原因、多头的作用。只说“点积”太浅,面试官会追问“为什么除√d”。
- ❌ “Self-attention 比 RNN 好,因为它能并行。” → ✅ 并行是结果,根本原因是 self-attention 没有序列依赖,但代价是 O(n²) 计算和位置信息缺失。只夸优点不提代价,显得不全面。
- ❌ “多头就是多个 attention 拼起来。” → ✅ 要解释每个头学习不同子空间(如语法 vs 语义),以及为什么需要线性投影降维。只说“拼起来”暴露没读过原始论文。
6️⃣ 简历呼应
- 如果你有 LLM 微调项目:从“self-attention 是长上下文瓶颈”切入,讲你用 FlashAttention 或 RoPE 扩展上下文到 32k 的实践,对比标准 attention 的显存和速度差异。
- 如果你只做过传统 NLP(如 LSTM 文本分类):用“RNN 的隐状态是固定窗口,self-attention 是全局窗口”做类比,强调 self-attention 如何解决长距离依赖问题,并提一句“但 O(n²) 复杂度让短序列任务(如情感分析)中 RNN 仍有优势”。
- 如果你是校招无项目:聚焦论文复现,说“我手写过单头 attention 的 PyTorch 代码,验证了除以 √d 对梯度的影响,并可视化注意力矩阵发现模型关注句法结构”,展示动手能力。
- 《Attention Is All You Need》(Vaswani et al., 2017)—— 原始论文,必读
- 《FlashAttention: Fast and Memory-Efficient Exact Attention》(Dao et al., 2022)—— 工程优化必读
- 《RoFormer: Enhanced Transformer with Rotary Position Embedding》(Su et al., 2021)—— RoPE 论文
- 《Efficient Transformers: A Survey》(Tay et al., 2020)—— 长序列优化综述
- 博客:Jay Alammar 的 “The Illustrated Transformer” —— 可视化入门