首页
/ models 仓库中的 BigBird:线性复杂度稀疏注意力与长序列训练的完整实践

models 仓库中的 BigBird:线性复杂度稀疏注意力与长序列训练的完整实践

2026-09-04 15:28:29作者:滕妙奇

本篇围绕 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 生成 3D band_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.keras API,分散在三个文件中,README 明确了它们的分工:

  • bigbird_attention.py:BigBird 稀疏注意力的层实现;
  • encoders.py:把 BigBird 注意力集成进 NLP 建模库的 EncoderScaffold
  • encoder.pyBigBirdEncoder,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.BigBirdAttentionattention_cfg 传入 from_block_sizeto_block_sizenum_rand_blocksmax_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.BigBirdAttentionmask_cls=layers.BigBirdMaskshidden_cls=layers.TransformerScaffoldattention_cfgkey_dimhidden_size // num_attention_heads,并用 layer_idx_as_attention_seed=True 保证每层随机块互不相同。源码中留有 TODO:后续计划把梯度检查点统一进 EncoderScaffold

对应的配置类 BigBirdEncoderConfigencoders.py#L156-L174)在 dataclass 层面暴露了同样的默认值:vocab_size=50358hidden_size=768num_layers=12num_attention_heads=12intermediate_size=3072max_position_embeddings=4096num_rand_blocks=3block_size=64type_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.pyrecompute_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.01exclude_from_weight_decay=['LayerNorm', 'layer_norm', 'bias'])、polynomial 学习率衰减(GLUE 初始 3e-5、SQuAD 初始 8e-5,终点 0.0)、polynomial warmup,并以 restrictions 约束 train_data.is_trainingvalidation_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=128query_length=64seq_length=1024question_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_eval

SQuAD(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_eval

GPU 使用方式:去掉 --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 = 4096bigbird_attention.py#L20)与 BigBirdEncodermax_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/gluebigbird/squad 两个配置工厂加上 YAML 覆盖机制,使 GLUE 与 SQuAD 微调只需一条 train.py 命令即可在 TPU/GPU 上运行。理解这条从配置到稀疏矩阵运算的链路,也为在仓库内复用 EncoderScaffold 开发其他长文本模型提供了模板。

登录后查看全文
热门项目推荐
相关项目推荐