为什么要用multi-head attention
1️⃣ 考察意图
面试官想考察你对Transformer核心设计动机的深度理解,而非简单背诵“多头注意力并行计算”。真正刁钻点在于:为什么单头不够?多头到底解决了什么表示学习上的根本问题? 答好了能展示你不仅会用Attention,还理解其表示瓶颈——单一子空间无法同时捕捉不同位置的不同语义关系(如语法、实体、长距离依赖)。这属于工程取舍+系统设计类问题,要求你从表示多样性、计算效率、消融实验三个层面展开。
2️⃣ 标准答
多头注意力(Multi-Head Attention)的核心动机是:单头注意力受限于单一表示空间,无法同时捕捉不同位置的不同语义子空间。具体来说:
- 表示多样性瓶颈:单头注意力中,Q、K、V的线性投影将输入映射到一个共享子空间。假设输入序列长度为n,每个位置需要关注多种关系(如主语-谓语语法、实体-实体共指、长距离语义依赖)。单头只能学习一种加权模式,比如偏向局部语法,就会牺牲全局语义。多头通过h个独立投影(h=8或16),每个头学习不同的注意力分布,覆盖不同语义维度。
- 具体实现:将d_model=512的输入分别投影到h个d_k=64的子空间(d_k = d_model / h)。每个头独立计算Attention(Q_i, K_i, V_i) = softmax(Q_i K_i^T / sqrt(d_k)) V_i,输出拼接后经W_O投影回d_model。关键取舍:每个头维度降低,计算量与单头相似(h * d_k^2 ≈ d_model^2),但通过并行实现高效。实际中,h=8时d_k=64,单头d_k=512,计算量从O(n^2 * 512)变为O(n^2 * 64 * 8) ≈ O(n^2 * 512),几乎持平。
- 实际落地的坑+解法:坑在于头数过多会导致冗余。比如在BERT-base中,h=12时部分头关注模式高度相似(如都聚焦于[CLS] token)。解法:使用注意力头剪枝(Head Pruning),如Michel et al. (2019)发现移除30%的头对下游任务影响极小。工程上,训练时加入L1正则化鼓励头稀疏,或通过Gumbel-Softmax学习头的重要性权重。
- 消融实验证据:Vaswani et al. (2017)在WMT14英德翻译任务中,将h从8减到1(单头),BLEU从28.4降至27.3(下降约4%)。在长距离依赖任务(如LAMBADA)中,单头性能下降更显著(约10%),因为单头难以同时建模局部语法和全局语义。
- 为什么不是更大单头? 如果增大单头维度(如d_k=1024),计算量从O(n^2 * 512)升至O(n^2 * 1024),且表示空间仍单一。多头通过低维子空间并行,在相同计算量下获得更丰富的表示。
3️⃣ 答题模板(30 秒电梯版)
“这个问题我从表示多样性、计算效率、消融实验三个层面回答。表示层面:单头注意力受限于单一子空间,无法同时捕捉语法、实体、长距离依赖等不同语义关系;多头通过h个低维投影学习不同注意力分布。效率层面:每个头维度降低,计算量与单头相似,但通过并行实现高效。消融实验:Vaswani et al.在WMT14上验证,单头比8头BLEU下降约4%,长距离任务下降更显著。总结一句:多头注意力在相同计算量下,通过子空间并行增强了表示多样性,是Transformer的核心设计。”
4️⃣ 高频追问 & 应对
追问1:多头注意力中,每个头真的学到不同模式吗?如何验证?
是的,但并非所有头都独立。可以通过注意力可视化验证:在BERT-base(12头)上,对句子“The cat sat on the mat”可视化,头1可能关注局部语法(cat→sat),头5关注实体关系(cat→mat),头9关注[CLS] token。更严谨的方法是计算头间相似度(如JS散度或余弦相似度),发现部分头高度冗余。工程上,使用Head Importance Score(如基于梯度或置信度)剪枝冗余头,可减少30%参数而不损失性能。
追问2:为什么选择d_k = d_model / h?如果d_k固定为64,h增加会怎样?
这是计算效率的权衡。如果d_k固定为64,h增加时总计算量线性增长(O(n^2 * 64 * h)),显存和延迟会飙升。标准设计d_k = d_model / h确保总计算量恒定(O(n^2 * d_model^2))。但h过大(如h=32)时,每个头维度太小(d_k=16),表示能力下降,导致信息瓶颈。实际中,h=8或12是经验最优,平衡了多样性和容量。
追问3:多头注意力在长序列(如8K tokens)中有什么问题?如何优化?
主要问题是计算复杂度O(n^2 * h * d_k)随n平方增长。优化方案:1)使用FlashAttention(Dao et al. 2022),通过分块和重计算减少显存,支持8K序列。2)采用稀疏注意力(如Longformer的滑动窗口+全局token),每个头只关注局部窗口或特定全局位置。3)在推理时,使用KV Cache复用历史K、V,但多头导致缓存量线性增长(h * d_k * n),需用MQA(Multi-Query Attention)或GQA(Grouped Query Attention)减少K、V头数。
5️⃣ 避坑 · 常见错误答法
- ❌ “多头注意力是为了并行计算,加快训练速度。” → ✅ 并行是结果而非动机。真正动机是表示多样性:单头无法同时捕捉多种语义关系。并行只是实现手段,且单头也能并行(通过矩阵乘法)。
- ❌ “多头注意力能捕捉长距离依赖,单头不行。” → ✅ 单头也能捕捉长距离依赖(如通过softmax权重),但多头通过不同子空间同时捕捉局部和全局,提升表示质量。长距离依赖是优势之一,非唯一原因。
- ❌ “头数越多越好,h=16比h=8强。” → ✅ 头数过多会导致冗余和过拟合。消融实验显示,h=8到h=16在WMT14上BLEU提升不到0.2,但计算量翻倍。实际中h=8或12是经验最优。
6️⃣ 简历呼应
- 如果你有RAG项目:从检索增强角度切入,说明多头注意力如何帮助模型在检索文档中同时关注实体匹配和语义相似性。例如,在RAG中,query通过多头注意力同时捕捉不同检索文档的局部和全局关系。
- 如果你只做过传统NLP(如LSTM):用类比迁移:LSTM的隐状态是单一表示,无法同时建模语法和语义;多头注意力类似多个LSTM并行,每个学习不同时间尺度依赖。
- 如果你是校招无项目:聚焦论文复现demo:在小型Transformer(6层,d_model=256)上对比单头、4头、8头,用IMDb情感分类任务,展示注意力热力图差异。产出:GitHub仓库+可视化报告。
- Vaswani et al., “Attention Is All You Need”, 2017 (Section 3.2: Multi-Head Attention)
- Michel et al., “Are Sixteen Heads Really Better than One?”, 2019 (Head Pruning)
- Dao et al., “FlashAttention: Fast and Memory-Efficient Exact Attention”, 2022
- Shazeer, “Fast Transformer Decoding: One Write-Head is All You Need”, 2019 (MQA)
- Ainslie et al., “GQA: Training Generalized Multi-Query Transformer Models”, 2023 (GQA)