尧图精选

fairseq 中的 Adaptive Span:基于可学习注意力跨度的 Transformer 语言模型训练与复现指南

🕒 发布时间:2026/9/19 22:39:09 📁 来源:尧图网络
fairseq 中的 Adaptive Span基于可学习注意力跨度的 Transformer 语言模型训练与复现指南【免费下载链接】fairseqFacebook AI Research Sequence-to-Sequence Toolkit written in Python.项目地址: https://gitcode.com/gh_mirrors/fa/fairseqAdaptive Span自适应注意力跨度是一种可学习的自注意力机制每个注意力头可以从数据中学习自己应该关注的上下文长度从而在显著扩展 Transformer 最大上下文范围的同时把内存占用与计算开销控制在可控范围内。本文基于 fairseq 仓库中的 examples/adaptive_span 完整示例讲解该机制的源码实现、Enwik8 数据预处理、12 层模型的训练命令、关键超参数含义以及如何用同一代码库复现 Transformer-XL 基线并完成评估。读完本文你将能够在 fairseq 中完整跑通 Adaptive Span → 训练 → 评估 的实战流程并理解可学习跨度背后的掩码学习原理。一、Adaptive Span 是什么Adaptive Span 由论文Adaptive Attention Span in TransformersarXiv:1905.07799提出发表时在语言建模任务上取得了当时最优的结果。它的核心思想是传统 Transformer 中每个注意力头都会看全部上下文而 Adaptive Span 允许每个头学习一个属于自己的注意力跨度attention span——有的头只需关注最近的少数 token有的头则可以关注很远的 token。这样模型既能利用长距离依赖又不会为所有头都付出满跨度的计算与显存代价。该实现采用 Truncated BPTT截断反向传播技术进行训练这与仓库中的 Transformer-XL 示例 一致长序列被切分成块语言模型按顺序逐块训练模型可以条件化于之前的块作为记忆 cache 传入但梯度只流经当前块。fairseq 中的这一实现尽量保留了原论文官方实现的思路并适配了 fairseq 的模型 / 任务 / 优化器 / 损失函数注册体系。二、仓库实现架构总览整个示例集中在 examples/adaptive_span 目录通过--user-dir examples/adaptive_span加载各文件职责如下文件职责注册名adaptive_span_model_wrapper.py模型包装层注册adaptive_span架构管理记忆 cache 状态register_model(adaptive_span)adaptive_span_model.py核心网络逐层 Transformer、顺序自注意力、可学习跨度—adaptive_span_attention.py可学习软掩码AdaptiveMask与跨度管理AdaptiveSpan—adaptive_span_loss.py语言模型损失 跨度正则项aux lossregister_criterion(adaptive_span_loss)adagrad_with_grad_clip.py带内建梯度裁剪的 Adagrad 优化器register_optimizer(adagrad_with_grad_clip)truncated_bptt_lm_task.pyTruncated BPTT 语言模型任务register_task(truncated_bptt_lm)目录下的 __init__.py 会自动导入目录内所有 Python 模块这也是 fairseq--user-dir插件的标准做法。模型侧的三层结构adaptive_span_model.pyTransformerSeq顶层语言模型包含 token 嵌入in_emb、可学习的位置嵌入参数key_pe形状为1 x d_model//n_head x attn_span、若干TransformerSeqLayer与输出 log-softmax 层提供get_aux_loss()、get_current_max_span()、get_current_avg_span()供损失函数与日志使用。TransformerSeqLayer单层 Transformer包含MultiHeadSeqAttention、两个LayerNorm和一个 ReLU 前馈网络d_model → d_inner → d_model并维护跨块的隐藏层 cache。SeqAttention / MultiHeadSeqAttention顺序自注意力——每个 token 只关注它之前的固定步数不包含自身先对 key/value 做trim_memory裁剪再通过_skew/_unskew张量移位技巧高效计算相对位置注意力。三、可学习跨度的底层原理跨度的学习完全依赖 adaptive_span_attention.py 中的两个模块。1. AdaptiveMask软掩码AdaptiveMask并不直接做硬截断而是构造一个从 1 渐变到 0的软掩码mask self.mask_template.float() self.current_val.float() * self._max_size mask mask / self._ramp_size 1 mask mask.clamp(0, 1)mask_template是一个从1 - max_size到0的等差模板current_val是一个可学习参数nn.Parameter初始值为init_val跨度初始比例掩码在边界处以ramp_size长度的斜坡从 1 平滑过渡到 0因此跨度长度可通过反向传播连续地学习每个更新步之后调用clamp_param()把current_val限制在[0, 1]。2. AdaptiveSpan逐头跨度管理AdaptiveSpan包装掩码并提供两种粒度adapt_span_layerTrue整层共享一个跨度shape(1,)adapt_span_layerFalse默认每个注意力头独立学习跨度shape(n_head, 1, 1)这是论文中最有价值的配置——不同头可以专注不同距离的模式。其关键方法方法作用get_trim_len()计算当前可以从记忆中裁剪掉多少历史步按 64 对齐便于内存管理从而减少计算量trim_memory(query, key, value, key_pe)在真正计算注意力前预先裁剪/补齐 cacheget_cache_size()决定 cache 应保留多长在裁剪基础上预留 64 步缓冲防止跨度后续增长get_loss()返回跨度正则项max_span * mean(current_val)用于抑制跨度无限增长get_current_max_span()/get_current_avg_span()报告当前最大 / 平均跨度供训练日志观测由此可见Adaptive Span 的收益是双重的记忆上cache 长度随学习到的跨度动态收缩计算上过长的历史在trim_memory阶段就被提前丢弃。这正对应 README 中在控制内存占用与计算时间的同时显著扩展最大上下文的表述。四、第 0 步Enwik8 数据预处理示例使用论文官方的预分词pre-tokenizedEnwik8 数据集。下载得到train.txt / valid.txt / test.txt后用 fairseq 的标准预处理命令生成 binarized 数据fairseq-preprocess --only-source --trainpref ~/data/enwik8/train.txt \ --validpref ~/data/enwik8/valid.txt --testpref ~/data/enwik8/test.txt \ --destdir ~/data/enwik8/data-bin/ --joined-dictionary --workers 20参数说明--only-source只处理源侧语言模型只有一个序列流--joined-dictionary训练/验证/测试集共享同一个词典Enwik8 是字符级任务词典极小--workers 20并行预处理进程数可按机器核数调整。预处理产物dict.txt会被 truncated_bptt_lm_task.py 中的Dictionary.load加载TokenBlockDataset按--tokens-per-sample把长流切成定长块。五、第 1 步训练 Adaptive Span 模型README 给出的训练命令遵循原论文的 12 层配置假设 4 块 GPU全局 batch size 为 64 条序列4 × 16在 4 块 V100 上约需 2–3 天CUDA_VISIBLE_DEVICES0,1,2,3 fairseq-train \ --user-dir examples/adaptive_span \ --data ~/data/enwik8/data-bin/ \ --fp16 --fp16-no-flatten-grads --max-update 600000 \ --task truncated_bptt_lm --tokens-per-sample 512 --arch adaptive_span \ --n-layer 12 --d-model 512 --n-head 8 --d-inner 2048 --dropout 0.3 \ --attn-span 8192 --optimizer adagrad_with_grad_clip --adagrad-clip 0.03 \ --validate-interval-updates 1000 \ --lr-scheduler fixed --warmup-updates 32000 --batch-size-valid 32 \ --lr 0.07 --criterion adaptive_span_loss --batch-size 16 --update-freq 1 \ --seed 2 --log-format json --log-interval 25 --aux-loss-scaler 5e-07README 报告该配置验证集可到约 1.05 bpc、测试集约 1.03 bpc相对 Transformer-XL 基线约有 ~0.03 bpc 的提升。几个关键点--fp16-no-flatten-grads该示例的优化器不支持扁平化梯度见 adagrad_with_grad_clip.py 中supports_flat_params返回False因此必须显式关闭--adagrad-clip 0.03启用优化器内部的梯度裁剪而不是 fairseq 全局的--clip-norm--aux-loss-scaler 5e-07跨度正则项的缩放系数调小它会让模型学习到更长的跨度性能可能更好但会相应增加计算/内存单卡训练时把--update-freq设为 4即可累积 4 倍梯度、模拟 4 卡训练。关键超参数速查以下配置由 adaptive_span_model_wrapper.py 中的AdaptiveSpanSmallConfig承接并传给底层模型参数训练示例值配置默认值含义--attn-span81921024最大注意力跨度上下文上限--n-layer128Transformer 层数--d-model512256模型隐藏维度--n-head84注意力头数--d-inner20481024前馈网络隐藏维度--dropout0.30.0注意力/前馈 dropout--aux-loss-scaler5e-072e-06跨度正则损失缩放--adapt-span-ramp未指定32掩码斜坡长度--adapt-span-init未指定0.0跨度初始比例--adapt-span-layer未指定FalseFalse 逐头学习True 整层共享训练日志中会通过 adaptive_span_loss.py 的reduce_metrics额外输出avg_span所有层所有头的平均当前跨度max_span全部头中的最大当前跨度total_loss语言模型损失 跨度正则项loss/ppl除以log(2)后以**比特/字符bpc**为单位报告README 中的 1.05 / 1.03 即 bpc。六、复现 Transformer-XL 基线同一代码库也可以复现 enwik8 上的 Transformer-XL 基线README 报告测试集约 1.06 bpc与原论文结果一致只需把--user-dir切换为 examples/truncated_bptt模型架构改为transformer_xlCUDA_VISIBLE_DEVICES0,1,2,3 fairseq-train \ --user-dir examples/truncated_bptt \ ~/data/enwik8/data-bin/ \ --task truncated_bptt_lm --fp16 --max-update 400000 \ --tokens-per-sample 512 --arch transformer_xl --n-layer 12 \ --d-model 512 --n-head 8 --d-head 64 --d-inner 2048 --dropout 0.1 \ --dropatt 0.0 --mem-len 512 --optimizer adam --clip-norm 0.25 \ --lr-scheduler cosine --warmup-updates 0 \ --lr 0.0 --lr 0.00025 --batch-size 15 \ --update-freq 1 --seed 2 --log-format json --log-interval 25 \ --fp16Transformer-XL 使用的训练设置与 Adaptive Span 形成鲜明对照使用adam--clip-norm 0.25全局梯度裁剪而 Adaptive Span 使用adagrad_with_grad_clip--adagrad-clip内部裁剪使用cosine学习率调度--lr 0.00025而 Adaptive Span 使用fixed调度 --warmup-updates 32000--lr 0.07依赖固定长度的记忆--mem-len 512跨度不可学习。该基线正是衡量 Adaptive Span 收益的对照组两者都用 Truncated BPTT 训练区别仅在于注意力跨度是否可学习。七、第 2 步评估Adaptive Span 评估fairseq-eval-lm ~/data/enwik8/data-bin/ --path model/checkpoint_best.pt \ --user-dir examples/adaptive_span \ --task truncated_bptt_lm --batch-size 8 --tokens-per-sample 512 --gen-subset testTransformer-XL 评估fairseq-eval-lm ~/data/enwik8/data-bin/ --path model/checkpoint_best.pt \ --user-dir examples/truncated_bptt/ --task truncated_bptt_lm --batch-size 8 \ --tokens-per-sample 80 \ --model-overrides {mem_len:2100,clamp_len:820,same_length:True} \ --gen-subset valid评估设置的两点说明README 原文强调Adaptive Span 训练时看到 512 token 上下文--tokens-per-sample512评估 batch size 为 8这些设置与论文实验脚本保持一致Transformer-XL 评估时通过--model-overrides把记忆长度扩展到 2100、并设置clamp_len与same_length——这是 Transformer-XL 评估的标准做法因为评估时可以使用比训练时更长的记忆与 examples/truncated_bptt 示例WikiText-103 上用mem_len:640的评估思路一致。另外注意评估语言模型时以整个序列为粒度推进隐藏层 cache 由模型包装器 adaptive_span_model_wrapper.py 的init_hid_cache按各层get_cache_size()动态初始化并逐块续传验证/测试过程中无需重新训练即可复现 README 报告的结果。八、延伸如何切换跨度学习粒度如果你想实验整层共享跨度与逐头独立跨度的差异只需在训练命令中追加--adapt-span-layer True从 adaptive_span_attention.py 的构造函数可以看到adapt_span_layerTrue时使用shape(1,)的单一可学习参数False默认时使用shape(n_head, 1, 1)每个头独立。前者更节省参数、更稳定后者表达能力更强、通常性能更好——这也正是论文的核心发现之一不同注意力头会自发学习出不同长度的关注范围。九、总结通过本指南你可以在 fairseq 中完整复现 Adaptive Span 在 Enwik8 上的语言建模实验数据预分词 Enwik8 →fairseq-preprocess生成 binarized 数据训练--user-dir examples/adaptive_span--arch adaptive_span--criterion adaptive_span_loss--optimizer adagrad_with_grad_clip12 层 / 8192 最大跨度配置约 2–3 天4×V100达到 1.05valid/ 1.03testbpc对比同一代码库切换examples/truncated_bptt可复现 1.06 bpc 的 Transformer-XL 基线评估fairseq-eval-lm--task truncated_bptt_lm注意 Transformer-XL 需--model-overrides扩大记忆长度。从源码层面看Adaptive Span 的成功来自两个精巧设计软掩码让跨度长度可以通过反向传播平滑学习adaptive_span_attention.py动态 cache 裁剪让长跨度模型的显存与计算开销随学习结果自适应收缩。若需进一步调整超参数组合可参照论文实验脚本中对应的 sweep 配置在本文给出的命令骨架基础上替换数值即可。【免费下载链接】fairseqFacebook AI Research Sequence-to-Sequence Toolkit written in Python.项目地址: https://gitcode.com/gh_mirrors/fa/fairseq创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
上一篇/下一篇内容由系统自动关联 返回资讯列表 →