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

发布时间:2026/9/14 22:22:47
基于 fairseq 的分层神经故事生成实战指南:WritingPrompts 数据预处理、卷积模型训练与采样生成 基于 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 GenerationFan et al., 2018ACL 2018的工作使用卷积 seq2seq 模型在 WritingPrompts 数据集上训练故事生成模型并进一步训练融合fusion模型将预训练语言模型的表示与训练的生成器融合从而提升生成故事的连贯性与质量。该示例在本仓库位于kosmos-2/fairseq/examples/stories/目录配套的模型实现为 fairseq 内置的fconv_self_att系列架构fconv_self_att.py任务类型为标准的序列到序列翻译式任务source → target即 prompt → story可通过fairseq-preprocess、fairseq-train、fairseq-generate三个 CLI 完成从数据到模型再到生成的全流程。二、预训练模型与示例输出原文档提供了论文官方发布的预训练模型与测试集对应信息如下表下载地址见原文档此处不再展开外部链接说明数据集模型测试集Stories with Convolutional ModelFan et al., 2018WritingPrompts官方 checkpoint 包官方测试集包官方还提供了两类模型的示例生成结果卷积 seq2seq 模型的示例故事seq2seq_stories融合fusion模型的示例故事fusion_stories以及对应的输入 promptfusion_prompts需要特别注意的是官方示例文件中存在unk标记。这是因为论文采用小型完整词表建模small full vocabulary没有使用 BPE 分词或预训练词嵌入。原文档明确说明这些含unk的 prompt 未用于人工评估。本仓库源码同样为这两个入口提供了便捷的 torch.hub 注册配置见 fconv_self_att.pyclassmethod 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_targetvalid.wp_source/valid.wp_targettest.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_sourceprompt 侧无需裁剪因为 prompt 本身很短。裁剪完成后wp_source与wp_target文件就可以交给 fairseq 做词表构建与二值化。四、数据二值化fairseq-preprocessfairseq 训练需要先把文本数据转换为索引化的二进制格式。原文档给出的二值化命令如下# Binarize the dataset: export TEXTexamples/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解码器输出投影维为 256decoder_out_embed_dim自注意力self_attention True4 头multihead_self_attention_nheads 4且开启project_input、gated_attention、downsample三个增强选项。基础架构base_architecturefconv_self_att.py的默认值为dropout 0.1、encoder_embed_dim 512、encoder_layers [(512, 3)] * 3、decoder_embed_dim 512、decoder_layers [(512, 3)] * 8、decoder_out_embed_dim 256、decoder_attention True、self_attention False、encoder_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_modelfconv_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.pypretrained eval(args.pretrained) if pretrained: logger.info(loading pretrained model) # 若 pretrained_checkpoint 不是绝对路径尝试拼接 data 目录 trained_model checkpoint_utils.load_model_ensemble( filenames[args.pretrained_checkpoint], tasktask, )[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训练时只更新新增的融合参数融合模型中新旧两个编码器被包装进CompositeEncoderfconv_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.pytrained_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 门控卷积 残差缩放FConvEncoderfconv_self_att.py的流程token embedding 与位置 embedding 相加后过 dropout线性层fc1投影到卷积输入维度逐层执行时间卷积每层用ConvTBC输出out_channels * 2个通道再经F.glu(x, dim2)按通道维度做 GLU 门控fconv_self_att.py若该层启用注意力则附加SelfAttention残差连接后乘sqrt(0.5)保持方差稳定最后fc2投影回 embedding 维度并用GradMultiply按注意力层数缩放梯度fconv_self_att.py输出 (x, y) 两路表示供解码器注意力使用。编码器默认attentionFalse因此故事任务的编码器是纯卷积编码器负责把 prompt 编码为上下文表示。8.2 解码器LinearizedConvolution 增量解码FConvDecoderfconv_self_att.py中的卷积使用LinearizedConvolutionlinearized_convolution.py。这是一个关键优化训练时退化为标准的ConvTBC时间维度一维卷积一次处理整个序列推理时利用incremental_state缓存输入缓冲区把卷积重写为线性层F.linear每次只接收新生成的 1 个 tokenlinearized_convolution.py实现 O(1) 的逐 token 自回归生成线性化权重在 checkpoint 中不持久化state_dict中剔除_linearized_weight避免冗余存储。解码器的逐层流程为卷积 GLU → encoder 注意力DownsampledMultiHeadAttention→ 自注意力SelfAttention→ 残差缩放。其中自注意力的实现fconv_self_att.py把 Q/K/V 分别线性投影后送入注意力模块并强制mask_future_timestepsTrue保证生成第 t 个 token 时只能看到前 t-1 个 token。8.3 下采样多头注意力与门控投影DownsampledMultiHeadAttentiondownsampled_multihead_attention.py是fconv_self_att_wp中downsampleTrue、gated_attentionTrue两个开关的底层实现GatingGLUGatedLinear用Linear → GLU → Linear → GLU → Linear的级联替代普通线性投影downsampled_multihead_attention.py为注意力注入更强的非线性DownsamplingDownsample模块每隔 head_index1 个元素取一个downsampled_multihead_attention.py每个注意力头在不同步长上降采样从而以不同粒度观察序列——这符合故事生成中同时需要局部与全局上下文的需求自注意力中还会叠加scalar_biasuse_scalar_biasTrue为注意力权重引入可学习的标量偏置fconv_self_att.py。8.4 损失函数标签平滑交叉熵训练命令使用--criterion label_smoothed_cross_entropy对应实现位于 label_smoothed_cross_entropy.py。其核心公式label_smoothed_cross_entropy.pyloss (1 - epsilon - eps_i) * nll_loss eps_i * smooth_loss eps_i epsilon / (vocab_size - 1)其中nll_loss为标准负对数似然smooth_loss -sum(lprobs)是平滑项label_smoothing0时epsilon0退化为标准交叉熵。该 criterion 还支持--report-accuracy报告准确率指标与--ignore-prefix-size忽略前 N 个 token 的损失等配置可用于进一步实验。九、复现路径小结与扩展建议完整的复现链路可归纳为四步下载与裁剪cd examples/stories curl ... | tar xvzf -再用 1000 词裁剪脚本处理*.wp_target二值化fairseq-preprocess生成data-bin/writingPrompts与词典词频阈值 10训练fairseq-train -a fconv_self_att_wp训练基础卷积 seq2seq 模型追加--pretrained True --pretrained-checkpoint path训练 fusion 模型生成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),仅供参考