十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

fairseq 实战教程:用字符级 RNN 在 fairseq 中实现分类任务(Model/Task 注册全流程)

fairseq 实战教程:用字符级 RNN 在 fairseq 中实现分类任务(Model/Task 注册全流程) fairseq 实战教程用字符级 RNN 在 fairseq 中实现分类任务Model/Task 注册全流程【免费下载链接】fairseqFacebook AI Research Sequence-to-Sequence Toolkit written in Python.项目地址: https://gitcode.com/gh_mirrors/fa/fairseq本篇教程将完整演示如何把 fairseq 从一个序列到序列Sequence-to-Sequence工具包扩展为支持分类任务的框架以 PyTorch 官方Classifying Names with a Character-Level RNN教程为蓝本在 fairseq 中重新实现一个基于字符级 RNN 的姓名国籍分类器。读者学完后将掌握 fairseq 扩展的四大核心能力用fairseq-preprocess预处理分类数据、用register_model注册自定义模型、用register_task注册自定义任务以及基于现有命令行工具完成训练和交互式评估。整篇教程的每一步都对应 fairseq 源码中的注册机制与数据接口既是上手教程也是理解 fairseq 插件化架构的最佳入口。教程概览分类任务如何复用序列到序列框架fairseq 本身是面向翻译、语言建模等序列任务的工具包但它的设计具有很强的可扩展性模型、任务、准则criterion、优化器都可以通过装饰器注册进全局注册表随后被统一的命令行工具调用。本教程的思路是——把分类伪装成目标序列长度为 1 的序列到序列任务从而最大化复用 fairseq 现有的数据管线、批处理逻辑和训练循环。教程共分五个步骤预处理数据创建输入与标签的字典注册新模型用简单 RNN 编码输入句子并预测输出标签注册新任务加载字典与数据集训练模型直接使用现有的fairseq-train命令行工具编写评估脚本以 import fairseq 的方式交互式评估新输入。在动手之前建议先快速浏览 PyTorch 官方的字符级 RNN 分类教程理解其网络结构与训练逻辑本教程与其最大区别在于我们的实现需要支持批量输入与GPU 张量。第一步预处理数据创建字典PyTorch 官方教程提供的是原始数据这里使用一份已经按字符分词、并划分好 train/valid/test 三个子集的修改版数据。下载并解压数据后使用fairseq-preprocess命令行工具创建字典。该工具主要为序列到序列问题设计但通过把标签视作长度为 1 的 target 序列即可直接复用它同时通过--dataset-impl raw让预处理产物以原始文本格式输出便于阅读 fairseq-preprocess \ --trainpref names/train --validpref names/valid --testpref names/test \ --source-lang input --target-lang label \ --destdir names-bin --dataset-impl raw参数说明--trainpref/--validpref/--testpref指定三个子集的文件前缀names/train实际会读取names/train.input与names/train.label--source-lang input、--target-lang label声明源语言为input输入句子、目标语言为label类别标签--destdir names-bin预处理产物输出目录--dataset-impl raw以原始文本格式存储数据集。在 fairseq_cli/preprocess.py 中可以看到当args.dataset_impl raw时工具不再生成二进制的 indexed dataset而是直接把原始文本文件复制到目标目录这对于理解数据内容、调试自定义任务非常友好其他可选值如cached等则用于加速加载。命令执行后names-bin/目录中会生成inputs和labels两套字典文件dict.input.txt与dict.label.txt以及对应的数据文件例如train.input-label.input、train.input-label.label等。后续任务代码将直接从该目录加载这些字典与数据。第二步注册新模型rnn_classifier2.1 从 PyTorch 教程复刻基础 RNN 模块在fairseq/models/rnn_classifier.py中创建文件先放入 PyTorch 教程中的简单 RNN 模块。该模块在每个时间步将当前字符的 one-hot 向量与上一时刻的隐藏状态拼接通过两个线性层分别计算新隐藏状态与输出import torch import torch.nn as nn class RNN(nn.Module): def __init__(self, input_size, hidden_size, output_size): super(RNN, self).__init__() self.hidden_size hidden_size self.i2h nn.Linear(input_size hidden_size, hidden_size) self.i2o nn.Linear(input_size hidden_size, output_size) self.softmax nn.LogSoftmax(dim1) def forward(self, input, hidden): combined torch.cat((input, hidden), 1) hidden self.i2h(combined) output self.i2o(combined) output self.softmax(output) return output, hidden def initHidden(self): return torch.zeros(1, self.hidden_size)其中input_size为输入字符的 one-hot 维度即输入字典大小hidden_size为隐藏层维度output_size为标签类别数即标签字典大小。注意这里与 PyTorch 原教程一致softmax 沿dim1即 batch 维之后计算返回的是对数概率。2.2 用register_model把模型注册进 fairseq仅仅定义一个nn.Module还不够fairseq 通过register_model装饰器把模型类登记进全局注册表之后才能被现有命令行工具识别。在 fairseq/models/init.py 中可以看到实现装饰器会把类写入MODEL_REGISTRY并且强制校验类必须是BaseFairseqModel的子类否则抛出ValueError(Model (...) must extend BaseFairseqModel)。因此需要再写一个包装类FairseqRNNClassifier继承BaseFairseqModel并以rnn_classifier为名注册from fairseq.models import BaseFairseqModel, register_model # Note: the register_model decorator should immediately precede the # definition of the Model class. register_model(rnn_classifier) class FairseqRNNClassifier(BaseFairseqModel): staticmethod def add_args(parser): # Models can override this method to add new command-line arguments. # Here well add a new command-line argument to configure the # dimensionality of the hidden state. parser.add_argument( --hidden-dim, typeint, metavarN, helpdimensionality of the hidden state, ) classmethod def build_model(cls, args, task): # Fairseq initializes models by calling the build_model() # function. This provides more flexibility, since the returned model # instance can be of a different type than the one that was called. # In this case well just return a FairseqRNNClassifier instance. # Initialize our RNN module rnn RNN( # Well define the Task in the next section, but for now just # notice that the task holds the dictionaries for the source # (i.e., the input sentence) and target (i.e., the label). input_sizelen(task.source_dictionary), hidden_sizeargs.hidden_dim, output_sizelen(task.target_dictionary), ) # Return the wrapped version of the module return FairseqRNNClassifier( rnnrnn, input_vocabtask.source_dictionary, ) def __init__(self, rnn, input_vocab): super(FairseqRNNClassifier, self).__init__() self.rnn rnn self.input_vocab input_vocab # The RNN module in the tutorial expects one-hot inputs, so we can # precompute the identity matrix to help convert from indices to # one-hot vectors. We register it as a buffer so that it is moved to # the GPU when cuda() is called. self.register_buffer(one_hot_inputs, torch.eye(len(input_vocab))) def forward(self, src_tokens, src_lengths): # The inputs to the forward() function are determined by the # Task, and in particular the net_input key in each # mini-batch. Well define the Task in the next section, but for # now just know that *src_tokens* has shape (batch, src_len) and # *src_lengths* has shape (batch). bsz, max_src_len src_tokens.size() # Initialize the RNN hidden state. Compared to the original PyTorch # tutorial well also handle batched inputs and work on the GPU. hidden self.rnn.initHidden() hidden hidden.repeat(bsz, 1) # expand for batched inputs hidden hidden.to(src_tokens.device) # move to GPU for i in range(max_src_len): # WARNING: The inputs have padding, so we should mask those # elements here so that padding doesnt affect the results. # This is left as an exercise for the reader. The padding symbol # is given by self.input_vocab.pad() and the unpadded length # of each input is given by *src_lengths*. # One-hot encode a batch of input characters. input self.one_hot_inputs[src_tokens[:, i].long()] # Feed the input to our RNN. output, hidden self.rnn(input, hidden) # Return the final output state for making a prediction return output这段代码的关键点值得展开add_args是模型向命令行暴露参数的约定入口。BaseFairseqModel基类见 fairseq/models/fairseq_model.py定义了add_args与build_model两个类方法接口子类可按需覆写。这里新增了--hidden-dim参数控制隐藏层维度。build_model是 fairseq 创建模型的统一入口。在 fairseq/models/init.py 的build_model顶层函数中fairseq 先从命令行/配置中解析出model_type查找注册表拿到模型类再调用model.build_model(cfg, task)。之所以采用类方法返回实例而非直接实例化是为了保留灵活性——返回的实例类型可以不同于被调用的类。这里我们直接返回FairseqRNNClassifier实例其内部 RNN 的输入/输出维度取自任务持有的两个字典task.source_dictionary与task.target_dictionary。register_buffer保证 one-hot 矩阵跟随设备迁移。torch.eye(len(input_vocab))注册为 buffer 后调用model.cuda()时会被自动移动到 GPU从而支持 GPU 训练。forward的入参由任务与批处理管线共同决定。它接收的是每个 mini-batch 中net_input键下的src_tokens形状(batch, src_len)与src_lengths形状(batch,)。循环按时间步取每个 batch 中所有样本的第i个字符索引通过查表self.one_hot_inputs[...]完成 one-hot 编码后喂给 RNN隐藏状态通过repeat(bsz, 1)扩展为批量形式并to(src_tokens.device)显式迁移到 GPU。遗留的 padding 问题是一个刻意练习批内样本长度不一短样本会被 pad 符号填充而上述循环会把这些 pad 位置也送入 RNN。读者应利用self.input_vocab.pad()拿到 pad 符号索引、结合src_lengths给出的每个样本真实长度在循环中屏蔽 padding 位置避免污染隐藏状态与输出。顺带一提BaseFairseqModel.get_normalized_probs_scriptable见 fairseq/models/fairseq_model.py对这类没有 decoder、直接输出张量的简单模型做了语法糖支持torch.is_tensor(net_output)时直接对 logits 做 log-softmax/softmax。源码注释里明确写到了the classification tutorial这一场景说明本教程的模型形态正是该分支设计所服务的典型用例。2.3 注册命名架构pytorch_tutorial_rnn最后用register_model_architecture注册一个命名架构把参数默认值固化下来这样就能用--arch命令行参数一键选用from fairseq.models import register_model_architecture # The first argument to register_model_architecture() should be the name # of the model we registered above (i.e., rnn_classifier). The function we # register here should take a single argument *args* and modify it in-place # to match the desired architecture. register_model_architecture(rnn_classifier, pytorch_tutorial_rnn) def pytorch_tutorial_rnn(args): # We use getattr() to prioritize arguments that are explicitly given # on the command-line, so that the defaults defined below are only used # when no other value has been specified. args.hidden_dim getattr(args, hidden_dim, 128)在 fairseq/models/init.py 中可以看到register_model_architecture要求model_name必须已存在于MODEL_REGISTRY否则报错并将架构函数登记进ARCH_CONFIG_REGISTRY当命令行指定--arch pytorch_tutorial_rnn时build_model顶层函数会先调用该函数原地修改配置对象再实例化模型。getattr(args, hidden_dim, 128)的写法保证了命令行显式传入的--hidden-dim优先未指定时才回退到默认值 128。第三步注册新任务simple_classification模型负责计算任务Task负责喂数据。任务持有字典、负责加载数据集并决定批处理方式。本教程中批处理直接复用fairseq.data.LanguagePairDataset。在fairseq/tasks/simple_classification.py中创建文件import os import torch from fairseq.data import Dictionary, LanguagePairDataset from fairseq.tasks import LegacyFairseqTask, register_task register_task(simple_classification) class SimpleClassificationTask(LegacyFairseqTask): staticmethod def add_args(parser): # Add some command-line arguments for specifying where the data is # located and the maximum supported input length. parser.add_argument(data, metavarFILE, helpfile prefix for data) parser.add_argument(--max-positions, default1024, typeint, helpmax input length) classmethod def setup_task(cls, args, **kwargs): # Here we can perform any setup required for the task. This may include # loading Dictionaries, initializing shared Embedding layers, etc. # In this case well just load the Dictionaries. input_vocab Dictionary.load(os.path.join(args.data, dict.input.txt)) label_vocab Dictionary.load(os.path.join(args.data, dict.label.txt)) print(| [input] dictionary: {} types.format(len(input_vocab))) print(| [label] dictionary: {} types.format(len(label_vocab))) return SimpleClassificationTask(args, input_vocab, label_vocab) def __init__(self, args, input_vocab, label_vocab): super().__init__(args) self.input_vocab input_vocab self.label_vocab label_vocab def load_dataset(self, split, **kwargs): Load a given dataset split (e.g., train, valid, test). prefix os.path.join(self.args.data, {}.input-label.format(split)) # Read input sentences. sentences, lengths [], [] with open(prefix .input, encodingutf-8) as file: for line in file: sentence line.strip() # Tokenize the sentence, splitting on spaces tokens self.input_vocab.encode_line( sentence, add_if_not_existFalse, ) sentences.append(tokens) lengths.append(tokens.numel()) # Read labels. labels [] with open(prefix .label, encodingutf-8) as file: for line in file: label line.strip() labels.append( # Convert label to a numeric ID. torch.LongTensor([self.label_vocab.add_symbol(label)]) ) assert len(sentences) len(labels) print(| {} {} {} examples.format(self.args.data, split, len(sentences))) # We reuse LanguagePairDataset since classification can be modeled as a # sequence-to-sequence task where the target sequence has length 1. self.datasets[split] LanguagePairDataset( srcsentences, src_sizeslengths, src_dictself.input_vocab, tgtlabels, tgt_sizestorch.ones(len(labels)), # targets have length 1 tgt_dictself.label_vocab, left_pad_sourceFalse, # Since our target is a single class label, theres no need for # teacher forcing. If we set this to True then our Models # forward() method would receive an additional argument called # *prev_output_tokens* that would contain a shifted version of the # target sequence. input_feedingFalse, ) def max_positions(self): Return the max input length allowed by the task. # The source should be less than *args.max_positions* and the target # has max length 1. return (self.args.max_positions, 1) property def source_dictionary(self): Return the source :class:~fairseq.data.Dictionary. return self.input_vocab property def target_dictionary(self): Return the target :class:~fairseq.data.Dictionary. return self.label_vocab # We could override this method if we wanted more control over how batches # are constructed, but its not necessary for this tutorial since we can # reuse the batching provided by LanguagePairDataset. # # 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, epoch1, # data_buffer_size0, disable_iterator_cacheFalse, # ): # (...)结合源码理解任务注册与各接口register_task与setup_task的联动。register_task见 fairseq/tasks/init.py把任务类登记进TASK_REGISTRY并校验其必须是FairseqTask子类。顶层函数setup_task见 fairseq/tasks/init.py根据配置中的task名称从注册表取出类调用类的setup_task(cfg)完成初始化——这正是本任务中加载两本字典的时机。继承LegacyFairseqTask而非FairseqTask。LegacyFairseqTask是 fairseq/tasks/fairseq_task.py 中为兼容旧的 argparse 风格参数而保留的基类它直接把args存为self.args适合教程这类以Namespace传参的写法新版 Hydra 风格任务则直接继承FairseqTask。add_args定义位置参数data与--max-positions。data是位置参数数据目录前缀--max-positions默认 1024控制允许的最大输入长度。Dictionary.load加载字典。Dictionary.load见 fairseq/data/dictionary.py读取symbol count格式的文本文件重建索引与词频表并自动附加s、pad、/s、unk等特殊符号。load_dataset按 split 加载并组装数据集。输入侧用encode_line(sentence, add_if_not_existFalse)见 fairseq/data/dictionary.py把按空格切好的字符序列编码为索引张量标签侧用add_symbol(label)见 fairseq/data/dictionary.py把字符串标签转成数值 ID不在字典中的标签会被动态加入。复用LanguagePairDataset实现目标长度为 1的分类批处理。LanguagePairDataset见 fairseq/data/language_pair_dataset.py原本为翻译设计这里把句子当 source、单标签当 targettgt_sizestorch.ones(len(labels))声明每个目标长度恒为 1left_pad_sourceFalse表示源序列右对齐补 pad左侧不补。关键参数是input_feedingFalse开启时默认值collate 会为每个 target 构造一个右移一位的prev_output_tokens见 fairseq/data/language_pair_dataset.py 中move_eos_to_beginning的逻辑模型的forward就会多收到一个prev_output_tokens参数由于分类目标只有 1 个符号、不需要教师强制teacher forcing这里显式关闭让forward保持我们第二步定义的签名。max_positions声明长度上限源长度受args.max_positions限制目标长度上限为 1。source_dictionary/target_dictionary属性这是模型build_model中读取字典的约定接口分别返回输入字典与标签字典从而把任务与模型接线起来。注释中的get_batch_iterator展示了高级定制点若想完全控制 mini-batch 的构造方式可以覆写该方法本教程因复用LanguagePairDataset的 batching 而无需覆写。第四步训练模型数据、模型、任务三者就绪后直接用现有的fairseq-train命令行工具开始训练只需显式指定新任务与架构 fairseq-train names-bin \ --task simple_classification \ --arch pytorch_tutorial_rnn \ --optimizer adam --lr 0.001 --lr-shrink 0.5 \ --max-tokens 1000 (...) | epoch 027 | loss 1.200 | ppl 2.30 | wps 15728 | ups 119.4 | wpb 116 | bsz 116 | num_updates 3726 | lr 1.5625e-05 | gnorm 1.290 | clip 0% | oom 0 | wall 32 | train_wall 21 | epoch 027 | valid on valid subset | valid_loss 1.41304 | valid_ppl 2.66 | num_updates 3726 | best 1.41208 | done training in 31.6 seconds参数解读names-bin对应任务add_args中定义的位置参数data指向预处理产物目录--task simple_classification选择刚注册的任务--arch pytorch_tutorial_rnn选择刚注册的命名架构--optimizer adam --lr 0.001 --lr-shrink 0.5配置优化器与学习率衰减策略--max-tokens 1000限制每个 mini-batch 的最大 token 数。提示可以通过--hidden-dim参数向fairseq-train传递自定义隐藏层维度例如--hidden-dim 256未指定时回退到架构默认值 128。训练日志中可以看到每轮 epoch 打印loss交叉熵损失、ppl困惑度、wps每秒词数、ups每秒更新次数、num_updates等训练指标以及验证集上的valid_loss/valid_ppl与历史最优值。上述示例中模型在第 27 个 epoch 收敛验证损失约 1.41整个训练耗时约 32 秒该数字与硬件、--max-tokens配置相关仅作量级参考。训练结束后模型文件会出现在checkpoints/目录下如checkpoint_best.pt、checkpoint_last.pt。第五步编写评估脚本交互式预测训练完成后编写一个独立脚本eval_classifier.py以 import fairseq 的方式加载模型并对新输入做交互式预测from fairseq import checkpoint_utils, data, options, tasks # Parse command-line arguments for generation parser options.get_generation_parser(default_tasksimple_classification) args options.parse_args_and_arch(parser) # Setup task task tasks.setup_task(args) # Load model print(| loading model from {}.format(args.path)) models, _model_args checkpoint_utils.load_model_ensemble([args.path], tasktask) model models[0] while True: sentence input(\nInput: ) # Tokenize into characters chars .join(list(sentence.strip())) tokens task.source_dictionary.encode_line( chars, add_if_not_existFalse, ) # Build mini-batch to feed to the model batch data.language_pair_dataset.collate( samples[{id: -1, source: tokens}], # bsz 1 pad_idxtask.source_dictionary.pad(), eos_idxtask.source_dictionary.eos(), left_pad_sourceFalse, input_feedingFalse, ) # Feed batch to the model and get predictions preds model(**batch[net_input]) # Print top 3 predictions and their log-probabilities top_scores, top_labels preds[0].topk(k3) for score, label_idx in zip(top_scores, top_labels): label_name task.target_dictionary.string([label_idx]) print(({:.2f})\t{}.format(score, label_name))脚本各段与源码的对应关系options.get_generation_parser(default_tasksimple_classification)见 fairseq/options.py构建一个面向生成/评估的命令行解析器并把默认任务设为我们的分类任务options.parse_args_and_arch见 fairseq/options.py负责真正的解析——它会先加载注册表含模型、任务、架构等再按注册的add_args动态注入参数。tasks.setup_task(args)走的就是第三步中register_task登记的入口加载names-bin下的两本字典脚本运行时会打印| [input] dictionary: 64 types与| [label] dictionary: 24 types。checkpoint_utils.load_model_ensemble([args.path], tasktask)见 fairseq/checkpoint_utils.py加载--path指定的 checkpoint 文件可传入多个文件构成模型集成返回模型列表这里取第一个模型用于推理。交互循环中先把用户输入字符串按字符空格化 .join(list(sentence.strip()))再用source_dictionary.encode_line(..., add_if_not_existFalse)编码为索引序列。data.language_pair_dataset.collate见 fairseq/data/language_pair_dataset.py把单个样本组装成 mini-batch构造{id: -1, source: tokens}形式的样本列表传入pad_idx、eos_idx并保持与训练一致的left_pad_sourceFalse、input_feedingFalse。collate 返回的字典中net_input键下包含src_tokens与src_lengths正好对应模型forward的签名因此直接model(**batch[net_input])即可前向。输出侧用preds[0].topk(k3)取对数概率最高的 3 个标签target_dictionary.string([label_idx])把索引还原为字符串标签名。运行评估注意要传入原始数据路径names-bin因为任务需要据此加载字典 python eval_classifier.py names-bin --path checkpoints/checkpoint_best.pt | [input] dictionary: 64 types | [label] dictionary: 24 types | loading model from checkpoints/checkpoint_best.pt Input: Satoshi (-0.61) Japanese (-1.20) Arabic (-2.86) Italian Input: Sinbad (-0.30) Arabic (-1.76) English (-4.08) Russian从示例输出可以看到输入日本名字 Satoshi 时模型首选 Japanese输入阿拉伯风格名字 Sinbad 时首选 Arabic模型已学到字符序列与国籍标签之间的对应关系且返回的是对数概率负值越接近 0 置信度越高。扩展阅读注册机制一览本教程背后是 fairseq 的插件化注册体系几个关键点可以帮助你把同样的套路迁移到其他任务如情感分类、文本蕴含等注册表即插件接口模型登记在MODEL_REGISTRY、架构登记在ARCH_MODEL_REGISTRY/ARCH_CONFIG_REGISTRY、任务登记在TASK_REGISTRY见 fairseq/models/init.py 与 fairseq/tasks/init.py。只要模块被 import注册即生效fairseq 的options.parse_args_and_arch在解析前会先加载--user-dir指定的自定义模块因此你既可以把新文件直接放进fairseq/models/、fairseq/tasks/下随包自动导入也可以通过--user-dir以外部插件方式引入不动仓库源码。net_input是模型与数据的契约模型forward的入参完全由 collate 产物中net_input键决定。想改输入形态比如加入prev_output_tokens做序列生成、或加入额外的特征字段只需让任务与 collate 保持一致。面向后续优化本教程的forward循环对每个时间步逐个前向未做 padding mask也刻意没有并行化如果你希望把它改造成更高效或更健壮的实现可以从self.input_vocab.pad()与src_lengths入手补齐掩码逻辑再考虑用pack_padded_sequence或矩阵化运算加速。至此你已经完整走通了 fairseq 的数据预处理 → 注册模型 → 注册任务 → 训练 → 交互式评估全链路并理解了注册表、net_input契约、字典与LanguagePairDataset这些核心组件如何协同工作。以此为模板你可以为任意分类/标注类任务快速搭建 fairseq 扩展。【免费下载链接】fairseqFacebook AI Research Sequence-to-Sequence Toolkit written in Python.项目地址: https://gitcode.com/gh_mirrors/fa/fairseq创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表