注意力机制:那一击突破


文档摘要

注意力机制:那一击突破 本节摘要:解码器不再眯眼盯着一个被压缩的摘要,而是开始打量整个源句。此后的所有进展,都是注意力加工程。第 09 节以一个有据可查的失败收尾:在玩具复制任务上,GRU 编码器-解码器从长度 5 的 89% 准确率掉到长度 80 的近随机。原因是结构性的,不是训练 bug——编码器榨出的每一点信息都得塞进一个定长隐状态,解码器看不到别的。Bahdanau、Cho、Bengio 在 2014 年给出了三行修复:不再只给解码器最终编码器状态,而是保留每一个编码器状态;每个解码步,算一个编码器状态的加权平均,权重回答「解码器此刻该多看编码器位置 一眼吗?」这个加权平均就是上下文,且每步都变。这就是全部想法。

注意力机制:那一击突破

本节摘要:解码器不再眯眼盯着一个被压缩的摘要,而是开始打量整个源句。此后的所有进展,都是注意力加工程。第 09 节以一个有据可查的失败收尾:在玩具复制任务上,GRU 编码器-解码器从长度 5 的 89% 准确率掉到长度 80 的近随机。原因是结构性的,不是训练 bug——编码器榨出的每一点信息都得塞进一个定长隐状态,解码器看不到别的。Bahdanau、Cho、Bengio 在 2014 年给出了三行修复:不再只给解码器最终编码器状态,而是保留每一个编码器状态;每个解码步,算一个编码器状态的加权平均,权重回答「解码器此刻该多看编码器位置 i 一眼吗?」这个加权平均就是上下文,且每步都变。这就是全部想法。Transformer 延伸了它,自注意力把它用到单序列,多头并行跑——但 2014 版已经打破了瓶颈,有了它,转向 Transformer 是工程而非概念。

对应原课程:Phase 5 · Lesson 10 · attention-mechanism(原英文 phases/05-nlp-foundations-to-advanced/10-attention-mechanism/docs/en.md)。前置依赖:第 09 节(序列到序列)。

学习目标

阅读完本节,你应当能够:

  1. 说清注意力如何用每步可变的加权平均打破 seq2seq 的定长上下文瓶颈。
  2. 从零实现加性(Bahdanau)与乘性(Luong dot/general)三种打分,并对照形状表排查每个张量。
  3. 把经典注意力的术语翻译成 Q/K/V,看清从 Bahdanau 到缩放点积注意力主要是记号之别。
  4. 警惕「注意力权重即解释」陷阱,知道何时仍该读 2014~2017 论文里的注意力图。

一、问题与直觉

第 09 节以一个有据可查的失败收尾。在玩具复制任务上训的 GRU 编码器-解码器,从长度 5 的 89% 准确率掉到长度 80 的近随机。原因是结构性的,不是训练 bug:编码器榨出的每一点信息都得塞进一个定长隐状态,解码器看不到别的。

Bahdanau、Cho、Bengio 在 2014 年给出了三行修复。不再只给解码器最终编码器状态,而是保留每一个编码器状态。每个解码步,算一个编码器状态的加权平均,权重回答「解码器此刻该多看编码器位置 i 一眼吗?」这个加权平均就是上下文,且每步都变。

这就是全部想法。Transformer 延伸了它,自注意力把它用到单序列,多头并行跑。但 2014 版已经打破了瓶颈,有了它,转向 Transformer 是工程,不是概念。

每个解码步 t:

  1. 把上一步解码器隐状态 s_{t-1}查询
  2. 对每个编码器隐状态 h_1, ..., h_T 打分,每个编码器位置一个标量。
  3. 对分数做 softmax,得到和为 1 的注意力权重 α_{t,1}, ..., α_{t,T}
  4. 上下文向量 c_t = Σ α_{t,i} * h_i,即编码器状态的加权平均。
  5. 解码器拿 c_t 加上一步输出 token,产出下一个 token。

加权平均是关键。当解码器要把 "Je" 译成 "I" 时,它给 "Je" 上的编码器状态高权重、其他低权重;要 "not" 时,给 "pas" 高权重。上下文向量每步都在重塑。

二、形状(每个人都栽过的地方)

每个注意力实现第一次都会在这里出错,慢慢读。

对象 形状 备注
编码器隐状态 H (T_enc, d_h) 若 BiLSTM,d_h = 2 * d_hidden
解码器隐状态 s_{t-1} (d_s,) 单个向量
注意力分数 e_{t,i} 标量 每个编码器位置一个
注意力权重 α_{t,i} 标量 对所有 i 做 softmax 后
上下文向量 c_t (d_h,) 与编码器状态同形

Bahdanau(加性)分数:e_{t,i} = v_α^T * tanh(W_a * s_{t-1} + U_a * h_i)

  • s_{t-1} 形状 (d_s,),h_i 形状 (d_h,)
  • W_a 形状 (d_attn, d_s),U_a 形状 (d_attn, d_h)
  • tanh 内的求和形状 (d_attn,)
  • v_α 形状 (d_attn,),与它的内积坍缩成标量。这就是 v_α 的作用,它不神秘,就是把注意力维向量投影成标量分数的那个投影。

Luong(乘性)分数,三种变体:

  • dot:e_{t,i} = s_t^T * h_i,要求 d_s == d_h。硬约束,编码器双向就跳过它。
  • general:e_{t,i} = s_t^T * W * h_i,W 形状 (d_s, d_h),解除等维约束。
  • concat:本质是 Bahdanau 形式,因前两者更省而少用。

⚠️ 一个 Bahdanau/Luong 坑值得一说:Bahdanau 用 s_{t-1}(生成当前词之前的解码器状态),Luong 用 s_t(之后的状态)。搞混会产生微妙错误的梯度,极难调试。选定一篇论文,死守它的约定。

三、从零实现

第 1 步:加性(Bahdanau)注意力

import numpy as np def additive_attention(decoder_state, encoder_states, W_a, U_a, v_a): projected_dec = W_a @ decoder_state projected_enc = encoder_states @ U_a.T combined = np.tanh(projected_enc + projected_dec) scores = combined @ v_a weights = softmax(scores) context = weights @ encoder_states return context, weights def softmax(x): x = x - np.max(x) e = np.exp(x) return e / e.sum()

对照形状表核对:encoder_states(T_enc, d_h),projected_enc(T_enc, d_attn),projected_dec(d_attn,) 并广播,combined(T_enc, d_attn),scores(T_enc,),weights(T_enc,),context(d_h,)。齐活。

第 2 步:Luong 的 dot 与 general

def dot_attention(decoder_state, encoder_states): scores = encoder_states @ decoder_state weights = softmax(scores) return weights @ encoder_states, weights def general_attention(decoder_state, encoder_states, W): projected = W.T @ decoder_state scores = encoder_states @ projected weights = softmax(scores) return weights @ encoder_states, weights

各三行。这就是 Luong 论文落地的原因:多数任务精度相当,代码少得多。

第 3 步:一个数值例

给三个编码器状态(大致是 "cat"、"sat"、"mat")和一个与第一个最对齐的解码器状态,注意力集中在位置 0;把解码器状态挪近第三个编码器状态,注意力就移到位置 2,上下文向量随之变。

H = np.array([ [1.0, 0.0, 0.2], [0.5, 0.5, 0.1], [0.1, 0.9, 0.3], ]) s_close_to_cat = np.array([0.9, 0.1, 0.2]) ctx, w = dot_attention(s_close_to_cat, H) print("weights:", w.round(3))
weights: [0.464 0.305 0.231]

第一行赢。再把解码器状态挪近第三个编码器状态,看权重迁移。就这些。注意力就是显式对齐。

第 4 步:为何这是通往 Transformer 的桥

把上面的语言翻成 Q/K/V:

  • 查询 = 解码器状态 s_{t-1}
  • = 编码器状态(我们打分的对象)
  • = 编码器状态(我们加权求和的对象)

经典注意力里,键和值是同一个东西。自注意力把它们分开:你可以让一个序列查询自己,K 和 V 用不同的学习投影。多头注意力用不同投影并行跑。Transformer 把整段堆很多次,丢掉 RNN。

数学是一样的,形状是一样的。从 Bahdanau 注意力到缩放点积注意力的教学跳跃,主要是记号。

四、框架对比

PyTorch 和 TensorFlow 直接提供注意力。

import torch import torch.nn as nn mha = nn.MultiheadAttention(embed_dim=128, num_heads=8, batch_first=True) query = torch.randn(2, 5, 128) key = torch.randn(2, 10, 128) value = torch.randn(2, 10, 128) output, weights = mha(query, key, value) print(output.shape, weights.shape)
torch.Size([2, 5, 128]) torch.Size([2, 5, 10])

这就是一个 Transformer 注意力层:查询批 5 个位置,键/值批 10 个位置,各 128 维,8 头。output 是上下文增强后的新查询,weights 是 5×10 的对齐矩阵,可以可视化。

经典注意力仍重要的场景

  • 教学:单头、单层、基于 RNN 的版本让每个概念都看得见。
  • Transformer 装不下的设备端序列任务
  • 2014~2017 的任何论文:不懂 Bahdanau 约定就会误读。
  • 机器翻译的细粒度对齐分析:原始注意力权重即便在 Transformer 上也是可解释性工具,读懂它需要知道它是什么。

注意力权重即解释的陷阱

注意力权重看着可解释:它是跨位置和为 1 的权重,能画图,高就是「看了这里」,审稿人爱它。

它没看上去那么可解释。Jain 与 Wallace(2019)表明,在某些任务上,注意力分布可以被置换、替换成任意替代品而不改变模型预测。没有消融或反事实检验,别把注意力权重当推理的证据上报

五、可复用产物

保存为 outputs/prompt-attention-shapes.md:

--- name: attention-shapes description: Debug shape bugs in attention implementations. phase: 5 lesson: 10 --- Given a broken attention implementation, you identify the shape mismatch. Output: 1. Which matrix has the wrong shape. Name the tensor. 2. What its shape should be, derived from (d_s, d_h, d_attn, T_enc, T_dec, batch_size). 3. One-line fix. Transpose, reshape, or project. 4. A test to catch regressions. Typically: assert `output.shape == (batch, T_dec, d_h)` and `weights.shape == (batch, T_dec, T_enc)` and `weights.sum(dim=-1) close to 1`. Refuse to recommend fixes that silently broadcast. Broadcast-hiding bugs surface later as silent accuracy degradation, the worst kind of attention bug. For Bahdanau confusion, insist the decoder input is `s_{t-1}` (pre-step state). For Luong, `s_t` (post-step state). For dot-product, flag dimension mismatch between query and key as the most common first-time error.

六、练习

  1. 基础:实现带掩码的 softmax,让编码器里的 padding token 注意力权重为零。在变长序列的批次上测试。
  2. 进阶:给 Luong general 加多头注意力。把 d_h 拆成 n_heads 组,每头跑注意力再拼接,验证单头情形与之前的实现一致。
  3. 挑战:在第 09 节的玩具复制任务上训一个带 Bahdanau 注意力的 GRU 编码器-解码器。画准确率对序列长度曲线,对比无注意力基线,应看到 gap 随长度拉大——证实注意力抬掉了瓶颈。

本节要点回顾

  1. 瓶颈是结构性的:定长上下文向量装不下长输入,长度 5 的 89% 掉到长度 80 的近随机。
  2. 三行修复:保留每个编码器状态,每解码步算加权平均,上下文每步变。
  3. 加权平均是关键:译 "Je" 给 "Je" 高权重、译 "not" 给 "pas" 高权重。
  4. 形状表是排查命门:v_α 把注意力维向量投影成标量分数,不神秘。
  5. Bahdanau 加性:v^T tanh(W s + U h);Luong 乘性:dot/general/concat,更省精度相当。
  6. Bahdanau 用 s_{t-1}、Luong 用 s_t,搞混产生难调的错梯度。
  7. Q/K/V 翻译:查询=解码器状态,键/值=编码器状态,自注意力把键值分开。
  8. 从 Bahdanau 到缩放点积主要是记号:数学与形状相同,Transformer 只是堆叠+去 RNN。
  9. nn.MultiheadAttention 一行调:output 增强查询,weights 是对齐矩阵。
  10. 注意力权重非充分解释:可被置换而不改预测,上报前要消融或反事实。

下一节,我们把注意力塞进真正的翻译系统——进入「机器翻译」,看 BLEU 怎么算、平行语料怎么喂、领域术语怎么保。


发布者: 作者: Rohit Gupta 转发
评论区 (0)
U