首页
/ TensorFlow 版 Longformer 实战:滑动窗口注意力实现解析与 MNLI 微调全流程

TensorFlow 版 Longformer 实战:滑动窗口注意力实现解析与 MNLI 微调全流程

2026-09-04 21:59:52作者:伍霜盼Ellen

本文基于 TensorFlow Model Garden 中 official/projects/longformer 目录的实现,讲解 Longformer 这一长文档 Transformer 在 TensorFlow/TPU 环境下的移植要点、与 HuggingFace PyTorch 版本的关键差异(固定 global_attention_size、去掉 tf.cond 的静态全局注意力索引),并完整给出从权重转换、数据准备到 MNLI 微调训练的可复制命令与配置参数说明,帮助读者掌握在 TF2 训练框架中运行长序列 Transformer 的完整流程。

一、目录结构与文件职责

Longformer 模块位于 official/projects/longformer,核心文件分工如下:

文件 职责
longformer.py 定义 LongformerEncoderConfig 与编码器工厂函数 get_encoder
longformer_encoder.py LongformerEncoder:嵌入、按窗口补长、逐层 Transformer 堆叠
longformer_encoder_block.py LongformerEncoderBlock:单层的注意力 + FFN + 残差/归一化
longformer_attention.py LongformerAttention:滑动窗口局部注意力 + 全局注意力拼接
longformer_experiments.py 注册 longformer/pretraininglongformer/glue 两个实验配置
train.py 训练入口,复用 Model Garden 的 train_lib.run_experiment
experiments/ MNLI 微调与预训练的 YAML 配置
utils/ PyTorch 权重转换、tokenized 数据转 TFRecord
README.md 移植说明与 MNLI 微调步骤(本文骨架)

二、与 HuggingFace 实现的两个核心差异

README 首先说明了本实现相对 HuggingFace transformers 版本的两处修改,理解它们是后续一切配置的前提。

2.1 强制固定的 global_attention_size

HuggingFace 版本允许为每个句子指定不同的全局注意力 token 集合,而本实现要求所有模型在配置中指定一个统一的 global_attention_size任何句子的前 global_attention_size 个 token 都参与全局注意力,不支持逐句子差异化的全局注意力大小。README 给出的理由是:TPU 上张量尺寸必须能够静态确定,动态的全局注意力索引会破坏这一点。

这一设计在源码中体现得很直接。longformer_attention.py 的 _get_global_attn_indices 不再从 is_index_global_attn 张量中动态搜索非零位置,而是直接按固定的 global_attention_size 构造索引:

# All global attention size are fixed through global_attention_size
max_num_global_attn_indices = global_attention_size
row_indices = tf.range(batch_size)
...
col_indices = tf.range(global_attention_size)

即在每个 batch 内,全局注意力 token 恒为位置 0 .. global_attention_size - 1。相应地,longformer_encoder.py 的 call 在构造输入时,把 attention mask 的前 global_attention_size 行强制置为 2(永不被屏蔽),并把 is_index_global_attn 构造为"前 N 个位置为 True、其余为 False"的静态布尔矩阵。

2.2 用 Python 分支替代 tf.cond

README 指出,由于全局注意力现在在"开头"就以静态配置确定,longformer_attention.py 中原来的 tf.cond 全部改成了普通的 Python if 条件。对应源码中的写法是:

# this function is only relevant for global attention
if self.global_attention_size > 0:
  attn_scores = self._concat_with_global_key_attn_probs(...)
else:
  pass

call 函数中所有与全局注意力相关的路径(拼接全局分数、计算全局输出、构造掩码)都由 if self.global_attention_size > 0: 这类图构建期的静态分支控制,而非运行时条件。这让整个计算图在 TPU/XLA 下形状完全可静态推断。

三、模型架构要点(源码级解读)

3.1 编码器配置:LongformerEncoderConfig

longformer.py 中,LongformerEncoderConfig 继承自 BERT 编码器配置,额外增加三个字段:

  • attention_window: List[int]:每一层的滑动窗口大小(列表长度需等于层数);
  • global_attention_size: int:参与全局注意力的前缀 token 数,默认 0;
  • pad_token_id: int:补长用的 pad token id,默认 1。

工厂函数 get_encoder 将其映射为 LongformerEncoder,其余超参(vocab_sizehidden_sizenum_layersnum_attention_headsintermediate_size、激活与 dropout 等)沿用 BERT 配置字段,因此 Longformer 可以直接挂载到 Model Garden 的 EncoderConfig(type='any', any=LongformerEncoderConfig()) 任务体系里。

3.2 滑动窗口局部注意力:O(L) 复杂度的关键

LongformerAttentionlongformer_attention.py#L107-L130)在初始化时强制校验:窗口大小必须为正偶数(attention_window % 2 == 0),并取单侧窗口 _one_sided_attn_window_size = window // 2

局部注意力的核心是 _sliding_chunks_query_key_matmulL445 起):

  1. 把 query/key 序列切成大小为 2w、重叠为 w 的滑动块(_chunk 方法借助 tf.signal.frame 实现,等价于步长为 w 的窗口卷积);
  2. 块内做稠密点积 tf.einsum("bcxd,bcyd->bcxy"),再通过对角化重组(_pad_and_transpose_last_two_dims_pad_and_diagonalize)拼回"每个 token 只看到左右各 w 个邻居"的带状注意力分数,形状为 [B, T, H, 2w+1]
  3. _mask_invalid_locations 用三角掩码把超出窗口的无效位置置为 -inf

前向流程(call 函数 L276-L324)是:局部 QK 分数 + 全局 key 分数拼接 → softmax → 分别对全局/局部 Value 求加权和。全局注意力使用独立的 _global_query_dense/_global_key_dense/_global_value_dense 投影(在 _build_from_signature 中创建,L180-L218),分数拼接后每行维度为 2w + 1 + global_attention_size

3.3 按窗口尺寸自动补长

由于分块要求序列长度是 2w 的整数倍,LongformerEncoder.call 会先调用 _pad_to_window_size:取各层最大窗口 attention_window,把 input_word_idspad_token_id 补齐、input_mask 用 0 补齐(pad 位置不参与注意力)、type_ids 用 0 补齐;前向结束后再把补上的尾部切掉(L270-L271),保证输出 sequence_output 与输入等长。输出为包含 sequence_outputpooled_output(首 token 过 tanh Dense 池化)和逐层 encoder_outputs 的字典。

3.4 编码器块与归一化

LongformerEncoderBlock 是标准 Transformer 块:注意力 → dropout → 残差+LayerNorm(或 norm_first=True 时先归一化),随后 FFN(inner_dim,默认 4 倍隐层)→ dropout → 残差+LayerNorm。两处实现细节值得注意:LayerNorm 固定用 dtype=tf.float32 保证混合精度下的数值稳定(L158-L166);混合精度策略为 mixed_bfloat16 时中间激活层回退到 float32(L176-L181),源码注释说明 bfloat16 下部分优化器收敛性未经充分验证。

四、准备预训练权重

微调前需要一份 TF 格式的 Longformer checkpoint,README 给出两条路径。

Option 1:直接拉取仓库作者保存的 allenai/longformer-base-4096 转换结果:

gsutil cp -r gs://model-garden-ucsd-zihan/longformer-4096 .

Option 2:用转换脚本从 HuggingFace 预训练模型自行生成:

python3 utils/convert_pretrained_pytorch_checkpoint_to_tf.py

该脚本位于 utils/convert_pretrained_pytorch_checkpoint_to_tf.py。从源码看,它通过 transformers.AutoModel.from_pretrained("allenai/longformer-base-4096") 加载 PyTorch 参数(L29-L34),再构造一个 TF LongformerEncodervocab_size=50265max_position_embeddings=4098global_attention_size=1L37-L62),逐组件 set_weights 完成映射——词嵌入、嵌入 LayerNorm、类型嵌入以及每层注意力的普通/全局 QKV 投影、中间层与输出层等,从而把 AllenAI 命名空间的变量名对齐到本仓库的变量结构,产出的 checkpoint 即 longformer-4096/longformer 目录(对应训练命令中的 INIT_CHECKPOINT=longformer-4096/longformer)。

五、可选:准备输入文件

若不使用云存储中现成的 MNLI TFRecord,可用 utils/longformer_tokenizer_to_tfrecord.py 自行生成:

python3 utils/longformer_tokenizer_to_tfrecord.py

脚本逻辑(见 L23-L72):加载 GLUE 的 mnli 数据集,用 allenai/longformer-base-4096 的 fast tokenizer 对 premise/hypothesis 双句做 tokenize,padding="max_length"max_seq_length=512(需与模型输入尺寸一致),train 集与 validation_matched 集分别写成训练/验证用的 TFRecord。

六、在 MNLI 上微调:完整命令与参数说明

README 给出的训练命令(注意需以仓库根目录的 official 包为 PYTHONPATH 基准):

TRAIN_DATA=task.train_data.input_path=gs://model-garden-ucsd-zihan/longformer_allenai_mnli_train.tf_record,task.validation_data.input_path=gs://model-garden-ucsd-zihan/longformer_allenai_mnli_eval.tf_record
INIT_CHECKPOINT=longformer-4096/longformer
PYTHONPATH=/path/to/model/garden \
    python3 train.py \
    --experiment=longformer/glue \
    --config_file=experiments/glue_mnli_allenai.yaml \
    --params_override="${TRAIN_DATA},runtime.distribution_strategy=tpu,task.init_checkpoint=${INIT_CHECKPOINT}" \
    --tpu=local \
    --model_dir=/path/to/outputdir \
    --mode=train_and_eval

按 README 的说明,该配置在 TPU 上运行约 3 小时,可取得约 86 的 MNLI 精度(此为 README 原始声明,实际以运行环境为准)。各参数含义:

  • --experiment=longformer/glue:使用 longformer_experiments.py 注册的句子预测任务——SentencePredictionConfig + LongformerEncoderConfig,AdamW 优化器(weight decay 0.01,对 LayerNorm/bias 不衰减),polynomial 学习率 3e-5 + polynomial warmup;
  • --config_file=experiments/glue_mnli_allenai.yaml:YAML 覆盖实验默认值;
  • --params_override=...:命令行再覆盖数据路径、分布策略与初始 checkpoint,优先级最高;
  • --tpu=local --mode=train_and_eval:本机/集群 TPU 上训练并同时评估。

6.1 配置文件 glue_mnli_allenai.yaml 关键参数

experiments/glue_mnli_allenai.yaml 中与 4096 长度预训练模型对齐的关键项:

task:
  model:
    num_classes: 3          # MNLI 三分类
    encoder:
      type: any
      any:
        max_position_embeddings: 4098
        attention_window: [128, 128, 128, 128, 128, 128, 128, 128, 128, 128, 128, 128]
        global_attention_size: 1   # 前缀 1 个 token(CLS)做全局注意力
        vocab_size: 50265          # AllenAI 词表
  metric_type: 'accuracy'
  train_data:
    global_batch_size: 32
    seq_length: 512                # 微调序列长度
trainer:
  optimizer_config:
    learning_rate:
      polynomial:
        decay_steps: 61359
        initial_learning_rate: 3.0e-05
        end_learning_rate: 0.0
        power: 1.0
    optimizer: {type: adamw}
    warmup:
      polynomial: {warmup_steps: 6136, power: 1}
  # Training data size 392,702 examples, 5 epochs.
  train_steps: 61359
  validation_interval: 2000
  validation_steps: 307

要点:train_steps=61359 由 392,702 个训练样本、batch 32、5 个 epoch 推出;validation_steps=307 覆盖整个 matched 验证集;steps_per_loop=1000summary_interval=1000 控制 TPU step 批处理与 TensorBoard 汇总频率。

仓库还另外提供了两个配置供参考:

  • experiments/glue_mnli.yaml:面向较短序列(seq_length: 128attention_window 全 32、max_position_embeddings: 512steps_per_loop: 100)的变体,学习率/步数与 allenai 版相同,适合从更小的 BERT 类权重起步;
  • experiments/pretraining_512.yaml:对应 longformer/pretraining 实验的 MaskedLM 预训练配置——BERT-base 规模(hidden 768 / 12 层 / 12 头 / inner 3072),attention_window 全 32、global_attention_size: 1seq_length: 512max_predictions_per_seq: 76、全局 batch 256、训练 100 万步、初始学习率 1e-4、warmup 1 万步,数据指向 gs://tf_model_garden/nlp/data/research_data/bert_pretrain/wikipedia.tfrecord-*

6.2 训练入口的底层调用链

train.pymain 是标准的 Model Garden 流水线:解析 gin 与 params_overridetrain_utils.parse_configuration → 按 runtime.mixed_precision_dtype 设置混合精度 → distribute_utils.get_distribution_strategy 建立 TPU/GPU 分布策略 → task_factory.get_task 实例化任务(即第 6 节所述 glue 任务)→ train_lib.run_experiment 执行训练/评估循环,最后 save_gin_config 归档配置。理解这一链条后,--params_override 中任何 section.key=value 都能直接对应到 ExperimentConfig 的字段上。

七、验证实现的单元测试

仓库自带测试可作为改配置后的回归验证手段。longformer_encoder_test.pycombinationsattention_window ∈ {32, 128}global_attention_size ∈ {0, 1, 2} 做参数化组合,检查输出形状 (batch, seq_len, hidden_size) 正确;另有 norm_first 与全局注意力大小的组合测试,以及 longformer_attention_test.pylongformer_encoder_test.py 等文件覆盖注意力层与编码器的细节。运行方式(以仓库 official 包可导入为前提):

python -m pytest official/projects/longformer/longformer_encoder_test.py
# 或
python -m pytest official/projects/longformer/longformer_attention_test.py

八、小结与适用边界

  • 本实现的核心取舍是以静态化的全局注意力(固定前缀 N 个 token、去掉 tf.cond)换取 TPU/XLA 可运行性,代价是不支持逐句定制全局 token 集合;attention_window 必须是正偶数,序列会被自动补齐到 max(attention_window) 的整数倍;
  • 微调 4096 预训练模型时,务必保持 vocab_size=50265max_position_embeddings=4098attention_window 12 层全 128 与 glue_mnli_allenai.yaml 一致,并用 task.init_checkpoint 指向转换后的 checkpoint;
  • 从源码结构看,滑动窗口实现基于 tf.signal.frame 的分块技巧,作者代码注释中也坦承对角重组部分"most likely not very efficient",若追求极致吞吐可结合 tf.profiler 定位热点;
  • 本文所有性能数字(约 3 小时、精度约 86)均来自 README 的原始声明,实际结果取决于硬件规格与数据版本。
登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.12 K
2.72 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
527
590
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
904
1.82 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
854
1.34 K
docsdocs
暂无描述
Markdown
889
5.78 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.52 K
1.01 K
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.33 K
1.45 K
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
982
502
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
540
384