基于 fairseq 的分层神经故事生成实战指南:WritingPrompts 数据预处理、卷积模型训练与采样生成
2026/9/14 22:22:39 网站建设 项目流程

基于 fairseq 的分层神经故事生成实战指南:WritingPrompts 数据预处理、卷积模型训练与采样生成

【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm

导读

本文基于 kosmos-2/fairseq/examples/stories/README.md,完整讲解如何复现 Fan et al. (2018) 的 Hierarchical Neural Story Generation(分层神经故事生成)实验:从 WritingPrompts 数据集下载与 1000 词裁剪,到fairseq-preprocess二值化、fairseq-train训练卷积 seq2seq 模型与 fusion 模型,再到fairseq-generate采样生成完整故事。同时结合本仓库 fairseq 源码(fconv_self_att.py、downsampled_multihead_attention.py、linearized_convolution.py 等)逐层拆解模型结构与参数含义,读完你可以直接在本地复现论文实验,并理解卷积自注意力模型、融合门控与增量解码的底层原理。

一、任务背景:故事生成与 WritingPrompts 数据集

故事生成(Story Generation)要求模型在给定一句话"prompt"(故事开头)的条件下,续写连贯、有情节的完整故事。本指南复现的是 Hierarchical Neural Story Generation(Fan et al., 2018,ACL 2018)的工作:使用卷积 seq2seq 模型在 WritingPrompts 数据集上训练故事生成模型,并进一步训练融合(fusion)模型,将预训练语言模型的表示与训练的生成器融合,从而提升生成故事的连贯性与质量。

该示例在本仓库位于kosmos-2/fairseq/examples/stories/目录,配套的模型实现为 fairseq 内置的fconv_self_att系列架构(fconv_self_att.py),任务类型为标准的序列到序列翻译式任务(source → target,即 prompt → story),可通过fairseq-preprocessfairseq-trainfairseq-generate三个 CLI 完成从数据到模型再到生成的全流程。

二、预训练模型与示例输出

原文档提供了论文官方发布的预训练模型与测试集,对应信息如下表(下载地址见原文档,此处不再展开外部链接):

说明数据集模型测试集
Stories with Convolutional Model(Fan et al., 2018)WritingPrompts官方 checkpoint 包官方测试集包

官方还提供了两类模型的示例生成结果

  • 卷积 seq2seq 模型的示例故事(seq2seq_stories)
  • 融合(fusion)模型的示例故事(fusion_stories),以及对应的输入 prompt(fusion_prompts)

需要特别注意的是:官方示例文件中存在unk标记。这是因为论文采用小型完整词表建模(small full vocabulary),没有使用 BPE 分词或预训练词嵌入。原文档明确说明这些含unk的 prompt 未用于人工评估。

本仓库源码同样为这两个入口提供了便捷的 torch.hub 注册配置(见 fconv_self_att.py):

@classmethod def hub_models(cls): return { "conv.stories.pretrained": { "path": ".../stories_checkpoint.tar.gz", "checkpoint_file": "pretrained_checkpoint.pt", "tokenizer": "nltk", }, "conv.stories": { "path": ".../stories_checkpoint.tar.gz", "checkpoint_file": "fusion_checkpoint.pt", "tokenizer": "nltk", "pretrained": "True", "pretrained_checkpoint": "./pretrained_checkpoint.pt", }, # Test set containing dictionaries "data.stories": ".../stories_test.tar.bz2", }

从中可以确认两个事实:官方发布包内含pretrained_checkpoint.pt(预训练模型)与fusion_checkpoint.pt(融合模型)两个 checkpoint;融合模型的pretrained_checkpoint默认指向本地相对路径./pretrained_checkpoint.pt,这与后文生成时需要--model-overrides的原因直接相关。

三、数据集下载与 1000 词裁剪

3.1 下载与结构

原文档给出数据集下载命令(在本仓库内执行,注意先进入示例目录):

cd examples/stories curl https://dl.fbaipublicfiles.com/fairseq/data/writingPrompts.tar.gz | tar xvzf -

解压后得到 WritingPrompts 数据集的train / test / valid 三个划分,数据格式为并行文件对:

  • train.wp_source/train.wp_target
  • valid.wp_source/valid.wp_target
  • test.wp_source/test.wp_target

其中wp_source是故事 prompt(输入),wp_target是完整故事(输出)。该数据集来自 Reddit 的 r/WritingPrompts 社区,论文描述见 Fan et al., 2018。

3.2 为什么裁剪到前 1000 词

原文档明确说明:论文只对每篇故事的前 1000 个词建模(包括一个换行 token),而数据集发行版本是完整数据。因此需要在训练前把每个故事裁剪为前 1000 词。原文档给出的裁剪脚本如下:

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")

这段代码的逻辑:对 train/test/valid 三个划分分别读取*.wp_target文件,按空白切分后仅保留前 1000 个 token,再用单个空格重新拼接并写回原文件。注意只裁剪wp_target(故事侧)wp_source(prompt 侧)无需裁剪,因为 prompt 本身很短。

裁剪完成后,wp_sourcewp_target文件就可以交给 fairseq 做词表构建与二值化。

四、数据二值化: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>@register_model_architecture("fconv_self_att", "fconv_self_att_wp") def fconv_self_att_wp(args): args.encoder_embed_dim = getattr(args, "encoder_embed_dim", 256) args.encoder_layers = getattr( args, "encoder_layers", "[(128, 3)] * 2 + [(512,3)] * 1" ) args.decoder_embed_dim = getattr(args, "decoder_embed_dim", 256) args.decoder_layers = getattr( args, "decoder_layers", "[(512, 4)] * 4 + [(768, 4)] * 2 + [(1024, 4)] * 1" ) args.decoder_out_embed_dim = getattr(args, "decoder_out_embed_dim", 256) args.self_attention = getattr(args, "self_attention", "True") args.multihead_self_attention_nheads = getattr( args, "multihead_self_attention_nheads", 4 ) args.project_input = getattr(args, "project_input", "True") args.gated_attention = getattr(args, "gated_attention", "True") args.downsample = getattr(args, "downsample", "True") base_architecture(args)

由此可确认 WritingPrompts 特化架构的实际形态:

  • 编码器:embedding 维 256,卷积层为[(128, 3)] * 2 + [(512, 3)] * 1,即 3 层时间卷积,前三层隐藏维 128、最后层 512,卷积核大小均为 3;
  • 解码器:embedding 维 256,卷积层为[(512, 4)] * 4 + [(768, 4)] * 2 + [(1024, 4)] * 1,共 7 层,隐藏维依次从 512 增长到 1024,卷积核大小为 4;解码器输出投影维为 256(decoder_out_embed_dim);
  • 自注意力self_attention = True,4 头(multihead_self_attention_nheads = 4),且开启project_inputgated_attentiondownsample三个增强选项。

基础架构(base_architecture,fconv_self_att.py)的默认值为:dropout 0.1encoder_embed_dim 512encoder_layers [(512, 3)] * 3decoder_embed_dim 512decoder_layers [(512, 3)] * 8decoder_out_embed_dim 256decoder_attention Trueself_attention Falseencoder_attention False、多头注意力头数默认 1。可见fconv_self_att_wp是论文针对长文故事任务专门调参后的变体。

5.3 模型可配置参数全表

fconv_self_att.py 中add_args定义了该模型可覆盖的全部命令行参数,训练时均可通过--key value传入:

参数类型说明
--dropoutfloat各层 dropout 概率
--encoder-embed-dimint编码器 embedding 维数
--encoder-layersstr编码器卷积层配置,形如[(dim, kernel_size), ...]
--decoder-embed-dimint解码器 embedding 维数
--decoder-layersstr解码器卷积层配置
--decoder-out-embed-dimint解码器输出 embedding 维数
--decoder-attentionstr解码器 encoder 注意力层开关列表,如[True, ...]
--self-attentionstr解码器自注意力层开关,如[True] + [False]*5
--multihead-attention-nheadsintencoder 注意力头数
--multihead-self-attention-nheadsint自注意力头数
--encoder-attentionstr编码器注意力层开关
--encoder-attention-nheadsint编码器注意力头数
--project-inputstr自注意力是否先投影输入,如[True, ...]
--gated-attentionstr自注意力投影中是否使用 GLU 门控层
--downsamplestr自注意力是否使用下采样
--pretrained-checkpointstr预训练模型 checkpoint 路径
--pretrainedstr训练时是否加载预训练模型(fusion 模式开关)

注意其中多个布尔参数以字符串形式传入并在源码中通过eval()解析(见build_model,fconv_self_att.py),因此既可以传单个True/False,也可以传 Python 表达式列表(如[True] + [False]*5)实现逐层精细控制。

六、训练 Fusion(融合)模型

原文档指出:训练融合模型只需在基础训练命令上追加两个参数

# 在 5.1 的命令基础上追加: --pretrained True --pretrained-checkpoint path/to/checkpoint

融合模型的核心思想:先用 5.1 的命令训练好一个基础卷积 seq2seq 模型,然后把该模型的参数冻结,作为额外的"预训练编码器/解码器"接入新模型中,由新模型学习两组表示的融合方式。

源码层面的实现证据(fconv_self_att.py):

pretrained = eval(args.pretrained) if pretrained: logger.info("loading pretrained model") # 若 pretrained_checkpoint 不是绝对路径,尝试拼接 data 目录 trained_model = checkpoint_utils.load_model_ensemble( filenames=[args.pretrained_checkpoint], task=task, )[0][0] trained_decoder = list(trained_model.children())[1] trained_encoder = list(trained_model.children())[0] # freeze pretrained model for param in trained_decoder.parameters(): param.requires_grad = False for param in trained_encoder.parameters(): param.requires_grad = False

关键点:

  • 加载预训练 checkpoint 后,其 encoder 与 decoder 的全部参数被冻结requires_grad = False),训练时只更新新增的融合参数;
  • 融合模型中,新旧两个编码器被包装进CompositeEncoder(fconv_self_att.py),前向时两者并行计算、在解码器中合并;
  • 解码器侧新增了融合门控模块(fconv_self_att.py):两个独立的 Sigmoid 门gate1/gate2分别作用于新模型输出x与预训练模型输出pretrained_outputs["out"],再经一个由多层Linear + LayerNorm + GLU构成的joining网络合并后接输出层fc3
  • 预训练模型的隐藏状态通过注册在fc2上的 forward hook 捕获(self.pretrained_decoder.fc2.register_forward_hook(save_output())),因为预训练模型自带输出层而融合发生在隐藏状态层面。

融合门控的前向逻辑(fconv_self_att.py):

trained_x, _ = self.pretrained_decoder.forward(prev_output_tokens, trained_encoder_out) y = torch.cat([x, self.pretrained_outputs["out"]], dim=-1) gate1 = self.gate1(y) # Sigmoid 门,控制新模型贡献 gate2 = self.gate2(y) # Sigmoid 门,控制预训练模型贡献 gated_x1 = gate1 * x gated_x2 = gate2 * self.pretrained_outputs["out"] fusion = torch.cat([gated_x1, gated_x2], dim=-1) fusion = self.joining(fusion) fusion_output = self.fc3(fusion)

七、生成故事:fairseq-generate 与采样参数

7.1 生成命令与 model-overrides

原文档给出的生成命令(完整继承):

fairseq-generate>assert ( not cfg.generation.sampling or cfg.generation.nbest == cfg.generation.beam ), "--sampling requires --nbest to be equal to --beam"

启用--sampling时,--nbest必须等于--beam。原文档命令中--beam 1 --nbest 1正是满足此约束的标准组合。如果你改成束搜索(如--beam 5),则需同步将--nbest设为 5,并移除--sampling相关参数。

八、模型内部原理:卷积、自注意力与增量解码

8.1 卷积编码器:GLU 门控卷积 + 残差缩放

FConvEncoder(fconv_self_att.py)的流程:

  1. token embedding 与位置 embedding 相加后过 dropout;
  2. 线性层fc1投影到卷积输入维度;
  3. 逐层执行时间卷积:每层用ConvTBC输出out_channels * 2个通道,再经F.glu(x, dim=2)按通道维度做 GLU 门控(fconv_self_att.py);
  4. 若该层启用注意力则附加SelfAttention
  5. 残差连接后乘sqrt(0.5)保持方差稳定;
  6. 最后fc2投影回 embedding 维度,并用GradMultiply按注意力层数缩放梯度(fconv_self_att.py),输出 (x, y) 两路表示供解码器注意力使用。

编码器默认attention=False,因此故事任务的编码器是纯卷积编码器,负责把 prompt 编码为上下文表示。

8.2 解码器:LinearizedConvolution 增量解码

FConvDecoder(fconv_self_att.py)中的卷积使用LinearizedConvolution(linearized_convolution.py)。这是一个关键优化:

  • 训练时:退化为标准的ConvTBC(时间维度一维卷积),一次处理整个序列;
  • 推理时:利用incremental_state缓存输入缓冲区,把卷积重写为线性层F.linear),每次只接收新生成的 1 个 token(linearized_convolution.py),实现 O(1) 的逐 token 自回归生成;
  • 线性化权重在 checkpoint 中不持久化(state_dict中剔除_linearized_weight),避免冗余存储。

解码器的逐层流程为:卷积 + GLU → encoder 注意力(DownsampledMultiHeadAttention)→ 自注意力(SelfAttention)→ 残差缩放。其中自注意力的实现(fconv_self_att.py)把 Q/K/V 分别线性投影后送入注意力模块,并强制mask_future_timesteps=True,保证生成第 t 个 token 时只能看到前 t-1 个 token。

8.3 下采样多头注意力与门控投影

DownsampledMultiHeadAttention(downsampled_multihead_attention.py)是fconv_self_att_wpdownsample=Truegated_attention=True两个开关的底层实现:

  • Gating(GLU)GatedLinear用"Linear → GLU → Linear → GLU → Linear"的级联替代普通线性投影(downsampled_multihead_attention.py),为注意力注入更强的非线性;
  • DownsamplingDownsample模块"每隔 head_index+1 个元素取一个"(downsampled_multihead_attention.py),每个注意力头在不同步长上降采样,从而以不同"粒度"观察序列——这符合故事生成中同时需要局部与全局上下文的需求;
  • 自注意力中还会叠加scalar_biasuse_scalar_bias=True),为注意力权重引入可学习的标量偏置(fconv_self_att.py)。

8.4 损失函数:标签平滑交叉熵

训练命令使用--criterion label_smoothed_cross_entropy,对应实现位于 label_smoothed_cross_entropy.py。其核心公式(label_smoothed_cross_entropy.py):

loss = (1 - epsilon - eps_i) * nll_loss + eps_i * smooth_loss eps_i = epsilon / (vocab_size - 1)

其中nll_loss为标准负对数似然,smooth_loss = -sum(lprobs)是平滑项;label_smoothing=0epsilon=0,退化为标准交叉熵。该 criterion 还支持--report-accuracy(报告准确率指标)与--ignore-prefix-size(忽略前 N 个 token 的损失)等配置,可用于进一步实验。

九、复现路径小结与扩展建议

完整的复现链路可归纳为四步:

  1. 下载与裁剪cd examples/stories && curl ... | tar xvzf -,再用 1000 词裁剪脚本处理*.wp_target
  2. 二值化fairseq-preprocess生成data-bin/writingPrompts与词典(词频阈值 10);
  3. 训练fairseq-train -a fconv_self_att_wp训练基础卷积 seq2seq 模型;追加--pretrained True --pretrained-checkpoint <path>训练 fusion 模型;
  4. 生成fairseq-generate配合--beam 1 --sampling --sampling-topk 10 --temperature 0.8 --nbest 1采样,fusion 模型需补--model-overrides "{'pretrained_checkpoint': ...}"

如果你希望进一步实验,可以围绕以下方向扩展(均有源码支撑):

  • 调整fconv_self_att_wp中的卷积层配置(--encoder-layers/--decoder-layers)、自注意力头数(--multihead-self-attention-nheads)或关闭--downsample/--gated-attention观察对生成质量的影响;
  • 修改词频阈值(--thresholdsrc/--thresholdtgt)以控制词表大小与unk比例;
  • 使用--criterion label_smoothed_cross_entropy --label-smoothing 0.1开启标签平滑;
  • 采样阶段调整--temperature--sampling-topk在多样性与连贯性之间权衡。

十、引用

若在研究中引用该方法,原文档给出的 BibTeX 如下:

@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),仅供参考

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

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

立即咨询