models 仓库中的 BigBird:线性复杂度稀疏注意力与长序列训练的完整实践
本篇围绕 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.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)、polynomial warmup,并以 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_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 = 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 开发其他长文本模型提供了模板。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0623
Hy4-previewHy4 preview 是由腾讯混元团队研发的新一代混合专家(MoE)旗舰模型。模型总参数量 770B,每个 token 激活 49B,主干共包含78层,第一层采用标准 FFN,其余 77 层均为 MoE 结构,每层包含 256 个路由专家与 1 个共享专家,每个 token 激活 top-8 路由专家及共享专家。主干之外原生内置 1 层 MTP(总参数量 10B,激活 0.7B)以支持投机解码。Python00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
GLM-5.3-FlashGLM-5.3-Flash (320B-A18B),是GLM-5系列的首个原生多模态模型。320B总参数,能力超过GLM-5.2Jinja00
Spark-X2.5-4BSpark-X2.5-4B 旨在让强大的 AI 更实用、更高效、更易获得。在广泛日常任务中表现强劲,涵盖对话、写作、翻译、推理、编码、工具调用以及智能体工作流,并在同等规模的开源模型中取得领先成绩。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00