FlagEmbedding 重排序模型微调之 AbsRerankerRunner 抽象运行器:从参数解析到训练流程的源码级解读
FlagEmbedding 重排序模型微调之 AbsRerankerRunner 抽象运行器从参数解析到训练流程的源码级解读【免费下载链接】FlagEmbeddingRetrieval and Retrieval-augmented LLMs项目地址: https://gitcode.com/GitHub_Trending/fl/FlagEmbedding本篇技术指南聚焦 FlagEmbedding 的abcAbstract Base Class抽象层中AbsRerankerRunner这一重排序Reranker微调运行器。它以统一的参数注入 → 资源加载 → 训练执行三段式流程屏蔽了 encoder-only 与 decoder-only 两类重排序模型的差异是后续DecoderOnlyRerankerRunner与EncoderOnlyRerankerRunner的公共骨架。读完本文你将掌握该抽象类各方法的职责、参数含义与默认值、数据集与 Collator 的调度逻辑并能基于仓库源码与官方示例独立编写一套可运行的 reranker 微调脚本。一、AbsRerankerRunner 在 FlagEmbedding 中的定位在 FlagEmbedding 仓库中重排序模型的微调能力按抽象层 具体实现两层组织抽象层abc位于 FlagEmbedding/abc/finetune/reranker/定义了与具体模型架构无关的通用流程包括AbsArguments参数、AbsDataset数据集与 Collator、AbsModeling模型基类、AbsTrainerTrainer 基类与AbsRunner运行器基类。具体实现层finetune位于 FlagEmbedding/finetune/reranker/分为decoder_only含base与layerwise和encoder_only含base两套体系各自继承抽象层组件并落地具体模型。本篇文章围绕的 API 文档 AbsRunner.rst 所描述的AbsRerankerRunner正是这条流水线的总导演——它在构造函数中依次完成输出目录校验、日志初始化、随机种子设置然后串行调用load_tokenizer_and_model、load_train_dataset、load_data_collator、load_trainer四个加载步骤最后通过run()启动训练。# FlagEmbedding/abc/finetune/reranker/AbsRunner.py class AbsRerankerRunner(ABC): def __init__(self, model_args, data_args, training_args): self.model_args model_args self.data_args data_args self.training_args training_args # 1. 输出目录保护性校验 if (os.path.exists(training_args.output_dir) and os.listdir(training_args.output_dir) and training_args.do_train and not training_args.overwrite_output_dir): raise ValueError( fOutput directory ({training_args.output_dir}) already exists and is not empty. Use --overwrite_output_dir to overcome. ) # 2. 日志与训练环境信息 logging.basicConfig( format%(asctime)s - %(levelname)s - %(name)s - %(message)s, datefmt%m/%d/%Y %H:%M:%S, levellogging.INFO if training_args.local_rank in [-1, 0] else logging.WARN, ) # 3. 固定随机种子transformers.set_seed set_seed(training_args.seed) # 4. 串行加载各组件 self.tokenizer, self.model self.load_tokenizer_and_model() self.train_dataset self.load_train_dataset() self.data_collator self.load_data_collator() self.trainer self.load_trainer()从__init__的执行顺序可以看出该抽象类对派生类的约束非常明确load_tokenizer_and_model与load_trainer是抽象方法abstractmethod必须由具体实现而load_train_dataset、load_data_collator与run提供了默认实现子类可按需覆盖。二、构造函数中的三件事目录校验、日志与随机种子2.1 输出目录防覆盖保护构造函数首先检查training_args.output_dir当目录已存在、非空、且开启了do_train但未设置overwrite_output_dir时直接抛ValueError提示使用--overwrite_output_dir。这一设计防止了误覆盖已有微调结果是运行脚本前最常见的报错来源之一。2.2 分布式环境感知的日志级别日志级别依据training_args.local_rank判断主进程local_rank in [-1, 0]输出INFO其余分布式子进程降为WARN避免多卡训练时日志刷屏。随后打印进程秩、设备、GPU 数、是否分布式训练、是否 16-bit 训练fp16以及三组参数对象的完整内容方便训练前核对配置。2.3 种子固定保证可复现通过transformers.set_seed(training_args.seed)一次性固定 PyTorch、NumPy 与 Python 随机模块的种子。随机种子同时影响数据采样如正例随机选取、负例采样与模型初始化是复现实验的关键。三、五个核心方法逐一拆解3.1load_tokenizer_and_model抽象方法签名要求返回Tuple[PreTrainedTokenizer, AbsRerankerModel]即加载分词器与重排序模型。该方法是唯一的强约束抽象方法之一具体实现因模型类型而异encoder 实现encoder_only/base/runner.py使用AutoTokenizerAutoConfignum_labels1AutoModelForSequenceClassification再包装为CrossEncoderModeldecoder 实现decoder_only/base/runner.py使用AutoTokenizer并补齐pad_token优先取unk_token其次eod兜底eos_token且显式设置padding_sideleft再通过get_model加载基座并包装为CrossDecoderModel。值得注意的是两个具体实现的load_tokenizer_and_model都读取self.training_args.gradient_checkpointing若开启则调用model.enable_input_require_grads()确保梯度检查点下输入嵌入仍可计算梯度这是 LoRA 微调 LLM 重排序模型时的常见坑。3.2load_train_dataset默认实现按model_type分发def load_train_dataset(self): if self.model_args.model_type encoder: train_dataset AbsRerankerTrainDataset(argsself.data_args, tokenizerself.tokenizer) else: train_dataset AbsLLMRerankerTrainDataset(argsself.data_args, tokenizerself.tokenizer) return train_dataset从源码可以看出数据集选择完全由model_args.model_type决定model_typeencoder→AbsRerankerTrainDataset交叉编码器使用的简单 query-passage 对其他取值如decoder→AbsLLMRerankerTrainDataset该子类在构造时预计算了sep_token的 token idsep_token默认\n用于分隔 query 与 passage并在__getitem__中拼接[BOS] query sep passage sep prompt的完整 LLM 提示模板。默认 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.且每条数据可用自有字段prompt覆盖。在AbsRerankerTrainDataset.__init__AbsDataset.py中train_data既支持单个/多个.json、.jsonl文件路径也支持目录路径自动遍历目录内所有 json/jsonl 文件最后通过datasets.concatenate_datasets合并并对每个数据集按max_example_num_per_dataset默认100000000随机抽样控制规模。3.3load_data_collator默认实现同样按model_type分发def load_data_collator(self): if self.model_args.model_type encoder: RerankerCollator AbsRerankerCollator else: RerankerCollator AbsLLMRerankerCollator data_collator RerankerCollator( tokenizerself.tokenizer, query_max_lenself.data_args.query_max_len, passage_max_lenself.data_args.passage_max_len, pad_to_multiple_ofself.data_args.pad_to_multiple_of, paddingTrue, return_tensorspt ) return data_collatorCollator 的职责是摊平 填充由于每个样本包含train_group_size条 passage1 正 n 负Collator 将 list 展平后统一tokenizer.pad并把返回结果组织为{pair: collated, teacher_scores: teacher_scores}字典——teacher_scores是知识蒸馏KD场景下教师模型打分列表非 KD 时为None。AbsLLMRerankerCollator额外继承DataCollatorForSeq2Seq的处理逻辑若样本含labels会先按批次内最大长度补齐支持pad_to_multiple_of对齐再交给tokenizer.pad保证 seq2seq 类标签可以正确组装。3.4load_trainer抽象方法第二个抽象方法负责装配AbsRerankerTrainer。该 Trainer 基类AbsTrainer.py继承自transformers.Trainer核心逻辑有两处抽象方法_save交由子类实现用于自定义 checkpoint 保存行为compute_loss直接调用model(**inputs)并取outputs.loss将损失计算完全下沉到模型层。而模型层的损失计算定义在 AbsModeling.py 的AbsRerankerModel.forward中将train_batch_size * train_group_size条打分重排为(train_batch_size, -1)用CrossEntropyLoss让正例得分排第一若启用知识蒸馏还会叠加一项 KL 散度项-mean(sum(log_softmax(logits) * softmax(teacher_scores)))。3.5run默认实现训练入口def run(self): Path(self.training_args.output_dir).mkdir(parentsTrue, exist_okTrue) self.trainer.train(resume_from_checkpointself.training_args.resume_from_checkpoint) self.trainer.save_model()默认流程即建目录 → 训练支持从 checkpoint 恢复→ 保存模型。decoder 子类DecoderOnlyRerankerRunner在 runner.py 中覆盖了run在默认流程之后追加了合并 LoRA 权重步骤当save_merged_lora_modelTrue且为主进程时调用save_merged_model将 LoRA 适配器合并回基座模型并保存完整权重。四、参数体系三组 Arguments 的字段全景运行器依赖三组参数对象全部定义在 AbsArguments.py 中命令行通过HfArgumentParser解析见 decoder_only/base/main.py。4.1AbsRerankerModelArguments模型参数参数默认值说明model_name_or_path必填初始化模型的 checkpoint 名或路径config_nameNone与模型名不同时的预训练 config 名/路径tokenizer_nameNone与模型名不同时的分词器名/路径cache_dirNone预训练模型下载缓存目录trust_remote_codeFalse是否信任远程代码model_typeencoder微调类型可选[encoder, decoder]决定数据集与 Collator 的分发use_fast_tokenizerTrue是否使用 fast tokenizertokenos.getenv(HF_TOKEN)访问受限模型时使用的 Hugging Face tokendecoder 专属子类RerankerModelArgumentsdecoder_only/base/arguments.py额外提供 LoRA 相关参数use_loraTrue、lora_rank64、lora_alpha16、lora_dropout0.1、target_modules默认[v_proj,q_proj,k_proj,gate_proj,down_proj,o_proj,up_proj]、use_flash_attnFalse、from_peftNone、raw_peftNone、save_merged_lora_modelFalse。4.2AbsRerankerDataArguments数据参数参数默认值说明train_dataNonenargs一个或多个训练数据路径要求每条数据含query: str、pos: List[str]、neg: List[str]cache_pathNone数据缓存目录传给datasets.load_datasettrain_group_size8每组训练的 passage 数量1 正 7 负负例不足时自动重复采样补齐query_max_len32query 分词后的最大长度passage_max_len128passage 分词后的最大长度max_len512总体最大序列长度pad_to_multiple_ofNone若设置则把序列 padding 到该值的整数倍便于高效推理示例中常用 8max_example_num_per_dataset100000000每个数据集最大样本数超出随机抽样query_instruction_for_rerankNonequery 侧指令query_instruction_format{}{}query 指令拼接格式__post_init__会把\\n替换为换行passage_instruction_for_rerankNonepassage 侧指令passage_instruction_format{}{}passage 指令拼接格式knowledge_distillationFalse开启后要求数据含pos_scores: List[float]与neg_scores: List[float]shuffle_ratio0.0对长度 100 的文本按比例随机打乱分块顺序sep_token\nLLM reranker 中区分 query 与 passage 的分隔 token其中两个行为值得注意__post_init__会校验train_data中每个路径真实存在否则抛FileNotFoundError从根上避免静默加载失败开启knowledge_distillation时_load_dataset会强制检查pos_scores/neg_scores列缺失即抛ValueError未开启时则会主动移除这两列。4.3AbsRerankerTrainingArguments训练参数该类继承transformers.TrainingArguments因此天然拥有output_dir、learning_rate、num_train_epochs、per_device_train_batch_size、gradient_accumulation_steps、warmup_ratio、weight_decay、fp16/bf16、gradient_checkpointing、deepspeed、resume_from_checkpoint、overwrite_output_dir、seed等全套 HF 训练参数额外新增的仅有sub_batch_size默认None源码注释标注尚未实现。per_device_train_batch_size会透传给AbsRerankerModel的train_batch_size用于损失函数中的分组视角。五、端到端实战两类重排序模型的标准微调命令5.1 Encoder-only交叉编码器微调参考 encoder_only/base.sh以BAAI/bge-reranker-base为例export WANDB_MODEdisabled train_data../example_data/normal/examples.jsonl num_train_epochs4 per_device_train_batch_size2 gradient_accumulation_steps1 train_group_size8 num_gpus2 torchrun --nproc_per_node $num_gpus \ -m FlagEmbedding.finetune.reranker.encoder_only.base \ --model_name_or_path BAAI/bge-reranker-base \ --train_data $train_data \ --train_group_size $train_group_size \ --query_max_len 256 \ --passage_max_len 256 \ --pad_to_multiple_of 8 \ --knowledge_distillation True \ --output_dir ./test_encoder_only_base_bge-reranker-base \ --overwrite_output_dir \ --learning_rate 6e-5 \ --fp16 \ --num_train_epochs $num_train_epochs \ --per_device_train_batch_size $per_device_train_batch_size \ --gradient_accumulation_steps $gradient_accumulation_steps \ --warmup_ratio 0.1 \ --gradient_checkpointing \ --weight_decay 0.01 \ --deepspeed ../../ds_stage0.json \ --logging_steps 1 \ --save_steps 10005.2 Decoder-onlyLLM 重排序器微调参考 decoder_only/base.sh以BAAI/bge-reranker-v2-gemma为例额外开启 LoRA 与知识蒸馏指令torchrun --nproc_per_node $num_gpus \ -m FlagEmbedding.finetune.reranker.decoder_only.base \ --model_name_or_path BAAI/bge-reranker-v2-gemma \ --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 \ --train_data ../example_data/prompt_based/examples.jsonl \ --train_group_size 8 \ --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 {}{} \ --output_dir ./test_decoder_only_base_bge-reranker-v2-gemma \ --learning_rate 2e-4 \ --bf16 \ --num_train_epochs 1 \ --per_device_train_batch_size 2 \ --warmup_ratio 0.1 \ --gradient_checkpointing \ --weight_decay 0.01 \ --deepspeed ../../ds_stage0.json5.3 训练数据格式训练数据为 JSONL每条包含query、pos正例列表、neg负例列表可选pos_scores/neg_scoresKD、query_prompt/passage_prompt覆盖默认指令、promptLLM reranker 的提示词。示例数据可查看 examples/finetune/reranker/example_data/其中normal/为 encoder 格式prompt_based/为 decoder 提示词格式。六、设计要点小结模板方法模式AbsRerankerRunner通过构造时固定流程 抽象方法留给子类实现统一训练管线新模型只需继承它并实现load_tokenizer_and_model与load_trainer即可复用全部数据加载、Collator 与训练逻辑。model_type单一开关数据集与 Collator 的选择完全由model_type驱动这是抽象层保持一次分发、处处生效的关键设计。知识蒸馏的一等公民从数据列校验、teacher_scores透传到模型层 KD 损失项整条链路在抽象层即已打通。防御性校验输出目录防覆盖、训练数据路径存在性检查、KD 分数列存在性检查将常见配置错误前置到启动阶段。若需在阅读本文后进一步深入可继续查看 AbsDataset.py 中正负例采样与指令拼接的细节、AbsModeling.py 中AbsRerankerModel.forward的损失计算以及文档树 docs/source/API/abc/finetune/reranker/ 下对应的 API 参考页面。【免费下载链接】FlagEmbeddingRetrieval and Retrieval-augmented LLMs项目地址: https://gitcode.com/GitHub_Trending/fl/FlagEmbedding创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考