
BART 在 fairseq 中的实战指南去噪 Seq2Seq 预训练、推理与微调全流程【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilmBARTBidirectional and Auto-Regressive Transformer是 fairseq 中实现的双向自回归 Transformer 序列到序列模型其核心思想是采用去噪自编码器目标进行预训练破坏原文后学习重建原文从而同时具备自然语言生成、翻译与理解能力。本文以 edgelm/examples/bart/README.md 为骨架结合仓库内fairseq/models/bart的源码实现完整讲解 BART 的模型体系、预训练权重、Python 推理 API编码、特征提取、分类、掩码填充、CNN-DM 摘要评估以及 GLUE 与 CNN-DM 两条微调实战链路读完即可在 fairseq 环境下复现 BART 的全部典型用法。BART 简介用去噪目标统一生成与理解BART 是一个以去噪denoising作为预训练目标的序列到序列模型。与传统自编码器或自回归语言模型不同BART 在预训练阶段对输入文本施加多种噪声如 token 掩码、文本删除、打乱句序等再训练模型从被破坏的文本中重建原始文本。这一目标更为通用论文中展示BART 在 SQuAD 与 GLUE 上能够追平 RoBERTa 的结果同时在摘要生成XSum、CNN/Daily Mail、长文生成式问答ELI5与对话回复生成ConvAI2上取得当时最优的效果。从源码结构看BART 模型在仓库中的实现位于 fairseq/models/bart/model.py其核心类BARTModel直接继承自TransformerModel通过register_model(bart)注册进 fairseq 的模型注册表model.py#L26-L27。这解释了为什么 BART 可以复用 fairseq 的translation、sentence_prediction等任务框架——它本质上是带去噪预训练目标的 Transformer。模型初始化遵循 BERT 的随机权重初始化策略init_bert_params并维护一个classification_heads模块字典用于挂载各类句子级分类头model.py#L43-L48这是后续 MNLI 等下游任务的架构基础。预训练模型一览官方为 BART 发布了 5 个预训练检查点覆盖通用预训练与下游微调两种形态模型描述参数量bart.base6 层 encoder 6 层 decoder 的 BART140Mbart.large12 层 encoder 12 层 decoder 的 BART400Mbart.large.mnli在 MNLI 上微调后的bart.large400Mbart.large.cnn在 CNN-DM 上微调后的bart.large400Mbart.large.xsum在 XSum 上微调后的bart.large400M上述模型均以.tar.gz压缩包形式发布内含model.pt权重文件与dict.txt等词典文件。这些预训练模型的注册信息同样固化在源码中BARTModel.hub_models()方法直接声明了各权重对应的下载地址model.py#L30-L38from_pretrained通过archive_mapcls.hub_models()将模型名映射到实际压缩包这也是torch.hub.load与from_pretrained两种加载方式能共享同一份权重清单的原因。各模型架构的超参数在 model.py#L315-L366 中定义bart_large使用 1024 维隐藏层、16 注意力头、FFN 维度 4096bart_base使用 768 维隐藏层、12 注意力头、FFN 维度 3072均开启 LayerNorm embedding、共享全部 embedding 与 decoder 输入输出 embedding。可通过--arch bart_large/--arch bart_base在训练与微调命令中直接引用。已报告实验结果GLUEdev 集单模型、单任务微调模型MNLIQNLIQQPRTESST-2MRPCCoLASTS-Broberta.large90.294.792.286.696.490.968.092.4bart.large89.994.992.587.096.690.462.891.2SQuADdev 集不使用额外数据模型SQuAD 1.1 EM/F1SQuAD 2.0 EM/F1roberta.large88.9 / 94.686.5 / 89.4bart.large88.8 / 94.686.1 / 89.2CNN/Daily Mailtest 集不使用额外数据模型R1R2RLBERTSUMEXTABS42.1319.6039.18bart.large44.1621.2840.90上述结果说明在理解类任务GLUE、SQuAD上 BART 与 RoBERTa 旗鼓相当而在生成类任务摘要上优势明显。快速上手加载 BART 模型方式一通过 torch.hub 加载PyTorch 1.1import torch bart torch.hub.load(pytorch/fairseq, bart.large) bart.eval() # 关闭 dropout若想微调则保留 train 模式方式二从本地检查点加载PyTorch 1.0 或自定义模型# 下载 bart.large 模型并解压 wget https://dl.fbaipublicfiles.com/fairseq/models/bart.large.tar.gz tar -xzvf bart.large.tar.gzfrom fairseq.models.bart import BARTModel bart BARTModel.from_pretrained(/path/to/bart.large, checkpoint_filemodel.pt) bart.eval()两种方式最终都会返回 BARTHubInterface 封装对象。从 model.py#L116-L138 可以看到from_pretrained内部走hub_utils.from_pretrained默认使用 GPT-2 BPEbpegpt2、sample_break_modeeos并开启load_checkpoint_headsTrue——这意味着微调时挂载的分类头如 MNLI head会随权重一起加载。值得注意的细节是当从预训练权重切到翻译任务微调时upgrade_state_dict_named 会自动删除 embedding 矩阵中对应mask的最后一行避免词典维度不匹配。文本编码与解码GPT-2 BPEBART 使用与 GPT-2 相同的 BPE 编码。encode会为序列添加s开头符号与/s结束符号hub_interface.py#L33-L63例如tokens bart.encode(Hello world!) assert tokens.tolist() [0, 31414, 232, 328, 2] bart.decode(tokens) # Hello world!编码遵循几个约定单个句子形如s a b c /s句子对形如s d e f /s 1 2 3 /sGPT-2 BPE 对前导空格敏感Hello world与 world、world编码结果不同。decode则会去除s并按连续的/s将多句拆分成列表返回hub_interface.py#L65-L78。从 BART 提取特征extract_features支持取最后一层特征或全部层的隐藏状态层 0 为 embedding 层# 提取最后一层特征 last_layer_features bart.extract_features(tokens) assert last_layer_features.size() torch.Size([1, 5, 1024]) # 提取 decoder 所有层的特征 all_layers bart.extract_features(tokens, return_all_hiddensTrue) assert len(all_layers) 13 assert torch.all(all_layers[-1] last_layer_features)实现上hub_interface.py#L119-L151extract_features会通过右移构造prev_output_tokens以features_onlyTrue模式前向一次return_all_hiddensTrue时把inner_states从T x B x C转置为B x T x C返回。句子对分类MNLI与自定义分类头使用官方 MNLI 微调模型直接预测# 加载已在 MNLI 上微调好的 BART bart torch.hub.load(pytorch/fairseq, bart.large.mnli) bart.eval() # 评估时关闭 dropout # 编码句子对并预测 tokens bart.encode(BART is a seq2seq model., BART is not sequence to sequence.) bart.predict(mnli, tokens).argmax() # 0: contradiction矛盾 tokens bart.encode(BART is denoising autoencoder., BART is version of autoencoder.) bart.predict(mnli, tokens).argmax() # 2: entailment蕴含predict的底层逻辑hub_interface.py#L160-L171是取extract_features的结果按/s位置聚合出句子表示喂给对应的classification_heads[head]后返回 log-softmax 概率。注册新的随机初始化分类头bart.register_classification_head(new_task, num_classes3) logprobs bart.predict(new_task, tokens)在源码中register_classification_head会构造一个BARTClassificationHeadmodel.py#L140-L164该头由dense线性层 激活函数默认 tanh 两层 dropout out_proj输出层组成并支持通过--pooler-dropout、--pooler-activation-fn、--spectral-norm-classification-head调节model.py#L284-L312。这就是把 BART 复用到任意分类任务的入口。批量预测import torch from fairseq.data.data_utils import collate_tokens bart torch.hub.load(pytorch/fairseq, bart.large.mnli) bart.eval() batch_of_pairs [ [BART is a seq2seq model., BART is not sequence to sequence.], [BART is denoising autoencoder., BART is version of autoencoder.], ] batch collate_tokens( [bart.encode(pair[0], pair[1]) for pair in batch_of_pairs], pad_idx1 ) logprobs bart.predict(mnli, batch) print(logprobs.argmax(dim1)) # tensor([0, 2])collate_tokens位于 fairseq/data/data_utils.py负责将变长 token 序列补齐pad_idx1对应/s为 batch 张量。使用 GPUbart.cuda() bart.predict(new_task, tokens)掩码填充BART 的填空能力BART 可以同时填充输入中的多个masktokenbart torch.hub.load(pytorch/fairseq, bart.base) bart.eval() bart.fill_mask([The cat mask on the mask.], topk3, beam10) # [[(The cat was on the ground., tensor(-0.6183)), (The cat was on the floor., tensor(-0.6798)), (The cat sleeps on the couch., tensor(-0.6830))]]注意默认强制输出长度与输入长度一致match_source_lenTrue关闭该约束后可生成更长文本bart.fill_mask([The cat mask on the mask.], topk3, beam10, match_source_lenFalse) # [[(The cat was on the ground., tensor(-0.6185)), (The cat was asleep on the couch., tensor(-0.6276)), (The cat was on the floor., tensor(-0.6800))]]批量填空GPU 上运行bart.cuda() bart.fill_mask([The cat mask on the mask., The dog mask on the mask.], topk3, beam10)从实现看hub_interface.py#L173-L208fill_mask会校验输入必须包含mask将文本按mask切分后逐段做 BPE 再拼接同时会把beam强制设为不小于topk并透传match_source_len给生成器返回(解码文本, score)的列表。评估 bart.large.mnliMNLI dev_matched 准确率以下代码逐行读取 MNLIdev_matched集统计准确率预期输出约0.9010label_map {0: contradiction, 1: neutral, 2: entailment} ncorrect, nsamples 0, 0 bart.cuda() bart.eval() with open(glue_data/MNLI/dev_matched.tsv) as fin: fin.readline() for index, line in enumerate(fin): tokens line.strip().split(\t) sent1, sent2, target tokens[8], tokens[9], tokens[-1] tokens bart.encode(sent1, sent2) prediction bart.predict(mnli, tokens).argmax().item() prediction_label label_map[prediction] ncorrect int(prediction_label target) nsamples 1 print(| Accuracy: , float(ncorrect)/float(nsamples)) # Expected output: 0.9010评估 bart.large.cnnCNN-DM 摘要生成与 ROUGE 计算数据准备参考abisee/cnn-dailymail的指引下载并处理 CNN/Daily Mail 数据得到每行一个未分词样本的test.source与test.target。也可下载cnn_dm_v2.tgz简化预处理分数可能有微小差异。huggingface/transformers提供更简单的单卡/多卡 beam search 评估接口对应模型路径为facebook/bart-large-cnn与facebook/bart-large-xsum。在 fairseq 中生成摘要cp>export CLASSPATH/path/to/stanford-corenlp-full-2016-10-31/stanford-corenlp-3.7.0.jar # 对假设和目标文件分词 cat test.hypo | java edu.stanford.nlp.process.PTBTokenizer -ioFileList -preserveLines test.hypo.tokenized cat test.target | java edu.stanford.nlp.process.PTBTokenizer -ioFileList -preserveLines test.hypo.target files2rouge test.hypo.tokenized test.hypo.target # Expected output: (ROUGE-2 Average_F: 0.21238)微调实战一GLUE 句子对分类任务1) 下载 GLUE 数据使用download_glue_data.py脚本可从 GLUE 官网相关 gist 获取下载全部任务数据wget https://gist.githubusercontent.com/W4ngatang/60c2bdb54d156a41194446737ce03e2e/raw/17b8dd0d724281ed7c3b2aeeda662b92809aadd5/download_glue_data.py python download_glue_data.py --data_dir glue_data --tasks all2) 预处理 GLUE 数据与 RoBERTa 相同./examples/roberta/preprocess_GLUE_tasks.sh glue_data glue_task_nameglue_task_name取值为{ALL, QQP, MNLI, QNLI, MRPC, RTE, STS-B, SST-2, CoLA}使用ALL一次处理全部任务。该脚本edgelm/examples/roberta/preprocess_GLUE_tasks.sh会自动下载 GPT-2 的encoder.json、vocab.bpe与dict.txt按任务抽取输入列与标签列如 QQP 取第 4、5 列输入与第 6 列标签经multiprocessing_bpe_encoder做 BPE 编码最后用fairseq-preprocess --only-source分别对 input0/input1/label 建 binMNLI 额外处理 dev/test 的 matched/mismatched 两个切分STS-B 则将标签除以 5.0 归一化到[0.0, 1.0]。3) 在 GLUE 任务上微调以RTE为例的完整命令TOTAL_NUM_UPDATES2036 # RTE 上 bsz 16 跑 10 个 epoch WARMUP_UPDATES61 # 更新步数的 6% LR1e-05 # polynomial LR 调度器的峰值学习率 NUM_CLASSES2 MAX_SENTENCES16 # Batch size BART_PATH/path/to/bart/model.pt CUDA_VISIBLE_DEVICES0,1 fairseq-train RTE-bin/ \ --restore-file $BART_PATH \ --batch-size $MAX_SENTENCES \ --max-tokens 4400 \ --task sentence_prediction \ --add-prev-output-tokens \ --layernorm-embedding \ --share-all-embeddings \ --share-decoder-input-output-embed \ --reset-optimizer --reset-dataloader --reset-meters \ --required-batch-size-multiple 1 \ --init-token 0 \ --arch bart_large \ --criterion sentence_prediction \ --num-classes $NUM_CLASSES \ --dropout 0.1 --attention-dropout 0.1 \ --weight-decay 0.01 --optimizer adam --adam-betas (0.9, 0.98) --adam-eps 1e-08 \ --clip-norm 0.0 \ --lr-scheduler polynomial_decay --lr $LR --total-num-update $TOTAL_NUM_UPDATES --warmup-updates $WARMUP_UPDATES \ --fp16 --fp16-init-scale 4 --threshold-loss-scale 1 --fp16-scale-window 128 \ --max-epoch 10 \ --find-unused-parameters \ --best-checkpoint-metric accuracy --maximize-best-checkpoint-metric;各 GLUE 任务需要不同的命令行参数模型MNLIQNLIQQPRTESST-2MRPCCoLASTS-B--num-classes32222221--lr5e-61e-51e-51e-55e-62e-52e-52e-5bsz128323232128646432--total-num-update309683311211327210185233114813341799--warmup-updates185819866796613146880107对STS-B还需追加--regression-target --best-checkpoint-metric loss并移除--maximize-best-checkpoint-metric该任务为回归任务优化目标是最小化 loss。注意事项--total-num-updates供polynomial_decay调度器使用按--max-epoch10与对应 batch size32/64/128计算得出。上述命令与超参在32GB显存的 Nvidia V100 上验证显存不足时可增大--update-freq并减小--batch-size。4) GLUE 推理微调完成后用checkpoints/目录中的检查点做推理from fairseq.models.bart import BARTModel bart BARTModel.from_pretrained( checkpoints/, checkpoint_filecheckpoint_best.pt, data_name_or_pathRTE-bin ) label_fn lambda label: bart.task.label_dictionary.string( [label bart.task.label_dictionary.nspecial] ) ncorrect, nsamples 0, 0 bart.cuda() bart.eval() with open(glue_data/RTE/dev.tsv) as fin: fin.readline() for index, line in enumerate(fin): tokens line.strip().split(\t) sent1, sent2, target tokens[1], tokens[2], tokens[3] tokens bart.encode(sent1, sent2) prediction bart.predict(sentence_classification_head, tokens).argmax().item() prediction_label label_fn(prediction) ncorrect int(prediction_label target) nsamples 1 print(| Accuracy: , float(ncorrect)/float(nsamples))微调实战二CNN-Dailymail 摘要生成1) 数据下载与预处理CNN/Daily Mail按abisee/cnn-dailymail指引下载原始数据确保样本未分词、未做 BPE。XSum按 EdinburghNLP/XSum 指引下载原始 Extreme Summarization 数据集同样保持原始不做 tokenization 与 BPE。2) BPE 预处理wget -N https://dl.fbaipublicfiles.com/fairseq/gpt2_bpe/encoder.json wget -N https://dl.fbaipublicfiles.com/fairseq/gpt2_bpe/vocab.bpe wget -N https://dl.fbaipublicfiles.com/fairseq/gpt2_bpe/dict.txt TASKcnn_dm for SPLIT in train val do for LANG in source target do python -m examples.roberta.multiprocessing_bpe_encoder \ --encoder-json encoder.json \ --vocab-bpe vocab.bpe \ --inputs $TASK/$SPLIT.$LANG \ --outputs $TASK/$SPLIT.bpe.$LANG \ --workers 60 \ --keep-empty; done done3) 二值化数据集fairseq-preprocess \ --source-lang source \ --target-lang target \ --trainpref ${TASK}/train.bpe \ --validpref ${TASK}/val.bpe \ --destdir ${TASK}-bin/ \ --workers 60 \ --srcdict dict.txt \ --tgtdict dict.txt;4) 在 CNN-DM 上微调TOTAL_NUM_UPDATES20000 WARMUP_UPDATES500 LR3e-05 MAX_TOKENS2048 UPDATE_FREQ4 BART_PATH/path/to/bart/model.pt CUDA_VISIBLE_DEVICES0,1,2,3,4,5,6,7 fairseq-train cnn_dm-bin \ --restore-file $BART_PATH \ --max-tokens $MAX_TOKENS \ --task translation \ --source-lang source --target-lang target \ --truncate-source \ --layernorm-embedding \ --share-all-embeddings \ --share-decoder-input-output-embed \ --reset-optimizer --reset-dataloader --reset-meters \ --required-batch-size-multiple 1 \ --arch bart_large \ --criterion label_smoothed_cross_entropy \ --label-smoothing 0.1 \ --dropout 0.1 --attention-dropout 0.1 \ --weight-decay 0.01 --optimizer adam --adam-betas (0.9, 0.999) --adam-eps 1e-08 \ --clip-norm 0.1 \ --lr-scheduler polynomial_decay --lr $LR --total-num-update $TOTAL_NUM_UPDATES --warmup-updates $WARMUP_UPDATES \ --fp16 --update-freq $UPDATE_FREQ \ --skip-invalid-size-inputs-valid-test \ --find-unused-parameters;上述配置预期在1个节点、8 张 32GB-V100上运行训练时长约5 小时改用4个节点分布式训练并设--update-freq 1可缩短训练时间。XSum 任务改用TOTAL_NUM_UPDATES15000、UPDATE_FREQ2。5) 摘要推理cp>cp>article{lewis2019bart, title {BART: Denoising Sequence-to-Sequence Pre-training for Natural Language Generation, Translation, and Comprehension}, author {Mike Lewis and Yinhan Liu and Naman Goyal and Marjan Ghazvininejad and Abdelrahman Mohamed and Omer Levy and Veselin Stoyanov and Luke Zettlemoyer }, journal{arXiv preprint arXiv:1910.13461}, year {2019}, }更详细的微调说明见 Finetuning on GLUE 与 Finetuning on CNN-DM。【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考