Skip to content
🔗 分享本题
查看我的学习进度 →

Scaled Dot-Product Multi-Head Attention 动漫知识图:输入投影为 Q/K/V,QK 转置点积经缩放、掩码和 Softmax 得到权重,再加权汇总 V,多头并行后拼接投影

🧠 图解记忆:QK 算关注,Softmax 分权重,再汇总 V,多头看不同关系;点击图片可查看原图。

Attention = 让模型在处理每个 Token 时,动态关注输入中所有其他 Token 的相关程度。

Self-Attention 核心公式(面试手撕级)

Attention(Q, K, V) = softmax(Q · K^T / √d_k) · V
  • Q(Query):当前要处理的 Token 在"查询什么信息"
  • K(Key):序列中每个 Token 的"被查询内容"
  • V(Value):匹配后实际取到的"价值信息"
  • √d_k:缩放因子,防止点积过大导致 Softmax 梯度消失

直观理解

句子: "苹果发布了新款手机,它的销量很好"

处理"它"时:
  - 与"苹果"的注意力权重高 → "它"指"苹果"
  - 与"销量"的注意力权重中等 → 上下文关联
  - 与"新款"的注意力权重较低

Multi-Head Attention 为什么更好?

维度Single-Head AttentionMulti-Head Attention
原理一组 Q/K/V 全局关注多组 Q/K/V 并行,各自关注不同方面
捕捉能力只能学到一种依赖模式语法、语义、长距离依赖等可同时捕捉
计算复杂度O(n²·d)O(n²·d·h),但可并行
python
# PyTorch 伪代码
num_heads = 8
head_dim = d_model // num_heads

# 线性变换得到 Q,K,V
Q = linear_Q(x)  # [batch, seq_len, d_model]
K = linear_K(x)
V = linear_V(x)

# 分割成多头
Q = Q.view(batch, seq_len, num_heads, head_dim).transpose(1, 2)  # [batch, heads, seq, dim]
K = K.view(batch, seq_len, num_heads, head_dim).transpose(1, 2)
V = V.view(batch, seq_len, num_heads, head_dim).transpose(1, 2)

# 对每个头做 Scaled Dot-Product Attention
attn_output = scaled_dot_product_attention(Q, K, V, mask)

# 拼接所有头的输出
attn_output = attn_output.transpose(1, 2).contiguous().view(batch, seq_len, d_model)
output = linear_out(attn_output)  # 最终投影

面试高频追问

  • 为什么要缩放 √d_k? 当 d_k 较大时,点积值分布方差变大,Softmax 会趋于 one-hot(梯度消失)。除以 √d_k 让方差保持在 ~1。
  • Single-Head 和 Multi-Head 哪个强? Multi-Head 几乎总是更强——相当于让模型"多角度观察同一件事"。除非极端资源受限场景。
  • Attention 时间复杂度? O(n² · d)。n 是序列长度,d 是隐藏维。这也是为什么需要 Flash Attention(见进阶题)。

面试话术:

"Attention 的核心思想是'让每个词都能看到并关注其他所有词'。公式上就是 Q·K^T 算相似度,Softmax 归一化,再加权求和 V。Multi-Head 相当于多个专家各看一个角度,最后拼接起来。它是 Transformer 取代 RNN 的关键——RNN 只能串行序列化地看前面,Attention 可以一次性全局关注,既快又准。"


📚 参考:The Illustrated Transformer(图解注意力机制) · Attention Is All You Need(原论文)