
在 fairseq 中使用截断 BPTT 训练 Transformer-XL 长序列语言模型从任务重写到多卡评测全流程【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm导读本篇文章基于 decoding/IAD/fairseq/examples/truncated_bptt 目录下的官方示例文档与配套实现系统讲解如何在 fairseq 中实现并运行截断反向传播Truncated Backpropagation Through TimeTruncated BPTT面对超长序列将数据按顺序切成 chunk 训练语言模型、让梯度只流过当前 chunk同时让模型通过显存记忆memory条件化于前序 chunk。你将掌握 Truncated BPTT 的核心原理、通过覆写FairseqTask::get_batch_iterator与自定义数据集实现顺序迭代的技巧并得到一个可直接复现的 WikiText-103 Transformer-XL 训练与评测实战方案。一、为什么需要 Truncated BPTT超越固定长度上下文的语言建模标准的 BPTT 要求把整个序列一次性送入模型反向传播时梯度跨越全部时间步。对于 WikiText-103 这类文档级语料单篇文章动辄数千 token直接整段训练既显存爆炸又让梯度过长、训练不稳定。Truncated BPTT 的思路是把一条长序列顺序切分成若干 chunk模型按顺序逐个 chunk 训练。语言模型可以条件化condition于前面已经看过的 chunk——这些历史信息以「记忆」形式保留在模型中——但梯度只流经当前 chunk前序 chunk 的隐藏状态不参与本轮反向传播。这一技术正是 Transformer-XL: Attentive Language Models Beyond a Fixed-Length ContextTransformer-XL 论文的核心训练基础该论文发表时在语言建模任务上取得了当时最优的结果。Transformer-XL 通过引入分段级循环机制segment-level recurrence和相对位置编码把固定长度上下文的限制打破训练时上文的隐藏状态被缓存为 memory 供当前段使用推理时记忆可以在多个段之间持续传递。在 fairseq 中实现的难点在 fairseq 中实现 Truncated BPTT 并不直接fairseq 默认的 batch 采样会随机打乱数据而 Truncated BPTT 要求严格按原始顺序迭代数据否则上一 chunk 与当前 chunk 相邻这一前提就失效了。示例文档明确指出实现策略覆写FairseqTask::get_batch_iterator让数据迭代完全顺序化并且同时支持 batching 与多 GPU数据并行训练——这正是本示例最有价值的工程部分。二、示例代码结构两个核心文件整个示例由三个文件组成均位于 decoding/IAD/fairseq/examples/truncated_bptt文件作用truncated_bptt_lm_task.py注册truncated_bptt_lm任务数据加载、顺序 batch 迭代、collate、评测 dataloadertransformer_xl_model.py注册transformer_xl模型封装 HuggingFaceTransfoXLLMHeadModel管理 memory 状态init.py导入以上两个模块供--user-dir机制发现使用方式训练/评测命令中通过--user-dir examples/truncated_bptt把该目录注册进 fairseq再通过--task truncated_bptt_lm和--arch transformer_xl引用其中实现。2.1 任务侧TruncatedBPTTLMTask的顺序迭代实现truncated_bptt_lm_task.py 中的TruncatedBPTTLMTask通过register_task(truncated_bptt_lm, dataclassTruncatedBPTTLMConfig)注册。任务配置类TruncatedBPTTLMConfig的关键字段如下配置项默认值说明data必填???数据目录路径tokens_per_sample1024每个样本chunk的最大 token 数batch_size取自dataset.batch_size每张 GPU 每个 forward 处理的序列数max_target_positions取task.tokens_per_sample位置嵌入上限默认与 chunk 长度一致data_parallel_rank/data_parallel_size自动填充未提供时从torch.distributed推断单卡时为 0/1实现要点逐一解读数据加载与分块。load_dataset通过data_utils.load_indexed_dataset读取已二值化的数据行为类似于open(split_path).readlines()再用TokenBlockDataset以block_sizetokens_per_sample、break_modenone将整条数据流切块注释明确指出这等价于data.view(-1).split(tokens_per_sample)——即不按句子边界、纯顺序硬切这正是 BPTT 语义。顺序迭代器。get_batch_iterator返回iterators.EpochBatchIterator其中两个设置至关重要batch_sampler[[i] for i in range(len(dataset))]——注释说明我们不用 EpochBatchIterator 的 batching 功能dataset 里每一项本身就是完整的一个 batch即数据集预先已按 batch 组织好disable_shufflingTrue——显式关闭 shuffle保证 chunk 顺序。数据并行分片。TruncatedBPTTDataset的batchify把全部 chunk 均匀切分成bsz_per_shard * num_shards份num_shards即 GPU 数每张卡通过shard_id * bsz_per_shard : (shard_id1) * bsz_per_shard拿到自己的那几条序列。代码注释用 16 个 item、bsz_per_shard2、num_shards3的例子演示了索引如何按 GPU 分片确保每个 GPU 都按原始顺序处理自己子集中的 chunk从而支持多卡数据并行训练。collate 与 target 构造。_collate_fn断言len(items) 1每个 item 就是整批对B x T的输入做collate_tokens后把输入右移一位作为 targettarget pad(item[:, 1:], ...)即标准的下一 token 预测next-token prediction末位补 pad。评测路径。eval_lm_dataloader明确指出 Transformer-XL 不需要--context-window而是通过--model-overrides {mem_len:42}控制上下文记忆长度build_dataset_for_inference按eos断句并在生成时去掉结尾 eos 以支持前缀生成。2.2 模型侧TransformerXLLanguageModel与 memory 管理transformer_xl_model.py 中的TransformerXLLanguageModel通过register_model(transformer_xl, dataclassTransformerXLConfig)注册内部直接复用 HuggingFacetransformers的TransfoXLLMHeadModel兼容新旧版本 import 路径。模型配置TransformerXLConfig默认值均取自原始 Transformer-XL 代码配置项默认值说明cutoffs[20000, 40000, 200000]自适应 softmax 的 cutoff 点超过词表大小的会被自动剔除d_model500隐藏维度n_head10注意力头数d_head50每头维度d_inner1000FFN 中间维度div_val1自适应嵌入的维度缩放因子n_layer12层数mem_len0memory 长度0 表示不记忆clamp_len-1位置编码截断上限-1 表示不截断same_lengthFalse是否让每位置只看到相同长度的上下文dropout/dropatt0.0普通 dropout / 注意力 dropoutcheckpoint_activationsFalse是否对每层使用激活检查点gradient checkpointing省显存模型侧有两个值得注意的工程细节HuggingFace bug 规避ProjectedAdaptiveLogSoftmax会把None塞进nn.ParameterListPyTorch 不支持示例代码在检测到out_projs参数全为None时将其替换为支持None的nn.ModuleList。激活检查点checkpoint_activationsTrue时对每一层transformer.layers[i]包装checkpoint_wrapper用计算换显存。memory 状态管理是模型的核心训练时incremental_state is Noneforward 后把输出里的 mems 存入self._mems供下一个 chunk 使用——这正是模型条件化于前序 chunk的落点推理时incremental_state非空只取src_tokens[:, -1:]最近一个 token并借助get/set_incremental_state保存 mems支持增量生成reorder_incremental_state在 beam search 等导致输入顺序变化的场景下用mems_i.index_select(1, new_order)对 memory 做同步重排。max_positions()返回max_target_positions保证位置嵌入上限与 chunk 长度匹配。三、环境准备WikiText-103 数据预处理按原文档指引预处理请参考 语言建模 README。完整流程如下1下载并解压原始语料运行 prepare-wikitext-103.sh仓库中真实存在cd examples/language_model/ bash prepare-wikitext-103.sh cd ../..该脚本从官方地址下载wikitext-103-v1.zip若文件已存在则跳过下载再按后缀解压支持.tgz/.tar/.zip得到wiki.train.tokens、wiki.valid.tokens、wiki.test.tokens。2额外 Python 依赖预处理还需要来自语言建模 READMEpip install fastBPE sacremoses3二值化TEXTexamples/language_model/wikitext-103 fairseq-preprocess \ --only-source \ --trainpref $TEXT/wiki.train.tokens \ --validpref $TEXT/wiki.valid.tokens \ --testpref $TEXT/wiki.test.tokens \ --destdir>CUDA_VISIBLE_DEVICES0,1,2,3 fairseq-train \ --user-dir examples/truncated_bptt \ >fairseq-eval-lm contenteditable="false">【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考