news 2026/9/15 13:05:09

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
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同时继承自ABCtransformers.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

其调用链在源码中清晰可见:

  1. main.py 使用HfArgumentParser将三类参数解析为 dataclass:RerankerModelArguments(模型/LoRA)、AbsRerankerDataArguments(数据)、AbsRerankerTrainingArguments(训练),然后实例化DecoderOnlyRerankerRunner并调用runner.run()
  2. 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 )
  1. 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是整个训练流程的收尾环节,其保存逻辑分为三步:

  1. 保存模型:先校验模型对象是否具备save接口,若不存在则抛出NotImplementedError;否则调用self.model.save(output_dir)。对于CrossDecoderModel而言,其save实现在 AbsModeling.py 中——先将参数克隆到 CPU,再通过save_pretrained(output_dir, state_dict=...)写出;
  2. 保存 tokenizer:当 tokenizer 存在且为全局 0 号进程时,执行tokenizer.save_pretrained(output_dir),确保推理时可复用一致的词表与特殊 token 配置;
  3. 保存训练参数:通过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_scoressoftmax后作为软标签,叠加一项 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_loraTrue是否使用 LoRA 参数高效微调
lora_rank64LoRA 秩
lora_alpha16LoRA 缩放参数
lora_dropout0.1LoRA 模块 dropout
target_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_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_dataNone一个或多个训练数据路径(.json/.jsonl
train_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 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

  1. 重新加载基础模型与 config;
  2. output_dir加载 PEFT adapter(若根目录下找不到,则通过find_largest_checkpoint自动定位步骤号最大的checkpoint-*目录);
  3. 执行merge_and_unload()将 LoRA 权重合并回基座;
  4. 将合并后的全量模型与 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 重排微调的训练器,通过继承AbsRerankerTrainertransformers.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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/15 13:02:47

三调图斑尖锐角与小缝隙的FME自动化检测修复实战

前几年做三调数据入库和质检的时候,最磨人的不是地类认定,也不是权属争议,反而是看起来特别不起眼的图形质量问题。尖锐角、小缝隙这两个词,干过三调的兄弟应该都不陌生——你辛辛苦苦把图斑矢量化完,一跑质检软件&…

作者头像 李华
网站建设 2026/9/15 13:02:22

EPS中用DOM与DSM创建垂直模型全流程解析

上周一个做测绘的朋友给我打电话,说手上接了个城市更新项目,甲方给的数据很朴素——一套0.2米分辨率的DOM加一套1米分辨率的DSM,要求在一个月内把核心区域的建筑垂直模型拉起来,用于方案比选和指标量算。他问我:在EPS里…

作者头像 李华
网站建设 2026/9/15 12:58:49

JWT认证:现代API安全与性能优化实践

1. 为什么现代API需要JWT保护?三年前我接手过一个电商项目,当时采用传统的Session-Cookie机制做用户认证。某个促销日系统崩溃后排查发现,服务器内存被海量Session数据撑爆。那次惨痛经历让我彻底转向了JWT(JSON Web Token&#x…

作者头像 李华