ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

FlagEmbedding 解码器重排模型微调实战:DecoderOnlyRerankerTrainer 核心机制与 LoRA 训练全流程

2026/9/15 13:05:23 拓冰建站 浏览量
FlagEmbedding 解码器重排模型微调实战:DecoderOnlyRerankerTrainer 核心机制与 LoRA 训练全流程 FlagEmbedding 解码器重排模型微调实战DecoderOnlyRerankerTrainer 核心机制与 LoRA 训练全流程【免费下载链接】FlagEmbeddingRetrieval and Retrieval-augmented LLMs项目地址: https://gitcode.com/GitHub_Trending/fl/FlagEmbedding本篇技术指南围绕 FlagEmbedding 中 decoder-only 重排模型如BAAI/bge-reranker-v2-gemma的微调训练器DecoderOnlyRerankerTrainer展开系统讲解其在整个训练管线中的定位、_save持久化机制、损失计算方式、配套 LoRA 参数与数据格式并结合仓库内的完整启动脚本给出可直接复用的实战方案。读完本文你将掌握从命令行入口到模型落盘、从 LoRA 微调到合并全量模型的一整套 decoder-only reranker 微调方法论。DecoderOnlyRerankerTrainer 在微调框架中的定位DecoderOnlyRerankerTrainer是 FlagEmbedding 中解码器decoder-only重排模型微调的训练器类定义于 trainer.py其继承体系如下transformers.Trainer └── AbsRerankerTrainer # 抽象训练器见 abc/finetune/reranker/AbsTrainer.py └── DecoderOnlyRerankerTrainer # decoder-only base 重排训练器在抽象层 AbsTrainer.py 中AbsRerankerTrainer同时继承自ABC与transformers.Trainer声明了抽象方法_save要求子类实现自定义保存逻辑实现了compute_loss将 Hugging Face Trainer 的损失计算统一为“模型前向输出中的loss字段”并支持return_outputsTrue时额外返回模型输出。DecoderOnlyRerankerTrainer在此基础上仅需专注于_save——即“训练完成后如何把模型、tokenizer 与训练参数正确落盘”其余训练循环、梯度累积、学习率调度、分布式通信等能力全部继承自transformers.Trainer。这也是该模块代码精简但功能完整的设计思路。一条完整的训练调用链从命令行到 Trainer以官方示例脚本 base.sh 为例训练入口为torchrun --nproc_per_node $num_gpus \ -m FlagEmbedding.finetune.reranker.decoder_only.base \ $model_args $data_args $training_args其调用链在源码中清晰可见main.py 使用HfArgumentParser将三类参数解析为 dataclassRerankerModelArguments模型/LoRA、AbsRerankerDataArguments数据、AbsRerankerTrainingArguments训练然后实例化DecoderOnlyRerankerRunner并调用runner.run()runner.py 的load_tokenizer_and_model()完成 tokenizer 与CrossDecoderModel的构建load_trainer()则创建本文主角trainer DecoderOnlyRerankerTrainer( modelself.model, argsself.training_args, train_datasetself.train_dataset, data_collatorself.data_collator, tokenizerself.tokenizer )run()依次执行trainer.train(resume_from_checkpoint...)、trainer.save_model()并在开启--save_merged_lora_model且为主进程process_index 0时调用save_merged_model()输出合并后的全量模型。值得注意的是Runner 在构造阶段继承自 AbsRunner.py还会做输出目录非空校验要求--overwrite_output_dir、日志初始化与随机种子设置随后依次加载 tokenizer、模型、训练数据集、data collator 与 trainer。_save 方法模型、分词器与训练参数的完整落盘DecoderOnlyRerankerTrainer._save是整个训练流程的收尾环节其保存逻辑分为三步保存模型先校验模型对象是否具备save接口若不存在则抛出NotImplementedError否则调用self.model.save(output_dir)。对于CrossDecoderModel而言其save实现在 AbsModeling.py 中——先将参数克隆到 CPU再通过save_pretrained(output_dir, state_dict...)写出保存 tokenizer当 tokenizer 存在且为全局 0 号进程时执行tokenizer.save_pretrained(output_dir)确保推理时可复用一致的词表与特殊 token 配置保存训练参数通过torch.save(self.args, os.path.join(output_dir, training_args.bin))将完整训练参数序列化便于后续复现实验或断点分析。源码中保留了 DeepSpeed ZeRO-3 场景下单独导出 LoRA adapteradapter_model.bin的注释代码表明在常规流程中训练产物即为save_pretrained风格的完整 checkpoint 目录。损失计算Cross-Entropy 与知识蒸馏项训练器的compute_loss直接复用AbsRerankerTrainer的实现outputs model(**inputs) loss outputs.loss return (loss, outputs) if return_outputs else loss而真正的损失计算发生在CrossDecoderModel的前向逻辑中见 AbsModeling.py模型以train_batch_size为单位将同组 query 的多个 passage 打分view(train_batch_size, -1)重组以正样本位置batch 内第 0 位为 target 计算CrossEntropyLoss若启用知识蒸馏knowledge_distillationTrue则用教师模型给出的pos_scores/neg_scores做softmax后作为软标签叠加一项 KL 散度损失loss -mean(sum(log_softmax(logits) * teacher_targets))。评分来源则是 modeling.py 中CrossDecoderModel.encode的实现取序列最后一个位置 logits 中Yestoken 的得分self.yes_loc在 AbsModeling.py 由tokenizer(Yes, add_special_tokensFalse)求得将“最后一个 token 预测 Yes 的概率”作为相关性分数。配套参数详解模型、LoRA、数据与训练模型与 LoRA 参数RerankerModelArguments定义于 arguments.py核心字段如下参数默认值说明model_name_or_path必填初始化的模型 checkpointmodel_typeencoder微调类型decoder 重排需设为decoderuse_loraTrue是否使用 LoRA 参数高效微调lora_rank64LoRA 秩lora_alpha16LoRA 缩放参数lora_dropout0.1LoRA 模块 dropouttarget_modules[v_proj,q_proj,k_proj,gate_proj,down_proj,o_proj,up_proj]注入 LoRA 的目标模块modules_to_saveNone需要额外保存在最终 checkpoint 中的模块use_flash_attnFalse是否启用 Flash Attention 2 加速训练from_peftNone加载已有 PEFT adapter 继续训练raw_peftNone多个原始 PEFT 路径先合并再训练save_merged_lora_modelFalse训练结束后合并 LoRA 并保存全量模型在 load_model.py 中use_loraTrue时会构建LoraConfig(task_typeTaskType.CAUSAL_LM, r..., target_modules..., modules_to_save..., lora_alpha..., lora_dropout...)并调用get_peft_model同时打印可训练参数量use_flash_attnTrue时加载AutoModelForCausalLM会传入attn_implementationflash_attention_2。数据参数AbsRerankerDataArguments定义于 AbsArguments.py关键项参数默认值说明train_dataNone一个或多个训练数据路径.json/.jsonltrain_group_size8每组 query 对应的 passage 数量1 正 N 负query_max_len32query 最大长度passage_max_len128passage 最大长度max_len512拼接后总序列最大长度pad_to_multiple_ofNone填充对齐到的倍数如 8利于算子优化knowledge_distillationFalse是否使用pos_scores/neg_scores做蒸馏query_instruction_for_rerankNonequery 侧指令前缀如A: passage_instruction_for_rerankNonepassage 侧指令前缀如B: query/passage_instruction_format{}{}指令拼接格式shuffle_ratio0.0长文本100 字符随机分块打乱的比例sep_token\n区分 query 与 passage 的分隔符数据加载时AbsDataset.py会校验文件存在性开启蒸馏但数据缺少pos_scores/neg_scores列时会抛出ValueError提示。负样本不足train_group_size - 1时会通过重复采样补齐保证每组样本数量一致。训练参数AbsRerankerTrainingArguments直接继承transformers.TrainingArguments并仅新增sub_batch_size预留字段当前未实现因此学习率、bf16/fp16、梯度累积、warmup、DeepSpeed 等全部复用 HF 标准能力。数据格式与 prompt 构造训练数据为逐行 JSON字段如下{query: str, pos: [...], neg: [..., ...], pos_scores: [96.0], neg_scores: [90.5, ...], prompt: str}query查询文本pos正样本列表每轮随机取 1 条neg负样本列表随机采样补齐 group sizepos_scores/neg_scores教师打分仅在knowledge_distillationTrue时必须提供且必须是数值prompt控制提示词若不提供则使用默认提示Given a query A and a passage B, determine whether the passage contains an answer to the query by providing a prediction of either Yes or No.对于 decoder 重排模型AbsDataset.py 中AbsLLMRerankerTrainDataset会按如下形式构造输入序列[BOS] query [sep] passage [sep] prompt其中sep由--sep_token指定默认\n并按query_max_len passage_max_len对 passage 部分做only_second截断。真实样例可参考 examples.jsonl其中pos_scores/neg_scores即为通过教师重排模型如BAAI/bge-reranker-v2-m3打分得到。完整可运行的训练脚本以下是仓库自带示例 base.sh 的完整内容以BAAI/bge-reranker-v2-gemma为例export WANDB_MODEdisabled train_data../example_data/prompt_based/examples.jsonl num_train_epochs1 per_device_train_batch_size2 gradient_accumulation_steps1 train_group_size8 num_gpus2 model_args\ --model_name_or_path BAAI/bge-reranker-v2-gemma \ --cache_dir $HF_HUB_CACHE \ --use_lora True \ --lora_rank 32 \ --lora_alpha 64 \ --use_flash_attn True \ --target_modules q_proj k_proj v_proj o_proj \ --save_merged_lora_model True \ --model_type decoder \ data_args\ --train_data $train_data \ --cache_path ~/.cache \ --train_group_size $train_group_size \ --query_max_len 512 \ --passage_max_len 512 \ --pad_to_multiple_of 8 \ --knowledge_distillation True \ --query_instruction_for_rerank A: \ --query_instruction_format {}{} \ --passage_instruction_for_rerank B: \ --passage_instruction_format {}{} \ training_args\ --output_dir ./test_decoder_only_base_bge-reranker-v2-gemma \ --overwrite_output_dir \ --learning_rate 2e-4 \ --bf16 \ --num_train_epochs $num_train_epochs \ --per_device_train_batch_size $per_device_train_batch_size \ --gradient_accumulation_steps $gradient_accumulation_steps \ --dataloader_drop_last True \ --warmup_ratio 0.1 \ --gradient_checkpointing \ --weight_decay 0.01 \ --deepspeed ../../ds_stage0.json \ --logging_steps 1 \ --save_steps 1000 \ cmdtorchrun --nproc_per_node $num_gpus \ -m FlagEmbedding.finetune.reranker.decoder_only.base \ $model_args $data_args $training_args \ eval $cmd要点说明LoRA 配置示例将lora_rank设为 32、lora_alpha设为 64仓库默认值分别为 64 与 16可按显存与效果调整target_modules仅注入q/k/v/o_proj四个注意力投影指令格式query 与 passage 分别加A: 、B: 前缀对应数据集中query [sep] passage的判别式任务训练精度decoder 模型示例统一使用--bf16需要 Ampere 及以上架构 GPU--gradient_checkpointing与--dataloader_drop_last True配合降低显存占用DeepSpeed通过--deepspeed ../../ds_stage0.json启用 ZeRO 优化配置见 ds_stage0.json指令微调入口等价命令也可参考 README.md 中bge-reranker-v2-gemma一节直接以torchrun运行。训练产物与模型合并训练结束后输出目录中包含模型权重与 config、tokenizer 文件、training_args.bin以及 HF Trainer 默认产出的checkpoint-{step}中间检查点。若设置了--save_merged_lora_model True主进程会调用 load_model.py 中的save_merged_model重新加载基础模型与 config从output_dir加载 PEFT adapter若根目录下找不到则通过find_largest_checkpoint自动定位步骤号最大的checkpoint-*目录执行merge_and_unload()将 LoRA 权重合并回基座将合并后的全量模型与 tokenizer 保存到output_dir/merged_model子目录。合并后的merged_model可直接被推理侧如FlagEmbedding.inference.reranker加载使用无需额外加载 adapter。若继续使用--from_peft参数则可基于已训练 adapter 进行二次微调--raw_peft则允许先合并一个或多个历史 adapter 再开始训练。断点续训与多卡分布式断点续训run()中通过trainer.train(resume_from_checkpointself.training_args.resume_from_checkpoint)支持从指定checkpoint-*目录恢复配合--save_steps定期保存即可实现中断恢复多卡训练示例使用torchrun --nproc_per_node N启动tokenizer 统一设置为padding_sideleft见 runner.py数据加载与梯度同步由 HF Trainer 与 DeepSpeed 自动处理梯度检查点兼容性开启--gradient_checkpointing时Runner 会调用model.enable_input_require_grads()确保输入嵌入层获得梯度与 LoRA 训练正确协同。小结DecoderOnlyRerankerTrainer作为 FlagEmbedding decoder-only 重排微调的训练器通过继承AbsRerankerTrainer与transformers.Trainer在保留完整 HF 训练生态的同时只暴露了最核心的_save持久化逻辑。搭配CrossDecoderModel的“最后一个 token 预测 Yes”打分机制、LoRA 参数高效微调、知识蒸馏软标签损失以及save_merged_lora_model一键合并能力构成了从数据准备、指令构造、多卡训练到可部署模型产出的完整闭环。相关源码与示例可继续参阅 trainer.py、runner.py、base.sh 与 README.md。【免费下载链接】FlagEmbeddingRetrieval and Retrieval-augmented LLMs项目地址: https://gitcode.com/GitHub_Trending/fl/FlagEmbedding创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考