fast_abs_rl源码解读②:带Copy机制的Seq2Seq摘要器——注意力+复制如何兼顾生成与忠实?
【免费下载链接】fast_abs_rlCode for ACL 2018 paper: "Fast Abstractive Summarization with Reinforce-Selected Sentence Rewriting. Chen and Bansal"项目地址: https://gitcode.com/gh_mirrors/fa/fast_abs_rl
在上一篇我们了解了 fast_abs_rl 的整体架构,这一篇深入其核心:带 Copy 机制的 Seq2Seq 摘要器。fast_abs_rl 是 ACL 2018 论文《Fast Abstractive Summarization with Reinforce-Selected Sentence Rewriting》的官方代码,它的摘要器(abstractor)用注意力(attention)+ 复制(copy)机制同时解决两个痛点:既能"改写"原文生成流畅句子,又能"照抄"原文保证内容忠实。下面用尽量少的代码,带你把这个模型拆开看。
🧩 为什么摘要器需要"生成 + 复制"两种能力?
纯神经生成式摘要有一个顽疾:训练词表有限,遇到原文里的实体(人名、地名、数字)容易写出错。Copy 机制(源自 Gu et al. 的 copy 机制)给模型开了一条"后门"——
- 生成(Generate):像正常 Seq2Seq 一样,从固定词表里预测下一个词,负责改写和润色;
- 复制(Copy):直接从源文章里把词"搬"过来,负责忠实引用实体和事实。
在 fast_abs_rl 里,这两条路被融合成经典的pointer-generator 结构:
P(w) = (1 - p_copy) · P_gen(w) + p_copy · P_attn(w)其中P_gen是词表上的生成概率,P_attn是注意力分布落到源词上的概率,p_copy是模型自己学出来的"复制门"。下面逐块拆解源码实现。
🔍 先看基础版:Seq2SeqSumm 的三件套
基础版摘要器在 Seq2SeqSumm 中定义(类Seq2SeqSumm,第 14 行起),由三部分组成:
| 组件 | 源码位置 | 说明 |
|---|---|---|
| 编码器 | model/summ.py(_enc_lstm,双向 LSTM) | 把源句编码为序列表示,支持 pack/unpack 变长序列(见model/rnn.py的lstm_encoder) |
| 解码器 | model/summ.py(AttentionalLSTMDecoder,第 139 行起) | 逐词解码,每步输入 = 上一词嵌入 + 上一步输出 |
| 注意力 | model/attention.py(step_attention,第 22 行起) | 点积打分 → mask 掉 padding → softmax → 加权聚合出 context 向量 |
解码一步的核心流程在AttentionalLSTMDecoder._step(model/summ.py第 158-173 行):
- 上一词嵌入和上一步输出拼接,喂给多层 LSTM(
model/rnn.py中手写的MultiLayerLSTMCells,方便逐 step 解码); - 用 LSTM 输出乘投影矩阵
_attn_w得到 query,调用step_attention对源句算注意力,得到 context; - 把 LSTM 输出和 context 拼接后过
_projection,再复用词嵌入矩阵的转置做线性投影,得到词表上的 logit——这是 2016 年后 NMT 模型常用的权重共享技巧。
到这里是一个标准的注意力 Seq2Seq,但词表外的词只能输出<unk>。真正的升级在下一节。
✨ CopySumm:指针-生成器的三步流水线
带 Copy 的版本是model/copy_summ.py中的CopySumm类(第 38 行起),它继承自Seq2SeqSumm,只多了一个_copy打分网络和自定义解码器CopyLSTMDecoder(第 175 行起)。解码每一步在_step(第 180-206 行)中分三步走:
第一步:算生成概率 P_gen
_compute_gen_prob(model/copy_summ.py第 251-262 行)先像基础版一样投影出词表 logit,然后有一个关键细节:如果本 batch 的"扩展词表"比模型词表大,就在 logit 后面补一段常数(eps = 1e-6)拉齐长度,再做 softmax。
为什么要补?因为 copy 机制的总词表 = 基础词表 ∪ 当前源句里的所有词(源句特有的实体词会分配新 id)。生成概率必须在"扩展词表"上归一化,否则和复制概率不在同一概率空间里。补 eps 就是给这些新词一个极小的生成先验,防止分母算错。
第二步:算复制门 p_copy
_CopyLinear(model/copy_summ.py第 15-35 行)是一个可学习的小打分器:它对context 向量、LSTM 状态、解码器输入三个来源分别做向量点积再相加,过 sigmoid 得到 0~1 之间的copy_prob。直观理解:
- context 里原文信息多、状态倾向于"照抄"时 →
p_copy升高; - 需要改写、润色时 →
p_copy降低,更多概率留给生成路。
第三步:融合两条概率
最终 log 概率用一行scatter_add完成(第 199-205 行,逻辑如下):
- 先取
(-copy_prob + 1) * gen_prob,即整个分布整体乘以"不复制"的比例; - 再沿着注意力分数
score的下标(也就是源词在扩展词表中的位置)做scatter_add,把score * copy_prob累加回去——注意力分布被缩放后直接"注入"到对应源词的格子中; - 加 1e-8 取 log,数值上更稳。
这一行代码正是P(w) = (1-p_copy)·P_gen(w) + p_copy·P_attn(w)的张量化实现,也是全模型最精妙的一行。
💡 解码时的一个小技巧:输出 token id 若 ≥ 基础词表大小(
vsize),说明是"复制来的源词",代码会临时把它映射回unk继续喂给下一 step(model/copy_summ.py的decode/batch_decode),同时把真实 id 保留在outputs里,最终按 id 查扩展词表还原出原文的词。
📦 扩展词表是怎么来的?
词表是每个 batch 动态构建的,相关逻辑在data/batcher.py:
convert_batch_copy(第 67-80 行):扫描 batch 内所有源句,把没出现过的新词依次分配新 id,形成ext_word2id;batchify_fn_copy(第 139-158 行):把ext_src(源句按扩展词表编码)和ext_vsize一起打包给模型,ext_vsize = 扩展词表中最大 id + 1。
也就是说,模型词表是"固定底座 + 每个 batch 的动态扩展",这正是_compute_gen_prob里要做长度对齐的原因。
🏋️ 它是怎么被训练出来的?
摘要器由 train_abstractor.py 训练,目标不是"整篇文章 → 整篇摘要",而是单句到单句的改写:
MatchDataset(第 38-50 行):用抽取模型预先选出的源句(extracts)与摘要句一一对齐,构成"源句 → 摘要句"的训练对;- 损失函数是标准序列交叉熵
sequence_loss(model/util.py第 29 行起),按 pad 位置做 mask; - 数据管线由
BucketedGenerater(data/batcher.py第 206 行起)按长度分桶 + 多进程预取,减少 padding 浪费。
训练完成后,摘要器会作为子模块被 RL 阶段(train_full_rl.py)复用——RL 策略从候选句里"选句子",选中的句交给摘要器"重写",这就是论文标题中 "Reinforce-Selected Sentence Rewriting" 的含义。
🚀 解码时它如何工作?
推理入口是 decode_full_model.py,摘要器走CopySumm.batched_beamsearch(model/copy_summ.py第 97-172 行),配合model/beam_search.py的 diverse beam search:
- 对每个 batch 内所有 beam 打包,调用
topk_step取 top-k 候选(注意代码里把 beam 维折进 batch 维,再在 copy 分支重新展开,因为_CopyLinear不支持 beam 广播——第 133 行附近的注释 "copy mechanism is not beamable" 说的就是这个); - beam_size=1 即贪心解码,beam_size=5 时论文中用抽取模型打分做 rerank;
- 每一步同步记录注意力分数,便于事后可视化"这个词是从哪抄的"。
📝 小结与关键文件速查
| 想深入了解 | 看这里 |
|---|---|
| 基础注意力 Seq2Seq | model/summ.py:Seq2SeqSumm、AttentionalLSTMDecoder |
| 点积注意力打分与 mask | model/attention.py:step_attention |
| Copy 门与概率融合 | model/copy_summ.py:_CopyLinear、CopyLSTMDecoder._step |
| 扩展词表构建 | data/batcher.py:convert_batch_copy、batchify_fn_copy |
| 摘要器训练脚本 | train_abstractor.py:MatchDataset、configure_net |
| 束搜索解码 | model/copy_summ.py:batched_beamsearch、model/beam_search.py |
一句话总结:注意力负责"看哪里",Copy 门负责"抄还是写",生成路径负责"怎么润色"——三者共用一套词嵌入权重,用scatter_add一行完成概率融合,让 fast_abs_rl 的摘要器既流畅又忠实,这也是它在 CNN/DailyMail 上取得当年 SOTA 的关键设计之一。
【免费下载链接】fast_abs_rlCode for ACL 2018 paper: "Fast Abstractive Summarization with Reinforce-Selected Sentence Rewriting. Chen and Bansal"项目地址: https://gitcode.com/gh_mirrors/fa/fast_abs_rl
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考