models 仓库中的 BigBird:线性复杂度稀疏注意力与长序列训练的完整实践
【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/models
本篇围绕official/projects/bigbird项目,讲解 BigBird 稀疏注意力机制如何在 TensorFlow 2 中将 Transformer 的二次方注意力开销降为线性,并结合仓库中的BigBirdEncoder实现、EncoderScaffold集成方式、GLUE/SQuAD 两份完整 YAML 实验配置以及official/nlp/train.py的训练命令,帮助你在 TPU 或 GPU 上复现 BigBird 的长文本微调流程,并理解从块掩码构建到分块稀疏矩阵乘法的底层调用链。
BigBird:从二次方注意力到线性稀疏注意力
标准 Transformer 的自注意力对序列长度呈二次方(O(n²))依赖,这使得处理长文档(数千 token)在计算与显存上都不可行。BigBird 是一种稀疏注意力机制,将这一二次依赖降为线性,同时理论分析表明:BigBird 是序列函数的通用逼近器(universal approximator),并且保持图灵完备性(Turing complete)——即保留了二次方全注意力模型的核心表达性质。分析中还揭示了一个实践性结论:机制中 O(1) 个能“看到”整条序列的全局 token(如 CLS)本身就有理论收益,这正是 BigBird 在全局块设计上保留 CLS 类 token 全局注意力的原因。
从源码结构看,BigBird 的稀疏模式由三部分构成,均在 bigbird_attention.py 中实现:
- 局部窗口(banded attention):中间块只关注自身及前后共 3 个块构成的带状窗口;
- 全局块(global tokens):序列的首块与末块(通常承载 CLS 等特殊 token)对整条序列做全量注意力;
- 随机块(random blocks):每个块额外随机关注
num_rand_blocks个不相邻的块,提供跨长距离的信息通路。
核心函数bigbird_block_sparse_attention将 Q/K/V 按block_size分块后,把注意力拆成 first / second / middle / second_last / last 五段分别计算(见 bigbird_attention.py#L136-L368):首块与末块直接对整条 K/V 做全量einsum;中间块仅与3*block_size的窗口加上随机块做点积;被掩码屏蔽的位置统一加上-10000的偏移再取 softmax。由于每块只与常量个数的块交互,整体复杂度从 O(n²) 降为 O(n)。
随机块掩码是如何生成的
随机注意力的“邻接表”由bigbird_block_rand_mask生成:它以 numpy 随机数在middle_seq(去除首尾块后的块索引序列)中做置换,为每个中间块行挑选num_rand_blocks个候选块,并刻意排除自身窗口内的块(middle_seq[:start]与middle_seq[end+1:last]拼接后采样),从而保证随机块与局部窗口不重叠。BigBirdAttention.__init__在构建层时就为每个注意力头固定生成随机掩码(last_idx=1024约束随机块最多取到 1024 token 之前),并以seed=层索引区分各层,随后在_compute_attention中按实际序列长度截取前from_seq_length // block_size - 2行使用(见 bigbird_attention.py#L405-L459)。
掩码准备由BigBirdMasks层完成:它把[B, L]的输入 mask 重排成blocked_encoder_mask([B, L//block_size, block_size]),并用create_band_mask_from_inputs通过 einsum 生成 3Dband_mask,最终返回四元组[band_mask, encoder_from_mask, encoder_to_mask, blocked_encoder_mask]供注意力层消费(见 bigbird_attention.py#L371-L390)。
环境要求
BigBird 示例代码依赖 TensorFlow:仓库在开发时针对 TensorFlow 2.5.0 进行了测试,并声明后续将持续跟进最新的已发布 TensorFlow 版本。运行前请先确认 Python 与 TensorFlow 版本:
python --version python -c 'import tensorflow as tf; print(tf.__version__)'要求 Python 3.6+ 与 TensorFlow 2.5.0 或更高版本。由于训练入口official/nlp/train.py位于仓库内部而非已安装的 pip 包,官方说明要求把models目录(即本仓库根目录)加入 Python path 后再调用该脚本。
网络实现:从配置文件到编码器
BigBird 的编码器与层基于tf.kerasAPI,分散在三个文件中,README 明确了它们的分工:
- bigbird_attention.py:BigBird 稀疏注意力的层实现;
- encoders.py:把 BigBird 注意力集成进 NLP 建模库的
EncoderScaffold; - encoder.py:
BigBirdEncoder,README 特别注明梯度检查点(gradient checkpointing)目前在这一实现中提供。
BigBirdEncoder 的关键参数
BigBirdEncoder.__init__(encoder.py#L91-L107)定义了完整的超参空间,默认值对应 BigBird-base 规格:
| 参数 | 默认值 | 含义 |
|---|---|---|
vocab_size | 必填 | 词表大小 |
hidden_size | 768 | Transformer 隐层宽度 |
num_layers | 12 | Transformer 层数 |
num_attention_heads | 12 | 注意力头数(hidden_size须被其整除) |
max_position_embeddings | 4096 | 位置嵌入的最大长度,决定编码器可消费的最大序列 |
type_vocab_size | 16 | type_ids可取类型数 |
intermediate_size | 3072 | 前馈层中间维度 |
block_size | 64 | BigBird 注意力的块大小(from/to 序列各分块) |
num_rand_blocks | 3 | 每行随机关注的块数 |
activation | gelu | 激活函数 |
dropout_rate/attention_dropout_rate | 0.1 | 常规 dropout 与注意力 dropout |
embedding_width | None | 词嵌入宽度;小于hidden_size时嵌入被分解为两个矩阵 |
use_gradient_checkpointing | False | 用额外计算换显存的梯度检查点开关 |
构建过程:词嵌入(OnDeviceEmbedding)+ 位置嵌入 + 类型嵌入相加,经 LayerNorm 与 dropout 后,若embedding_width != hidden_size再经EinsumDense投影到hidden_size;随后BigBirdMasks生成四元组掩码,逐层堆叠TransformerScaffold(其attention_cls指定为layers.BigBirdAttention,attention_cfg传入from_block_size、to_block_size、num_rand_blocks、max_rand_mask_length=max_position_embeddings,并以seed=i用层索引做随机种子),输出{'sequence_output', 'encoder_outputs'}字典(encoder.py#L139-L220)。
两条构建路径:EncoderScaffold 与 BigBirdEncoder
encoders.py 中的EncoderScaffold工厂按encoder.type == 'bigbird'分发构建逻辑,并且存在一条与 README 表述一致的分支(encoders.py#L467-L531):
- 当
use_gradient_checkpointing为 True 时,直接返回official/projects/bigbird/encoder.py里的BigBirdEncoder——因为该版本通过RecomputeTransformerLayer在反向传播中重算前向,实现显存换计算的检查点; - 否则返回通用
networks.EncoderScaffold,其attention_cls=layers.BigBirdAttention、mask_cls=layers.BigBirdMasks、hidden_cls=layers.TransformerScaffold,attention_cfg中key_dim取hidden_size // num_attention_heads,并用layer_idx_as_attention_seed=True保证每层随机块互不相同。源码中留有 TODO:后续计划把梯度检查点统一进EncoderScaffold。
对应的配置类BigBirdEncoderConfig(encoders.py#L156-L174)在 dataclass 层面暴露了同样的默认值:vocab_size=50358、hidden_size=768、num_layers=12、num_attention_heads=12、intermediate_size=3072、max_position_embeddings=4096、num_rand_blocks=3、block_size=64、type_vocab_size=16,以及use_gradient_checkpointing=False。
梯度检查点的实现方式
encoder.py 顶部的RecomputeTransformerLayer继承layers.TransformerScaffold,其call把嵌套输入[emb, mask]展开成 5 个张量参数(emb、band_mask、encoder_from_mask、encoder_to_mask、blocked_encoder_mask),用recompute_grad.recompute_grad(f)包装后调用,从而在反向传播时重算整层前向。配套的两个工具模块:
- recompute_grad.py:
recompute_grad装饰器,前向不保存中间激活、反向重算; - recomputing_dropout.py + stateless_dropout.py:当开启检查点时,
BigBirdEncoder会把tf_keras.layers.Dropout替换为RecomputingDropout——因为普通 dropout 的随机噪声在重算前向时无法复现,必须使用无状态(seeded)dropout 才能保证重算结果一致(encoder.py#L111-L116)。
训练流程:基于 YAML 配置的运行方式
实验配置注册
experiment_configs.py 用@exp_factory.register_config_factory注册了两个实验类型:
bigbird/glue:任务为sentence_prediction.SentencePredictionConfig,训练/验证数据用SentencePredictionDataConfig(验证侧is_training=False, drop_remainder=False),并将task.model.encoder.type强制置为'bigbird';bigbird/squad:任务为question_answering.QuestionAnsweringConfig,数据用QADataConfig。
两者的TrainerConfig共用一组默认优化器配置:AdamW(weight_decay_rate=0.01,exclude_from_weight_decay=['LayerNorm', 'layer_norm', 'bias'])、polynomial学习率衰减(GLUE 初始 3e-5、SQuAD 初始 8e-5,终点 0.0)、polynomialwarmup,并以restrictions约束train_data.is_training与validation_data.is_training必须显式给出(experiment_configs.py#L26-L99)。
GLUE 实验配置
experiments/glue_mnli_matched.yaml 是 MNLI-matched 的完整示例,三段结构如下:
task: hub_module_url: '' model: num_classes: 3 # MNLI 三分类 encoder: type: bigbird bigbird: use_gradient_checkpointing: false # hidden_size: 768 # 未注释即覆盖 BigBirdEncoderConfig 默认值 # num_layers: 12 # num_attention_heads: 12 # intermediate_size: 3072 init_checkpoint: 'TODO' # 用预训练权重初始化 metric_type: 'accuracy' train_data: drop_remainder: true global_batch_size: 32 input_path: 'TODO' is_training: true seq_length: 1024 label_type: 'int' validation_data: drop_remainder: false global_batch_size: 32 input_path: 'TODO' is_training: false seq_length: 1024 label_type: 'int' trainer: checkpoint_interval: 3000 optimizer_config: learning_rate: polynomial: decay_steps: 36813 # 100% of train_steps end_learning_rate: 0.0 initial_learning_rate: 3.0e-05 power: 1.0 type: polynomial optimizer: type: adamw warmup: polynomial: power: 1 warmup_steps: 3681 # ~10% of train_steps type: polynomial steps_per_loop: 1000 summary_interval: 1000 # Training data size 392,702 examples, 3 epochs. train_steps: 36813 validation_interval: 6135 # Eval data size = 9815 examples. validation_steps: 307 best_checkpoint_export_subdir: 'best_ckpt' best_checkpoint_eval_metric: 'cls_accuracy' best_checkpoint_metric_comp: 'higher'要点:train_steps=36813按 392,702 条训练样本、3 个 epoch 计算,warmup 取其 10%(3681 步);每 6135 步验证一次(验证集 9815 条约 307 步),并按cls_accuracy越高越好导出best_ckpt最优检查点。YAML 中注释掉的hidden_size/num_layers/num_attention_heads/intermediate_size展示了如何用配置覆盖BigBirdEncoderConfig的默认结构。
SQuAD 实验配置
experiments/squad_v1.yaml 在相同骨架上增加了问答任务特有字段:
task: model: encoder: type: bigbird bigbird: use_gradient_checkpointing: false max_answer_length: 30 # 答案最大 token 数 n_best_size: 20 # 取前 n 个候选答案 null_score_diff_threshold: 0.0 # v2 中允许空答案的阈值 init_checkpoint: 'TODO' train_data: global_batch_size: 48 is_training: true seq_length: 1024 validation_data: do_lower_case: true doc_stride: 128 # 文档滑窗步长 global_batch_size: 48 is_training: false query_length: 64 seq_length: 1024 tokenization: SentencePiece # 使用 SentencePiece 分词 version_2_with_negative: false # SQuAD v1(含负例的 v2 设为 true) vocab_file: 'TODO' trainer: max_to_keep: 5 optimizer_config: # 与 glue 相同结构:adamw + polynomial LR(8e-5→0) + warmup train_steps: 3699 validation_steps: 226 best_checkpoint_eval_metric: 'final_f1'其中doc_stride=128、query_length=64、seq_length=1024由question_answering_dataloader.QADataConfig消费,tokenization: SentencePiece表明 SQuAD 流程使用 SentencePiece 词表而非 WordPiece。
数据准备
训练数据脚本与 BERT 的官方流程一致:先按 NLP 文档中的微调数据制备方式生成tf.train.Example序列,再按 README 指引下载官方提供的 SentencePiece 词表文件vocab_sp.model(BigBird 官方词表,随 checkpoint 一起发布)。GLUE 数据对应sentence_prediction数据加载器(label_type: 'int'表示整型标签),SQuAD 数据对应question_answering数据加载器并需额外传入vocab_file。
训练命令
训练入口是 train.py,代码支持train / train_and_eval / eval三种模式,通过--mode指定。
GLUE(TPU):
INIT_CKPT=??? TRAIN_FILE=??? EVAL_FILE=??? python3 official/nlp/train.py \ --experiment_type=bigbird/glue \ --config_file=experiments/glue_mnli_matched.yaml \ --params_override=task.init_checkpoint=${INIT_CKPT} \ --params_override=runtime.distribution_strategy=tpu \ --params_override=task.train_data.input_path=${TRAIN_FILE},task.validation_data.input_path=${EVAL_FILE} \ --tpu=??? \ --mode=train_and_evalSQuAD(README 中给出的配置路径是 bazel 风格路径,在本仓库内实际对应 experiments/squad_v1.yaml,使用时请以仓库内路径为准):
VOCAB_FILE=??? TRAIN_FILE=??? EVAL_FILE=??? python3 official/nlp/train.py \ --experiment_type=bigbird/squad \ --config_file=official/projects/bigbird/experiments/squad_v1.yaml \ --params_override=task.init_checkpoint=${INIT_CKPT} \ --params_override=task.train_data.input_path=${TRAIN_FILE},task.validation_data.input_path=${EVAL_FILE},task.validation_data.vocab_file=${VOCAB_FILE} \ --params_override=runtime.distribution_strategy=tpu \ --tpu=??? \ --mode=train_and_evalGPU 使用方式:去掉--tpu标志,并把runtime.distribution_strategy通过--params_override设为mirrored,即使用tf.distribute.MirroredStrategy做多卡数据并行。--params_override的写法是section.field=value,同一标志内用逗号分隔多个键值对,这是 Orbit 配置系统(official/core/config_definitions.py)覆盖 YAML 的统一方式。
官方 Checkpoint
README 给出了 BigBird 官方发布模型的规格与参考指标:
| 模型 | 配置 | 预训练数据 | Checkpoint | 参考指标 |
|---|---|---|---|---|
| BigBird base | 12 层,序列长度 1024 ≤ L ≤ 4096 | Wiki + Books + CC-News + Stories(Common Crawl 一部分) | bigbird_base(发布在 TensorFlow 模型库的 BigBird 存储桶,bigbird.etc.base.keras.tar.gz) | SQuAD v1 F1 91.3,TriviaQA F1 79.8 |
表中序列长度范围 1024–4096 正对应实现中的MAX_SEQ_LEN = 4096(bigbird_attention.py#L20)与BigBirdEncoder的max_position_embeddings=4096默认值;use_gradient_checkpointing与长序列(如 4096)配合使用可以显著降低激活内存占用,适合在有限显存设备上处理长文档。
小结
official/projects/bigbird用一套小而完整的模块展示了长序列 Transformer 的工程范式:稀疏注意力层(BigBirdAttention+BigBirdMasks+ 分块einsum计算)负责把复杂度从 O(n²) 降到 O(n);EncoderScaffold集成与独立的BigBirdEncoder(含recompute_grad梯度检查点与无状态 dropout)覆盖常规与显存受限两条训练路径;experiment_configs.py注册的bigbird/glue、bigbird/squad两个配置工厂加上 YAML 覆盖机制,使 GLUE 与 SQuAD 微调只需一条train.py命令即可在 TPU/GPU 上运行。理解这条从配置到稀疏矩阵运算的链路,也为在仓库内复用EncoderScaffold开发其他长文本模型提供了模板。
【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/models
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考