- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
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.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_MEAN | BERT 词片嵌入的掩码均值 | 768 |
BERT_REDUCE_WEIGHTED_MEAN | 按词频倒数加权的均值(论文中 a=0.0001) | 768 |
BERT_REDUCE_MIN_MAX | min 与 max 拼接 | 768×2 |
BERT_REDUCE_MIN_MAX_MEAN | mean/min/max 三者拼接 | 768×3 |
BERT_CLASS_TOKEN | 取 BERT 的 CLS token 输出 | 768 |
UNIVERSAL_SENTENCE | Universal 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.shbuild_tfds_dataset.sh 内部实际执行三步:
- 下载并解压 BERT 模型
cased_L-12_H-768_A-12.zip; - 通过
tensorflow_datasets.scripts.download_and_prepare从 ROC Stories 官方 Google Sheets 拉取原始 CSV(train2016/2017、valid2016/2018、test2016/2018); - 以
--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_cl | Story Cloze 任务最优 MLP(使用 CSLoss) |
mlp_best_story_cloze_nocl | Story Cloze 任务最优 MLP(不使用 CSLoss) |
resmlp_best_largescale_cl | 大尺度重排任务最优残差 MLP(使用 CSLoss) |
resmlp_best_largescale_nocl | 大尺度重排任务最优残差 MLP(不使用 CSLoss) |
resmlp_best_story_cloze_cl | Story Cloze 任务最优残差 MLP(使用 CSLoss) |
resmlp_best_story_cloze_nocl | Story 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_dir | TFDS 数据集目录,如tfds_datasets/ |
--gin_config | Gin 配置文件路径(必填) |
--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 展示了数据集如何被划分为多份:
train:train[2%:],用于主训练并构建嵌入矩阵;valid_nolabel:train[:2%],用于"无标签"评估(在 2000 个 distractor 中选出正确续句);train_nolabel:train[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-research
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 示例,系统
人工智能深度学习机器学习
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考