ReadTwice:面向超长文档阅读的“读两遍“BERT 模型实战指南
2026/9/21 15:29:34 网站建设 项目流程
  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

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

导读

ReadTwice 是 Google Research 提出的"带记忆的超长文档阅读"(Reading Very Large Documents with Memories)模型实现,核心思想是让模型对文档进行两次阅读:第一次读取生成全局摘要(记忆),第二次阅读时携带这些摘要信息重新处理原文,从而突破 Transformer 输入长度限制、更好地处理跨段长距离依赖。本文以readtwice/README.md为主干,结合仓库源码,完整讲解该模型的架构原理、依赖安装、预训练权重加载,以及 HotpotQA、TriviaQA、NarrativeQA 三个阅读理解任务从数据预处理、TFRecord 生成到 TPU 微调的全流程实操,并附上配置项与关键源码的对照解析。

模型概述与仓库结构

ReadTwice 仓库对应论文ReadTwice: Reading Very Large Documents with Memories,部分函数最初为 ETC(Enhanced Transformer Capture,同样位于本仓库etcmodel目录)项目实现,在此基础上演进出"二次阅读 + 全局记忆"的机制。

仓库核心目录结构如下:

  • readtwice/data_utils:数据工具,包含 Beam 流水线工具、TF 示例构造(data_utils.py)、SentencePiece 分词(tokenization.py)及对应测试;
  • readtwice/layers:模型层实现,如侧输入注意力 attention.py、Transformer 主体 transformer.py、嵌入层、重计算梯度、TPU 工具(tpu_utils.py)等;
  • readtwice/models:模型主体(modeling.py、config.py)、优化器、损失函数,以及 HotpotQA / TriviaQA / NarrativeQA 三个下游任务的预处理与微调脚本;
  • readtwice/run.sh:一键脚本,创建虚拟环境、安装依赖并运行全部单元测试。

依赖安装与环境要求

安装命令

README 推荐的获取与安装方式如下:

svn export https://github.com/google-research/google-research/trunk/readtwice pip install -r readtwice/requirements.txt

仓库为只读镜像,实际使用中也可直接以本仓库readtwice目录为工作目录,其依赖清单见 requirements.txt。

依赖清单(readtwice/requirements.txt)主要包括:tensorflow>=1.15.0(使用 TF1 兼容 API)、absl-pyapache-beam(预处理阶段用于并行生成 TFRecord)、nltk(句子切分与答案标注)、numpysentencepiece(RoBERTa 同款分词)、intervaltreesortedcontainers以及bert-tensorflow

重要硬件前提

WARNING:当前代码依赖部分自定义算子(特别是cross_replica_concat),要求TPU(或 CPU)环境。从源码看,该算子位于 readtwice/layers/tpu_utils.py,其实现通过xla.replica_id()获取当前 TPU 核编号,用tf.scatter_nd把本地张量散落到"复制副本"维度,再经tpu.cross_replica_sum跨核求和、tf.reshape展平,从而把分布在num_replicas个 TPU 核上的张量拼接成首维扩大num_replicas倍的完整张量(详见 losses.py 中cross_batch_softmax等跨批操作的使用)。README 也提示:可以通过调整代码使其适配 GPU,但需要自行改造。

此外还需安装:

pip install cloud-tpu-client

运行单元测试

单元测试通过 unittest 发现机制运行,前提是当前工作目录必须为readtwice文件夹的父目录(即仓库根目录),否则模块导入会失败:

python -m unittest discover -s readtwice -p '*_test.py'

该命令会扫描readtwice下所有*_test.py,覆盖数据工具、各网络层、模型、损失、优化器与三个任务的评估代码。readtwice/run.sh则等价地封装了"虚拟环境 + 安装依赖 + 跑测试"的完整流程。

预训练模型与分词器

官方发布了一个预训练 checkpoint(readtwice.tar.gz,存放于公共存储桶),可用于复现论文实验。模型词汇表与 RoBERTa 完全一致,SentencePiece 分词器由本仓库的 bertseq2seq 项目提供(vocab_gpt.model)。

微调前需要准备:

export PRETRAINED_MODEL_DIR=gs://path/to/directory/with/pretrained/model export CONFIG_PATH=${PRETRAINED_MODEL_DIR}/read_it_twice_bert_config.json export PRETRAINED_MODEL_CHECKPOINIT=${PRETRAINED_MODEL_DIR}/model.ckpt-1000000 export SPM_MODEL_PATH=/path/to/vocab_gpt.model export NLTK_DATA_PATH=/tmp/nltk_dir

其中read_it_twice_bert_config.json是模型配置文件,model.ckpt-1000000是预训练 checkpoint(README 原文如此拼写CHECKPOINIT,实际为 checkpoint 之意),NLTK_DATA_PATH用于 nltk 句子切分所需数据的存放位置。

配置文件的完整字段

配置类ReadItTwiceBertConfig定义在 readtwice/models/config.py,由from_json_file从 JSON 加载、to_json_string序列化。完整字段及含义如下:

字段默认值说明
vocab_size(必填)token_ids的词汇表大小
use_sparse_memory_attention(必填)是否允许非实体 token 关注基于实体的摘要
max_seq_length512token_ids的最大长度
max_num_blocks_per_document256单个文档内最大块数,即block_pos的最大值
cross_attention_pos_emb_modeNone是否基于block_pos添加位置嵌入
embedding_sizeNonetoken 嵌入维度;None时等于hidden_size(同原始 BERT),可设小值(如 128)类似 ALBERT
hidden_size768编码器与池化层维度
num_hidden_layers12Transformer 编码器层数
num_attention_heads12注意力头数
intermediate_size3072FFN 中间层维度
hidden_act"gelu"激活函数
share_kv_projectionsFalse主-主与主-侧注意力是否共享 K/V 投影(每层由 2 组 K/V 变为 1 组)
hidden_dropout_prob0.1全连接层 dropout
attention_probs_dropout_prob0.1注意力概率 dropout
initializer_range0.02权重初始化截断正态标准差
grad_checkpointing_period0激活重计算间隔;0表示全部保存,大于 0 时以重计算换显存
second_read_type"from_scratch"第二次读取方式,见下文"两次读取"详解
second_read_num_new_layersNonesecond_read_type=new_layers时新增的 Transformer 层数
second_read_num_cross_attention_headsNone第二次读取跨注意力的头数
second_read_enable_default_side_inputFalse是否加入默认侧输入(类似 no-op 注意力,允许注意力权重之和小于 1)
summary_mode"cls"摘要提取方式(如clstext_blockentity等)
summary_postprocessing_type"none"摘要后处理:none/linear/transformer
summary_postprocessing_num_layersNone摘要后处理 Transformer 层数
cross_attention_top_kNone计算摘要注意力前是否做 Top-K 截断(仅支持cross_attend_once
text_block_extract_every_xNone文本块摘要抽取间隔

模型配置加载逻辑get_model_config(config.py)支持三种来源:优先读取model_dir下的read_it_twice_bert_config.json;若不存在,可从source_file或 Base64 编码的source_base64读取,并(默认)把源配置写入模型目录以便后续复用。

核心机制:两次读取(Read-It-Twice)架构

从 readtwice/models/modeling.py 的实现看,ReadItTwiceBertModel的处理流程为:

  1. 第一次读取transformer_with_side_inputsTransformerWithSideInputLayers)对块内 token 做标准自注意力,同时通过FusedSideAttention(layers/attention.py)把侧输入(第一次读取时可无)作为额外的 K/V 参与注意力。该层不使用相对注意力,是标准 Transformer 层的直接推广,因此可以方便地从预训练 BERT/RoBERTa 直接迁移权重
  2. 摘要提取SummaryExtraction从第一次读取的隐藏状态中提取块级摘要(Summary结构体含statesprocessed_statesblock_idsblock_poslabels),并通过get_cross_block_att(modeling.py)依据文档 ID 计算块间注意力掩码——cross_block_attention_mode决定了不同块摘要之间的交互范围(block/doc/batch/other_blocks)。必要时借助cross_replica_concat在 TPU 多核间汇聚全局摘要,实现"全局记忆"。
  3. 第二次读取:依据second_read_type决定如何利用摘要:
    • from_scratch(默认):把第一次读取的输出从零重新处理,将全局摘要的processed_states作为侧输入,att_mask_with_side_input同时包含 token-token 掩码与 token-摘要映射掩码(modeling.py);
    • new_layers/new_layers_cross_attention:在第一次读取结果之上堆叠second_read_num_new_layers个新 Transformer 层;
    • cross_attend_once:先用SideAttention残差块让 token 一次性关注全局摘要(支持cross_attention_top_kTop-K 截断与块位置嵌入),再送入新层(modeling.py)。

这种"两遍 + 记忆"设计使模型无需把整个长文档硬塞进单一上下文窗口,而是以摘要为媒介实现跨块、跨文档的信息传递,这正是应对超长文档阅读的核心思路。模型测试 modeling_test.py 中的test_modeltest_get_cross_block_att等用例覆盖了不同second_read_typecross_block_attention_mode与摘要后处理的组合,可作为理解各配置行为的参考。

下游任务微调实战

三个任务的通用流程均为:① 下载数据并设置路径 → ② 用preprocess脚本生成 TFRecord 示例并拷贝到 GCS → ③ 用run_finetuning在 TPU 上微调(先训练后评估)

HotpotQA(多跳抽取式问答)

下载hotpot_train_v1.1.jsonhotpot_dev_distractor_v1.json${HOTPOTQA_DATA_DIR}后:

export HOTPOTQA_DATA_DIR=/path/to/HotpotQA export HOTPOTQA_EXAMPLE_DIR=${HOTPOTQA_DATA_DIR}/examples export HOTPOTQA_EXAMPLE_GCP_BUCKET=gs://path/to/gcp/bucket export HOTPOTQA_OUTPUT_FOLDER=gs://path/to/HotpotQA/output/folder mkdir -p ${HOTPOTQA_EXAMPLE_DIR}

生成 TFRecord(验证集不生成答案标注,训练集加--generate_answers):

python -m readtwice.models.hotpot_qa.preprocess \ --spm_model_path=${SPM_MODEL_PATH} \ --input_file=${HOTPOTQA_DATA_DIR}/hotpot_dev_distractor_v1.json \ --output_prefix=${HOTPOTQA_EXAMPLE_DIR}/valid \ --nltk_data_path=${NLTK_DATA_PATH} python -m readtwice.models.hotpot_qa.preprocess \ --spm_model_path=${SPM_MODEL_PATH} \ --input_file=${HOTPOTQA_DATA_DIR}/hotpot_train_v1.1.json \ --output_prefix=${HOTPOTQA_EXAMPLE_DIR}/train \ --generate_answers \ --nltk_data_path=${NLTK_DATA_PATH} gcloud storage cp ${HOTPOTQA_EXAMPLE_DIR}/* ${HOTPOTQA_EXAMPLE_GCP_BUCKET} gcloud storage cp ${HOTPOTQA_DATA_DIR}/hotpot_dev_distractor_v1.json ${HOTPOTQA_EXAMPLE_GCP_BUCKET}

TPU 微调(训练阶段):

python -m readtwice.models.hotpot_qa.run_finetuning \ --read_it_twice_bert_config_file=${CONFIG_PATH} \ --input_file=${HOTPOTQA_EXAMPLE_GCP_BUCKET}/train.tfrecord-* \ --output_dir=${HOTPOTQA_OUTPUT_FOLDER} \ --init_checkpoint=${PRETRAINED_MODEL_CHECKPOINIT} \ --enable_side_inputs \ --cross_block_attention_mode=doc \ --do_train \ --nodo_eval \ --optimizer=adamw \ --learning_rate=3e-05 \ --num_train_epochs=6 \ --warmup_proportion=0.1 \ --learning_rate_schedule=inverse_sqrt \ --poly_power=1 \ --start_warmup_step=0 \ --save_checkpoints_steps=5000 \ --iterations_per_loop=1000 \ --nouse_one_hot_embeddings \ --use_tpu \ --tpu_job_name=??? \ --num_tpu_cores=16 \ --num_tpu_tasks=1 \ --decode_top_k=40 \ --decode_max_size=10 \ --tpu_name=??? \ --cross_attention_top_k=100

评估阶段仅把--do_train改为--nodo_train --do_eval--nodo_eval改为--do_eval,其余参数不变(注意 README 中nodo_eval/nodo_train为原文写法,其语义即关闭对应开关)。

WARNING:对输出结果的正式评估还需要论文附录中的额外步骤(HotpotQA 的完整打分包含 yes/no、支持事实等,仓库 hotpot_qa/evaluation.py 与 hotpot_qa/losses.py 分别实现了评估指标与含cross_replica_concat的跨核损失)。

TriviaQA(开放域长文档问答)

下载 TriviaQA 官方数据(wikipedia/web 证据 + QA json)后:

export TRIVIAQA_DATA_DIR=/path/to/TriviaQA export TRIVIAQA_EXAMPLE_DIR=${TRIVIAQA_DATA_DIR}/examples export TRIVIAQA_EXAMPLE_GCP_BUCKET=gs://path/to/gcp/bucket export TRIVIAQA_OUTPUT_FOLDER=gs://path/to/TriviaQA/output/folder mkdir -p ${TRIVIAQA_EXAMPLE_DIR}

生成 TFRecord 时需额外指定证据语料目录:

python -m readtwice.models.trivia_qa.preprocess \ --spm_model_path=${SPM_MODEL_PATH} \ --input_file=${TRIVIAQA_DATA_DIR}/qa/wikipedia-dev.json \ --wikipedia_dir=${TRIVIAQA_DATA_DIR}/evidence/wikipedia \ --web_dir=${TRIVIAQA_DATA_DIR}/evidence/web \ --output_prefix=${TRIVIAQA_EXAMPLE_DIR}/valid \ --nltk_data_path=${NLTK_DATA_PATH} python -m readtwice.models.trivia_qa.preprocess \ --spm_model_path=${SPM_MODEL_PATH} \ --input_file=${TRIVIAQA_DATA_DIR}/qa/wikipedia-train.json \ --wikipedia_dir=${TRIVIAQA_DATA_DIR}/evidence/wikipedia \ --web_dir=${TRIVIAQA_DATA_DIR}/evidence/web \ --output_prefix=${TRIVIAQA_EXAMPLE_DIR}/train \ --generate_answers \ --nltk_data_path=${NLTK_DATA_PATH} gcloud storage cp ${TRIVIAQA_EXAMPLE_DIR}/* ${TRIVIAQA_EXAMPLE_GCP_BUCKET} gcloud storage cp ${TRIVIAQA_DATA_DIR}/qa/wikipedia-dev.json ${TRIVIAQA_EXAMPLE_GCP_BUCKET}

微调命令与 HotpotQA 的差异点:学习率更低(--learning_rate=1e-05)、学习率调度为--learning_rate_schedule=poly_decay--save_checkpoints_steps=3000--iterations_per_loop=200--decode_top_k=8 --decode_max_size=20,并新增--eval_json_path--eval_data_split=valid

python -m readtwice.models.trivia_qa.run_finetuning \ --read_it_twice_bert_config_file=${CONFIG_PATH} \ --input_file=${TRIVIAQA_EXAMPLE_GCP_BUCKET}/train.tfrecord-* \ --eval_json_path=${TRIVIAQA_EXAMPLE_GCP_BUCKET}/wikipedia-dev.json \ --output_dir=${TRIVIAQA_OUTPUT_FOLDER} \ --init_checkpoint=${PRETRAINED_MODEL_CHECKPOINIT} \ --enable_side_inputs \ --cross_block_attention_mode=doc \ --do_train \ --nodo_eval \ --optimizer=adamw \ --learning_rate=1e-05 \ --num_train_epochs=6 \ --warmup_proportion=0.1 \ --learning_rate_schedule=poly_decay \ --poly_power=1 \ --start_warmup_step=0 \ --save_checkpoints_steps=3000 \ --iterations_per_loop=200 \ --nouse_one_hot_embeddings \ --use_tpu \ --tpu_job_name=??? \ --num_tpu_cores=16 \ --num_tpu_tasks=1 \ --decode_top_k=8 \ --decode_max_size=20 \ --eval_data_split=valid \ --spm_model_path=${SPM_MODEL_PATH} \ --tpu_name=??? \ --cross_attention_top_k=100

评估阶段同样切换为--nodo_train --do_eval。TriviaQA 的评估实现(evaluate_triviaqa)位于 trivia_qa/evaluation.py,包含答案归一化(去冠词、标点、下划线、统一大小写与空白)以及基于 ground-truth 集合的 EM/F1 计算。

NarrativeQA(整本故事书阅读)

从 NarrativeQA 官网下载后,NarrativeQA 特殊之处在于复用trivia_qa.run_finetuning入口,且预处理需要 qaps(问答对)与 documents(故事文本)两份 CSV:

export NARRATIVEQA_DATA_DIR=/path/to/NarrativeQA export NARRATIVEQA_EXAMPLE_DIR=${NARRATIVEQA_DATA_DIR}/examples export NARRATIVEQA_EXAMPLE_GCP_BUCKET=gs://path/to/gcp/bucket export NARRATIVEQA_OUTPUT_FOLDER=gs://path/to/NarrativeQA/output/folder mkdir -p ${NARRATIVEQA_EXAMPLE_DIR}
python -m readtwice.models.narrative_qa.preprocess \ --spm_model_path=${SPM_MODEL_PATH} \ --input_qaps=${NARRATIVEQA_DATA_DIR}/qaps.csv \ --input_documents=${NARRATIVEQA_DATA_DIR}/documents.csv \ --data_split=valid \ --stories_dir=${NARRATIVEQA_DATA_DIR}/tmp/ \ --output_prefix=${NARRATIVEQA_EXAMPLE_DIR}/valid \ --nltk_data_path=${NLTK_DATA_PATH} python -m readtwice.models.narrative_qa.preprocess \ --spm_model_path=${SPM_MODEL_PATH} \ --input_qaps=${NARRATIVEQA_DATA_DIR}/qaps.csv \ --input_documents=${NARRATIVEQA_DATA_DIR}/documents.csv \ --data_split=train \ --stories_dir=${NARRATIVEQA_DATA_DIR}/tmp/ \ --output_prefix=${OUTPUT_NARRATIVE_QA}/train \ --generate_answers \ --nltk_data_path=${NLTK_DATA_PATH} gcloud storage cp ${NARRATIVEQA_DATA_DIR}qaps.csv ${NARRATIVEQA_EXAMPLE_DIR}

NarrativeQA 的预处理实现了基于 ROUGE-L oracle 的抽取式答案搜索(extractive_oracle.py)与故事文本解析(Gutenberg/电影剧本格式),评估逻辑见 narrative_qa/evaluation.py。

微调时使用trivia_qa.run_finetuning,学习率进一步降至5e-06,并关闭默认侧输入:

python -m readtwice.models.trivia_qa.run_finetuning \ --read_it_twice_bert_config_file=${CONFIG_PATH} \ --input_file=${NARRATIVEQA_EXAMPLE_DIR}/train.tfrecord-* \ --input_qaps=${NARRATIVEQA_EXAMPLE_DIR}/qaps.csv \ --eval_data_split=valid \ --output_dir=${NARRATIVEQA_OUTPUT_FOLDER} \ --init_checkpoint=${PRETRAINED_MODEL_CHECKPOINIT} \ --enable_side_inputs \ --cross_block_attention_mode=doc \ --do_train \ --nodo_eval \ --optimizer=adamw \ --learning_rate=5e-06 \ --nosummary_enable_default_side_input \ --num_train_epochs=6 \ --warmup_proportion=0.1 \ --learning_rate_schedule=inverse_sqrt \ --poly_power=1 \ --start_warmup_step=0 \ --save_checkpoints_steps=15000 \ --spm_model_path=${SPM_MODEL_PATH} \ --iterations_per_loop=1000 \ --nouse_one_hot_embeddings \ --use_tpu \ --tpu_job_name=??? \ --num_tpu_cores=16 \ --num_tpu_tasks=1 \ --decode_top_k=40 \ --decode_max_size=10 \ --tpu_name=??? \ --cross_attention_top_k=100

评估阶段切换为--nodo_train --do_eval,参数与训练基本一致。

微调命令行参数速查

综合三个任务的run_finetuning命令,核心参数及其作用如下(实现见各任务的 run_finetuning.py 与 trivia_qa/run_finetuning.py):

参数示例值作用
--read_it_twice_bert_config_file${CONFIG_PATH}模型架构配置 JSON
--input_filegs://.../train.tfrecord-*预处理产物(支持 glob)
--init_checkpoint${PRETRAINED_MODEL_CHECKPOINIT}预训练 checkpoint 初始化
--output_dirgs://...checkpoint 与结果输出目录
--enable_side_inputs(布尔)开启 Read-It-Twice 侧输入机制;关闭则退化为标准 Transformer
--cross_block_attention_modedoc摘要跨块交互范围:block/doc/batch/other_blocks
--do_train/--do_eval(布尔)训练 / 评估开关
--optimizeradamw优化器:adamwlamb(lamb_optimizer.py)
--learning_rate3e-05/1e-05/5e-06初始学习率(各任务不同)
--num_train_epochs6训练轮数
--warmup_proportion0.1预热步数占比
--learning_rate_scheduleinverse_sqrt/poly_decay学习率调度(optimization.py 中实现inverse_sqrt_learning_rate_schedule与多项式衰减)
--poly_power1多项式衰减幂次
--start_warmup_step0预热起始步
--save_checkpoints_steps5000/3000/15000保存 checkpoint 间隔
--iterations_per_loop1000/200TPU 每次 loop 迭代数
--nouse_one_hot_embeddings(布尔)关闭 one-hot 嵌入查找
--use_tpu(布尔)启用 TPU
--tpu_name/--tpu_job_name???TPU 名称与作业名(按实际环境填写)
--num_tpu_cores/--num_tpu_tasks16/1TPU 核数与任务数
--decode_top_k40/8解码候选 top-k
--decode_max_size10/20解码候选最大数量
--cross_attention_top_k100摘要注意力 Top-K 截断
--eval_json_path/--eval_data_split(TriviaQA/NarrativeQA)评估数据指定

训练目标方面,微调采用 masked language model 与跨度预测等损失(models/losses.py 提供LanguageModelLoss、批量共指消解损失等;HotpotQA 另有 yes/no 与 supporting-fact 损失)。预训练阶段的核心配置(如mlm_fraction_to_mask=0.15mention_mask_modemlm_use_whole_wordnum_replicas_concat等)可在预训练 demo 脚本 run_pretraining_demo.py 中查看。

预训练说明与 demo

README 明确说明:预训练代码尚未完整发布,主要有两个待解决事项:

  1. 预训练依赖自定义 TF 算子,用于在训练过程中动态执行词与实体掩码(对应run_pretraining_demo.py中的mention_mask_modemlm_use_whole_word等掩码策略,以及input_utils.pymask_same_entity_mentions等函数);
  2. 数据预处理目前依赖内部专有基础设施,无法直接开源。

作为替代,仓库发布了预训练 demo(run_pretraining_demo.py):该脚本虽不可直接执行,但完整展示了核心实现细节,包括:MLM 与共指消解(coreference resolution)损失的组合、source_model_config_file/source_model_config_base64两种模型配置注入方式、num_replicas_concat跨副本摘要汇聚、以及cross_block_attention_mode的块交互策略等,是理解 ReadTwice 预训练目标与数据流的关键入口。

小结

ReadTwice 通过"两次阅读 + 块级全局记忆"的机制,把超长文档阅读理解转化为可并行的块级处理:第一次读取生成摘要记忆,第二次读取携带记忆精读原文。本文从架构原理(ReadItTwiceBertModelFusedSideAttentionSummaryExtraction的配合)、配置字段(ReadItTwiceBertConfig全参数)、依赖与测试,到 HotpotQA / TriviaQA / NarrativeQA 三大任务的数据预处理与 TPU 微调命令,完整复现了官方 README 的实操路径,并补充了对应的源码级证据。实际复现时请注意:环境需 TPU(或 CPU)、tpu_name等参数需按集群实际填写、评估还需论文附录的额外步骤。

  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

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

相关推荐

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

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

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

立即咨询