在 fairseq 中复现分层神经故事生成:WritingPrompts 数据集从预处理、训练到融合模型生成的完整实战指南
【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm
本指南以unilm仓库中 infoxlm/fairseq/examples/stories/README.md 为核心,系统讲解如何在 fairseq 框架下复现 Fan et al. (2018) 提出的分层神经故事生成(Hierarchical Neural Story Generation)方案:包括 WritingPrompts 数据集的下载与 1000 词截断、fairseq-preprocess二值化、基于卷积 seq2seq 模型fconv_self_att的训练、融合模型(fusion model)的两阶段训练,以及基于 top-k 采样的故事生成。读完本文,你将掌握这套故事生成模型在 fairseq 中的完整可运行链路,并理解其底层卷积注意力架构与模型融合机制。
1. 任务背景与模型概览
分层神经故事生成(Hierarchical Neural Story Generation)由 Fan, Lewis 和 Dauphin 在 2018 年 ACL 论文中提出。其核心思路是:给定一句故事提示(prompt),让模型续写出完整的故事情节。该任务使用WritingPrompts数据集——一个从 Reddit r/WritingPrompts 社区爬取的大规模"提示-故事"对数据集。
在 fairseq 中,该任务通过卷积 seq2seq 模型族fconv_self_att实现,包含两种模型形态:
- 卷积模型(Convolutional Model):以卷积网络为编码器和解码器主体,叠加自注意力与多头注意力,实现提示到故事的映射;
- 融合模型(Fusion Model):在已训练好的卷积模型基础上,再训练一个融合了解码器,通过可学习的门控机制结合预训练模型与新模型的隐状态,从而在训练数据较少的情况下获得更好的故事生成质量。
在本仓库中,该模型的完整实现位于 infoxlm/fairseq/fairseq/models/fconv_self_att.py,通过@register_model('fconv_self_att')注册(见该文件第 31 行),并预置了面向 WritingPrompts 任务的模型架构配置fconv_self_att_wp。
2. 预训练模型与样例故事
原 README 提供了如下与论文对应的预训练资源:
| 说明 | 数据集 | 模型 | 测试集 |
|---|---|---|---|
| 卷积模型故事生成(Fan et al., 2018) | WritingPrompts | 预训练 checkpoint | 测试集数据(含词典) |
注意:上述资源的下载地址仅存在于原 README 及源码中,本文不重复列出外部链接。其中模型与数据的下载路径被固化在 fconv_self_att.py 的
hub_models类方法中,注册了三个可加载条目:conv.stories.pretrained(预训练卷积模型)、conv.stories(融合模型,加载时依赖pretrained_checkpoint)与data.stories(含词典的测试集)。
原 README 还提到,论文提供了卷积 seq2seq 模型和融合模型生成的样例故事文件,以及融合模型对应的提示文件。需要特别说明的是:这些样例文件中存在unk标记,因为该实验建模的是一个小型完整词表(未使用 BPE 或预训练),并且这些带unk的提示未用于人工评估。这一细节对理解后续--thresholdtgt 10 --thresholdsrc 10的词频截断设置(见第 4 节)很重要——低频词被映射为unk正是小词表策略的体现。
3. 数据集:WritingPrompts 的下载与裁剪
3.1 下载与解压
原 README 给出的下载方式是在examples/stories目录下执行(下载地址见 README 原文,本文以占位符表示):
cd infoxlm/fairseq/examples/stories curl <WritingPrompts 数据包下载地址(见 README 原文)> | tar xvzf -解压后得到的数据集包含train、test、valid三个划分。数据集由原论文(arXiv: 1805.04833)描述,其格式为成对的.wp_source(提示)与.wp_target(故事)文件。
3.2 裁剪到前 1000 词
原 README 明确指出:数据集发布的是完整数据,但论文只对每篇故事的前 1000 个词进行建模(包含一个换行 token)。因此训练前必须先将每个故事裁剪到前 1000 词。原文档提供了如下 Python 裁剪脚本:
data = ["train", "test", "valid"] for name in data: with open(name + ".wp_target") as f: stories = f.readlines() stories = [" ".join(i.split()[0:1000]) for i in stories] with open(name + ".wp_target", "w") as o: for line in stories: o.write(line.strip() + "\n")这段脚本对三个划分的.wp_target文件逐一处理:按空白切分取前 1000 个 token 后重新拼接写入。i.split()会丢弃原始换行,因此 1000 词限制是硬性的;脚本末行line.strip() + "\n"保证每篇故事独占一行,便于后续 fairseq 的按行文本数据读取。提示侧(.wp_source)无需裁剪。
4. 数据预处理:fairseq-preprocess 与参数解析
数据裁剪完成后,需要将文本二值化为 fairseq 的二进制数据集。原文档命令如下:
# Binarize the dataset: export TEXT=examples/stories/writingPrompts fairseq-preprocess --source-lang wp_source --target-lang wp_target \ --trainpref $TEXT/train --validpref $TEXT/valid --testpref $TEXT/test \ --destdir>fairseq-train># Train a fusion model: # add the arguments: --pretrained True --pretrained-checkpoint path/to/checkpoint即在fairseq-train命令末尾加上--pretrained True --pretrained-checkpoint path/to/checkpoint,其中 checkpoint 指向第一阶段(普通卷积模型)训练好的模型文件。
从 fconv_self_att.py 的build_model可以看出融合机制的实现细节:
- 当
--pretrained True时,源码会通过checkpoint_utils.load_model_ensemble加载预训练 checkpoint,并将预训练编码器与解码器的全部参数requires_grad置为 False(冻结); - 随后将预训练编码器与本次训练的编码器一起包装进
CompositeEncoder(见该文件第 53-63 行),两者前向结果在解码器中合并; - 模型类
FConvModelSelfAtt还会把编码器注意力层数计入num_attention_layers,用于梯度缩放。
融合发生在解码器的隐状态层面(见 fconv_self_att.py 与第 450-466 行的前向逻辑):模型为预训练解码器注册了一个前向 hook(register_forward_hook),捕获其fc2输出作为预训练隐状态;随后:
- 用两个可学习的 Sigmoid 门
gate1、gate2分别对"新模型隐状态"和"预训练模型隐状态"做逐元素门控; - 将两个门控结果拼接,送入
joining模块——一个由 线性层 → LayerNorm → GLU → 线性层 → LayerNorm → GLU → 线性层 → LayerNorm 组成的多层门控单元; - 最后经
fc3映射到词表得到融合后的 logits。
这种"冻结预训练 + 门控融合"的设计正是融合模型能用更少训练数据生成更连贯故事的关键。
仓库中的单元测试 infoxlm/fairseq/tests/test_binaries.py 完整复现了这一两阶段流程:先以fconv_self_att_wp架构、小规模配置(如[(128, 3)] * 2层、嵌入维度 8)训练普通模型,随后把checkpoint_last.pt改名为pretrained.pt,再以--pretrained True --pretrained-checkpoint <pretrained.pt>续训出融合模型并存到独立目录。这为理解上述命令提供了可直接对照的自动化验证样例。
7. 生成故事:fairseq-generate 与采样参数
7.1 生成命令
训练完成后,使用fairseq-generate进行故事生成。原文档命令如下:
fairseq-generate>--model-overrides "{'pretrained_checkpoint':'/path/to/pretrained/model/checkpoint'}"--model-overrides在 options.py 中定义,默认值为空字典"{}",用于在加载 checkpoint 时以 JSON 字典形式覆盖存档中的模型参数。需要特别留意的是:如果是从非融合模型(普通卷积模型)生成,则完全不需要--model-overrides参数。
8. 全流程速查
将上述步骤串成一条完整的可运行链路(目录均相对仓库根):
- 下载并解压 WritingPrompts 数据集(下载地址见 stories/README.md 原文);
- 用第 3.2 节的 Python 脚本将
train/test/valid的故事侧裁剪到前 1000 词; fairseq-preprocess二值化数据到data-bin/writingPrompts(--padding-factor 1 --thresholdtgt 10 --thresholdsrc 10);fairseq-train以fconv_self_att_wp训练普通卷积模型(--pretrained False);- (可选)追加
--pretrained True --pretrained-checkpoint训练融合模型; fairseq-generate以 top-k 采样(--beam 1 --sampling --sampling-topk 10 --temperature 0.8)生成故事,融合模型记得传--model-overrides。
各阶段命令均可对照源码验证参数行为:预处理参数见 options.py,模型架构与融合逻辑见 fconv_self_att.py,端到端流程见 test_binaries.py 的自动化测试。
9. 引用
本指南对应的原始研究工作,请引用:
@inproceedings{fan2018hierarchical, title = {Hierarchical Neural Story Generation}, author = {Fan, Angela and Lewis, Mike and Dauphin, Yann}, booktitle = {Conference of the Association for Computational Linguistics (ACL)}, year = 2018, }【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考