ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

fairseq Tasks 任务体系详解:以 InfoXLM 仓库为例掌握 fairseq 的任务注册、数据集加载与损失计算全流程

fairseq Tasks 任务体系详解:以 InfoXLM 仓库为例掌握 fairseq 的任务注册、数据集加载与损失计算全流程 fairseq Tasks 任务体系详解以 InfoXLM 仓库为例掌握 fairseq 的任务注册、数据集加载与损失计算全流程【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm本文基于 InfoXLM 仓库中内嵌的 fairseq 文档 tasks.rst 展开系统讲解 fairseq “任务Task”这一核心抽象Task 负责持有词典Dictionary、提供数据集的加载与迭代辅助方法、负责构建模型与损失函数Criterion并计算损失。读完后你将能够理解--task命令行参数背后的注册与分发机制掌握setup_task → build_model / build_criterion → load_dataset → get_batch_iterator → train_step的完整调用链并知道如何基于register_task装饰器为自己的项目扩展新任务。Task 是什么fairseq 中数据集与训练流程的“中枢”官方文档对 Task 的定义非常凝练引自 tasks.rstTasks store dictionaries and provide helpers for loading/iterating over Datasets, initializing the Model/Criterion and calculating the loss.也就是说Task 在 fairseq 架构中承担四个职责持有词典源/目标语言的Dictionary翻译任务或单语词典语言模型任务数据集加载与迭代load_dataset(split)加载指定 splitget_batch_iterator(...)生成可按 epoch 复用、可分片的批量迭代器构建模型与损失build_model(args)与build_criterion(args)分别实例化BaseFairseqModel与FairseqCriterion计算损失train_step/valid_step封装前向、求 loss、反向传播的完整一步。任务通过--task命令行参数选择。选定任务后该任务还会向全局解析器注入自己的附加参数——例如语言模型任务会追加--tokens-per-sample、--sample-break-mode等参数。这一“动态注入参数”的机制是理解整套任务体系的关键。注册与分发机制register_task、TASK_REGISTRY 与自动导入任务选择与分发逻辑位于 fairseq/tasks/init.py。其核心结构如下全局注册表模块级字典TASK_REGISTRY {}保存“任务名 → 任务类”的映射TASK_CLASS_NAMES集合用于防止类名重复入口函数def setup_task(args, **kwargs): return TASK_REGISTRY[args.task].setup_task(args, **kwargs)即fairseq.tasks.setup_task(args)本质上是根据args.task从注册表中取出任务类调用其类方法setup_task。注册装饰器register_task(name)要求被装饰的类必须继承自FairseqTask并会拒绝重复注册任务名或重复类名def register_task_cls(cls): if name in TASK_REGISTRY: raise ValueError(Cannot register duplicate task ({}).format(name)) if not issubclass(cls, FairseqTask): raise ValueError(Task ({}: {}) must extend FairseqTask.format(name, cls.__name__)) ... TASK_REGISTRY[name] cls自动导入与参数注入模块加载时会遍历tasks/目录下所有.py文件并逐一importlib.import_module从而触发各文件中register_task(...)装饰器的执行。随后对每个已注册任务代码会创建一个argparse.ArgumentParser把--task任务名参数与该任务的add_args中声明的所有“附加命令行参数”都挂到对应参数组中for file in os.listdir(os.path.dirname(__file__)): if file.endswith(.py) and not file.startswith(_): task_name file[:file.find(.py)] importlib.import_module(fairseq.tasks. task_name) if task_name in TASK_REGISTRY: parser argparse.ArgumentParser(add_helpFalse) group_task parser.add_argument_group(Task name) group_task.add_argument(--task, metavartask_name, ...) group_args parser.add_argument_group(Additional command-line arguments) TASK_REGISTRY[task_name].add_args(group_args) globals()[task_name _parser] parser这意味着你只要在tasks/目录新增一个 Python 文件并写好register_task(my_task)无需修改任何注册代码任务就会被自动发现并暴露其命令行参数。这也是 tasks.rst 末尾 “Adding new tasks” 一节所承诺的扩展方式。在当前 InfoXLM 仓库中该目录实际注册了 translation、language_modeling、denoising、masked_lm、multilingual_translation、semisupervised_translation、sentence_prediction、sentence_ranking 等十余个任务见 tasks 目录。标准使用流程从 setup_task 到 batch 循环tasks.rst 给出的官方示例流程是理解 Task 接口的最佳入口# setup the task (e.g., load dictionaries) task fairseq.tasks.setup_task(args) # build model and criterion model task.build_model(args) criterion task.build_criterion(args) # load datasets task.load_dataset(train) task.load_dataset(valid) # iterate over mini-batches of data batch_itr task.get_batch_iterator( task.dataset(train), max_tokens4096, ) for batch in batch_itr: # compute the loss loss, sample_size, logging_output task.get_loss( model, criterion, batch, ) loss.backward()对照基类实现 fairseq_task.py各环节的实际行为如下setup_task 与词典管理基类的setup_task默认实现就是构造任务实例return cls(args, **kwargs)各具体任务会覆写它以加载词典、解析额外参数。词典相关的两个类方法同样定义在基类中load_dictionary(filename)调用Dictionary.load(filename)从磁盘加载词典build_dictionary(filenames, workers1, threshold-1, nwords-1, padding_factor8)从原始文本文件构建词典。其中threshold控制最小词频、nwords限定最终词典大小、padding_factor默认 8会把词典大小补齐到 8 的倍数——文档注释指出这对某些硬件如 Nvidia Tensor Cores很重要。build_model 与 build_criterion两者都遵循 fairseq 的统一“registry 构建”模式def build_model(self, args): from fairseq import models return models.build_model(args, self) def build_criterion(self, args): from fairseq import criterions return criterions.build_criterion(args, self)注意第二个参数self即当前 task被传入因此模型与损失函数在构建时可以访问任务持有的词典、数据配置等上下文。load_dataset 与 dataset 缓存load_dataset(split, combineFalse, **kwargs)在基类中抛出NotImplementedError必须由具体任务实现——它负责把磁盘上的索引化数据包装成FairseqDataset并存入self.datasets[split]。加载完成后通过dataset(split)取回该方法会做两道校验未加载则抛KeyError不是FairseqDataset类型则抛TypeError。任务实例的__init__中同时维护了self.datasets {}与self.dataset_to_epoch_iter {}两个缓存字典后者用于跨 epoch 复用批量迭代器。get_batch_iterator批量迭代器的构建细节get_batch_iterator是数据流水线的核心其完整签名携带了丰富的控制项def get_batch_iterator( self, dataset, max_tokensNone, max_sentencesNone, max_positionsNone, ignore_invalid_inputsFalse, required_batch_size_multiple1, seed1, num_shards1, shard_id0, num_workers0, epoch0, ):从 fairseq_task.py 的实现看其内部依次完成跨 epoch 复用若该 dataset 已有迭代器缓存dataset_to_epoch_iter直接返回——注释说明“对于默认的 fairseq task数据集不是动态的可以跨 epoch 返回同一迭代器”任务可覆写此行为设置起始 epochdataset.set_epoch(epoch)按样本长度排序在data_utils.numpy_seed(seed)的随机种子下调用dataset.ordered_indices()过滤超长样本若给定max_positions用data_utils.filter_by_size过滤ignore_invalid_inputsTrue时静默丢弃否则抛异常按大小约束切分 mini-batchdata_utils.batch_by_size(indices, dataset.num_tokens, max_tokens..., max_sentences..., required_batch_size_multiple...)这就是示例中max_tokens4096生效的位置构造可分片、可复用的迭代器iterators.EpochBatchIterator(dataset..., collate_fndataset.collater, batch_sampler..., seed..., num_shards..., shard_id..., num_workers..., epoch...)collate_fn直接复用 dataset 自己的collater。值得留意的是当前 InfoXLM 仓库对基类的这段实现加入了调试打印如print(| At task.get_batch_iterator ..., flushTrue)这与上游 fairseq 原版不同属于该仓库的本地修改从源码结构看这些打印用于观察数据迭代各阶段是否被触发。train_step / valid_step / inference_step损失的“三个现场”文档示例中的task.get_loss(model, criterion, batch)对应的完整一步在基类中由train_step实现def train_step(self, sample, model, criterion, optimizer, ignore_gradFalse): model.train() loss, sample_size, logging_output criterion(model, sample) if ignore_grad: loss * 0 optimizer.backward(loss) return loss, sample_size, logging_output关键约定sample_size是梯度分母的度量“which is used as the denominator for the gradient”ignore_gradTrue时把 loss 置零例如半监督任务中未带标注的样本。valid_step则在torch.no_grad()下只算 lossinference_step默认委托给 generatorgenerator.generate(models, sample, prefix_tokensprefix_tokens)。此外build_generator展示了 Task 与解码参数的耦合当args.score_reference为真时返回SequenceScorer否则根据print_alignment选择SequenceGeneratorWithAlignment或SequenceGenerator并把beam默认 5、lenpen、unkpen、sampling、sampling_topk、sampling_topp、temperature、diverse_beam_groups等解码超参数一一传入。还有max_positions()、source_dictionary、target_dictionary等接口分别返回任务允许的最大输入长度与源/目标词典。内置任务一TranslationTask翻译任务定义于 translation.py以register_task(translation)注册。其文档说明与官方文档一致Translate from one (source) language to another (target) language.The translation task is compatible withfairseq-train,fairseq-generateandfairseq-interactive.它向命令行暴露的主要参数包括参数默认值说明data必填冒号分隔的数据目录列表训练期间按轮询round-robin方式依次使用-s/--source-langNone源语言代码-t/--target-langNone目标语言代码--lazy-load关惰性加载数据集--raw-text关加载原始文本数据集--load-alignments关加载分词后的词对齐文件--left-pad-sourceTrue源序列左侧填充--left-pad-targetFalse目标序列左侧填充--max-source-positions1024源序列最大 token 数--max-target-positions1024目标序列最大 token 数其数据加载核心是模块级函数load_langpair_dataset(...)从源码可以读出翻译任务的数据约定文件命名协议数据目录下的索引文件形如split.src-tgt.src如train.en-de.en函数会尝试src-tgt与tgt-src两种方向来推断语言代码找不到则抛FileNotFoundError分片合并combineTrue时通过itertools.count()循环探测train.1、train.2… 等编号分片并用ConcatDataset按sample_ratios拼接其中主分片的采样比例为upsample_primary可选截断truncate_sourceTrue时用StripTokenDataset → TruncateDataset(max_source_positions - 1) → AppendTokenDataset(eos)把过长的源句截到限长并补回 eos对齐加载load_alignmentsTrue时读取split.align.src-tgt对齐文件最终所有包装都落在LanguagePairDataset上透传left_pad_source/target、max_source_positions、max_target_positions等参数。内置任务二LanguageModelingTask语言模型任务定义于 language_modeling.py以register_task(language_modeling)注册。类文档说明它持有 input 词典dictionary与输出词典output_dictionary二者通常相同除非使用--output-dictionary-size截断输出词表以及目标类型列表targets——可取self、future、past三种默认为[future]。文档注明该任务兼容fairseq-train、fairseq-generate、fairseq-interactive与fairseq-eval-lm四个工具。命令行参数从add_args中可整理出完整的参数表参数默认值说明data必填数据目录路径--sample-break-modenone取值none/complete/complete_doc/eosnone时每样本固定填tokens-per-sample个 tokencomplete只在句末切分一个样本可含多句complete_doc类似但尊重文档边界eos时每个样本只含一句--tokens-per-sample1024LM 数据集每个样本的最大 token 数--lazy-load关惰性加载已废弃建议改用--dataset-impllazy--raw-text关加载原始文本已废弃建议改用--dataset-implraw--output-dictionary-size-1限制输出词表大小--self-target关加入“预测自身”目标--future-target关加入“预测未来”目标--past-target关加入“预测过去”目标--add-bos-token关在句首插入stoken--max-target-positions无目标序列最大 token 数setup_task词典加载与目标解析setup_task的源码展示了几个值得注意的细节废弃参数迁移--raw-text/--lazy-load会触发utils.deprecation_warning并把args.dataset_impl改写为raw/lazy多路径数据args.data支持冒号分隔的多个目录词典从第一个路径的dict.txt读取Dictionary.load(os.path.join(paths[0], dict.txt))并打印词表规模| dictionary: {} types输出词表截断若output_dictionary_size 0用TruncatedDictionary(dictionary, size)包装仅允许预测词表内的高频词旧 checkpoint 兼容若 args 中存在历史参数exclude_self_target会转换为self_target not exclude_self_target目标集合推导按self_target / future_target / past_target三个开关收集targets全部未设置时回落到标准自回归设定[future]。load_datasetTokenBlockDataset MonolingualDataset 的组合LM 任务的load_dataset(split, epoch0, ...)按如下步骤组装数据集依据epoch % len(paths)在多数据路径间轮询选取当前 epoch 使用的data_pathdata_utils.load_indexed_dataset(split_path, dictionary, dataset_impl, combinecombine)加载索引化分词文件找不到即抛FileNotFoundError用TokenBlockDataset(...)把连续 token 流切成长度约为tokens_per_sample的块break_mode由--sample-break-mode控制include_targetsTrue只有当sample_break_mode不是none时才为其它目标补 eosadd_eos_for_other_targets最终封装为MonolingualDataset(dataset, sizes, dictionary, output_dictionary, add_eos_for_other_targets..., shuffleTrue, targetsself.targets, add_bos_token...)。此外build_model会覆写基类以校验模型支持的任务目标若self.targets中存在model.supported_targets之外的目标直接抛ValueError——这保证了--self-target之类开关不会与不支持的模型静默错配。build_dataset_for_inference则用TokenBlockDataset MonolingualDataset TransformEosDataset构造推理输入remove_eos_from_srcTrue因为该输入将作为生成的 prefix。其inference_step的覆写逻辑是若外部未提供prefix_tokens就用样本中的src_tokens作为 prefix并先剥离开头的 eos再调用generator.generate(models, sample, prefix_tokensprefix_tokens)完成续写式生成。source_dictionary/target_dictionary两个 property 分别返回self.dictionary与self.output_dictionary。如何添加一个新任务FairseqTask 接口清单tasks.rst 的最后一节 “Adding new tasks” 指向fairseq.tasks.register_task与fairseq.tasks.FairseqTask的完整成员列表。结合 fairseq_task.py 的实现一个新任务需要关注的接口清单如下接口类型作用与约定add_args(parser)静态方法向解析器追加任务专属参数基类默认空实现load_dictionary(filename)类方法加载词典默认Dictionary.loadbuild_dictionary(filenames, ...)类方法从原始文本构建词典支持词频阈值与 padding_factorsetup_task(args, **kwargs)类方法任务入口默认构造实例通常覆写以加载词典/解析参数load_dataset(split, combineFalse, **kwargs)实例方法加载指定 split基类抛NotImplementedError必须实现dataset(split)实例方法取回已加载 split带类型校验get_batch_iterator(dataset, ...)实例方法生成EpochBatchIterator可覆写以实现动态数据build_model(args)实例方法经models.build_model(args, self)构建模型build_criterion(args)实例方法经criterions.build_criterion(args, self)构建损失build_generator(args)实例方法构建 beam search / 采样解码器train_step(sample, model, criterion, optimizer, ignore_grad)实例方法前向 criterion(model, sample)optimizer.backward返回(loss, sample_size, logging_output)valid_step(sample, model, criterion)实例方法no_grad下计算验证损失inference_step(generator, models, sample, prefix_tokens)实例方法委托generator.generateupdate_step(num_updates)/grad_denom/aggregate_logging_outputs实例方法训练步更新钩子、梯度分母与日志聚合默认委托给 criterionmax_positions()实例方法任务允许的最大输入长度默认None不限制source_dictionary/target_dictionaryproperty返回对应Dictionary按任务语义实现文档给出的最小示例形态是register_task(classification) class ClassificationTask(FairseqTask): (...)即继承FairseqTask→ 用register_task装饰并命名 → 放在tasks/目录下 → 实现load_dataset等必要方法。由__init__.py的自动导入与参数注入逻辑可知一旦满足这三点--taskclassification即可在命令行中被识别且其add_args中声明的参数会自动出现在帮助信息中。小结Task 是 fairseq 中连接“数据—模型—损失”的枢纽fairseq.tasks.setup_task(args)按--task参数从TASK_REGISTRY分发到对应任务类tasks/__init__.py新任务“零注册成本”文件名 register_task 继承FairseqTask目录级自动导入会完成注册与参数暴露批量迭代器get_batch_iterator的max_tokens/max_sentences/num_shards/epoch等参数直接决定数据切分与分布式采样行为实现细节见 fairseq_task.py内置的 TranslationTask 与 LanguageModelingTask 分别展示了“双词典 语言对文件协议”和“词典截断 多目标self/future/past TokenBlockDataset”两类典型任务设计模式translation.py、language_modeling.py需要自定义任务如 InfoXLM 的跨语言/跨模态场景时按上文接口清单实现load_dataset与词典 property并覆写train_step以适配自定义损失的分母约定即可。【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表