FlagEmbedding BGE-M3 微调建模源码解析:EncoderOnlyEmbedderM3Model 的三路表征与统一微调实现
【免费下载链接】FlagEmbeddingRetrieval and Retrieval-augmented LLMs项目地址: https://gitcode.com/GitHub_Trending/fl/FlagEmbedding
本篇技术指南以 FlagEmbedding 官方 API 文档 modeling.rst 为骨架,结合仓库内 M3 微调模块 的完整源码,深入剖析 EncoderOnlyEmbedderM3Model 及其推理子类 EncoderOnlyEmbedderM3ModelForInference 的架构设计、三路表征(Dense / Sparse / ColBERT)的生成与打分机制、统一微调(Unified Fine-tuning)与自蒸馏损失的计算细节,并给出基于 examples/finetune/embedder/encoder_only/m3.sh 的可复现实战配置。读者读完即可掌握 M3 模型在 FlagEmbedding 框架中的建模原理、关键参数含义及二次开发切入点。
一、模块定位:M3 微调建模在 FlagEmbedding 中的角色
FlagEmbedding 将 BGE-M3 的多向量混合检索能力封装在FlagEmbedding.finetune.embedder.encoder_only.m3包中,该包仅包含 5 个文件:
- modeling.py:本文主角,定义
EncoderOnlyEmbedderM3Model与EncoderOnlyEmbedderM3ModelForInference; - arguments.py:模型参数与训练参数的数据类;
- runner.py:负责模型加载、tokenizer 加载与 Trainer 装配;
- trainer.py:自定义 Trainer 的保存逻辑;
- main.py:命令行入口,支持
torchrun -m FlagEmbedding.finetune.embedder.encoder_only.m3直接启动训练。
从继承关系看,EncoderOnlyEmbedderM3Model继承自抽象基类AbsEmbedderModel(AbsModeling.py),因此它天然具备基类提供的能力:in-batch negatives 损失、跨设备 negatives 损失(negatives_cross_device)、知识蒸馏损失分发(kd_loss_type)、MRL 支持等。但 M3 是一个多向量、三路表征的模型,其核心差异体现在:除了常规 dense 向量外,还通过两个额外线性层产出 sparse 词权重与 ColBERT token 向量,并据此设计了专用的m3_kd_loss蒸馏方式与四路损失加权方案。这正是本模块区别于 base embedder 建模的精华所在。
二、模型构造与关键超参
EncoderOnlyEmbedderM3Model.__init__的完整签名(modeling.py):
| 参数 | 默认值 | 含义 |
|---|---|---|
base_model | 必填 | 由 Runner 装配好的 dict,含model、colbert_linear、sparse_linear三个组件 |
tokenizer | None | 训练所用 tokenizer |
negatives_cross_device | False | 是否启用跨设备负样本(需要先初始化分布式环境) |
temperature | 1.0 | 打分时的温度系数,控制分数缩放 |
sub_batch_size | -1 | 编码时的子批次大小,为负则不拆分子批次(省显存技巧) |
kd_loss_type | 'm3_kd_loss' | 蒸馏损失类型,M3 默认专用m3_kd_loss,也可用kl_div |
use_mrl | False | 是否启用 MRL 训练——注意 M3 构造函数中若置 True 会直接raise NotImplementedError |
mrl_dims | [] | MRL 各层维度(M3 当前不支持) |
sentence_pooling_method | 'cls' | dense 向量池化方式 |
normalize_embeddings | False | 是否对 dense / colbert 向量做 L2 归一化 |
unified_finetuning | True | 是否启用统一微调(训练 sparse 与 colbert 头) |
use_self_distill | False | 是否启用自蒸馏 |
self_distill_start_step | -1 | 自蒸馏起始步数 |
关键实现细节:
- 组件装配:
unified_finetuning=True时,self.model / self.colbert_linear / self.sparse_linear全部取自base_model;为False时只保留 backbone,两个线性头置为None,模型退化为纯 dense 微调。 - sparse 头是
Linear(hidden_size, 1),输出经relu得到非负 token 权重;colbert 头是Linear(hidden_size, hidden_size)(维度可通过colbert_dim覆盖,见 runner.py)。 - 构造时即持有
self.cross_entropy = torch.nn.CrossEntropyLoss(reduction='mean'),用于最终的对比学习损失;同时记录self.vocab_size供 sparse 向量 one-hot 展开使用。
配套参数类(arguments.py):
EncoderOnlyEmbedderM3ModelArguments:继承AbsEmbedderModelArguments,额外只有colbert_dim: int = -1(colbert 线性层输出维度,≤0 时用 hidden_size);EncoderOnlyEmbedderM3TrainingArguments:继承AbsEmbedderTrainingArguments,新增unified_finetuning、use_self_distill、fix_encoder(冻结 backbone、只训练两个头)、self_distill_start_step四个开关。
三、三路表征的生成:_dense_embedding/_sparse_embedding/_colbert_embedding
三个下划线私有方法分别实现 Dense、Sparse、ColBERT 三种表征,统一在_encode与encode中串联调用。
3.1 Dense:池化得到句向量
_dense_embedding从 backbone 的last_hidden_state提取句向量,支持三种池化:
cls:直接取第 0 个 token 的隐状态last_hidden_state[:, 0];mean:按 attention_mask 加权求和后除以有效 token 数;last_token:根据 padding 方向取最后一个有效 token——先判断是否为左侧 padding(attention_mask[:, -1].sum() == attention_mask.shape[0]),左侧 padding 取[:, -1],否则用attention_mask.sum(dim=1) - 1定位每个样本的末尾位置索引后 gather。
其余池化方式会抛出NotImplementedError,这是推理端与训练端保持一致性的基础。
3.2 Sparse:词级权重展开为词表向量
_sparse_embedding是理解 M3 稀疏检索的关键:
- 将 last_hidden_state 过
sparse_linear后接relu,得到每个 token 的非负权重token_weights; return_embedding=False时直接返回 token 权重(推理时可配合 BM25 风格加权);- 训练与推理分支使用不同实现:
- 训练态:初始化
(batch, seq_len, vocab_size)零张量,用torch.scatter把 token 权重按input_ids写入对应词表位置,再沿序列维度取max汇聚; - 推理态:采用
scatter_reduce(..., reduce="amax")直接在(batch, vocab_size)上取每个词的最大权重。源码注释明确指出这一优化源自 issue #1364,并强调训练态不能使用该路径,否则会触发RuntimeError: one of the variables needed for gradient computation has been modified by an inplace operation;
- 训练态:初始化
- 最后将
cls / eos / pad / unk四个特殊 token 对应的词表维度权重清零(sparse_embedding[:, unused_tokens] *= 0.)。
得到的 sparse 向量可直接与 dense 向量拼接或独立建稀疏索引,实现词级精确匹配能力。
3.3 ColBERT:token 级向量
_colbert_embedding将 last_hidden_state 去掉[CLS]位([:, 1:])后过colbert_linear投影,再按 attention_mask 去掉 padding token(colbert_vecs * mask[:, 1:][:, :, None]),得到(batch, seq_len, dim)的 token 向量序列,用于后期交互(late interaction)式的细粒度匹配。
3.4 统一入口:_encode与encode
_encode单次前向:
last_hidden_state = self.model(**features, return_dict=True).last_hidden_state dense_vecs = self._dense_embedding(last_hidden_state, features['attention_mask']) sparse_vecs = self._sparse_embedding(last_hidden_state, features['input_ids']) # unified_finetuning 时 colbert_vecs = self._colbert_embedding(last_hidden_state, features['attention_mask']) # unified_finetuning 时 if self.normalize_embeddings: dense_vecs = F.normalize(dense_vecs, dim=-1) colbert_vecs = F.normalize(colbert_vecs, dim=-1)encode则处理两种输入形态:dict时按sub_batch_size切分(避免 OOM)或整体编码;list[dict]时逐个编码后torch.cat。统一微调开启时返回(dense_vecs, sparse_vecs, colbert_vecs)三元组,关闭时返回(dense_vecs, None, None),并全部调用.contiguous()保证内存布局。
四、三种打分与集成打分:compute_*_score/ensemble_score
4.1 Dense / Sparse 打分
compute_dense_score与compute_sparse_score逻辑一致:先由_compute_similarity计算内积相似度(二维输入走matmul(q, p.T),更高维走matmul(q, p.transpose(-2,-1))),再除以temperature并view(q_reps.size(0), -1)还原为(batch, batch*group_size)的得分矩阵。
4.2 ColBERT 打分(后期交互)
compute_colbert_score实现 ColBERT 的 MaxSim 后期交互:
token_scores = torch.einsum('qin,pjn->qipj', q_reps, p_reps) # 全量 token 对点积 scores, _ = token_scores.max(-1) # 每个 query token 取最相似 passage token scores = scores.sum(1) / q_mask[:, 1:].sum(-1, keepdim=True) # 求和并按 query 有效 token 数归一 scores = scores / self.temperature其中q_mask来自_get_queries_attention_mask:当 queries 是list[dict]时,会按 padding 方向把各子批的 attention_mask 右侧或左侧 pad 到同一长度再拼接,保证归一化分母正确。
4.3 三路集成打分
compute_score提供带权重的即时集成:
dense_score * dense_weight + sparse_score * sparse_weight + colbert_score * colbert_weight默认权重为dense=1.0, sparse=0.3, colbert=1.0。而ensemble_score则要求三路分数作为参数传入,硬编码dense + 0.3*sparse + colbert,三个分数缺一即抛ValueError。这组权重与官方 BGE-M3 推理侧的0.3稀疏权重设计保持一致,在训练中也被复用为集成损失。
五、前向传播与统一微调损失:forward/compute_loss
5.1 forward 主流程
forward是训练核心,先编码 query 与 passage(passage 形状为(batch*group_size, dim)),随后仅在self.training分支计算损失:
- teacher_targets 构造:若传入
teacher_scores,reshape 为(batch, group_size)后detach()并softmax化为分布;否则为None,走纯对比学习。 - 损失函数选择:
no_in_batch_neg_flag为 True 时用_compute_no_in_batch_neg_loss;否则按negatives_cross_device选择跨设备或批内负样本实现(这些实现在基类 AbsModeling.py 中)。 - 三路损失 + 集成损失(
unified_finetuning开启时):- dense loss:由
compute_dense_score计算; - sparse loss:由
compute_sparse_score计算,加权系数 0.1; - colbert loss:由
compute_colbert_score计算,并传入q_mask=self._get_queries_attention_mask(queries); - ensemble loss:用
ensemble_score融合三路分数后再算对比损失。 - 最终
loss = (loss + ensemble_loss + 0.1 * sparse_loss + colbert_loss) / 4。
- dense loss:由
- 跨设备截取:开启
negatives_cross_device时按process_rank从全局得分矩阵中截取本进程的 dense 得分参与集成(源码注释引用了 issue #1410 的 bug 修复:no_in_batch_neg_flag下需先按group_size取对角线 passage)。 - 自蒸馏:
use_self_distill且self.step > self_distill_start_step时,以集成得分softmax作为教师分布,分别对 dense / sparse / colbert 三路计算kl_div蒸馏损失并累加,最后整体再除 2。
5.2 compute_loss 与 m3_kd_loss
compute_loss即对得分矩阵与目标索引求交叉熵。而 M3 专用的蒸馏损失m3_kd_loss实现在基类 AbsModeling.py:它以group_size为步长构造目标位置,逐组对得分加掩码后计算逐样本交叉熵,并用教师分布teacher_targets[:, i]对每组损失加权求和——这使得「同一 query 的多个正负样本」各自获得与教师置信度匹配的梯度贡献,比单纯 KL 散度更贴合 BGE-M3 的蒸馏训练。
5.3 内存优化钩子
gradient_checkpointing_enable与enable_input_require_grads分别透传 backbone 的梯度检查点与输入梯度要求。Runner 在加载模型后会根据训练参数调用它们(见 runner.py),并支持通过fix_position_embedding/fix_encoder冻结特定参数。
六、模型保存:save与 Trainer 落盘
save先以 CPU 克隆方式保存 backbone 权重(save_pretrained),随后在unified_finetuning开启时,将两个头的状态字典分别落盘为colbert_linear.pt与sparse_linear.pt。这正好与加载逻辑对称——Runner 的get_model会检查模型目录下是否同时存在这两个文件,存在则恢复权重,否则提示「参数为新初始化,确保是训练而非推理加载」。
Trainer 的_save则调用self.model.save(output_dir),同时保存 tokenizer 与training_args.bin,保证 checkpoint 可完整恢复训练。
七、推理子类:EncoderOnlyEmbedderM3ModelForInference
EncoderOnlyEmbedderM3ModelForInference继承训练类并覆写 forward,将训练语义切换为按需产出表征的推理语义:
| 参数 | 默认值 | 说明 |
|---|---|---|
return_dense | True | 返回 dense 句向量 |
return_sparse | False | 返回 sparse 向量(词表维度) |
return_colbert_vecs | False | 返回 colbert token 向量 |
return_sparse_embedding | False | 透传给_sparse_embedding:True 返回展开的词表向量,False 只返回 token 权重 |
truncate_dim | None | 截断输出维度(兼容 Matryoshka 表征学习类模型) |
实现要点:
- 入口断言
return_dense or return_sparse or return_colbert_vecs至少一项为 True; - 强制
self.training = False,从而让 sparse 走推理态的高效scatter_reduce(amax)路径(再次呼应 issue #1364 的优化); - 只做一次 backbone 前向,按需组合输出字典
{'dense_vecs', 'sparse_vecs', 'colbert_vecs'};normalize_embeddings开启时对 dense 与 colbert 做 L2 归一化。
该子类与推理侧FlagEmbedding.inference.embedder.decoder_only/encoder_only中的封装共同构成 M3 从训练到部署的完整链路。
八、实战:统一微调完整命令与参数说明
仓库提供了可直接运行的示例脚本 examples/finetune/embedder/encoder_only/m3.sh,其核心配置如下(路径均相对该脚本所在目录):
train_data="\ ../example_data/retrieval \ ../example_data/sts/sts.jsonl \ ../example_data/classification-no_in_batch_neg \ ../example_data/clustering-no_in_batch_neg " num_train_epochs=4 per_device_train_batch_size=2 num_gpus=2 model_args="--model_name_or_path BAAI/bge-m3 --cache_dir $HF_HUB_CACHE" data_args="\ --train_data $train_data \ --cache_path ~/.cache \ --train_group_size 8 \ --query_max_len 512 \ --passage_max_len 512 \ --pad_to_multiple_of 8 \ --knowledge_distillation False \ " training_args="\ --output_dir ./test_encoder_only_m3_bge-m3 \ --overwrite_output_dir \ --learning_rate 1e-5 \ --fp16 \ --num_train_epochs $num_train_epochs \ --per_device_train_batch_size $per_device_train_batch_size \ --dataloader_drop_last True \ --warmup_ratio 0.1 \ --gradient_checkpointing \ --deepspeed ../../ds_stage0.json \ --logging_steps 1 \ --save_steps 1000 \ --negatives_cross_device \ --temperature 0.02 \ --sentence_pooling_method cls \ --normalize_embeddings True \ --kd_loss_type m3_kd_loss \ --unified_finetuning True \ --use_self_distill True \ --fix_encoder False \ --self_distill_start_step 0 \ " cmd="torchrun --nproc_per_node $num_gpus \ -m FlagEmbedding.finetune.embedder.encoder_only.m3 \ $model_args $data_args $training_args"关键参数逐条解读(对应前文源码):
--unified_finetuning True:开启 sparse / colbert 头训练,此时 forward 才会计算四路损失;--use_self_distill True --self_distill_start_step 0:从第 0 步起启用集成分数自蒸馏;--kd_loss_type m3_kd_loss:使用上文 5.2 的 M3 专用蒸馏损失;--negatives_cross_device:跨设备负样本,需配合torchrun多卡;--temperature 0.02:打分温度,直接进入compute_*_score的分母;--sentence_pooling_method cls --normalize_embeddings True:与模型构造参数一一对应;--train_group_size 8:每个 query 对应 1 正 + 7 负,与基类get_local_score中的group_size计算一致;--gradient_checkpointing:训练时由 Runner 调用model.enable_input_require_grads()配合梯度检查点。
数据格式要求:每个训练样本需包含query: str、pos: List[str]、neg: List[str]字段(见 AbsArguments.py);混合 retrieval / sts / classification / clustering 四种任务时,classification-no_in_batch_neg这类数据集对应no_in_batch_neg_flag路径,源码在 forward 中会据此选择_compute_no_in_batch_neg_loss。若设置same_dataset_within_batch,Runner 还会注册数据刷新回调(runner.py)。
九、总结与延伸阅读
EncoderOnlyEmbedderM3Model的设计精髓可概括为三点:三路表征解耦(dense 池化 + sparse 词权重 + colbert 后期交互)、统一微调联合训练(四路损失加权 + 可选自蒸馏)、训练/推理实现分离(sparse 的高效scatter_reduce推理路径与显存友好的子批次编码)。理解这一建模层,是进一步阅读推理封装、评估脚本或二次开发的基础。
- 推理侧封装:见 FlagEmbedding/inference/embedder(M3 对应 encoder_only/m3.py);
- 基类损失与蒸馏:见 AbsModeling.py;
- 训练入口与参数:见main.py、arguments.py;
- 端到端示例:见 examples/finetune/embedder/encoder_only/m3.sh 与 m3_same_dataset.sh;
- 官方 API 文档:见 modeling.rst(本文章节一一对应其中的方法清单)。
【免费下载链接】FlagEmbeddingRetrieval and Retrieval-augmented LLMs项目地址: https://gitcode.com/GitHub_Trending/fl/FlagEmbedding
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考