TensorFlow Models LRA 项目实战训练 MEGA、Transformer 与 Linformer 长程序列建模基线【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/modelsofficial/projects/lra目录是 TensorFlow Model Gardenmodels 仓库中针对 Long Range ArenaLRA基准的 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_fileYAML 配置、--params_override参数覆盖、--tpu、--model_dir、--mode。以下三组命令完整继承自 README。在 ListOps 上训练 TransformerTRAIN_DATAtask.train_data.input_pathgs://model-garden-ucsd-zihan/lra_listops_train.tf_record,task.validation_data.input_pathgs://model-garden-ucsd-zihan/lra_listops_eval.tf_record PYTHONPATH[/PATH/TO/MODEL_GARDEN] \ python3 train.py \ --experimenttransformer/lra_listops \ --config_file../experiments/lra_listops.yaml \ --params_override${TRAIN_DATA},runtime.distribution_strategytpu \ --tpulocal \ --model_dir[OUTPUT_DIR] \ --modetrain_and_eval在 ListOps 上训练 LinformerTRAIN_DATAtask.train_data.input_pathgs://model-garden-ucsd-zihan/lra_listops_train.tf_record,task.validation_data.input_pathgs://model-garden-ucsd-zihan/lra_listops_eval.tf_record PYTHONPATH[/PATH/TO/MODEL_GARDEN] \ python3 train.py \ --experimentlinformer/lra_listops \ --config_file../experiments/lra_listops_linformer.yaml \ --params_override${TRAIN_DATA},runtime.distribution_strategytpu \ --tpulocal \ --model_dir[OUTPUT_DIR] \ --modetrain_and_eval在 TextIMDB-4096上训练 MEGAREADME 标注该配置为 Reproduced Acc 87.55即作者在该数据集上复现报告 87.55 的准确率TRAIN_DATAtask.train_data.input_pathgs://model-garden-ucsd-zihan/lra_imdb_4096_train.tf_record,task.validation_data.input_pathgs://model-garden-ucsd-zihan/lra_imdb_4096_eval.tf_record PYTHONPATH[/PATH/TO/MODEL_GARDEN] \ python3 train.py \ --experimentmega/lra_imdb \ --config_file../experiments/lra_imdb_mega.yaml \ --params_override${TRAIN_DATA},runtime.distribution_strategytpu \ --tpulocal \ --model_dir[OUTPUT_DIR] \ --modetrain_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_overridegin 风格的keyvalue覆盖列表此处覆盖数据路径并指定runtime.distribution_strategytpugin.parse_config_files_and_bindings(FLAGS.gin_file, FLAGS.gin_params)train.py 中gin_params与gin_file合并解析--tpuTPU 地址local表示本地 TPUdistribute_utils.get_distribution_strategy(..., tpu_addressparams.runtime.tpu)--model_dircheckpoint 与序列化配置的输出目录train_utils.serialize_config(params, model_dir)、train_lib.run_experiment--modetrain_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 上收益来自 float16TPU 上来自 bfloat16loss_scale仅在 float16 下生效→ 构建DistributionStrategy→ 在策略作用域内task_factory.get_task构建任务 →train_lib.run_experiment执行训练/评估 → 最后train_utils.save_gin_config落盘 gin 配置。命令行 flag 本身定义于 official/common/flags.pytfm_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_aanlinformer_experiments.py 与 mega_experiments.py 同构分别注册linformer/...与mega/...前缀的同名任务。Linformer 与 MEGA 两个实验文件结构完全对称唯一差异在于编码器配置类与默认学习率Transformer/Linformer 的_TRAINER默认初始学习率3e-5adamwweight_decay_rate0.01排除LayerNorm/layer_norm/bias的权重衰减MEGA 的_TRAINER默认初始学习率1e-7mega_experiments.py 中显式配置实际训练时几乎都依赖 YAML 覆盖例如 IMDB-MEGA 配置中覆盖为0.004。任务与数据加载器的组合也有规律ListOps / IMDB / CIFAR / Pathfinder 四个任务复用 NLP 通用的sentence_prediction.SentencePredictionConfigSentencePredictionDataConfig即双句预测框架下的序列分类任务AANApproximate Area Navigation任务则使用 LRA 项目自研的DualEncoderConfigDualEncoderDataConfig见下文双编码器任务小节。实验配置文件详解experiments/ 目录README 指出所有实验配置位于experiments子文件夹共 15 份 YAML命名为lra_任务[_模型].yaml无后缀为 Transformer_linformer、_mega为对应变体。所有 YAML 均为tasktrainer两级结构input_path留空为TODO运行时必须像 README 命令那样用params_override注入。下面按任务归纳关键默认值数值取自各 YAML 文件本身ListOps10 类序列长度 2000lra_listops.yamlTransformernum_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→ 0decay_steps: 5000power: 0.5 1000 步 warmup、AdamW。lra_listops_linformer.yaml 与 Transformer 版几乎相同唯一新增项是low_rank_features: 32——即 Linformer 低秩投影维度这是实现线性复杂度注意力的核心超参LinformerEncoderConfig默认值为 256见 linformer.py此处按 LRA 论文推荐压到 32。IMDB / Text2 类情感序列长度 1000lra_imdb.yamlTransformerembedding_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.yamlMEGA在 Transformer 结构参数之外额外包含 MEGA 专属超参zdim: 64门控隐维度、hdim: 256隐层维度、ndim: 16EMA 头数、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。AAN2 类序列长度 4000双编码器任务lra_aan.yamlnum_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*.yamlCIFAR-10 图像转序列任务与lra_pathfinder*.yaml路径规划任务同样按Transformer 默认 Linformer/MEGA 变体三份一组的方式组织可按任务名直接查文件。三种编码器配置类与关键实现三种模型都继承BertEncoderConfigofficial/nlp/configs/encoders.py通过base_config.bind(...)将配置 dataclass 绑定到编码器实例化函数gin 可绑定这是 Model Garden 统一的配置即对象模式。Transformer 基线transformer.py 中TransformerEncoderConfig未新增字段get_encoder()直接构造 transformer_encoder.py 的TransformerEncodervocab_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_accuracybest_checkpoint_metric_comp: higher呼应——训练循环会在每次验证后把该指标最优的 checkpoint 导出到best_ckpt子目录。AAN 任务使用自研的 lra_dual_encoder_task.py 中DualEncoderTaskbuild_model()支持从hub_module_urlTF Hub或本仓库编码器构建 backbone再包装为 lra_dual_encoder.py 的LRADualEncoder基于 LaBSE 论文的双编码器结构use_encoder_pooler为真时分类头直接接编码器池化输出否则额外插入inner_dim hidden_size * 2的稠密层损失num_classes 1时为 MSE回归否则为 sparse CCElogits评估指标由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指向自定义路径任务数据路径ListOpsgs://model-garden-ucsd-zihan/lra_listops_[train/eval/test].tf_recordIMDBgs://model-garden-ucsd-zihan/lra_imdb_[train/eval/test].tf_recordIMDB-4096gs://model-garden-ucsd-zihan/lra_imdb_4096_[train/eval/test].tf_recordAANgs://model-garden-ucsd-zihan/lra_aan_[train/eval/test].tf_recordCIFAR10gs://model-garden-ucsd-zihan/lra_cifar_[train/eval/test].tf_recordPathfindergs://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路径去读文件而报错。实验名与配置需匹配--experimentmega/lra_imdb注册的是MegaEncoderConfig若误配 Transformer 的 YAML 会缺少zdim/hdim/ndim等字段回落到配置类默认值 64/256/16反之亦然——实验名前缀决定编码器YAML 决定超参两者必须成对选择。分布式runtime.distribution_strategytpu是 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),仅供参考