序列到序列模型 本节摘要:两个 RNN 假装自己是翻译官。它们撞上的瓶颈,正是注意力存在的理由。分类把变长序列映到单个标签,翻译则把变长序列映到另一个变长序列——输入输出活在不同的词表里,可能是不同的语言,长度也未必对等。seq2seq 架构(Sutskever、Vinyals、Le, 2014)用一个刻意简单的配方破解了它:两个 RNN,一个读源句产出一个定长上下文向量,另一个读这个向量逐 token 生成目标句。第 08 节写的代码换个拼法而已。它值得学,有两个原因:其一,上下文向量瓶颈是 NLP 里教学价值最高的失败,它解释了注意力与 Transformer 一切所长;其二,它的训练配方(教师强制、计划采样、推理时束搜索)至今仍适用于包括 LLM 在内的每个现代生成系统。
本节摘要:两个 RNN 假装自己是翻译官。它们撞上的瓶颈,正是注意力存在的理由。分类把变长序列映到单个标签,翻译则把变长序列映到另一个变长序列——输入输出活在不同的词表里,可能是不同的语言,长度也未必对等。seq2seq 架构(Sutskever、Vinyals、Le, 2014)用一个刻意简单的配方破解了它:两个 RNN,一个读源句产出一个定长上下文向量,另一个读这个向量逐 token 生成目标句。第 08 节写的代码换个拼法而已。它值得学,有两个原因:其一,上下文向量瓶颈是 NLP 里教学价值最高的失败,它解释了注意力与 Transformer 一切所长;其二,它的训练配方(教师强制、计划采样、推理时束搜索)至今仍适用于包括 LLM 在内的每个现代生成系统。
对应原课程:Phase 5 · Lesson 09 ·
sequence-to-sequence(原英文phases/05-nlp-foundations-to-advanced/09-sequence-to-sequence/docs/en.md)。前置依赖:第 08 节(文本的 CNN 与 RNN)、Phase 3 · 11(PyTorch 入门)。
阅读完本节,你应当能够:
分类把变长序列映到单个标签,翻译把变长序列映到另一个变长序列。输入输出活在不同的词表里,可能不同语言,长度也未必对等。
seq2seq 架构(Sutskever、Vinyals、Le, 2014)用一个刻意简单的配方破解了它:两个 RNN。一个读源句,产出一个定长的上下文向量;另一个读这个向量,逐 token 生成目标句。第 08 节的代码换个拼法而已。
<EOS> 或撞到最大长度。t 步的输入是位置 t-1 的真值 token,而非解码器自己的上一步预测。这稳定训练;没有它,早期错误会雪崩,模型永远学不会。推理时只能用模型自己的预测,所以总有训练/推理分布 gap,这个 gap 叫暴露偏差。注意力(第 10 节)修掉了这个,让解码器看每一个编码器隐状态,而非只看最后一个。这就是它的全部卖点。
import torch import torch.nn as nn class Encoder(nn.Module): def __init__(self, src_vocab_size, embed_dim, hidden_dim): super().__init__() self.embed = nn.Embedding(src_vocab_size, embed_dim, padding_idx=0) self.gru = nn.GRU(embed_dim, hidden_dim, batch_first=True) def forward(self, src): e = self.embed(src) outputs, hidden = self.gru(e) return outputs, hidden
outputs 形状 [batch, seq_len, hidden_dim]——每个输入位置一个隐状态。hidden 形状 [1, batch, hidden_dim]——最终步。第 08 节说「对 outputs 池化做分类」,这里我们留最后隐状态当上下文向量,忽略逐步输出。
class Decoder(nn.Module): def __init__(self, tgt_vocab_size, embed_dim, hidden_dim): super().__init__() self.embed = nn.Embedding(tgt_vocab_size, embed_dim, padding_idx=0) self.gru = nn.GRU(embed_dim, hidden_dim, batch_first=True) self.fc = nn.Linear(hidden_dim, tgt_vocab_size) def forward(self, token, hidden): e = self.embed(token) out, hidden = self.gru(e, hidden) logits = self.fc(out) return logits, hidden
解码器一次调一步。输入是一批单 token 和当前隐状态,输出是下一 token 的词表 logits 和更新后的隐状态。
def train_batch(encoder, decoder, src, tgt, bos_id, optimizer, teacher_forcing_ratio=0.9): optimizer.zero_grad() _, hidden = encoder(src) batch_size, tgt_len = tgt.shape input_token = torch.full((batch_size, 1), bos_id, dtype=torch.long) loss = 0.0 loss_fn = nn.CrossEntropyLoss(ignore_index=0) for t in range(tgt_len): logits, hidden = decoder(input_token, hidden) step_loss = loss_fn(logits.squeeze(1), tgt[:, t]) loss += step_loss use_teacher = torch.rand(1).item() < teacher_forcing_ratio if use_teacher: input_token = tgt[:, t].unsqueeze(1) else: input_token = logits.argmax(dim=-1) loss.backward() optimizer.step() return loss.item() / tgt_len
两个旋钮值得一说。ignore_index=0 跳过 padding token 的损失。teacher_forcing_ratio 是每步用真值 token 还是模型预测的概率。从 1.0(全教师强制)起,训练中退火到约 0.5,以缩小暴露偏差 gap。
@torch.no_grad() def greedy_decode(encoder, decoder, src, bos_id, eos_id, max_len=50): _, hidden = encoder(src) batch_size = src.shape[0] input_token = torch.full((batch_size, 1), bos_id, dtype=torch.long) output_ids = [] for _ in range(max_len): logits, hidden = decoder(input_token, hidden) next_token = logits.argmax(dim=-1) output_ids.append(next_token) input_token = next_token if (next_token == eos_id).all(): break return torch.cat(output_ids, dim=1)
贪心解码每步选概率最高的 token,但它会跑偏:一旦选定一个 token,就收不回。束搜索保留前 k 个部分序列,最后挑得分最高的完整序列。束宽 3~5 是标准。
在一个玩具复制任务上训:源 [a, b, c, d, e],目标 [a, b, c, d, e]。增加序列长度,观察准确率。
seq_len=5 复制准确率: 98% seq_len=10 复制准确率: 91% seq_len=20 复制准确率: 62% seq_len=40 复制准确率: 23%
单个 GRU 隐状态无法无损记住 40 个 token 的输入。信息在编码器的每一步都在,但解码器只看到最后状态。注意力直接修掉这个。
PyTorch 有 nn.Transformer 和基于 nn.LSTM 的 seq2seq 模板。Hugging Face 的 transformers 提供在数十亿 token 上训练的完整编码器-解码器模型(BART、T5、mBART、NLLB)。
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM tok = AutoTokenizer.from_pretrained("facebook/bart-base") model = AutoModelForSeq2SeqLM.from_pretrained("facebook/bart-base") src = tok("Translate this to French: Hello, how are you?", return_tensors="pt") out = model.generate(**src, max_new_tokens=50, num_beams=4) print(tok.decode(out[0], skip_special_tokens=True))
现代编码器-解码器用 Transformer 替掉了 RNN,但高层形状(编码器、解码器、逐 token 生成)与 2014 年的 seq2seq 论文一模一样,变的只是每个块内部的机制。
对新项目,几乎不再用。特例:
三者至今仍适用于基于 Transformer 的生成。
保存为 outputs/prompt-seq2seq-design.md:
--- name: seq2seq-design description: Design a sequence-to-sequence pipeline for a given task. phase: 5 lesson: 09 --- Given a task (translation, summarization, paraphrase, question rewrite), output: 1. Architecture. Pretrained transformer encoder-decoder (BART, T5, mBART, NLLB) is the default. RNN-based seq2seq only for specific constraints. 2. Starting checkpoint. Name it (`facebook/bart-base`, `google/flan-t5-base`, `facebook/nllb-200-distilled-600M`). Match the checkpoint to task and language coverage. 3. Decoding strategy. Greedy for deterministic output, beam search (width 4-5) for quality, sampling with temperature for diversity. One sentence justification. 4. One failure mode to verify before shipping. Exposure bias manifests as generation drift on longer outputs; sample 20 outputs at the 90th-percentile length and eyeball. Refuse to recommend training a seq2seq from scratch for under a million parallel examples. Flag any pipeline that uses greedy decoding for user-facing content as fragile (greedy repeats and loops).
facebook/bart-base。对比微调模型与基线模型的束 4 输出,报 BLEU 并挑 10 个定性例子。teacher_forcing_ratio 从 1.0 退火到约 0.5 缩小暴露偏差。下一节,我们让解码器不再只盯一个向量——进入「注意力机制」,看它如何一击解决 seq2seq 的瓶颈,并为 Transformer 铺平道路。