2026/9/13 14:13:36

用 fairseq 微调 RoBERTa 完成自定义分类任务:以 IMDB 情感分类为例的完整实战指南

用 fairseq 微调 RoBERTa 完成自定义分类任务:以 IMDB 情感分类为例的完整实战指南 用 fairseq 微调 RoBERTa 完成自定义分类任务以 IMDB 情感分类为例的完整实战指南【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm导读本文以 infoxlm/fairseq/examples/roberta/README.custom_classification.md 为主线讲解如何使用 fairseq 在自定义分类任务上微调 RoBERTa 预训练模型。教程以 IMDB 影评情感二分类positive/negative为具体案例完整覆盖从原始数据获取、格式转换、GPT-2 BPE 编码、fairseq-preprocess二进制预处理到启动训练的五个核心步骤。读者学完后可以将同样的流程迁移到任意单句分类 / 句对分类 / 回归打分类型的任务上。文中所有参数说明均结合仓库源码任务与损失实现逐项展开确保命令可直接复制运行。说明本文涉及的 fairseq 代码快照位于 infoxlm/fairseq 目录下RoBERTa 相关示例集中在 infoxlm/fairseq/examples/roberta/预训练与下游任务示例GLUE、RACE、CommonsenseQA、WSC 等的完整介绍见 README.md。1. 总体流程与前置准备把 RoBERTa 微调到任意分类任务上本质是把任务数据组织成 fairseq 能读的格式然后在预训练 checkpoint 基础上继续训练一个随机初始化的分类头。整体分五步步骤目标产物1. 获取数据下载并解压原始语料原始文本目录2. 格式化数据每个样本一行文本与标签分文件存放train.input0、train.label、dev.input0、dev.label3. BPE 编码用 GPT-2 BPE 把文本行转成 token id 序列*.input0.bpe4. 预处理fairseq-preprocess生成二进制的 indexed datasetIMDB-bin/input0、IMDB-bin/label5. 启动训练train.py加载预训练权重并微调分类头微调后的 checkpoint预训练模型从 RoBERTa 官方发布中选择 checkpoint本仓库 README.md 列出了两个规模roberta.baseBERT-base 结构约 125M 参数roberta.largeBERT-large 结构约 355M 参数本文示例命令即使用该规格。教程示例在单张 NVIDIA V100 32GB 上完成训练roberta_large架构配合--max-tokens 4400的显存预算正是按该显卡规格设计的。2. 第 1 步获取数据以 IMDB 为例IMDB 数据集按训练集/测试集 × 情感类别分目录存放每个影评是一个独立的.txt文件# 下载原始压缩包并解压数据集下载地址见原文档此处以占位符代替 URL wget -O aclImdb_v1.tar.gz 数据集下载地址 tar zxvf aclImdb_v1.tar.gz解压后目录结构大致为aclImdb/train/pos/*.txt、aclImdb/train/neg/*.txt、aclImdb/test/pos/*.txt、aclImdb/test/neg/*.txt。任何每条样本一个文件格式的数据集都可以用同样的方式组织。3. 第 2 步格式化数据文本行与标签行分文件IMDB 原始格式是一个样本一个文件而 fairseq 希望一个样本一行。下面的 Python 脚本把pos/neg两类文件合并、打乱并分别写出文本文件.input0与标签文件.labelimport argparse import os import random from glob import glob random.seed(0) def main(args): for split in [train, test]: samples [] for class_label in [pos, neg]: fnames glob(os.path.join(args.datadir, split, class_label) /*.txt) for fname in fnames: with open(fname) as fin: line fin.readline() samples.append((line, 1 if class_label pos else 0)) random.shuffle(samples) out_fname train if split train else dev f1 open(os.path.join(args.datadir, out_fname .input0), w) f2 open(os.path.join(args.datadir, out_fname .label), w) for sample in samples: f1.write(sample[0] \n) f2.write(str(sample[1]) \n) f1.close() f2.close() if __name__ __main__: parser argparse.ArgumentParser() parser.add_argument(--datadir, defaultaclImdb) args parser.parse_args() main(args)要点说明训练/验证划分脚本把train拆出train.input0/train.label把test命名为dev.input0/dev.label便于 fairseq 的--trainpref/--validpref使用标签编码pos记为1neg记为0逐行写入标签文件文本长度脚本只读取每个文件的第一行fin.readline()因为 IMDB 影评是多行文本若你的数据本身就是单行直接整行写入即可。从后续任务实现 fairseq/tasks/sentence_prediction.py 看这种文本与标签各自成文件、逐行一一对应的布局是该任务的标准输入约定也是第 4 步能够被正确读取的前提。4. 第 3 步BPE 编码RoBERTa 使用 GPT-2 的 BPE 子词词表因此文本必须先经过 BPE 编码才能送入模型。仓库提供了多进程编码脚本 multiprocessing_bpe_encoder.py用法如下# 下载 BPE 词表文件encoder.json 与 vocab.bpe下载地址见脚本 docstring 与原文档此处以占位符代替 wget -N encoder.json 下载地址 wget -N vocab.bpe 下载地址 for SPLIT in train dev; do python -m examples.roberta.multiprocessing_bpe_encoder \ --encoder-json encoder.json \ --vocab-bpe vocab.bpe \ --inputs aclImdb/$SPLIT.input0 \ --outputs aclImdb/$SPLIT.input0.bpe \ --workers 60 \ --keep-empty done为什么单独一步做编码脚本 docstring 中说明使用多进程辅助编码原始文本把整份语料批量编码比在第 2 步逐样本编码更快--workers 60即并行度。结合脚本源码multiprocessing_bpe_encoder.py各参数含义如下参数默认值说明--encoder-json必填GPT-2 BPE 的encoder.jsontoken→id 映射路径--vocab-bpe必填合并规则文件vocab.bpe路径--inputs-stdin输入文件列表可同时传入多个文件多个文件的行将逐行打包一起编码--outputs-stdout输出文件列表数量必须与--inputs一致脚本中有assert校验--keep-empty关闭是否保留空行不指定时遇到空行会把整批过滤掉encode_lines返回EMPTY状态并丢弃--workers20编码进程数示例中用 60 加速编码核心逻辑encode_lines逐行strip()后调用fairseq.data.encoders.gpt2_bpe.get_encoder得到 BPE 编码器把文本转成以空格分隔的 token id 序列写入输出文件。这一格式正是后续fairseq-preprocess与词典匹配所需的输入形态。5. 第 4 步预处理为二进制数据fairseq-preprocessBPE 编码后仍是纯文本需要用fairseq-preprocess转成 fairseq 的二进制 indexed dataset。文本与标签需要分别预处理到两个独立目录# 下载 fairseq 词典GPT-2 BPE 对应的 dict.txt下载地址以占位符代替 wget -N dict.txt 下载地址 fairseq-preprocess \ --only-source \ --trainpref aclImdb/train.input0.bpe \ --validpref aclImdb/dev.input0.bpe \ --destdir IMDB-bin/input0 \ --workers 60 \ --srcdict dict.txt fairseq-preprocess \ --only-source \ --trainpref aclImdb/train.label \ --validpref aclImdb/dev.label \ --destdir IMDB-bin/label \ --workers 60参数说明--only-source单语source-only预处理不涉及目标语言--trainpref/--validpref训练/验证前缀脚本会补齐.input0.bpe等后缀读取--destdir输出目录。文本进入IMDB-bin/input0标签进入IMDB-bin/label这个input0 / label 子目录布局与任务实现的目录约定一一对应见第 7 节--srcdict dict.txt文本侧显式指定 GPT-2 BPE 词典保证与预训练词表一致标签侧不指定fairseq-preprocess会自动从标签数据生成 label 词典两个整数 0/1加上/s等特殊符号。6. 第 5 步启动训练预处理完成后即可微调。完整训练命令如下原文档命令参数含义见后文逐项说明TOTAL_NUM_UPDATES7812 # 10 epochs through IMDB for bsz 32 WARMUP_UPDATES469 # 6 percent of the number of updates LR1e-05 # Peak LR for polynomial LR scheduler. NUM_CLASSES2 MAX_SENTENCES8 # Batch size. ROBERTA_PATH/path/to/roberta/model.pt CUDA_VISIBLE_DEVICES0 python train.py IMDB-bin/ \ --restore-file $ROBERTA_PATH \ --max-positions 512 \ --max-sentences $MAX_SENTENCES \ --max-tokens 4400 \ --task sentence_prediction \ --reset-optimizer --reset-dataloader --reset-meters \ --required-batch-size-multiple 1 \ --init-token 0 --separator-token 2 \ --arch roberta_large \ --criterion sentence_prediction \ --num-classes $NUM_CLASSES \ --dropout 0.1 --attention-dropout 0.1 \ --weight-decay 0.1 --optimizer adam --adam-betas (0.9, 0.98) --adam-eps 1e-06 \ --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 \ --best-checkpoint-metric accuracy --maximize-best-checkpoint-metric \ --truncate-sequence \ --find-unused-parameters \ --update-freq 46.1 关键超参数逐项解读参数示例值作用与说明--task sentence_prediction-指定句子或句对分类/回归任务对应 SentencePredictionTask--criterion sentence_prediction-对应 SentencePredictionCriterion分类用 NLL 损失、回归用 MSE 损失并统计 accuracy--num-classes 22分类类别数任务层会据此注册分类头--restore-file-预训练 checkpoint 路径微调关键需配合下方三个--reset-*--reset-optimizer --reset-dataloader --reset-meters-丢弃 checkpoint 中的优化器状态、数据加载器状态与统计信息仅保留模型权重是从预训练模型开始微调的必要配置--init-token 00每个样本开头插入的 token id0 对应sRoBERTa 的句子起始符--separator-token 22句对输入之间插入的 token id2 对应/s。单句任务只用init-token句对任务两者都会用到见第 7 节源码分析--arch roberta_large-模型架构与 checkpoint 匹配也可用roberta_base--max-positions 512512序列最大长度同时用于--truncate-sequence的截断上限--max-sentences 8/--update-freq 48 / 4每步 8 个句子 × 梯度累积 4 次 有效 batch size 32这正是TOTAL_NUM_UPDATES7812的换算基准--max-tokens 44004400每批 token 上限与--max-sentences共同约束 batch 大小适配 32GB 显存--lr 1e-051e-05多项式衰减调度器的峰值学习率微调 RoBERTa 一般用远低于预训练的学习率--total-num-update 78127812总更新步数25,000 训练样本 ÷ 32 ≈ 781 步/epoch × 10 epoch ≈ 7812--warmup-updates 469469学习率 warmup 步数约为总更新数的 6%--lr-scheduler polynomial_decay-多项式衰减调度器--fp16系列-混合精度训练--fp16-init-scale 4初始动态损失缩放、--threshold-loss-scale 1阈值、--fp16-scale-window 128缩放窗口--max-epoch 1010最多训练 10 个 epoch--best-checkpoint-metric accuracy --maximize-best-checkpoint-metric-以验证集 accuracy 为保存 checkpoint 的选优指标并取最大化--truncate-sequence-超长序列截断到--max-positions配合--max-positions 512使用--find-unused-parameters-由于部分预训练参数在分类微调时不参与梯度计算需要开启以防 DDP 报错--clip-norm 0.00.0梯度裁剪阈值0 表示不裁剪预期效果按上述配置在单张 NVIDIA V100 32GB 上训练 10 个 epochbest-validation-accuracy约为~96.5%数据来自原文档的实测说明。7. 源码级原理训练背后发生了什么7.1 任务层SentencePredictionTask 如何组装数据sentence_prediction.pytask 注册了sentence_prediction任务并定义了任务级参数--num-classes默认 -1setup_task中会断言必须大于 0、--init-token、--separator-token、--regression-target、--no-shuffle、--truncate-sequence。数据加载逻辑load_dataset揭示了目录约定的底层实现分别读取{data}/input0/{split}与可选的{data}/input1/{split}两个 indexed dataset若指定--init-token用PrependTokenDataset在input0开头插入sid0若存在input1句对任务用PrependTokenDataset在句对第二句前插入--separator-tokenid2再用ConcatSentencesDataset拼接成s 句A /s 句B若开启--truncate-sequence用TruncateDataset截断到args.max_positions标签侧通过StripTokenDataset去掉/s与OffsetTokensDataset减去nspecial把 label 词典 id 还原为原始标签数值 0/1。此外build_model 在构建模型后会自动调用model.register_classification_head(sentence_classification_head, num_classes...)即分类头由任务层按--num-classes自动挂载到预训练模型之上。7.2 损失层SentencePredictionCriterion 如何计算 loss 与 accuracysentence_prediction.pycriterion 的forward逻辑以features_onlyTrue提取编码器特征并通过classification_head_namesentence_classification_head拿到分类 logits分类任务使用F.nll_loss(F.log_softmax(logits))回归任务--regression-target则退化为F.mse_loss分类时额外统计ncorrect预测与标签相等的个数聚合阶段aggregate_logging_outputs计算accuracy ncorrect / nsentences——这正是训练命令中--best-checkpoint-metric accuracy的来源。7.3 模型层RobertaClassificationHead 的注册与推理model.py 中的register_classification_head把分类头保存进self.classification_heads该头由RobertaClassificationHead实现一个全连接dense 池化激活 out_proj其中隐藏维度与 dropout 由--pooler-activation-fn、--pooler-dropout等预训练配置决定前向时若传入classification_head_name则先取编码器特征再过分类头forward。微调完成后可通过 hub_interface.py 暴露的接口做推理register_classification_head(name, num_classes)为模型注册分类头predict(head, tokens)提取特征并返回 log-softmax 概率参考用法见 README.md 的 Register a new classification head 一节。8. 扩展把流程迁移到其他分类任务8.1 单句分类如 SST-2 情感与 IMDB 完全一致一个input0目录存放编码后的句子一个label目录存放整数标签命令仅需调整--num-classes与数据路径。8.2 句对分类如 MNLI、QQP在格式化阶段额外生成第二个文本文件dev/input1以及对应的train.input1预处理时再补一条命令把input1也预处理到IMDB-bin/input1。训练命令保留--init-token 0 --separator-token 2任务层会自动完成s A /s B的拼接见 7.1 节。GLUE 全量任务的预处理脚本见 preprocess_GLUE_tasks.sh微调示例见 README.glue.md。8.3 回归任务在训练命令中加入--regression-target任务层参数见 sentence_prediction.py损失会自动切换为 MSE见 7.2 节标签文件写入浮点数即可任务层通过RawLabelDataset读取sentence_prediction.py。8.4 仓库中的其他微调范例Commonsense QA带选项的多选问答任务含自定义 task 实现与数据下载脚本WSCWinoGrande 指代消解自定义 criterion 与 task配合--user-dir加载扩展RACE阅读理解任务配套预处理脚本 preprocess_RACE.py。9. 常见问题排查--total-num-update如何计算训练样本数 ÷ 有效 batch size × epoch 数有效 batch size --max-sentences × --update-freq。更换数据集时务必重新计算否则--max-epoch与调度器步数不匹配换用roberta.base的注意点--arch改为roberta_base且 base 的encoder_embed_dim为 768large 为 1024分类头内维度会随之变化见 register_classification_head 中encoder_embed_dim的传入显存不足优先降低--max-tokens或--max-sentences并用--update-freq补偿 batch sizefp16 相关参数--fp16-init-scale等用于缓解低精度训练中的损失缩放问题分类头重复注册警告若 checkpoint 已含同名分类头且类别数不同模型会打印 re-register 警告并覆盖model.py属预期行为。至此一条任意分类任务 → RoBERTa 微调的完整流水线已经打通从原始语料出发经过格式化、BPE 编码、二进制预处理三个数据环节最后用一行可复现的训练命令完成微调并可通过--best-checkpoint-metric accuracy自动挑选验证集最优 checkpoint。【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考