news 2026/9/20 23:03:24

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
better_storylines 实战指南:用句子级语言模型在 ROC Stories 上复现 ACL 2020 故事续写实验
  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google 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 锁定的关键依赖为:

依赖版本
tensorflow>=2.1.0
tensorflow-datasets>=3.0
tensorflow-hub==0.8.0
gin-config==0.3.0
apache-beam>=2.20.0
absl-py==0.9.0
numpy>=1.18
scipy==1.4.1

需要特别注意的是:训练/评估用 TF2,而数据从零生成必须用 TF1(详见下节)。两种环境的依赖分别由requirements.train.txtrequirements.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 词片嵌入的掩码均值768
BERT_REDUCE_WEIGHTED_MEAN按词频倒数加权的均值(论文中 a=0.0001)768
BERT_REDUCE_MIN_MAXmin 与 max 拼接768×2
BERT_REDUCE_MIN_MAX_MEANmean/min/max 三者拼接768×3
BERT_CLASS_TOKEN取 BERT 的 CLS token 输出768
UNIVERSAL_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_dir='tfds_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.sh

build_tfds_dataset.sh 内部实际执行三步:

  1. 下载并解压 BERT 模型cased_L-12_H-768_A-12.zip
  2. 通过tensorflow_datasets.scripts.download_and_prepare从 ROC Stories 官方 Google Sheets 拉取原始 CSV(train2016/2017、valid2016/2018、test2016/2018);
  3. --module_import='src.rocstories_sentence_embeddings'注册自定义 Beam-based builder,为每条句子计算 BERT 嵌入并写出 TFDS 数据集。

从 rocstories_sentence_embeddings.py 源码可见,GenerateBERTEmbeddings在 TF2 环境下会直接抛出ValueError("Data 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(使用 CSLoss)
mlp_best_largescale_nocl大尺度重排任务最优 MLP(不使用 CSLoss)
mlp_best_story_cloze_clStory Cloze 任务最优 MLP(使用 CSLoss)
mlp_best_story_cloze_noclStory Cloze 任务最优 MLP(不使用 CSLoss)
resmlp_best_largescale_cl大尺度重排任务最优残差 MLP(使用 CSLoss)
resmlp_best_largescale_nocl大尺度重排任务最优残差 MLP(不使用 CSLoss)
resmlp_best_story_cloze_clStory Cloze 任务最优残差 MLP(使用 CSLoss)
resmlp_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_dir='tfds_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_dir=saved_checkpoints \ --data_dir='tfds_datasets/' \ --gin_config="configs/residual_best.gin" \ --gin_bindings="dataset.dataset_name = 'roc_stories_embeddings/bert_mean_emb'" \ --gin_bindings="train.learning_rate = 0.0001" \ --gin_bindings="ResidualModel.hparams.small_context_loss_weight = 1.0"

训练期间每个 epoch 都会做一次多任务评估,结果写入 TensorBoard summary;final_eval.tsv中保存最终各项指标;检查点按ep%04d_step%05d.ckpt格式定期保存(见 train.py)。

线性模型配置:configs/linear_best.gin

dataset.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.gin

dataset.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_b=True)得到对每个候选句的点积得分(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_acctrain_subset_accvalid_winter2018_accvalid_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-research

Google Research

项目地址:https://gitcode.com/gh_mirrors/go/google-research
点击查看免费下载

相关推荐

上一篇:Python程序执行可视化终极指南:使用Heartrate实时监控代码运行
下一篇:Spotify Web API错误处理与调试:开发者必须掌握的10个技巧

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/20 23:01:30

Python+飞书机器人:自动化消息推送实战指南

“发个通知而已&#xff0c;能有多麻烦&#xff1f;”说实话&#xff0c;我以前也这么想。直到有一天&#xff0c;我负责的一个数据同步任务在凌晨2点挂了&#xff0c;群里安安静静&#xff0c;第二天早上才被人发现。那一刻我才意识到&#xff1a;手动发通知这事&#xff0c;看…

作者头像 李华
网站建设 2026/9/20 23:01:14

淘宝评论情感分析:Python机器学习从分词到LIME

简介&#xff1a;面向毕业设计、课程大作业及机器学习入门实践&#xff0c;这是一份基于机器学习的电商淘宝商品评论情感分析项目源码与数据包。项目从Selenium模拟登录爬取淘宝评论入手&#xff0c;依次完成数据清理、jieba精确模式分词、词语索引与词向量构造&#xff0c;并对…

作者头像 李华
网站建设 2026/9/20 23:01:01

Hermes Agent 实战笔记:跟一条消息走完工具调用循环的每一轮

Hermes Agent 实战笔记&#xff1a;跟一条消息走完工具调用循环的每一轮 【免费下载链接】hermes-agent The agent that grows with you 项目地址: https://gitcode.com/GitHub_Trending/he/hermes-agent Hermes Agent 是一个可接多种模型后端的 AI 代理框架&#xff08…

作者头像 李华