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_outputs=True时额外返回模型输出。
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将三类参数解析为 dataclass:RerankerModelArguments(模型/LoRA)、AbsRerankerDataArguments(数据)、AbsRerankerTrainingArguments(训练),然后实例化DecoderOnlyRerankerRunner并调用runner.run(); - runner.py 的
load_tokenizer_and_model()完成 tokenizer 与CrossDecoderModel的构建,load_trainer()则创建本文主角:
trainer = DecoderOnlyRerankerTrainer( model=self.model, args=self.training_args, train_dataset=self.train_dataset, data_collator=self.data_collator, tokenizer=self.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 adapter(adapter_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_distillation=True),则用教师模型给出的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_tokens=False)求得),将“最后一个 token 预测 Yes 的概率”作为相关性分数。
配套参数详解:模型、LoRA、数据与训练
模型与 LoRA 参数(RerankerModelArguments)
定义于 arguments.py,核心字段如下:
| 参数 | 默认值 | 说明 |
|---|---|---|
model_name_or_path | 必填 | 初始化的模型 checkpoint |
model_type | 'encoder' | 微调类型,decoder 重排需设为'decoder' |
use_lora | True | 是否使用 LoRA 参数高效微调 |
lora_rank | 64 | LoRA 秩 |
lora_alpha | 16 | LoRA 缩放参数 |
lora_dropout | 0.1 | LoRA 模块 dropout |
target_modules | ['v_proj','q_proj','k_proj','gate_proj','down_proj','o_proj','up_proj'] | 注入 LoRA 的目标模块 |
modules_to_save | None | 需要额外保存在最终 checkpoint 中的模块 |
use_flash_attn | False | 是否启用 Flash Attention 2 加速训练 |
from_peft | None | 加载已有 PEFT adapter 继续训练 |
raw_peft | None | 多个原始 PEFT 路径,先合并再训练 |
save_merged_lora_model | False | 训练结束后合并 LoRA 并保存全量模型 |
在 load_model.py 中,use_lora=True时会构建LoraConfig(task_type=TaskType.CAUSAL_LM, r=..., target_modules=..., modules_to_save=..., lora_alpha=..., lora_dropout=...)并调用get_peft_model,同时打印可训练参数量;use_flash_attn=True时加载AutoModelForCausalLM会传入attn_implementation="flash_attention_2"。
数据参数(AbsRerankerDataArguments)
定义于 AbsArguments.py,关键项:
| 参数 | 默认值 | 说明 |
|---|---|---|
train_data | None | 一个或多个训练数据路径(.json/.jsonl) |
train_group_size | 8 | 每组 query 对应的 passage 数量(1 正 + N 负) |
query_max_len | 32 | query 最大长度 |
passage_max_len | 128 | passage 最大长度 |
max_len | 512 | 拼接后总序列最大长度 |
pad_to_multiple_of | None | 填充对齐到的倍数(如 8,利于算子优化) |
knowledge_distillation | False | 是否使用pos_scores/neg_scores做蒸馏 |
query_instruction_for_rerank | None | query 侧指令前缀,如'A: ' |
passage_instruction_for_rerank | None | passage 侧指令前缀,如'B: ' |
query/passage_instruction_format | '{}{}' | 指令拼接格式 |
shuffle_ratio | 0.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 size);pos_scores/neg_scores:教师打分,仅在knowledge_distillation=True时必须提供,且必须是数值;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_MODE=disabled train_data="../example_data/prompt_based/examples.jsonl" num_train_epochs=1 per_device_train_batch_size=2 gradient_accumulation_steps=1 train_group_size=8 num_gpus=2 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 \ " cmd="torchrun --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_checkpoint=self.training_args.resume_from_checkpoint)支持从指定checkpoint-*目录恢复,配合--save_steps定期保存即可实现中断恢复; - 多卡训练:示例使用
torchrun --nproc_per_node N启动,tokenizer 统一设置为padding_side='left'(见 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),仅供参考