fast_abs_rl源码解读②:带Copy机制的Seq2Seq摘要器——注意力+复制如何兼顾生成与忠实?
2026/8/25 8:55:03 网站建设 项目流程

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.pylstm_encoder
解码器model/summ.pyAttentionalLSTMDecoder,第 139 行起)逐词解码,每步输入 = 上一词嵌入 + 上一步输出
注意力model/attention.pystep_attention,第 22 行起)点积打分 → mask 掉 padding → softmax → 加权聚合出 context 向量

解码一步的核心流程在AttentionalLSTMDecoder._stepmodel/summ.py第 158-173 行):

  1. 上一词嵌入和上一步输出拼接,喂给多层 LSTM(model/rnn.py中手写的MultiLayerLSTMCells,方便逐 step 解码);
  2. 用 LSTM 输出乘投影矩阵_attn_w得到 query,调用step_attention对源句算注意力,得到 context;
  3. 把 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_probmodel/copy_summ.py第 251-262 行)先像基础版一样投影出词表 logit,然后有一个关键细节:如果本 batch 的"扩展词表"比模型词表大,就在 logit 后面补一段常数(eps = 1e-6)拉齐长度,再做 softmax

为什么要补?因为 copy 机制的总词表 = 基础词表 ∪ 当前源句里的所有词(源句特有的实体词会分配新 id)。生成概率必须在"扩展词表"上归一化,否则和复制概率不在同一概率空间里。补 eps 就是给这些新词一个极小的生成先验,防止分母算错。

第二步:算复制门 p_copy

_CopyLinearmodel/copy_summ.py第 15-35 行)是一个可学习的小打分器:它对context 向量、LSTM 状态、解码器输入三个来源分别做向量点积再相加,过 sigmoid 得到 0~1 之间的copy_prob。直观理解:

  • context 里原文信息多、状态倾向于"照抄"时 →p_copy升高;
  • 需要改写、润色时 →p_copy降低,更多概率留给生成路。

第三步:融合两条概率

最终 log 概率用一行scatter_add完成(第 199-205 行,逻辑如下):

  1. 先取(-copy_prob + 1) * gen_prob,即整个分布整体乘以"不复制"的比例;
  2. 再沿着注意力分数score的下标(也就是源词在扩展词表中的位置)做scatter_add,把score * copy_prob累加回去——注意力分布被缩放后直接"注入"到对应源词的格子中
  3. 加 1e-8 取 log,数值上更稳。

这一行代码正是P(w) = (1-p_copy)·P_gen(w) + p_copy·P_attn(w)的张量化实现,也是全模型最精妙的一行。

💡 解码时的一个小技巧:输出 token id 若 ≥ 基础词表大小(vsize),说明是"复制来的源词",代码会临时把它映射回unk继续喂给下一 step(model/copy_summ.pydecode/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_lossmodel/util.py第 29 行起),按 pad 位置做 mask;
  • 数据管线由BucketedGeneraterdata/batcher.py第 206 行起)按长度分桶 + 多进程预取,减少 padding 浪费。

训练完成后,摘要器会作为子模块被 RL 阶段(train_full_rl.py)复用——RL 策略从候选句里"选句子",选中的句交给摘要器"重写",这就是论文标题中 "Reinforce-Selected Sentence Rewriting" 的含义。

🚀 解码时它如何工作?

推理入口是 decode_full_model.py,摘要器走CopySumm.batched_beamsearchmodel/copy_summ.py第 97-172 行),配合model/beam_search.py的 diverse beam search:

  1. 对每个 batch 内所有 beam 打包,调用topk_step取 top-k 候选(注意代码里把 beam 维折进 batch 维,再在 copy 分支重新展开,因为_CopyLinear不支持 beam 广播——第 133 行附近的注释 "copy mechanism is not beamable" 说的就是这个);
  2. beam_size=1 即贪心解码,beam_size=5 时论文中用抽取模型打分做 rerank;
  3. 每一步同步记录注意力分数,便于事后可视化"这个词是从哪抄的"。

📝 小结与关键文件速查

想深入了解看这里
基础注意力 Seq2Seqmodel/summ.pySeq2SeqSummAttentionalLSTMDecoder
点积注意力打分与 maskmodel/attention.pystep_attention
Copy 门与概率融合model/copy_summ.py_CopyLinearCopyLSTMDecoder._step
扩展词表构建data/batcher.pyconvert_batch_copybatchify_fn_copy
摘要器训练脚本train_abstractor.pyMatchDatasetconfigure_net
束搜索解码model/copy_summ.pybatched_beamsearchmodel/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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询