
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.pyAttentionalLSTMDecoder第 139 行起逐词解码每步输入 上一词嵌入 上一步输出注意力model/attention.pystep_attention第 22 行起点积打分 → mask 掉 padding → softmax → 加权聚合出 context 向量解码一步的核心流程在AttentionalLSTMDecoder._stepmodel/summ.py第 158-173 行上一词嵌入和上一步输出拼接喂给多层 LSTMmodel/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_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 行逻辑如下先取(-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继续喂给下一 stepmodel/copy_summ.py的decode/batch_decode同时把真实 id 保留在outputs里最终按 id 查扩展词表还原出原文的词。 扩展词表是怎么来的词表是每个 batch 动态构建的相关逻辑在data/batcher.pyconvert_batch_copy第 67-80 行扫描 batch 内所有源句把没出现过的新词依次分配新 id形成ext_word2idbatchify_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对每个 batch 内所有 beam 打包调用topk_step取 top-k 候选注意代码里把 beam 维折进 batch 维再在 copy 分支重新展开因为_CopyLinear不支持 beam 广播——第 133 行附近的注释 copy mechanism is not beamable 说的就是这个beam_size1 即贪心解码beam_size5 时论文中用抽取模型打分做 rerank每一步同步记录注意力分数便于事后可视化这个词是从哪抄的。 小结与关键文件速查想深入了解看这里基础注意力 Seq2Seqmodel/summ.pySeq2SeqSumm、AttentionalLSTMDecoder点积注意力打分与 maskmodel/attention.pystep_attentionCopy 门与概率融合model/copy_summ.py_CopyLinear、CopyLSTMDecoder._step扩展词表构建data/batcher.pyconvert_batch_copy、batchify_fn_copy摘要器训练脚本train_abstractor.pyMatchDataset、configure_net束搜索解码model/copy_summ.pybatched_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),仅供参考