第 8 章 2014—2017:序列到序列与注意力机制
学习目标
- 理解 Seq2Seq(2014)的编码器-解码器框架
- 理解注意力机制为什么能解决「信息瓶颈」
- 对比 Bahdanau 注意力(2015)与 Luong 注意力(2015)
- 实现带注意力的序列到序列模型,并观察注意力对齐
8.1 时代背景:让机器翻译机器
第 4 章的 LSTM 能做语言模型,但很多任务需要把一个序列变成另一个序列:机器翻译(英文→中文)、摘要(长文→短文)、对话(问题→回答)。
2014 年,Sutskever、Vinyals 与 Le 提出 Seq2Seq:一个编码器把整个输入序列压缩成一个向量,一个解码器从这个向量逐字生成输出。同年 Cho 等人提出 GRU(第 4 章),也是在这个框架里诞生的。
8.2 Seq2Seq 与信息瓶颈
Seq2Seq 的结构:
输入 x₁ x₂ x₃ ... xₙ
↓ 编码器(GRU/LSTM)
→ 上下文向量 c(通常是最后隐藏状态)
↓ 解码器逐字生成
输出 y₁ y₂ y₃ ... yₘ问题很明显:所有信息都要挤进最后一个隐藏状态 c。句子越长,c 越像「把整本书塞进一句话的摘要」——细节必然丢失,这被称为信息瓶颈。
8.3 注意力:让解码器「回头看」
2015 年,Bahdanau 等人在机器翻译中提出注意力机制(attention):解码器每生成一个字,不是只看压缩的 c,而是回头看编码器的所有隐藏状态,按相关性加权求和,得到当前步的上下文向量:
scoreᵢ = 打分(h_dec, h_enc_i)
αᵢ = softmax(score₁..scoreₙ) # 对齐权重:第 i 个输入词对当前输出的重要性
context = Σ αᵢ·h_enc_i解码器每一步都有一个不同的 context——它可以根据当前要生成的词,主动「读」输入中相关的部分。这解决了信息瓶颈,还带来一个副产品:注意力权重 α 就是对齐关系,可以可视化「翻译第 j 个词时,模型在看输入的第 i 个词」。
8.3.1 Bahdanau 注意力(加性)
score = v·tanh(W₁·h_dec + W₂·h_enc)
用一个两层小网络打分,输入是解码器状态与编码器状态的组合。论文里叫 additive attention。
8.3.2 Luong 注意力(点积)
score = h_encᵀ·W·h_dec(general 版本),或直接点积 h_enc·h_dec。
Luong 等人的论文在同一时期提出点积/乘性注意力,计算更简单、更易并行。
examples/ch08_seq2seq.py 在「数字倒序」任务(输入 6 位数字,输出倒序)上同时实现两种注意力:
class BahdanauAttention(nn.Module):
def __init__(self):
super().__init__()
self.w1 = nn.Linear(64, 64)
self.w2 = nn.Linear(64, 64)
self.v = nn.Linear(64, 1)
def forward(self, h_dec, enc_out):
score = self.v(torch.tanh(self.w1(h_dec).unsqueeze(1) + self.w2(enc_out))).squeeze(-1)
alpha = score.softmax(dim=1)
return (alpha.unsqueeze(-1) * enc_out).sum(dim=1), alpha
class LuongAttention(nn.Module):
def __init__(self):
super().__init__()
self.w = nn.Linear(64, 64)
def forward(self, h_dec, enc_out):
score = torch.bmm(enc_out, self.w(h_dec).unsqueeze(-1)).squeeze(-1)
alpha = score.softmax(dim=1)
return (alpha.unsqueeze(-1) * enc_out).sum(dim=1), alpha运行结果:
Bahdanau(加性注意力):最终 loss = 0.0010,测试准确率 = 100.0%
Luong(点积注意力) :最终 loss = 0.0013,测试准确率 = 100.0%两种注意力在这个任务上都 100% 准确。注意训练时用教师强制(teacher forcing):解码器输入用真实的前一个目标词,而不是模型自己生成的词;推理时则用自己生成的上一个词,遇到 <EOS> 停止。
8.4 注意力为什么重要
注意力机制是序列建模史上最重要的思想之一,原因有三:
- 解决信息瓶颈:每一步都能访问全部输入;
- 可解释:对齐权重直接展示模型在「看什么」;
- 通往 Transformer:2017 年,研究者发现——既然注意力这么有用,为什么不只用注意力、去掉循环呢?这就是第 9 章。
完整代码
本章用到的完整示例代码:
examples/ch08_seq2seq.py
"""第 8 章示例(2014—2017):Seq2Seq 与两种注意力机制对比。
任务:把一串数字倒序输出(如 3 7 1 2 → 2 1 7 3)。
编码器把输入序列编码成向量序列;解码器用注意力对齐输入位置。
Bahdanau 注意力(2015)用加性打分,Luong 注意力(2015)用点积打分。
运行方式:
uv run python examples/ch08_seq2seq.py
"""
import torch
from torch import nn
torch.manual_seed(0)
device = "cuda" if torch.cuda.is_available() else "cpu"
SOS, EOS, VOCAB, SEQ_LEN = 10, 11, 12, 6
def make_batch(n, generator=None):
src = torch.randint(0, 10, (n, SEQ_LEN), generator=generator)
rev = src.flip(1)
tgt = torch.cat([torch.full((n, 1), SOS), rev, torch.full((n, 1), EOS)], dim=1)
return src, tgt
class Encoder(nn.Module):
def __init__(self):
super().__init__()
self.embed = nn.Embedding(VOCAB, 32)
self.gru = nn.GRU(32, 64, batch_first=True)
def forward(self, x):
return self.gru(self.embed(x))[0]
class BahdanauAttention(nn.Module):
"""加性注意力:score = v·tanh(W1·h_dec + W2·enc)。"""
def __init__(self):
super().__init__()
self.w1 = nn.Linear(64, 64)
self.w2 = nn.Linear(64, 64)
self.v = nn.Linear(64, 1)
def forward(self, h_dec, enc_out):
score = self.v(torch.tanh(self.w1(h_dec).unsqueeze(1) + self.w2(enc_out))).squeeze(-1)
alpha = score.softmax(dim=1)
return (alpha.unsqueeze(-1) * enc_out).sum(dim=1), alpha
class LuongAttention(nn.Module):
"""点积注意力:score = enc · (W·h_dec)。"""
def __init__(self):
super().__init__()
self.w = nn.Linear(64, 64)
def forward(self, h_dec, enc_out):
score = torch.bmm(enc_out, self.w(h_dec).unsqueeze(-1)).squeeze(-1)
alpha = score.softmax(dim=1)
return (alpha.unsqueeze(-1) * enc_out).sum(dim=1), alpha
class Decoder(nn.Module):
def __init__(self, attn):
super().__init__()
self.attn = attn
self.embed = nn.Embedding(VOCAB, 32)
self.gru = nn.GRU(32 + 64, 64, batch_first=True)
self.head = nn.Linear(64, VOCAB)
def forward(self, tgt, enc_out):
h = torch.zeros(1, tgt.shape[0], 64, device=tgt.device)
logits = []
for t in range(tgt.shape[1]):
emb = self.embed(tgt[:, t:t + 1])
context, _ = self.attn(h[-1], enc_out)
out, h = self.gru(torch.cat([emb, context.unsqueeze(1)], dim=-1), h)
logits.append(self.head(out))
return torch.cat(logits, dim=1)
class Seq2Seq(nn.Module):
def __init__(self, attn):
super().__init__()
self.encoder = Encoder()
self.decoder = Decoder(attn)
def forward(self, src, tgt):
return self.decoder(tgt, self.encoder(src))
def decode_greedy(model, src):
model.eval()
with torch.no_grad():
enc_out = model.encoder(src)
h = torch.zeros(1, src.shape[0], 64, device=src.device)
inp = torch.full((src.shape[0], 1), SOS, device=src.device)
out = []
for _ in range(SEQ_LEN + 1):
emb = model.decoder.embed(inp[:, -1:])
context, _ = model.decoder.attn(h[-1], enc_out)
out_step, h = model.decoder.gru(
torch.cat([emb, context.unsqueeze(1)], dim=-1), h
)
pred = model.decoder.head(out_step).argmax(-1)
out.append(pred)
inp = pred
if (pred == EOS).all():
break
return torch.cat(out, dim=1)
def train_and_eval(name, attn):
model = Seq2Seq(attn).to(device)
opt = torch.optim.Adam(model.parameters(), lr=1e-3)
g = torch.Generator().manual_seed(0)
for step in range(1200):
src, tgt = make_batch(64, g)
src, tgt = src.to(device), tgt.to(device)
opt.zero_grad()
# 教师强制:输入 tgt[:, :-1](SOS+前 6 位),预测 tgt[:, 1:](6 位+EOS)
loss = nn.functional.cross_entropy(
model(src, tgt[:, :-1]).transpose(1, 2), tgt[:, 1:]
)
loss.backward()
opt.step()
src_test, tgt_test = make_batch(256, g)
pred = decode_greedy(model, src_test.to(device))
correct = 0
for i in range(256):
gold = tgt_test[i, 1:-1].tolist()
got = pred[i, :SEQ_LEN].tolist()
correct += int(gold == got)
print(f"{name}:最终 loss = {loss.item():.4f},测试准确率 = {correct / 256 * 100:.1f}%")
train_and_eval("Bahdanau(加性注意力)", BahdanauAttention())
train_and_eval("Luong(点积注意力) ", LuongAttention())动手实践
- 运行
examples/ch08_seq2seq.py,把序列长度从 6 改成 12,观察两种注意力的准确率变化,解释为什么变长后困难。 - 在解码器里把
context去掉(只输入嵌入向量),训练并对比准确率,验证注意力对长序列的价值。 - 打印一组测试样本的注意力矩阵(解码每一步对编码各位置的 α),观察倒序任务的对齐是否近似「镜像」。
- 把编码器从 GRU 换成 LSTM,重新训练,对比。
常见错误
| 错误写法 | 现象 | 原因 |
|---|---|---|
| 教师强制目标错位 | loss 很低但推理全错 | 解码输入用 tgt[:, :-1],目标用 tgt[:, 1:] |
| 打分公式里维度没对齐 | mat1 and mat2 shapes cannot be multiplied | h_dec 是 (B, H),enc_out 是 (B, L, H),加性注意要先 unsqueeze 扩展 |
| softmax 维度写错 | 权重加和不为 1 | 对齐权重在序列维(axis=1)上 softmax |
| 推理时还用教师强制 | 结果与训练不符 | 推理必须用自己的生成序列自回归 |
| 生成不设停止条件 | 无限循环 | 遇到 <EOS> 或达到最大长度要停止 |
| 忽略 padding | 对齐权重被 padding 污染 | 真实任务要加 padding mask |
章末练习
基础
- Seq2Seq 的信息瓶颈是什么?
- 注意力中的 αᵢ 表示什么?
- Bahdanau 与 Luong 注意力的打分方式各是什么?
提高
- 把数字倒序任务改成「奇偶翻转」:偶数位置不变、奇数位置取反,训练并评估。
- 实现 Luong 的 dot 版本(无参数 W),与 general 版本对比。
- 在测试集上统计两种注意力在不同序列长度(4/6/8/10)下的准确率,画出趋势。
挑战
- 实现一个「加两个数」的 Seq2Seq:输入 "12+34",输出 "46"(注意对齐),训练并评估。
- 用注意力权重做一个对齐可视化:对一个翻译样本,画出输入词 × 输出词的热力图(用 matplotlib 或纯文本矩阵)。
章末自测
- Seq2Seq 的两个组成部分是?
- A. 生成器与判别器
- B. 编码器与解码器
- C. Q 与 V
- D. 编码器与池化
- 信息瓶颈指?
- A. 数据太少
- B. 所有信息被压缩进一个向量
- C. 显存不足
- D. 序列太短
- 注意力的核心操作是?
- A. 卷积
- B. 加权求和
- C. 池化
- D. 归一化
- αᵢ 经过什么得到?
- A. sigmoid
- B. softmax
- C. relu
- D. tanh
- Bahdanau 注意力又称?
- A. 点积注意力
- B. 加性注意力
- C. 卷积注意力
- D. 全局注意力
- 教师强制指?
- A. 用真实目标词作为解码输入
- B. 用模型自己的输出
- C. 用随机词
- D. 不输入任何词
- 推理时解码器输入来自?
- A. 真实目标
- B. 自己上一步的输出
- C. 编码器输出
- D. 随机
- 注意力权重可视化的意义是?
- A. 展示对齐关系
- B. 提高准确率
- C. 减少参数
- D. 加速训练
- Luong 注意力的打分比 Bahdanau 更?
- A. 复杂
- B. 简单(点积)
- C. 慢
- D. 不准确
- 注意力思想直接通往哪个架构?
- A. CNN
- B. Transformer
- C. GAN
- D. Hopfield
