better_storylines 实战指南:用句子级语言模型在 ROC Stories 上复现 ACL 2020 故事续写实验

发布时间:2026/9/20 23:03:30
better_storylines 实战指南:用句子级语言模型在 ROC Stories 上复现 ACL 2020 故事续写实验 人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载导读本指南围绕better_storylines/README.md展开系统讲解如何在 ROC Stories 数据集上复现 ACL 2020 论文Toward Better Storylines with Sentence-Level Language Models的完整实验流程。你将掌握基于 BERT 句子嵌入构建 TFDS 数据集的两种方式下载现成数据或从零生成、Story Cloze 与大规模重排large-scale reranking两类任务的评估方法、以及线性 MLP 与残差 MLP 两种模型的从零训练与 Gin 超参数配置。文中所有命令与脚本均取自仓库内真实文件可直接复制运行。项目概览句子级语言模型与故事续写任务better_storylines是 Google Research 中用于复现论文实验的代码库核心思路是不直接以 token 序列建模整个故事而是先用预训练编码器把每个句子映射为稠密嵌入向量再训练一个轻量 MLP 从前 4 句的嵌入预测第 5 句的嵌入。预测结果通过点积打分在候选句子集合中排序从而完成续写与评测。该仓库围绕两条主线组织训练与评估脚本scripts/ 目录下共 7 个 shell 脚本覆盖数据构建、训练、三类评估。核心实现src/ 目录下的 8 个 Python 文件其中 models.py 定义模型结构rocstories_sentence_embeddings.py 定义 TFDS 数据集构建器train.py 与 utils.py 实现训练与评估循环。README 中说明该代码可复现论文 Table 1、Table 2 以及 Figure 1 中的精度数字并提供了用于复现的预训练检查点下载地址。环境搭建Python 3 TensorFlow 2训练与评估代码基于 Python 3 与 TensorFlow 2。README 推荐在虚拟环境中安装依赖python3 -m venv pyenv_tf2 source pyenv_tf2/bin/activate pip install --upgrade pip pip3 install -R requirements.train.txt其中 requirements.train.txt 锁定的关键依赖为依赖版本tensorflow2.1.0tensorflow-datasets3.0tensorflow-hub0.8.0gin-config0.3.0apache-beam2.20.0absl-py0.9.0numpy1.18scipy1.4.1需要特别注意的是训练/评估用 TF2而数据从零生成必须用 TF1详见下节。两种环境的依赖分别由requirements.train.txt与requirements.datagen.txt管理不能混用。数据集ROC Stories 的句子级 BERT 嵌入数据形态数据集的每个样本是一条故事其中训练集完整 5 句话的故事含正确第 5 句验证集/测试集前 4 句 2 个候选第 5 句Story Cloze 格式标签指示哪个候选是正确的每条句子被替换为其BERT 平均词片wordpiece嵌入768 维并打包为一个 TFDS 数据集。从 rocstories_sentence_embeddings.py 的_DESCRIPTION可以看到该数据集正是为 Story Cloze 任务设计的训练集故事均为 5 句验证/测试集为前 4 句加 2 个候选结尾。从源码看数据集还支持多种嵌入变体由 EmbeddingType 枚举定义EmbeddingType 枚举值说明输出维度BERT_REDUCE_MEANBERT 词片嵌入的掩码均值768BERT_REDUCE_WEIGHTED_MEAN按词频倒数加权的均值论文中 a0.0001768BERT_REDUCE_MIN_MAXmin 与 max 拼接768×2BERT_REDUCE_MIN_MAX_MEANmean/min/max 三者拼接768×3BERT_CLASS_TOKEN取 BERT 的 CLS token 输出768UNIVERSAL_SENTENCEUniversal Sentence Encoder large 版本嵌入512论文与默认配置两个.gin文件使用的均是bert_mean_emb即 BERT 掩码均值嵌入。BERT 模型固定为cased_L-12_H-768_A-12见源码常量。方式一直接下载预计算数据集由于对约 40 万条句子逐句跑 BERT 嵌入耗时较长README 提供了论文使用的均值嵌入数据集下载方式在仓库根目录执行wget https://storage.googleapis.com/gresearch/better_storylines/roc_stories_embeddings.zip mkdir tfds_datasets unzip roc_stories_embeddings.zip -d tfds_datasets rm roc_stories_embeddings.zip解压后tfds_datasets/即为 TFDS 数据目录后续训练与评估脚本默认通过--data_dirtfds_datasets/引用。方式二从零生成数据集TF1 环境若需自行生成例如复现加权嵌入变体README 给出了完整流程注意必须切换到 TF1 虚拟环境python3 -m venv pyenv_tf1 source pyenv_tf1/bin/activate pip install --upgrade pip pip install -R requirements.datagen.txt # 仅当需要生成 frequency-weighted 嵌入时才需要以下一行 wget https://storage.googleapis.com/gresearch/better_storylines/vocab_frequencies sh scripts/build_tfds_dataset.shbuild_tfds_dataset.sh 内部实际执行三步下载并解压 BERT 模型cased_L-12_H-768_A-12.zip通过tensorflow_datasets.scripts.download_and_prepare从 ROC Stories 官方 Google Sheets 拉取原始 CSVtrain2016/2017、valid2016/2018、test2016/2018以--module_importsrc.rocstories_sentence_embeddings注册自定义 Beam-based builder为每条句子计算 BERT 嵌入并写出 TFDS 数据集。从 rocstories_sentence_embeddings.py 源码可见GenerateBERTEmbeddings在 TF2 环境下会直接抛出ValueErrorData generation can only be performed with TF1.这正是 README 要求单独建 TF1 环境的原因。整个生成流程基于 Apache Beam 的BeamBasedBuilder逐条处理句子并借助掩码聚合得到句级嵌入masked_mean 等函数。README 同时提醒本地运行无 Apache Beam 集群时耗时可能很长。预训练检查点一览为直接复现论文精度仓库提供 8 个可下载检查点。它们按两个维度划分模型架构MLP vs 残差 MLP与任务/损失大尺度重排任务 vs Story Cloze 任务是否使用 CSLoss。检查点名称说明mlp_best_largescale_cl大尺度重排任务最优 MLP使用 CSLossmlp_best_largescale_nocl大尺度重排任务最优 MLP不使用 CSLossmlp_best_story_cloze_clStory Cloze 任务最优 MLP使用 CSLossmlp_best_story_cloze_noclStory Cloze 任务最优 MLP不使用 CSLossresmlp_best_largescale_cl大尺度重排任务最优残差 MLP使用 CSLossresmlp_best_largescale_nocl大尺度重排任务最优残差 MLP不使用 CSLossresmlp_best_story_cloze_clStory Cloze 任务最优残差 MLP使用 CSLossresmlp_best_story_cloze_noclStory Cloze 任务最优残差 MLP不使用 CSLoss下载地址均为https://storage.googleapis.com/gresearch/better_storylines/检查点名.zip。以 README 中的 Story Cloze 评估示例为准wget https://storage.googleapis.com/gresearch/better_storylines/mlp_best_largescale_cl.zip mkdir trained_models unzip mlp_best_largescale_cl.zip -d trained_models rm mlp_best_largescale_cl.zip评估三个脚本覆盖三类任务评估体系围绕all_metrics.csv展开必须先运行评估所有检查点的脚本生成该文件后续脚本依赖它挑选最优检查点。1. 评估全部检查点基础步骤sh scripts/evaluate_all_checkpoints.sh trained_models/mlp_best_largescale_cl该脚本evaluate_all_checkpoints.sh调用src/evaluate_full.py以--base_dir指向检查点目录、--data_dirtfds_datasets指向数据目录对目录下每个检查点在验证集上计算精度结果写入all_metrics.csv。从 utils.py 的pick_best_checkpoint可以看到选择逻辑读取eval/all_metrics.csv默认以valid_spring2016_acc列为排序指标逐行比较找出最高精度对应的检查点再通过 glob 匹配*ckpt*index还原出完整检查点路径。2. Story Cloze 2016 测试集评估sh scripts/evaluate_best_story_cloze_test.sh trained_models/mlp_best_largescale_cl该脚本输出指定目录中最优检查点在Story Cloze 2016 测试集上的精度。README 特别说明2018 测试集只能通过提交 CodaLab 排行榜来评估代码内无法直接运行。底层由 evaluate_story_cloze_test.py 实现。3. 大规模重排任务评估sh scripts/evaluate_ranking_task.sh trained_models/mlp_best_largescale_cl输出最优检查点在大尺度重排任务上的accuracy 与 MRR两项指标由 evaluate_ranking_task.py 实现。4. 大规模重排的定性评估sh scripts/evaluate_ranking_qualitative.sh path/to/rocstories/csvs trained_models/mlp_best_largescale_cl此脚本输出大尺度重排任务中得分最高的候选下一句用于人工检查续写质量。前提是先向 ROC Stories 官网申请验证集与训练集的 CSV 文件并将目录路径作为第一个参数传入。对应实现为 evaluate_qualitative.py。从零训练两种模型与 Gin 配置启动训练README 给出的训练入口即残差模型脚本sh scripts/train_residual.sh若想训练线性 MLP 模型仓库另提供 train_linear.sh。两个脚本均调用src/train.py核心命令行参数如下来自 train.py 的 flags 定义参数说明--save_dir模型保存目录必填--data_dirTFDS 数据集目录如tfds_datasets/--gin_configGin 配置文件路径必填--gin_bindings额外的 Gin 参数绑定可多次传入以 train_residual.sh 为例实际执行的命令为python src/train.py \ --save_dirsaved_checkpoints \ --data_dirtfds_datasets/ \ --gin_configconfigs/residual_best.gin \ --gin_bindingsdataset.dataset_name roc_stories_embeddings/bert_mean_emb \ --gin_bindingstrain.learning_rate 0.0001 \ --gin_bindingsResidualModel.hparams.small_context_loss_weight 1.0训练期间每个 epoch 都会做一次多任务评估结果写入 TensorBoard summaryfinal_eval.tsv中保存最终各项指标检查点按ep%04d_step%05d.ckpt格式定期保存见 train.py。线性模型配置configs/linear_best.gindataset.dataset_name roc_stories_embeddings/bert_mean_emb dataset.shuffle_input_sentences False LinearModel.hparams.dropout_amount 0.5 LinearModel.hparams.relu_layers [1024, 1024, 1024] LinearModel.hparams.small_context_loss_weight 1.0 LinearModel.hparams.normalize_embeddings True LinearModel.hparams.final_dropout False train.learning_rate 0.0001 train.num_epochs 400 build_model.network_class LinearModel残差模型配置configs/residual_best.gindataset.dataset_name roc_stories_embeddings/bert_mean_emb dataset.shuffle_input_sentences False ResidualModel.hparams.dropout_amount 0.5 ResidualModel.hparams.num_residual_layers 1 ResidualModel.hparams.residual_layer_size 1024 train.learning_rate 0.0001 train.num_epochs 50 train.save_every_n_epochs 1 build_model.network_class ResidualModel关键超参数含义对应 models.py 源码relu_layers线性层 ReLU 层的维度列表。输入[batch, 4 句 × 768 维]被展平为[batch, 4*768]后依次经过这些全连接层LinearModel._build_network。linear_best.gin使用[1024, 1024, 1024]。dropout_amount每个隐藏层后的 dropout 比例默认为 0.5。normalize_embeddings是否对输入与预测嵌入做 LayerNormalization标准化为均值 0、单位方差。final_dropout是否在最后的嵌入层之后追加 dropout。线性最优配置关闭了它。small_context_loss_weight论文的核心创新点 CSLoss 的权重。大于 0 时在包含大量负样本distractor的主损失之外额外计算一个仅以 4 句上下文作为负样本的小上下文损失。其计算细节见 utils.py 的 train_step将预测嵌入与 4 个上下文句嵌入做点积再与真实第 5 句得分拼接构成 5 类分类损失。训练脚本中通过--gin_bindings显式设为 1.0。max_num_distractors大于等于 0 时随机截取真实第 5 句附近的一个 distractor 窗口参与损失计算用于控制训练时负样本数量models.py compute_loss。residual_layer_size/num_residual_layers残差模型专属残差块内部隐藏维度与残差块个数。每个块为Dense(ReLU) → Dropout → Dense(ReLU)后与输入相加ResidualModel._build_network。预测机制嵌入矩阵作为输出层两个模型的输出层并不是传统 softmax 分类头而是乘以上下文无关的嵌入矩阵由训练集所有第 5 句嵌入拼接而成。训练时标签是该故事第 5 句在嵌入矩阵中的行号见 build_train_style_dataset预测时用tf.matmul(embedding, embedding_matrix, transpose_bTrue)得到对每个候选句的点积得分models.py call。验证时同理只是候选从全量嵌入矩阵换成 Story Cloze 的两个候选结尾utils.py eval_step。训练数据划分细节prepare_datasets 展示了数据集如何被划分为多份traintrain[2%:]用于主训练并构建嵌入矩阵valid_nolabeltrain[:2%]用于无标签评估在 2000 个 distractor 中选出正确续句train_nolabeltrain[2%:4%]同类的训练子集评估valid2018/valid2016官方 Story Cloze 验证集每个样本只有 2 个候选。每轮训练结束后do_evaluation 会依次运行四组评估并记录valid_nolabel_acc、train_subset_acc、valid_winter2018_acc、valid_spring2016_acc四个指标。引用论文若在研究中使用了该代码库README 给出的引用格式为inproceedings{ippolito2020toward, title{Toward Better Storylines with Sentence-Level Language Models}, author{Ippolito, Daphne and Grangier, David and Eck, Douglas and Callison-Burch, Chris}, booktitle{Proceedings of the 58th Annual Meeting of the Association for Computational Linguistics}, year{2020} }小结better_storylines提供了一个端到端、可复现的句子级故事续写实验方案从 BERT 句嵌入数据集下载或自建出发用 Gin 配置驱动线性/残差 MLP 训练再通过统一的all_metrics.csv机制完成 Story Cloze、大尺度重排与定性评估。建议按以下顺序实践先搭建 TF2 环境并下载预计算嵌入与检查点 → 运行evaluate_all_checkpoints.sh生成基线 → 分别跑 Story Cloze 与重排任务评估 → 最后用train_linear.sh/train_residual.sh从零训练并对比 Gin 超参数的效果。赞分享人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载相关推荐Roc语言实战指南Roc语言实战指南 项目介绍 Roc 是一个新兴的编程语言旨在提供一种高效、安全且表达力强的开发体验尤其专注于并发和系统级编程。它采用现代编译器技术确保高在 Langfuse 前端编写高质量 Storybook Stories组件故事编写规范与 CSF Next 实践指南在 Langfuse 前端编写高质量 Storybook Stories组件故事编写规范与 CSF Next 实践指南 本文是 Langfuse 开源仓库A人工智能LLMOps可观测性AI 评测LLM 网关后端前端用 Flax 在 LM1B 上训练 Transformer 语言模型完整实战指南用 Flax 在 LM1B 上训练 Transformer 语言模型完整实战指南 导读 本指南基于 Flax 仓库中的 examples/lm1b 示例系统人工智能深度学习机器学习上一篇Python程序执行可视化终极指南使用Heartrate实时监控代码运行下一篇Spotify Web API错误处理与调试开发者必须掌握的10个技巧创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

关于本文作者

来自尧图内容编辑团队

尧图内容编辑团队 内容团队

尧图内容编辑团队

本文由尧图网络内容编辑团队执笔。团队由资深项目经理、前端工程师与设计师组成,所有内容均来自亲手交付的真实项目,先讲清问题、再给出可落地的解法。尧图深耕北京网站建设十年,服务过京华建材集团、智造科技等各行业客户,把一线经验沉淀为可复用的行业观察。

  • 十年建站经验,覆盖建材、制造、服务、文创等
  • 项目经理把关选题与事实准确性
  • 工程师与设计师联合撰写专业细节
  • 统一编辑规范,保证文风与排版一致
  • 每月复盘转化数据,迭代选题方向

延伸阅读

相关资讯与近期热门内容

深度阅读推荐

建站决策前值得细读的三篇

网站改版的5个关键决策
2024-08-12

网站改版的5个关键决策

什么时候该改版、改到什么程度、如何避免流量掉光,京华建材集团改版复盘给出答案。

获取专属建站方案

看完文章,把您的行业与预算告诉我们,免费获取一份量身定制的官网建设方案与报价。

立即免费咨询