ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

Fairseq Tasks 任务系统深度解析:从翻译、语言建模到自定义任务开发

2026/9/19 3:19:07 拓冰建站 浏览量
Fairseq Tasks 任务系统深度解析:从翻译、语言建模到自定义任务开发 Fairseq Tasks 任务系统深度解析从翻译、语言建模到自定义任务开发【免费下载链接】fairseqFacebook AI Research Sequence-to-Sequence Toolkit written in Python.项目地址: https://gitcode.com/gh_mirrors/fa/fairseq导读Tasks 是 fairseqFacebook AI Research 的序列到序列工具包中负责数据与模型之间的编排层的核心抽象它存储词典Dictionary、加载并迭代数据集、构建模型与损失函数Model/Criterion并计算训练损失。本文以 docs/tasks.rst 为主体结合 fairseq/tasks/ 下的源码实现系统讲解 Task 的选择机制、生命周期调用链、翻译与语言建模两个内置任务的完整配置以及如何通过register_task注册自定义任务帮助你掌握 fairseq 训练流程的枢纽环节。Task 是什么fairseq 训练流水线的编排中枢按照 docs/tasks.rst 的官方定义Task 承担四项核心职责存储词典维护源语言/目标语言或单语的Dictionary即词表与 token ↔ id 的映射提供数据加载与迭代的助手方法按 split如train、valid、test加载FairseqDataset并生成可复用的批次迭代器EpochBatchIterator初始化 Model 与 Criterion根据配置构建模型和损失函数计算损失对给定的 mini-batch 调用 criterion 计算loss、sample_size与logging_output。在基类 fairseq/tasks/fairseq_task.py#L50-L64 的 docstring 中可以看到Task 还具有有限状态性limited statefulness需要随 checkpoint 保存/恢复的状态必须存放在self.state一个StatefulContainer中。例如self.state.add_factory(dictionary, self.load_dictionary) print(self.state.dictionary) # 触发 self.load_dictionary()结果会被缓存这样设计是为了在加载 checkpoint 时可以先实例化 Task再正确地重建其内部状态。通过 --task 命令行参数选择任务Task 通过--task命令行参数进行选择。fairseq-cli下所有入口训练、生成、交互、评估都会在启动时调用fairseq.tasks.setup_task(cfg.task)来实例化任务例如 fairseq_cli/train.py#L87、fairseq_cli/generate.py#L83、fairseq_cli/interactive.py#L140、fairseq_cli/eval_lm.py#L249。选定任务后该任务可以额外暴露自己的命令行参数用于进一步配置例如翻译任务的--source-lang、语言建模任务的--tokens-per-sample。Task 的完整生命周期一个端到端调用示例docs/tasks.rst给出了最核心的编程式使用范例展示了 Task 在一个训练迭代中的全部关键方法# 1. 初始化任务例如加载词典 task fairseq.tasks.setup_task(args) # 2. 构建模型和损失函数 model task.build_model(args) criterion task.build_criterion(args) # 3. 加载数据集 task.load_dataset(train) task.load_dataset(valid) # 4. 遍历数据 mini-batch batch_itr task.get_batch_iterator( task.dataset(train), max_tokens4096, ) for batch in batch_itr: # 5. 计算损失 loss, sample_size, logging_output task.get_loss( model, criterion, batch, ) loss.backward()需要说明的是示例中的get_loss是文档中简化的示意方法名在真实训练流程中fairseq/trainer.py 调用的是task.train_step(sample, model, criterion, optimizer, update_num)见 fairseq/trainer.py#L843。其核心逻辑位于 fairseq/tasks/fairseq_task.py#L505-L537def train_step(self, sample, model, criterion, optimizer, update_num, ignore_gradFalse): model.train() model.set_num_updates(update_num) with torch.autograd.profiler.record_function(forward): with torch.cuda.amp.autocast(enabled(isinstance(optimizer, AMPOptimizer))): loss, sample_size, logging_output criterion(model, sample) if ignore_grad: loss * 0 with torch.autograd.profiler.record_function(backward): optimizer.backward(loss) return loss, sample_size, logging_output对应的验证路径则是valid_stepfairseq/tasks/fairseq_task.py#L539-L543它在torch.no_grad()下完成前向并返回同样的三元组。每个方法职责如下方法职责基类默认行为setup_task(cls, cfg)类方法解析配置、加载词典返回任务实例直接cls(cfg)load_dataset(split)加载指定 split 的FairseqDataset存入self.datasets抛NotImplementedError必须由子类实现dataset(split)取出已加载的 split 数据集校验其为FairseqDataset从self.datasets中读取get_batch_iterator(...)构造EpochBatchIterator支持按 token 数/句数分桶、分片、多进程内置完整实现build_model(cfg)调用models.build_model并做标量量化内置实现build_criterion(cfg)调用criterions.build_criterion内置实现train_step(...)前向 反向返回(loss, sample_size, logging_output)内置实现valid_step(...)无梯度前向验证内置实现max_positions()返回任务允许的最大输入长度返回Nonesource_dictionary/target_dictionary源/目标词典属性返回Noneget_batch_iterator 的底层逻辑get_batch_iteratorfairseq/tasks/fairseq_task.py#L208-L340是数据流水线的关键其内部做了四件事排序与索引在numpy_seed(seed epoch)的确定性随机种子下调用dataset.ordered_indices()得到按样本长度排序的索引过滤超长样本若指定了max_positions调用filter_indices_by_size过滤超出模型长度限制的样本默认会抛异常可通过ignore_invalid_inputsTrue或--skip-invalid-size-inputs-valid-test跳过构造 mini-batch调用dataset.batch_by_size按max_tokens每批最大 token 数与max_sentences每批最大句子数约束分桶required_batch_size_multiple可要求 batch size 为 N 的倍数包装为迭代器构造iterators.EpochBatchIterator支持num_shards/shard_id数据分片分布式、num_workers子进程加载、data_buffer_size预取缓冲、skip_remainder_batch丢弃尾批等。同时它遵循can_reuse_epoch_itr策略当数据集允许跨 epoch 复用迭代器时会把迭代器缓存到self.dataset_to_epoch_iter避免每个 epoch 重新建批次可通过disable_iterator_cacheTrue或配置rebuild_batches关闭。TranslationTask翻译任务详解docs/tasks.rst中通过autoclass引用了 fairseq/tasks/translation.py 中的TranslationTask。它负责从源语言到目标语言的翻译训练与推理兼容fairseq-train、fairseq-generate和fairseq-interactive三个 CLI 入口。词典加载与语言对推断TranslationTask.setup_taskfairseq/tasks/translation.py#L290-L321在初始化阶段完成通过utils.split_paths(cfg.data)拆分数据目录列表自动推断语言对当--source-lang/--target-lang未显式给出时调用data_utils.infer_language_pair从第一个数据目录中的dict.{lang}.txt文件名推断src-tgt分别加载dict.{source_lang}.txt与dict.{target_lang}.txt两个词典并断言两者的 pad/eos/unk 索引一致日志输出各语言词典的词汇量。TranslationConfig 核心参数TranslationTask通过 dataclassTranslationConfigfairseq/tasks/translation.py#L174-L265暴露命令行参数参数默认值说明--datacfg.data无冒号分隔的数据目录列表训练时按 epoch 轮转round-robinvalid/test 始终取第一个目录--source-lang/-s自动推断源语言--target-lang/-t自动推断目标语言--load-alignmentsFalse加载二值化的词对齐数据--left-pad-sourceTrue源序列左侧 padding--left-pad-targetFalse目标序列左侧 padding--max-source-positions1024源序列最大 token 数--max-target-positions1024目标序列最大 token 数--upsample-primary-1主数据集第一个的上采样倍数用于数据混合--truncate-sourceFalse将源序列截断到max-source-positions--num-batch-buckets0若 0将源/目标长度分桶并按桶 paddingTPU 上减少编译次数--eval-bleuFalse验证时计算 BLEU 分数--eval-bleu-detokspace计算 BLEU 前的去分词方式如mosesspace表示不去分词--eval-bleu-args{}BLEU 评测的生成参数 JSON如{beam: 4, lenpen: 0.6}--eval-tokenized-bleuFalse使用 tokenized BLEU 而非 sacrebleu--eval-bleu-remove-bpe无计算 BLEU 前移除 BPE--eval-bleu-print-samplesFalse验证时打印生成的样本其中不少字段通过II(...)Interpolation引用了全局配置例如train_subset来自dataset.train_subset、dataset_impl来自dataset.dataset_impl、required_seq_len_multiple来自dataset.required_seq_len_multiple——这正是 Hydra 配置系统在 fairseq 中的典型用法见 fairseq/config/config.yaml。数据集加载load_langpair_datasetTranslationTask.load_datasetfairseq/tasks/translation.py#L323-L358会按data_path paths[(epoch - 1) % len(paths)]轮转选择数据目录非训练 split 只用第一个目录然后委托给模块级函数load_langpair_datasetfairseq/tasks/translation.py#L40-L171尝试按{split}.{src}-{tgt}.{lang}的命名规则定位索引数据集支持combineTrue时合并多个分片可选地对源序列做TruncateDatasetAppendTokenDataset截断、PrependTokenDataset加 BOS、AppendTokenDataset追加[lang]语言标记多语言翻译场景若--load-alignments且存在{split}.align.{src}-{tgt}文件则加载对齐数据最终组装为LanguagePairDataset定义于 fairseq/data/language_pair_dataset.py。max_positions()返回(max_source_positions, max_target_positions)二元组source_dictionary/target_dictionary属性分别返回src_dict/tgt_dict。验证时的 BLEU 评测当--eval-bleu开启时build_model会额外构建 tokenizer 与SequenceGeneratorfairseq/tasks/translation.py#L369-L381valid_step在正常验证之外调用_inference_with_bleu生成译文并与参考译文计算 sacrebleu 分数fairseq/tasks/translation.py#L383-L395把 4-gram 的 counts/totals 拆分到独立日志字段以支持跨 worker 的高效聚合reduce_metrics再通过metrics.log_derived(bleu, ...)汇总出最终 BLEUfairseq/tasks/translation.py#L397-L448。LanguageModelingTask语言建模任务详解docs/tasks.rst中通过autoclass引用了 fairseq/tasks/language_modeling.py 中的LanguageModelingTask。它训练语言模型兼容fairseq-train、fairseq-generate、fairseq-interactive和fairseq-eval-lm。与TranslationTask继承FairseqTask不同它是LegacyFairseqTask使用 argparseNamespace而非 dataclass 配置的实例。预测目标targetsLanguageModelingTask支持三类预测目标fairseq/tasks/language_modeling.py#L139-L146future预测下一个 token标准语言建模默认self自目标--self-target用于去噪训练等past预测过去 token--past-target。当三个开关均未开启时默认退化为[future]。build_model会校验所选目标是否被模型支持若target not in model.supported_targets则抛出ValueErrorfairseq/tasks/language_modeling.py#L190-L198。LanguageModelingConfig 核心参数LanguageModelingTask的 dataclass 配置LanguageModelingConfigfairseq/tasks/language_modeling.py#L41-L106暴露如下参数参数默认值说明--data无数据目录路径--sample-break-modenone样本切分方式none填满tokens-per-sample、complete仅在句末切分可含多句、complete_doc尊重文档边界、eos每样本仅一句--tokens-per-sample1024每个 LM 样本的最大 token 数--output-dictionary-size-1限制输出词典大小-10 时使用TruncatedDictionary--self-targetFalse包含自目标--future-targetFalse包含未来目标--past-targetFalse包含过去目标--add-bos-tokenFalse在输入前加 BOSs--max-target-positionsNone目标序列最大 token 数--shorten-methodnone对超长序列的截断策略none/truncate/random_crop--shorten-data-split-list应用截断的 split 列表逗号分隔默认全部--pad-to-fixed-lengthFalse固定长度 padding--pad-to-fixed-bszFalse固定 batch size padding单语数据集组装LanguageModelingTask.load_datasetfairseq/tasks/language_modeling.py#L200-L267的组装链路清晰展示了 fairseq 数据 wrapper 的组合式设计load_indexed_dataset原始索引数据 → maybe_shorten_dataset按 shorten-method 截断/裁剪超长序列 → TokenBlockDataset按 tokens-per-sample 与 sample-break-mode 切块include_targetsTrue → MonolingualDataset包装 src_vocab/tgt_vocab、targets、BOS/padding 逻辑推理时build_dataset_for_inferencefairseq/tasks/language_modeling.py#L269-L312会构造一个NestedDictionaryDataset源序列前加 BOS或 EOS取决于--add-bos-token目标序列后补 pad并通过StripTokenDataset去掉序列末尾的 EOS。inference_step则把src_tokens作为prefix_tokens传给SequenceGenerator进行条件续写fairseq/tasks/language_modeling.py#L314-L338。此外eval_lm_dataloaderfairseq/tasks/language_modeling.py#L340-L371为fairseq-eval-lm提供了专用的评估迭代器支持context_window上下文窗口通过LMContextWindowDataset实现确保每个被评估的 token 在可能的情况下都能访问到足够大的上下文这也是困惑度perplexity评估数据流的关键所在。添加新任务注册机制与 FairseqTask 接口docs/tasks.rst明确指出扩展方式是register_task装饰器 继承FairseqTask。注册机制的实现在 fairseq/tasks/init.py#L50-L103它维护三个全局结构TASK_REGISTRY任务名 → 任务类TASK_DATACLASS_REGISTRY任务名 → dataclass 配置类TASK_CLASS_NAMES已注册的类名集合防止重复注册同名类。注册流程做了三重校验任务名不得重复类必须是FairseqTask的子类否则抛Task (...) must extend FairseqTask类名不得重复。若提供了 dataclass还会通过 Hydra 的ConfigStore把配置节点注册到task配置组provider 为fairseq从而使该任务可在 Hydra 配置中按task任务名选择。register_task的文档示例fairseq/tasks/init.py#L50-L64register_task(classification) class ClassificationTask(FairseqTask): (...)任务目录的自动导入fairseq/tasks/__init__.py的末尾通过import_tasks(tasks_dir, fairseq.tasks)fairseq/tasks/init.py#L110-L133自动导入tasks/目录下所有不以_或.开头的 Python 文件/子目录因此新增任务只需把文件放进该目录即可被自动发现并注册。该函数还会为每个已注册任务生成一个用于 Sphinx 文档的{task_name}_parser包含--task与任务自定义参数。setup_task 的分发逻辑fairseq.tasks.setup_taskfairseq/tasks/init.py#L24-L47根据配置形态分发若cfg.task是字符串legacy argparse 模式直接从TASK_REGISTRY查类若该任务注册了 dataclass则调用dc.from_namespace(cfg)把 Namespace 转成 dataclass若cfg是 Hydra 的DictConfig读取cfg._name作为任务名通过merge_with_parent把任务默认配置与用户配置合并from_checkpoint时允许移除缺失字段。两种路径最终都调用task_cls.setup_task(cfg, **kwargs)返回任务实例。最小自定义任务模板综合 fairseq/tasks/fairseq_task.py 的接口约束一个最小可注册的自定义任务如下from fairseq.tasks import FairseqTask, register_task register_task(my_task) class MyTask(FairseqTask): classmethod def setup_task(cls, cfg, **kwargs): # 1. 解析 cfg加载词典 # 2. 返回任务实例 return cls(cfg) def load_dataset(self, split, epoch1, combineFalse, **kwargs): # 构造 FairseqDataset 并存入 self.datasets[split] raise NotImplementedError property def source_dictionary(self): # 返回源词典 return self.src_dict一个完整的参考实现可以对照 fairseq/tasks/translation.py继承FairseqTask、dataclass 配置范式与 fairseq/tasks/language_modeling.pylegacy Namespace 范式来编写。更贴近实际的自定义任务还包括 fairseq/tasks/denoising.py、fairseq/tasks/masked_lm.py、fairseq/tasks/sentence_prediction.py 等它们展示了不同数据形态噪声注入、掩码、分类标签下如何实现load_dataset与词典属性。任务与训练主循环的协作在真实训练中Task 并非被单独调用而是被 fairseq/trainer.py 的Trainer统一驱动数据侧trainer.get_train_iterator/get_valid_iterator内部调用task.load_dataset带load_datasetTrue标志与task.get_batch_iterator见 fairseq/trainer.py#L697-L745训练侧每个 step 调用task.train_step(sample, model, criterion, optimizer, update_num)完成前向反向fairseq/trainer.py#L843验证侧调用task.valid_step并最终由task.reduce_metrics聚合各 worker 的日志指标如wpb词数/批、wps词数/秒、bsz批大小见 fairseq/tasks/fairseq_task.py#L579-L613。因此Task 接口的设计把数据形态差异翻译的平行语料、LM 的连续文本块、掩码 LM 的噪声注入等封装在了统一接口之下而 Trainer、优化器、Checkpoint 管理等基础设施无需关心具体任务类型——这正是 fairseq 能同时支持翻译、语言建模、语音、多模态等众多任务的关键架构决策。当你在 fairseq/tasks/ 目录中看到translation.py、language_modeling.py、audio_pretraining.py、hubert_pretraining.py、speech_to_text.py、speech_to_speech.py等二十余个任务文件时它们全部遵循了本文所述的生命周期与注册规范。【免费下载链接】fairseqFacebook AI Research Sequence-to-Sequence Toolkit written in Python.项目地址: https://gitcode.com/gh_mirrors/fa/fairseq创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考