FlagEmbedding 交叉编码器重排模型微调:CrossEncoderModel 源码级解析与实战指南
发布时间:2026/9/15 13:32:19 锦皓数字建站

FlagEmbedding 交叉编码器重排模型微调CrossEncoderModel 源码级解析与实战指南【免费下载链接】FlagEmbeddingRetrieval and Retrieval-augmented LLMs项目地址: https://gitcode.com/GitHub_Trending/fl/FlagEmbedding导读本文聚焦 FlagEmbedding 开源项目中 encoder-only 重排器Reranker微调链路的模型核心 ——FlagEmbedding.finetune.reranker.encoder_only.base.CrossEncoderModel完整解析其类定义、encode方法实现、底层前向传播与损失计算逻辑并结合 Runner、Trainer 与官方示例脚本给出可直接复跑的微调方案。读完本文你将掌握交叉编码器重排模型在 FlagEmbedding 中的训练原理、参数语义与端到端微调流程能够在自己的数据集上训练 BGE 系列重排模型。一、模块定位encoder-only 重排器微调链路中的建模层在 FlagEmbedding 的微调代码结构中重排器reranker被划分为 decoder-only 与 encoder-only 两大类其中 encoder-only 又细分为base基础路径。CrossEncoderModel正是 base/modeling.py 中定义的模型类它对外呈现为「交叉编码器」形态将 query 与 passage 拼接后一次性送入编码器直接输出相关性 logits而非像双塔嵌入模型那样分别编码再计算相似度。该模块位于微调链路的建模层与同目录下的runner.py加载模型与数据集、trainer.py训练与保存共同构成完整训练管线并统一遵循FlagEmbedding/abc/finetune/reranker/下的抽象基类规范。其在包中的导出关系可见 base/init.py三者一并导出from .modeling import CrossEncoderModel from .runner import EncoderOnlyRerankerRunner from .trainer import EncoderOnlyRerankerTrainer二、CrossEncoderModel类定义与构造参数modeling.py 中的CrossEncoderModel定义极为精简本质是对基类的「薄封装」class CrossEncoderModel(AbsRerankerModel): Model class for reranker. def __init__( self, base_model: PreTrainedModel, tokenizer: AutoTokenizer None, train_batch_size: int 4, ): super().__init__( base_model, tokenizertokenizer, train_batch_sizetrain_batch_size, )三个构造参数的含义与默认值如下表参数类型默认值说明base_modelPreTrainedModel必填底层预训练模型实际承担编码与打分任务。在 Runner 中由AutoModelForSequenceClassification加载得到tokenizerAutoTokenizerNone用于编码输入文本的分词器train_batch_sizeint4训练批次大小其语义是「每个训练样本内 query-passage 分组对应的样本数」直接影响 loss 的分组视角见下文 forward 解析类本身不引入新的可学习参数所有前向、损失、保存逻辑均由抽象基类AbsRerankerModel提供。因此理解该类关键在于理解其继承链。三、encode 方法从输入特征到相关性 logitsencode是本模型对外暴露的核心方法定义见 modeling.pydef encode(self, features): Encodes input features to logits. Args: features (dict): Dictionary with input features. Returns: torch.Tensor: The logits output from the model. return self.model(**features, return_dictTrue).logits实现要点输入features为字典包含input_ids、attention_mask、token_type_ids等由数据整理器collator产出、已被分词并拼接好的模型输入query passage 拼接后的序列处理将features以关键字参数形式直接透传给底层的PreTrainedModelSequenceClassification 模型并开启return_dictTrue获取结构化输出输出返回.logits即模型打出的相关性分数。由于 Runner 加载时固定num_labels1输出形状为(batch_size * group_size, 1)其中group_size即训练样本中 query 对应的文档数量正样本 负样本数见train_group_size参数。由于encode是抽象基类AbsRerankerModel中的抽象方法见 AbsModeling.pyCrossEncoderModel必须实现它这是子类唯一必须补齐的能力点。四、继承体系AbsRerankerModel 的初始化与前向逻辑CrossEncoderModel继承自FlagEmbedding.abc.finetune.reranker.AbsRerankerModel该抽象类实现了完整的训练语义源码位于 AbsModeling.py。4.1 初始化阶段的关键行为构造时基类会完成四件重要的事缓存模型与分词器将base_model存为self.model并将model.config同步为self.config补齐 pad_token若model.config.pad_token_id is None则自动用tokenizer.pad_token_id填充AbsModeling.py避免分组拼接与 pad 时报错内置交叉熵损失self.cross_entropy nn.CrossEntropyLoss(reductionmean)默认 reduction 为均值计算 Yes 的 token 位置self.yes_loc self.tokenizer(Yes, add_special_tokensFalse)[input_ids][-1]供 decoder-only 重排器使用encoder-only 链路不依赖此项。同时基类还透传了gradient_checkpointing_enable与enable_input_require_grads用于配合梯度检查点训练Runner 在开启gradient_checkpointing时即调用后者。4.2 forward 与损失计算forward方法AbsModeling.py定义了每个训练 step 的计算过程ranker_logits self.encode(pair) # (batch_size * group_size, 1) ... if self.training: grouped_logits ranker_logits.view(self.train_batch_size, -1) target torch.zeros(self.train_batch_size, ...) # 正样本永远排在第 0 位 loss self.compute_loss(grouped_logits, target) if teacher_scores is not None: # 知识蒸馏项以 teacher 的 softmax 分数为目标 loss -torch.mean(torch.sum(torch.log_softmax(grouped_logits, dim-1) * teacher_targets, dim-1))核心机制可归纳为三点分组视角ranker_logits被view(self.train_batch_size, -1)重整为「每个样本一行、组内各候选一列」的二维矩阵第 0 列恒为正样本目标构造target为全零向量即始终要求正样本分数最高损失即为组内 Softmax 交叉熵——这是「让正样本排在组内第一」的直接实现知识蒸馏当teacher_scores传入时对应数据参数knowledge_distillationTrue额外累加一项「log-softmax(logits) 与 teacher 概率的逐元素乘积均值」的负值等价于最小化 logits 分布与 teacher 分布的 KL 散度项。这一设计使训练可以兼容教师模型软标签训练数据中带pos_scores/neg_scores字段。compute_lossAbsModeling.py即封装了上述交叉熵。输出统一为RerankerOutputdataclass含loss与scores字段推理阶段loss为None仅返回scores。4.3 保存语义基类提供两种保存途径save(output_dir)将state_dict全部迁移到 CPU 后调用save_pretrained保存save_pretrained(*args, **kwargs)同时保存 tokenizer 与模型先 tokenizer 后 model保证产物可被from_pretrained完整恢复。训练器EncoderOnlyRerankerTrainer的_save正是走这条路径trainer.py并额外将training_args.bin与模型一同落盘。五、模型如何被加载Runner 中的组装逻辑CrossEncoderModel不在用户代码中直接实例化而是由EncoderOnlyRerankerRunner.load_tokenizer_and_model完成组装见 runner.pytokenizer AutoTokenizer.from_pretrained(self.model_args.model_name_or_path, ...) num_labels 1 config AutoConfig.from_pretrained( self.model_args.config_name if self.model_args.config_name else self.model_args.model_name_or_path, num_labelsnum_labels, ...) base_model AutoModelForSequenceClassification.from_pretrained( self.model_args.model_name_or_path, configconfig, ...) model CrossEncoderModel( base_model, tokenizertokenizer, train_batch_sizeself.training_args.per_device_train_batch_size, )值得注意的三个细节num_labels1重排打分是回归式单分数输出SequenceClassification 头只有 1 个输出单元对应encode返回(N, 1)的 logitstrain_batch_size与训练参数绑定模型构造时的train_batch_size直接取training_args.per_device_train_batch_size从而保证 forward 中view分组与数据加载批次严格一致——这是 loss 计算正确的隐含前提条件梯度检查点if self.training_args.gradient_checkpointing: model.enable_input_require_grads()配合--gradient_checkpointing使用。Runner 基类AbsRunner.py随后根据model_args.model_type encoder选择AbsRerankerTrainDataset与AbsRerankerCollator训练结束调用trainer.save_model()落盘。六、端到端实战完整微调命令与参数解读6.1 启动入口encoder-only base 重排器微调的命令行入口为 base/main.py通过HfArgumentParser解析三组参数模型参数、数据参数、训练参数实例化EncoderOnlyRerankerRunner并调用runner.run()。6.2 官方示例脚本仓库提供了可直接运行的示例 base.sh其核心命令为torchrun --nproc_per_node $num_gpus \ -m FlagEmbedding.finetune.reranker.encoder_only.base \ --model_name_or_path BAAI/bge-reranker-base \ --train_data ../example_data/normal/examples.jsonl \ --train_group_size 8 \ --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 \ --learning_rate 6e-5 \ --fp16 \ --num_train_epochs 4 \ --per_device_train_batch_size 2 \ --gradient_accumulation_steps 1 \ --warmup_ratio 0.1 \ --gradient_checkpointing \ --weight_decay 0.01 \ --deepspeed ../../ds_stage0.json6.3 关键参数语义表以下参数由AbsRerankerDataArgumentsAbsArguments.py定义直接决定数据如何喂给CrossEncoderModel参数默认值说明train_data必填训练数据路径可传多个。要求每条包含query: str、pos: List[str]、neg: List[str]若开启蒸馏还需pos_scores/neg_scorestrain_group_size8每个 query 参与打分的文档数正负样本合计决定 forward 中每组 logits 的列数query_max_len32query 截断长度passage_max_len128passage 截断长度max_len512拼接后的总序列最大长度encoder-only 场景下 querypassage 拼接后截断pad_to_multiple_ofNone将序列 pad 到该值的整数倍示例中为8利于算子加速knowledge_distillationFalse是否启用知识蒸馏损失项shuffle_ratio0.0文本洗牌比例query_instruction_for_rerankNone查询侧指令前缀数据加载阶段会对train_data做存在性校验__post_init__中FileNotFoundError并对\n转义做归一化处理。6.4 数据格式示例训练数据JSONL的字段结构如下{query: 什么是知识蒸馏, pos: [知识蒸馏是一种模型压缩技术], neg: [今天天气很好, 量子计算的原理]}若开启蒸馏则每个文档追加软标签{query: ..., pos: [{text: ..., score: 0.9}], neg: [{text: ..., score: 0.1}]}字段形式可参考 示例数据目录 与 数据参数定义。七、训练、保存与验证闭环训练控制training_args继承自 transformers 的TrainingArgumentsAbsArguments.py示例脚本展示了fp16、deepspeed、warmup_ratio、weight_decay、gradient_checkpointing、save_steps等常用配置的组合用法断点续训Runner 的run()支持resume_from_checkpointAbsRunner.py输出目录保护若output_dir已存在且非空、且未声明--overwrite_output_dirRunner 会在启动时抛出ValueError拦截避免误覆盖AbsRunner.py产物结构保存目录内含pytorch_model.bin、config.json、tokenizer文件与training_args.bin。微调后的模型既可用FlagEmbedding推理侧加载做重排也可继续作为CrossEncoderModel的base_model二次微调。八、总结CrossEncoderModel虽只是一个轻量封装类却是 FlagEmbedding encoder-only 重排器微调管线的建模枢纽它以「querypassage 拼接 → 单头打分 → 组内交叉熵」的方式定义了交叉编码器重排的训练范式同时通过AbsRerankerModel基类天然支持知识蒸馏、梯度检查点与标准化保存。理解encode返回 logits 的形状语义与train_batch_size的分组含义是正确调参和排查训练问题的关键。配合 runner.py、trainer.py 与官方 base.sh 示例你可以快速将 BGE 系列 encoder 模型微调为适配自身业务相关性打分的重排器。【免费下载链接】FlagEmbeddingRetrieval and Retrieval-augmented LLMs项目地址: https://gitcode.com/GitHub_Trending/fl/FlagEmbedding创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
锦
锦皓数字建站
深耕本土企业品牌数字化升级,专注原创端正雅致商务官网,从视觉设计到稳定运维全程保驾护航。