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

资讯详情

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

使用 fairseq 在 GLUE 上微调 BART:从数据预处理到分类微调与推理的完整实践指南

使用 fairseq 在 GLUE 上微调 BART:从数据预处理到分类微调与推理的完整实践指南 使用 fairseq 在 GLUE 上微调 BART从数据预处理到分类微调与推理的完整实践指南【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm本指南以 edgelm/examples/bart/README.glue.md 为骨架系统讲解如何在 unilm 仓库的 edgelmfairseq子项目中将预训练的 BART 序列到序列模型微调到 GLUE 自然语言理解基准的 8 个任务上。你将掌握 GLUE 数据下载与 BPE 预处理的完整流程、RTE 等任务的可复现微调命令与逐参数含义、各任务超参数总表以及基于 checkpoint 的离线推理与准确率评估方法并结合 sentence_prediction 任务实现 与 分类头源码 理解其底层运行原理。一、任务背景用去噪自编码器做句子对分类BARTDenoising Sequence-to-Sequence Pre-training是一个以去噪为预训练目标的序列到序列模型其 官方 README 指出该预训练目标比单向语言模型更通用在 SQuAD 与 GLUE 上可以匹配 RoBERTa 的结果如bart.large在 MNLI dev 集上为 89.9、RTE 上为 87.0同时在摘要生成XSum、CNN-DM、长文本生成式问答ELI5、对话回复生成ConvAI2等生成任务上取得当时的最佳效果。在 GLUE 上微调 BART 时核心思路是在其编码器末端挂接一个随机初始化的分类头classification head把「句子对分类/回归」建模为在预训练特征之上的浅层判别任务。整个流程分为三步下载 GLUE 原始数据 → 用与 RoBERTa 相同的脚本做 BPE 预处理 → 用 fairseq-train 微调。下面逐一展开。二、第 1 步下载 GLUE 原始数据使用社区维护的download_glue_data.py脚本GLUE 官网推荐的数据下载方式一次性获取全部任务数据wget download_glue_data.py 下载地址 python download_glue_data.py --data_dir glue_data --tasks all--data_dir glue_data指定数据存放目录脚本会在其中为每个任务创建子目录如glue_data/RTE/、glue_data/MNLI/等内含train.tsv、dev.tsv、test.tsv等原始 TSV 文件--tasks all下载 GLUE 全部任务也可以换成单个任务名如RTE只下载所需数据。下载完成后glue_data目录即作为第 2 步预处理脚本的输入目录。三、第 2 步预处理 GLUE 任务数据与 RoBERTa 完全一致预处理复用 edgelm/examples/roberta/preprocess_GLUE_tasks.sh命令格式为./examples/roberta/preprocess_GLUE_tasks.sh glue_data glue_task_name其中glue_task_name取值为{ALL, QQP, MNLI, QNLI, MRPC, RTE, STS-B, SST-2, CoLA}传ALL表示一次性预处理全部 8 个任务。3.1 脚本内部做了哪些事阅读 preprocess_GLUE_tasks.sh 源码可以看清完整的处理链路下载 GPT-2 BPE 资源脚本开头自动wget下载encoder.json、vocab.bpe与dict.txtfairseq 字典BPE 编码依赖这三个文件按任务解析列号脚本为每个任务硬编码了输入列与标签列。例如QQP输入列 4、5标签列 6测试输入列 2、3MNLI输入列 9、10训练标签列 12dev/test 标签列 16SST-2与CoLA是单句任务INPUT_COUNT1其余为双句任务清洗与分列去掉表头tail -n 2对 QQP 用awk -F \t过滤字段数不足的行再用cut分别抽取input0、input1与labelBPE 编码调用python -m examples.roberta.multiprocessing_bpe_encoder以 60 个 worker 并行对每个 split 的两个输入列做 GPT-2 BPE 编码保留空行--keep-emptyfairseq-preprocess 建 bin 数据对每个输入列用fairseq-preprocess --only-source生成TASK-bin/input0、TASK-bin/input1标签单独生成TASK-bin/label并指定--srcdict dict.txt作为字典。3.2 两个值得注意的特殊分支MNLI 的多 split 处理MNLI 的 dev/test 分为dev_matched、dev_mismatched、test_matched、test_mismatched四个子集脚本通过DEVPREF与TESTPREF用逗号拼接多个前缀传入fairseq-preprocessSTS-B 的标签缩放STS-B 是回归任务原始标签为 0~5 的相似度分数。脚本用awk {print $1 / 5.0}将训练与验证标签除以 5归一化到[0.0, 1.0]区间这与后面微调时--regression-target的 MSE 损失设计保持一致。预处理成功后会在当前目录生成RTE-bin/内含input0、input1、label三个子目录等任务 bin 目录供第 3 步fairseq-train直接读取。四、第 3 步微调 BART——以 RTE 为例以下是原文档给出的 RTE 任务完整微调命令10 个 epoch、bsz 16TOTAL_NUM_UPDATES2036 # 10 epochs through RTE for bsz 16 WARMUP_UPDATES61 # 6 percent of the number of updates LR1e-05 # Peak LR for polynomial LR scheduler. 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;4.1 关键参数逐项解析参数取值作用--task sentence_prediction固定启用句子/句子对分类任务对应 fairseq/tasks/sentence_prediction.py 中注册的SentencePredictionTask--restore-file $BART_PATH预训练模型加载bart.large400M 参数12 层 encoder/decoder权重--init-token 00在每个 batch 样本开头添加 token 0s实现句子对拼接时的分隔--add-prev-output-tokens开为 encoder-decoder 架构在 sample 中生成prev_output_tokens见任务配置中的同名选项--num-classes 2随任务变化分类头输出维度RTE 为二分类entailment/not entailment--criterion sentence_prediction固定分类/回归损失函数对应 fairseq/criterions/sentence_prediction.py--lr-scheduler polynomial_decay固定多项式衰减学习率调度器--lr为峰值学习率--best-checkpoint-metric accuracy--maximize-best-checkpoint-metric固定按验证准确率保存最优 checkpointcheckpoint_best.pt--find-unused-parameters开兼容 BART 编码器在分类任务下部分参数不参与反向传播的情况4.2sentence_prediction任务与准则的源码视角从 sentence_prediction.py 任务实现 可以看到SentencePredictionConfig暴露了num_classes、init_token、separator_token、add_prev_output_tokens、max_positions默认 512等配置项并在setup_task中断言cfg.num_classes 0即必须显式指定类别数。任务会加载input0/dict.txt作为输入字典并对每个输入样本执行加前缀 tokenPrependTokenDataset与偏移等操作。在 sentence_prediction 准则源码 中分类头名称默认sentence_classification_head可通过--classification-head-name覆盖非回归时计算F.log_softmaxF.nll_loss回归时--regression-target计算F.mse_loss训练日志中通过ncorrect / nsentences统计并上报准确率这正是--best-checkpoint-metric accuracy的数据来源。分类头本身的挂接逻辑在 edgelm/fairseq/models/bart/model.pyBARTModel 内部维护self.classification_heads nn.ModuleDict()register_classification_head会为每个新任务动态注册一个BARTClassificationHead含 dense 层 out_proj并在加载 checkpoint 时按classification_heads.name.*前缀区分哪些参数需要重新初始化、哪些是预训练权重从而保证「BART 骨干冻结复用、分类头从零训练」。五、各 GLUE 任务的微调超参数总表微调时需根据任务修改如下命令行参数原文档完整表格ModelMNLIQNLIQQPRTESST-2MRPCCoLASTS-B--num-classes32222221--lr5e-61e-51e-51e-55e-62e-52e-52e-5bsz128323232128646432--total-num-update309683311211327210185233114813341799--warmup-updates185819866796613146880107解读要点--num-classesMNLI 三分类contradiction/neutral/entailmentSTS-B 为回归输出 1 个连续值其余任务为二分类bsz即上节命令中的--batch-size如 RTE 用 16 时total-num-update相应调整为 2036表中 1018 对应 bsz 32学习率整体很小5e-6 ~ 2e-5因为只做判别式微调过大的学习率会破坏预训练表示。六、STS-B 回归任务的特殊处理对于STS-B除了按上表设置--num-classes 1、--lr 2e-5、--total-num-update 1799等参数外还需要追加--regression-target --best-checkpoint-metric loss移除--maximize-best-checkpoint-metric即改为最小化验证 loss 来选择最优 checkpoint。这对应源码中回归分支的行为SentencePredictionCriterion在regression_targetTrue时改用 MSE 损失criterions/sentence_prediction.py且不再记录ncorrect准确率因此最优模型只能依据 loss 挑选。此外预处理阶段 STS-B 标签已缩放至[0.0, 1.0]与 MSE 损失的数值尺度匹配。七、注意事项与调参建议原文档 Notea)--total-num-updates的由来该值由polynomial_decay调度器使用按--max-epoch10与--batch-size32/64/128视任务而定计算得出。换用其他 epoch 数或 batch size 时需要按公式「epoch × 训练样本数 ÷ batch size」重新计算并同步调整 warmup通常为总更新数的 6%。b)显存与--update-freq的组合上表超参数均在 NvidiaV100 32GB显存上验证通过。如果你的 GPU 显存更小可以通过增大--update-freq、减小--batch-size来保持等效 batch size梯度累积例如 batch size 减半、--update-freq 2保证训练总更新数与学习率曲线不变。c) 微调前务必使用--reset-optimizer --reset-dataloader --reset-meters避免把预训练阶段或之前任务的优化器状态、数据加载器状态与统计量带到新任务中。八、GLUE 任务推理与准确率评估训练完成后最优权重保存在checkpoints/checkpoint_best.pt。原文档给出的 RTE 推理评估代码如下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))代码要点BARTModel.from_pretrained同时加载模型权重checkpoint_filecheckpoint_best.pt与任务数据字典data_name_or_pathRTE-bin推理时bart.encode(sent1, sent2)会对句子对做 BPE 编码并拼接bart.predict(sentence_classification_head, tokens)返回 logitsargmax()得到类别索引label_fn通过label_dictionary.string将索引还原为字符串标签如entailment、not_entailment与dev.tsv中第 4 列的目标标签比对统计准确率。若已微调bart.large.mnli这类官方模型也可参考 edgelm/examples/bart/README.md 中 MNLI 的评估方式定义label_map {0: contradiction, 1: neutral, 2: entailment}遍历glue_data/MNLI/dev_matched.tsv句子在第 8、9 列标签在最后一列得到约 0.9010 的匹配准确率对于 MNLI 这种 3 分类任务还需先通过bart.register_classification_head(mnli, num_classes3)若模型未内置该头再执行bart.predict(mnli, tokens)。九、总结一条从数据到精度的完整链路在 unilm 仓库的 edgelm 子项目中微调 BART 到 GLUE 的完整链路为download_glue_data.py取数 → preprocess_GLUE_tasks.sh 完成列抽取、BPE 编码与 bin 化含 MNLI 四 split、STS-B 标签缩放两个特例→fairseq-train以sentence_prediction任务 多项式衰减学习率微调 →BARTModel.from_pretrained加载checkpoint_best.pt离线推理评估。其中任务tasks/sentence_prediction.py、准则criterions/sentence_prediction.py与分类头注册models/bart/model.py三段源码相互印证解释了--num-classes、--regression-target、--best-checkpoint-metric等参数在底层如何驱动训练与选点。按照上述命令与超参数表逐任务执行即可复现原文档所述的 BART GLUE 微调流程。【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表