TensorFlow Models LRA 项目实战:训练 MEGA、Transformer 与 Linformer 长程序列建模基线
【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/models
official/projects/lra目录是 TensorFlow Model Garden(models 仓库)中针对 Long Range Arena(LRA)基准的 TensorFlow 2.x 实现,包含 MEGA、Transformer 与 Linformer 三种长程序列建模基线,代码改编自 google-research/long-range-arena 官方仓库。本文以该目录的 README 为主线,完整继承其中的训练命令、数据集路径与实验配置,并结合仓库源码剖析训练入口、实验注册机制、YAML 配置结构与三种编码器的关键实现,帮助读者在 TPU/GPU 上复现并调优 LRA 基准任务。
项目定位与目录结构
LRA 是一组用于评测"高效 Transformer"在长序列(数千到上万长度)上建模能力的任务集。本项目的实现在official/projects/lra/下组织为三块:
- 训练入口:train.py —— 基于 Model Garden 通用
train_lib的定制训练脚本; - 实验注册:transformer_experiments.py、linformer_experiments.py、mega_experiments.py —— 通过装饰器注册
实验名 → 配置工厂的映射; - 模型与数据:三种编码器(
transformer_encoder.py、linformer_encoder.py、mega_encoder.py 及对应的 attention block 实现)、AAN 任务专用的双编码器任务类 lra_dual_encoder_task.py,以及experiments/子目录下的 15 份 YAML 实验配置(每种模型 × ListOps/IMDB/AAN/CIFAR/Pathfinder 五个任务)。
训练命令:README 原始工作流
README 给出的标准流程是:设置TRAIN_DATA覆盖训练/验证数据的 GCS 路径 → 以PYTHONPATH指向 Model Garden 根目录 → 调用train.py并指定--experiment(已注册的实验名)、--config_file(YAML 配置)、--params_override(参数覆盖)、--tpu、--model_dir、--mode。以下三组命令完整继承自 README。
在 ListOps 上训练 Transformer
TRAIN_DATA=task.train_data.input_path=gs://model-garden-ucsd-zihan/lra_listops_train.tf_record,task.validation_data.input_path=gs://model-garden-ucsd-zihan/lra_listops_eval.tf_record PYTHONPATH=[/PATH/TO/MODEL_GARDEN] \ python3 train.py \ --experiment=transformer/lra_listops \ --config_file=../experiments/lra_listops.yaml \ --params_override="${TRAIN_DATA},runtime.distribution_strategy=tpu" \ --tpu=local \ --model_dir=[OUTPUT_DIR] \ --mode=train_and_eval在 ListOps 上训练 Linformer
TRAIN_DATA=task.train_data.input_path=gs://model-garden-ucsd-zihan/lra_listops_train.tf_record,task.validation_data.input_path=gs://model-garden-ucsd-zihan/lra_listops_eval.tf_record PYTHONPATH=[/PATH/TO/MODEL_GARDEN] \ python3 train.py \ --experiment=linformer/lra_listops \ --config_file=../experiments/lra_listops_linformer.yaml \ --params_override="${TRAIN_DATA},runtime.distribution_strategy=tpu" \ --tpu=local \ --model_dir=[OUTPUT_DIR] \ --mode=train_and_eval在 Text(IMDB-4096)上训练 MEGA
README 标注该配置为 "Reproduced Acc = 87.55"(即作者在该数据集上复现报告 87.55 的准确率):
TRAIN_DATA=task.train_data.input_path=gs://model-garden-ucsd-zihan/lra_imdb_4096_train.tf_record,task.validation_data.input_path=gs://model-garden-ucsd-zihan/lra_imdb_4096_eval.tf_record PYTHONPATH=[/PATH/TO/MODEL_GARDEN] \ python3 train.py \ --experiment=mega/lra_imdb \ --config_file=../experiments/lra_imdb_mega.yaml \ --params_override="${TRAIN_DATA},runtime.distribution_strategy=tpu" \ --tpu=local \ --model_dir=[OUTPUT_DIR] \ --mode=train_and_eval参数说明(结合 train.py 源码):
| 参数 | 作用 | 源码对应 |
|---|---|---|
--experiment | 选择已注册实验(如mega/lra_imdb),决定编码器类型与任务/数据加载器组合 | 各*_experiments.py中的@exp_factory.register_config_factory(...) |
--config_file | 指向 YAML 配置文件,提供 task 与 trainer 的完整默认值(注意命令中../experiments/...是相对于 LRA 目录的 CLI 相对路径,需按 README 的目录约定执行) | train_utils.parse_configuration(FLAGS) |
--params_override | gin 风格的key=value覆盖列表,此处覆盖数据路径并指定runtime.distribution_strategy=tpu | gin.parse_config_files_and_bindings(FLAGS.gin_file, FLAGS.gin_params)(train.py 中gin_params与gin_file合并解析) |
--tpu | TPU 地址(local表示本地 TPU) | distribute_utils.get_distribution_strategy(..., tpu_address=params.runtime.tpu) |
--model_dir | checkpoint 与序列化配置的输出目录 | train_utils.serialize_config(params, model_dir)、train_lib.run_experiment |
--mode | train_and_eval/train/eval等;train.py 中仅当 mode 含train时才序列化 YAML 配置,以避免连续 eval 任务与训练任务写文件竞争 | if 'train' in FLAGS.mode: train_utils.serialize_config(...) |
main()的完整调用链为:解析 gin 配置 → 若启用runtime.mixed_precision_dtype则调用performance.set_mixed_precision_policy设置混合精度(GPU 上收益来自 float16,TPU 上来自 bfloat16,loss_scale仅在 float16 下生效)→ 构建DistributionStrategy→ 在策略作用域内task_factory.get_task构建任务 →train_lib.run_experiment执行训练/评估 → 最后train_utils.save_gin_config落盘 gin 配置。命令行 flag 本身定义于 official/common/flags.py(tfm_flags.define_flags())。
--experiment是怎么工作的:实验注册机制
每个*_experiments.py模块通过@exp_factory.register_config_factory('模型/任务')把"实验名"绑定到一个返回cfg.ExperimentConfig的工厂函数。以 transformer_experiments.py 为例,注册了五个实验:transformer/lra_listops、transformer/lra_imdb、transformer/lra_cifar、transformer/lra_pathfinder、transformer/lra_aan;linformer_experiments.py 与 mega_experiments.py 同构,分别注册linformer/...与mega/...前缀的同名任务。Linformer 与 MEGA 两个实验文件结构完全对称,唯一差异在于编码器配置类与默认学习率:
- Transformer/Linformer 的
_TRAINER默认初始学习率3e-5(adamw,weight_decay_rate=0.01,排除LayerNorm/layer_norm/bias的权重衰减); - MEGA 的
_TRAINER默认初始学习率1e-7(mega_experiments.py 中显式配置),实际训练时几乎都依赖 YAML 覆盖,例如 IMDB-MEGA 配置中覆盖为0.004。
任务与数据加载器的组合也有规律:ListOps / IMDB / CIFAR / Pathfinder 四个任务复用 NLP 通用的sentence_prediction.SentencePredictionConfig+SentencePredictionDataConfig(即"双句预测"框架下的序列分类任务);AAN(Approximate Area Navigation)任务则使用 LRA 项目自研的DualEncoderConfig+DualEncoderDataConfig(见下文"双编码器任务"小节)。
实验配置文件详解(experiments/ 目录)
README 指出所有实验配置位于experiments子文件夹,共 15 份 YAML,命名为lra_<任务>[_<模型>].yaml(无后缀为 Transformer,_linformer、_mega为对应变体)。所有 YAML 均为task+trainer两级结构,input_path留空为TODO,运行时必须像 README 命令那样用params_override注入。下面按任务归纳关键默认值(数值取自各 YAML 文件本身):
ListOps(10 类,序列长度 2000)
lra_listops.yaml(Transformer):num_classes: 10、vocab_size: 100、embedding_size: 512、hidden_size: 512、intermediate_size: 1024、num_attention_heads: 8、num_layers: 4、max_position_embeddings: 2000、dropout 均为 0.1、gelu 激活;数据global_batch_size: 64、seq_length: 2000;训练train_steps: 5000、学习率 polynomial 衰减(initial_learning_rate: 5e-5→ 0,decay_steps: 5000,power: 0.5)+ 1000 步 warmup、AdamW。
lra_listops_linformer.yaml 与 Transformer 版几乎相同,唯一新增项是low_rank_features: 32——即 Linformer 低秩投影维度,这是实现"线性复杂度注意力"的核心超参(LinformerEncoderConfig默认值为 256,见 linformer.py,此处按 LRA 论文推荐压到 32)。
IMDB / Text(2 类情感,序列长度 1000)
lra_imdb.yaml(Transformer):embedding_size/hidden_size: 256、4 头、4 层、vocab_size: 258、seq_length: 1000、batch 32、train_steps: 20000、初始学习率5e-5、warmup 8000 步、衰减 20000 步。
lra_imdb_mega.yaml(MEGA)在 Transformer 结构参数之外额外包含 MEGA 专属超参:zdim: 64(门控隐维度)、hdim: 256(隐层维度)、ndim: 16(EMA 头数)、activation: 'silu'、bidirectional: true(双向 EMA)、dropout: 0.1、hidden_dropout: 0.1,并设置use_encoder_pooler: true复用编码器池化输出;训练train_steps: 50000、初始学习率0.004、warmup 10000 步、power 1 衰减 25000 步。注意 README 中 MEGA 的示例命令将lra_imdb_mega.yaml搭配IMDB-4096数据集路径使用(lra_imdb_4096_*.tf_record)。
AAN(2 类,序列长度 4000,双编码器任务)
lra_aan.yaml:num_classes: 2、max_seq_length: 4000、embedding_size/hidden_size: 128、4 头 4 层、vocab_size: 258、seq_length: 4000、batch 32、train_steps: 5000、学习率5e-4、warmup 800 步;checkpoint_interval与validation_interval均为 500 步。_linformer与_mega变体结构同构。
其余任务:lra_cifar*.yaml(CIFAR-10 图像转序列任务)与lra_pathfinder*.yaml(路径规划任务)同样按"Transformer 默认 + Linformer/MEGA 变体"三份一组的方式组织,可按任务名直接查文件。
三种编码器:配置类与关键实现
三种模型都继承BertEncoderConfig(official/nlp/configs/encoders.py),通过@base_config.bind(...)将配置 dataclass 绑定到编码器实例化函数(gin 可绑定),这是 Model Garden 统一的"配置即对象"模式。
Transformer 基线
transformer.py 中TransformerEncoderConfig未新增字段,get_encoder()直接构造 transformer_encoder.py 的TransformerEncoder:vocab_size、hidden_size、num_layers、num_attention_heads、intermediate_size(映射为inner_dim)、hidden_activation(经tf_utils.get_activation转为 gelu 等)、max_position_embeddings(映射为max_sequence_length,决定位置嵌入形状)、embedding_size、initializer_range(截断正态初始化标准差)等,全部由 YAML 的encoder.any字段驱动。
Linformer:低秩线性注意力
Linformer 来自论文 "Linformer: Self-Attention with Linear Complexity",把 K/V 先投影到固定低秩维度再做注意力,把复杂度从 O(L²) 降到 O(L)(linformer_encoder_block.py 的类 docstring 中明确引用了该文与 LRA 基准论文)。配置类 LinformerEncoderConfig 相比 Transformer 新增两个字段:
pad_token_id: int = 0 # pad token 的 id low_rank_features: int = 256 # 低秩投影维度low_rank_features是调参重点:越大信息保留越多、越接近标准注意力,越小越省显存与算力;LRA ListOps 配置采用 32。
MEGA:移动平均门控注意力
MegaEncoder 实现 "Mega: Moving Average Equipped Gated Attention"(见 moving_average_gated_attention.py 中MovingAverageGatedAttention的 docstring),核心思想是用多组指数移动平均(EMA)替代 key 的软加权求和:exponential_moving_average.py 实现MultiHeadEMA层,把过去状态沿时间维递归平滑,从而得到线性于序列长度的注意力。MovingAverageGatedAttention内部由 Q/K/V/Z 四组投影构成(zdim为门控向量维度、hdim为隐状态维度、ndim为 EMA 头数),并对相对位置做可学习偏置(同文件的RelativePositionBias层,形状2*max_positions - 1,前向时裁剪为seq_len × (2*seq_len-1)的相对位置偏置矩阵);激活函数在silu与softmax间二选一(get_activation_fn)。
MegaEncoderConfig 新增字段及默认值(可直接用于对比 YAML 覆盖):
zdim: int = 64 # 门控隐维度 hdim: int = 256 # MEGA 隐状态维度 ndim: int = 16 # EMA 头数 activation: str = 'silu' # 门控激活 bidirectional: bool = False # 是否双向 EMA dropout: float = 0.0 hidden_dropout: float = 0.0MegaEncoder的inner_activation、attention_dropout、max_sequence_length等仍沿用父类字段,mega_encoder.py 中标注"Modified From huggingface/transformers",实现细节可对照该文件与其测试 mega_encoder_test.py。
任务侧:SentencePrediction 与 AAN 双编码器
除 AAN 外的任务复用 NLP 通用 sentence_prediction 任务:编码器输出(末位/池化)接分类头,二分类时评估cls_accuracy与 PR 曲线 AUC,与 YAML 中best_checkpoint_eval_metric: 'cls_accuracy'+best_checkpoint_metric_comp: 'higher'呼应——训练循环会在每次验证后把该指标最优的 checkpoint 导出到best_ckpt子目录。
AAN 任务使用自研的 lra_dual_encoder_task.py 中DualEncoderTask:
build_model()支持从hub_module_url(TF Hub)或本仓库编码器构建 backbone,再包装为 lra_dual_encoder.py 的LRADualEncoder(基于 LaBSE 论文的双编码器结构),use_encoder_pooler为真时分类头直接接编码器池化输出,否则额外插入inner_dim = hidden_size * 2的稠密层;- 损失:
num_classes == 1时为 MSE(回归),否则为 sparse CCE(logits); - 评估指标由
metric_type选择,合法集合为accuracy / f1 / matthews_corrcoef / pearson_spearman_corr,后三者会在reduce_aggregated_logs中把整验证集的 logits 与标签聚拢后用 sklearn/scipy 计算(如 pearson 与 spearman 相关系数取平均),这是accuracy之外的更严格评测口径; initialize()支持从init_checkpoint部分加载预训练编码器权重(status.expect_partial()断言已存在对象匹配)。
数据集路径
README 列出的预打包 TFRecord 位于 GCS bucketgs://model-garden-ucsd-zihan/,使用时需具备该 bucket 的读取权限,或自行按 LRA 原始格式转写 tf_record 后通过params_override指向自定义路径:
| 任务 | 数据路径 |
|---|---|
| ListOps | gs://model-garden-ucsd-zihan/lra_listops_[train/eval/test].tf_record |
| IMDB | gs://model-garden-ucsd-zihan/lra_imdb_[train/eval/test].tf_record |
| IMDB-4096 | gs://model-garden-ucsd-zihan/lra_imdb_4096_[train/eval/test].tf_record |
| AAN | gs://model-garden-ucsd-zihan/lra_aan_[train/eval/test].tf_record |
| CIFAR10 | gs://model-garden-ucsd-zihan/lra_cifar_[train/eval/test].tf_record |
| Pathfinder | gs://model-garden-ucsd-zihan/lra_pathfinder_[train/eval/test].tf_record |
复现要点与常见坑
- 执行目录:README 命令使用
--config_file=../experiments/xxx.yaml的相对路径,意味着命令预期在official/projects/lra/目录下执行(train.py也在该目录);PYTHONPATH需指向仓库根目录以导入official包。 - 数据必须覆盖:所有 YAML 的
input_path均为TODO,漏掉params_override中的数据项会直接以TODO路径去读文件而报错。 - 实验名与配置需匹配:
--experiment=mega/lra_imdb注册的是MegaEncoderConfig,若误配 Transformer 的 YAML 会缺少zdim/hdim/ndim等字段(回落到配置类默认值 64/256/16),反之亦然——实验名前缀决定编码器,YAML 决定超参,两者必须成对选择。 - 分布式:
runtime.distribution_strategy=tpu是 README 示例的默认策略;从 train.py 源码看该字段透传给distribute_utils.get_distribution_strategy,在纯 GPU 环境可覆盖为其他策略,并可配合runtime.mixed_precision_dtype开启混合精度。 - 指标口径:README 中 "Reproduced Acc = 87.55" 是作者用 IMDB-4096 数据集 +
lra_imdb_mega.yaml配置复现的报告值,属于该特定组合的结果,不能外推到其它任务或模型组合;验证阶段的validation_steps: 99999表示跑完整验证集,cls_accuracy即最终口径。 - 扩展新任务:若要在 LRA 框架上加一种新编码器,按既有模式操作即可——定义
XxxEncoderConfig(BertEncoderConfig)与@base_config.bind工厂(参照 linformer.py)、实现编码器(参照 linformer_encoder.py),再到xxx_experiments.py中用@exp_factory.register_config_factory('xxx/lra_<task>')注册五个实验,即可沿用本文全部训练命令模板。
综上,LRA 项目以"实验注册 + YAML 覆盖 + 通用 train_lib"三层解耦的方式,把三种不同注意力机制(全注意力、低秩注意力、EMA 门控注意力)的长程基准训练收敛到同一套命令形态,是研究高效 Transformer 时一个结构清晰、可直接复用的 TF2 参考实现。
【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/models
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考