Skip to content

第 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 位数字,输出倒序)上同时实现两种注意力:

python
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 注意力为什么重要 ​

注意力机制是序列建模史上最重要的思想之一,原因有三:

  1. 解决信息瓶颈:每一步都能访问全部输入;
  2. 可解释:对齐权重直接展示模型在「看什么」;
  3. 通往 Transformer:2017 年,研究者发现——既然注意力这么有用,为什么不只用注意力、去掉循环呢?这就是第 9 章。

完整代码 ​

本章用到的完整示例代码:

examples/ch08_seq2seq.py ​

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())

动手实践 ​

  1. 运行 examples/ch08_seq2seq.py,把序列长度从 6 改成 12,观察两种注意力的准确率变化,解释为什么变长后困难。
  2. 在解码器里把 context 去掉(只输入嵌入向量),训练并对比准确率,验证注意力对长序列的价值。
  3. 打印一组测试样本的注意力矩阵(解码每一步对编码各位置的 α),观察倒序任务的对齐是否近似「镜像」。
  4. 把编码器从 GRU 换成 LSTM,重新训练,对比。

常见错误 ​

错误写法现象原因
教师强制目标错位loss 很低但推理全错解码输入用 tgt[:, :-1],目标用 tgt[:, 1:]
打分公式里维度没对齐mat1 and mat2 shapes cannot be multipliedh_dec 是 (B, H),enc_out 是 (B, L, H),加性注意要先 unsqueeze 扩展
softmax 维度写错权重加和不为 1对齐权重在序列维(axis=1)上 softmax
推理时还用教师强制结果与训练不符推理必须用自己的生成序列自回归
生成不设停止条件无限循环遇到 <EOS> 或达到最大长度要停止
忽略 padding对齐权重被 padding 污染真实任务要加 padding mask

章末练习 ​

基础

  1. Seq2Seq 的信息瓶颈是什么?
  2. 注意力中的 αᵢ 表示什么?
  3. Bahdanau 与 Luong 注意力的打分方式各是什么?

提高

  1. 把数字倒序任务改成「奇偶翻转」:偶数位置不变、奇数位置取反,训练并评估。
  2. 实现 Luong 的 dot 版本(无参数 W),与 general 版本对比。
  3. 在测试集上统计两种注意力在不同序列长度(4/6/8/10)下的准确率,画出趋势。

挑战

  1. 实现一个「加两个数」的 Seq2Seq:输入 "12+34",输出 "46"(注意对齐),训练并评估。
  2. 用注意力权重做一个对齐可视化:对一个翻译样本,画出输入词 × 输出词的热力图(用 matplotlib 或纯文本矩阵)。

章末自测 ​

  1. Seq2Seq 的两个组成部分是?
    • A. 生成器与判别器
    • B. 编码器与解码器
    • C. Q 与 V
    • D. 编码器与池化
  2. 信息瓶颈指?
    • A. 数据太少
    • B. 所有信息被压缩进一个向量
    • C. 显存不足
    • D. 序列太短
  3. 注意力的核心操作是?
    • A. 卷积
    • B. 加权求和
    • C. 池化
    • D. 归一化
  4. αᵢ 经过什么得到?
    • A. sigmoid
    • B. softmax
    • C. relu
    • D. tanh
  5. Bahdanau 注意力又称?
    • A. 点积注意力
    • B. 加性注意力
    • C. 卷积注意力
    • D. 全局注意力
  6. 教师强制指?
    • A. 用真实目标词作为解码输入
    • B. 用模型自己的输出
    • C. 用随机词
    • D. 不输入任何词
  7. 推理时解码器输入来自?
    • A. 真实目标
    • B. 自己上一步的输出
    • C. 编码器输出
    • D. 随机
  8. 注意力权重可视化的意义是?
    • A. 展示对齐关系
    • B. 提高准确率
    • C. 减少参数
    • D. 加速训练
  9. Luong 注意力的打分比 Bahdanau 更?
    • A. 复杂
    • B. 简单(点积)
    • C. 慢
    • D. 不准确
  10. 注意力思想直接通往哪个架构?
    • A. CNN
    • B. Transformer
    • C. GAN
    • D. Hopfield