FlagEmbedding 交叉编码器重排模型微调:CrossEncoderModel 源码级解析与实战指南
2026/9/15 13:32:06 网站建设 项目流程

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, tokenizer=tokenizer, train_batch_size=train_batch_size, )

三个构造参数的含义与默认值如下表:

参数类型默认值说明
base_modelPreTrainedModel必填底层预训练模型,实际承担编码与打分任务。在 Runner 中由AutoModelForSequenceClassification加载得到
tokenizerAutoTokenizerNone用于编码输入文本的分词器
train_batch_sizeint4训练批次大小,其语义是「每个训练样本内 query-passage 分组对应的样本数」,直接影响 loss 的分组视角(见下文 forward 解析)

类本身不引入新的可学习参数,所有前向、损失、保存逻辑均由抽象基类AbsRerankerModel提供。因此理解该类,关键在于理解其继承链。

三、encode 方法:从输入特征到相关性 logits

encode是本模型对外暴露的核心方法,定义见 modeling.py:

def 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_dict=True).logits

实现要点:

  • 输入features为字典,包含input_idsattention_masktoken_type_ids等由数据整理器(collator)产出、已被分词并拼接好的模型输入(query + passage 拼接后的序列);
  • 处理:将features以关键字参数形式直接透传给底层的PreTrainedModel(SequenceClassification 模型),并开启return_dict=True获取结构化输出;
  • 输出:返回.logits,即模型打出的相关性分数。由于 Runner 加载时固定num_labels=1,输出形状为(batch_size * group_size, 1),其中group_size即训练样本中 query 对应的文档数量(正样本 + 负样本数,见train_group_size参数)。

由于encode是抽象基类AbsRerankerModel中的抽象方法(见 AbsModeling.py),CrossEncoderModel必须实现它,这是子类唯一必须补齐的能力点。

四、继承体系:AbsRerankerModel 的初始化与前向逻辑

CrossEncoderModel继承自FlagEmbedding.abc.finetune.reranker.AbsRerankerModel,该抽象类实现了完整的训练语义,源码位于 AbsModeling.py。

4.1 初始化阶段的关键行为

构造时基类会完成四件重要的事:

  1. 缓存模型与分词器:将base_model存为self.model,并将model.config同步为self.config
  2. 补齐 pad_token:若model.config.pad_token_id is None,则自动用tokenizer.pad_token_id填充(AbsModeling.py),避免分组拼接与 pad 时报错;
  3. 内置交叉熵损失self.cross_entropy = nn.CrossEntropyLoss(reduction='mean')(默认 reduction 为均值);
  4. 计算 "Yes" 的 token 位置self.yes_loc = self.tokenizer('Yes', add_special_tokens=False)['input_ids'][-1],供 decoder-only 重排器使用,encoder-only 链路不依赖此项。

同时基类还透传了gradient_checkpointing_enableenable_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_logitsview(self.train_batch_size, -1)重整为「每个样本一行、组内各候选一列」的二维矩阵,第 0 列恒为正样本;
  • 目标构造target为全零向量,即始终要求正样本分数最高,损失即为组内 Softmax 交叉熵——这是「让正样本排在组内第一」的直接实现;
  • 知识蒸馏:当teacher_scores传入时(对应数据参数knowledge_distillation=True),额外累加一项「log-softmax(logits) 与 teacher 概率的逐元素乘积均值」的负值,等价于最小化 logits 分布与 teacher 分布的 KL 散度项。这一设计使训练可以兼容教师模型软标签(训练数据中带pos_scores/neg_scores字段)。

compute_loss(AbsModeling.py)即封装了上述交叉熵。输出统一为RerankerOutput(dataclass,含lossscores字段),推理阶段lossNone,仅返回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.py:

tokenizer = 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_labels=num_labels, ...) base_model = AutoModelForSequenceClassification.from_pretrained( self.model_args.model_name_or_path, config=config, ...) model = CrossEncoderModel( base_model, tokenizer=tokenizer, train_batch_size=self.training_args.per_device_train_batch_size, )

值得注意的三个细节:

  1. num_labels=1:重排打分是回归式单分数输出,SequenceClassification 头只有 1 个输出单元,对应encode返回(N, 1)的 logits;
  2. train_batch_size与训练参数绑定:模型构造时的train_batch_size直接取training_args.per_device_train_batch_size,从而保证 forward 中view分组与数据加载批次严格一致——这是 loss 计算正确的隐含前提;
  3. 条件梯度检查点if self.training_args.gradient_checkpointing: model.enable_input_require_grads(),配合--gradient_checkpointing使用。

Runner 基类(AbsRunner.py)随后根据model_args.model_type == 'encoder'选择AbsRerankerTrainDatasetAbsRerankerCollator,训练结束调用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.json

6.3 关键参数语义表

以下参数由AbsRerankerDataArguments(AbsArguments.py)定义,直接决定数据如何喂给CrossEncoderModel

参数默认值说明
train_data必填训练数据路径,可传多个。要求每条包含query: strpos: List[str]neg: List[str];若开启蒸馏还需pos_scores/neg_scores
train_group_size8每个 query 参与打分的文档数(正负样本合计),决定 forward 中每组 logits 的列数
query_max_len32query 截断长度
passage_max_len128passage 截断长度
max_len512拼接后的总序列最大长度(encoder-only 场景下 query+passage 拼接后截断)
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 的TrainingArguments(AbsArguments.py),示例脚本展示了fp16deepspeedwarmup_ratioweight_decaygradient_checkpointingsave_steps等常用配置的组合用法;
  • 断点续训:Runner 的run()支持resume_from_checkpoint(AbsRunner.py);
  • 输出目录保护:若output_dir已存在且非空、且未声明--overwrite_output_dir,Runner 会在启动时抛出ValueError拦截,避免误覆盖(AbsRunner.py);
  • 产物结构:保存目录内含pytorch_model.binconfig.jsontokenizer文件与training_args.bin。微调后的模型既可用FlagEmbedding推理侧加载做重排,也可继续作为CrossEncoderModelbase_model二次微调。

八、总结

CrossEncoderModel虽只是一个轻量封装类,却是 FlagEmbedding encoder-only 重排器微调管线的建模枢纽:它以「query+passage 拼接 → 单头打分 → 组内交叉熵」的方式定义了交叉编码器重排的训练范式,同时通过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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询