FlagEmbedding 解码器专用嵌入模型详解:BiDecoderOnlyEmbedderModel 建模原理与实践
2026/9/15 17:41:33 网站建设 项目流程

FlagEmbedding 解码器专用嵌入模型详解:BiDecoderOnlyEmbedderModel 建模原理与实践

【免费下载链接】FlagEmbeddingRetrieval and Retrieval-augmented LLMs项目地址: https://gitcode.com/GitHub_Trending/fl/FlagEmbedding

导读

本文以 FlagEmbedding 项目文档 docs/source/API/finetune/embedder/decoder_only/base/modeling.rst 为核心,深入讲解基于 decoder-only 架构(如 Qwen、LLaMA 等因果语言模型)训练文本嵌入向量的核心建模类BiDecoderOnlyEmbedderModel。你将掌握该类的构造参数、编码与池化流程、相似度与损失计算、MRL 多维表示学习,以及配套的 LoRA 参数与模型加载/保存机制,并能在 examples/finetune/embedder/decoder_only 的示例脚本基础上独立完成基于 LLM 的嵌入模型微调。

一、建模文档指向哪个类:从 Sphinx autodoc 到源码

modeling.rst是一份 Sphinx autodoc 文档,通过autoclassautomethod指令直接绑定 Python 源码中的类与方法,属于"文档即接口声明"的 API 文档形态:

.. autoclass:: FlagEmbedding.finetune.reranker.decoder_only.base.CrossDecoderModel Methods ======= .. automethod:: FlagEmbedding.finetune.reranker.decoder_only.base.CrossDecoderModel.encode

值得说明的是,从当前仓库的文档结构看,该文件中的 autodoc 目标写为FlagEmbedding.finetune.reranker.decoder_only.base.CrossDecoderModel,而路径本身位于 embedder(嵌入器)的 API 目录下;对称地,docs/source/API/finetune/reranker/decoder_only/base/modeling.rst 中则绑定了BiDecoderOnlyEmbedderModel。从仓库源码对照来看,这两处文档的类名引用存在互换的痕迹:与 embedder 建模文档路径语义一致的核心类是BiDecoderOnlyEmbedderModel(位于 FlagEmbedding/finetune/embedder/decoder_only/base/modeling.py)。因此本文以该文档路径所归属的"decoder-only 嵌入器建模"主题展开,聚焦BiDecoderOnlyEmbedderModel的完整实现。

无论文档绑定写法如何,理解该 API 文档的关键都在于:autoclass声明的是训练侧嵌入模型类automethod列出的encodecompute_scorecompute_lossgradient_checkpointing_enableenable_input_require_gradssave_sentence_embedding_compute_similarity这八个方法,正是该类对外暴露的核心接口,下文逐一剖析。

二、类的定位:继承自抽象嵌入模型的 decoder-only 实现

BiDecoderOnlyEmbedderModel继承自 FlagEmbedding/abc/finetune/embedder/AbsModeling.py 中的AbsEmbedderModel(抽象基类),后者本身继承ABCtorch.nn.Module,并强制要求子类实现四个抽象方法:encodecompute_losscompute_scoresave。继承关系可以概括为:

  • AbsEmbedderModel(抽象层,负责训练循环中的前向、负样本与蒸馏逻辑)
    • BiDecoderOnlyEmbedderModel(decoder-only 具体实现,负责编码与池化)

构造函数签名如下:

def __init__( self, base_model: PreTrainedModel, tokenizer: PreTrainedTokenizer = None, negatives_cross_device: bool = False, temperature: float = 1.0, sub_batch_size: int = -1, kd_loss_type: str = 'kl_div', use_mrl: bool = False, mrl_dims: List[int] = [], sentence_pooling_method: str = 'last_token', normalize_embeddings: bool = False, ):

各参数含义与默认值整理如下:

参数默认值作用
base_model必填用于训练的 decoder-only 预训练模型,通常为AutoModel加载的因果 LM 骨干
tokenizerNone分词器,用于编码输入文本
negatives_cross_deviceFalse是否启用跨设备负样本(所有 GPU 上的 batch 互为负样本),启用后计算量随world_size线性增长
temperature1.0温度系数,缩放相似度得分后再进入损失计算,控制 softmax 分布的锐利程度
sub_batch_size-1编码时的子批次大小;为正数时将 batch 切分为子批次逐块编码再拼接,用于显存受限场景;负数表示不切分
kd_loss_type'kl_div'知识蒸馏损失类型,支持'kl_div''m3_kd_loss'
use_mrlFalse是否启用 Matryoshka Representation Learning(MRL)多维表示学习
mrl_dims[]MRL 层的维度列表,例如[512, 256, 128]use_mrl=True时必须非空,否则基类会抛出ValueError
sentence_pooling_method'last_token'句向量池化方式:'cls'/'mean'/'last_token'
normalize_embeddingsFalse是否对最终嵌入向量做 L2 归一化

此外类属性TRANSFORMER_CLS = AutoModel声明了骨干模型的加载类,模型主体保存在self.model中。

三、encode:从输入特征到嵌入向量的完整流程

encode是训练与推理共用的核心方法,接收模型输入特征(dict 或 dict 列表),返回嵌入向量或 MRL 向量列表。其执行逻辑分三步:

1)子批次编码(显存优化)

当输入为单个 dict 且sub_batch_size > 0时,按attention_mask的长度切分子批次,逐块前向获取last_hidden_state并池化,最后torch.cat拼接回完整 batch;当输入为 dict 列表(不同样本长度不同,无法统一 padding)时,则逐样本前向后再拼接:

if not isinstance(features, list): if self.sub_batch_size is not None and self.sub_batch_size > 0: for i in range(0, len(features['attention_mask']), self.sub_batch_size): ... # 切片子特征 -> 前向 -> 池化 else: for sub_features in features: # 逐样本编码 ...

2)池化得到句向量

对每个子批次的last_hidden_state调用_sentence_embedding(last_hidden_state, attention_mask),得到句子表示p_reps

3)MRL 分支与归一化

  • use_mrl=True:对每个mrl_dims维度截取前dim维(all_p_reps[:, :dim]),若dim超过原始维度会记录 warning 并退化为原始维度;每段子向量按normalize_embeddings决定是否归一化,最终返回列表(每个元素对应一个 MRL 维度)。
  • 否则返回完整的归一化(可选)嵌入向量all_p_reps.contiguous()

一个值得注意的细节:MRL 模式下normalize_embeddings只在截断子向量时生效;非 MRL 模式下则对全维度向量调用torch.nn.functional.normalize(dim=-1)。两种路径的归一化语义完全一致,只是作用在"当前使用的表示"上。

四、三种池化策略:_sentence_embedding 的实现细节

_sentence_embedding根据sentence_pooling_methodlast_hidden_state提取句向量,源码位于 FlagEmbedding/finetune/embedder/decoder_only/base/modeling.py:

  • cls:直接取序列首 token 的隐状态last_hidden_state[:, 0],与 encoder-only 架构的 [CLS] 池化一致。
  • mean:对last_hidden_stateattention_mask做掩码加权平均,即sum(hidden * mask) / sum(mask),避免 padding 位置稀释表示。
  • last_token(默认):取每个序列的最后一个有效 token(即attention_mask.sum(dim=1) - 1位置)的隐状态。这是 decoder-only 嵌入的主流做法——因果注意力下最后一个 token 聚合了前文全部信息。实现中还兼容左 padding 的情况(left_padding判断),此时直接取last_hidden_state[:, -1]
  • 其他取值抛出NotImplementedError

五、相似度与损失:compute_score / _compute_similarity / compute_loss

相似度计算。_compute_similarity使用内积(torch.matmul)计算 query 与 passage 表示之间的相似度矩阵:二维表示走q @ p^T,三维表示走 batch 内矩阵乘法。compute_score在此基础上除以温度temperature后展平为(batch_size, -1)

scores = self._compute_similarity(q_reps, p_reps) / self.temperature scores = scores.view(q_reps.size(0), -1)

温度越低,得分分布越尖锐,对正负样本的区分越强。

损失计算。compute_loss直接复用torch.nn.CrossEntropyLoss(reduction='mean'),在训练时由基类的损失函数调用,目标为 batch 内每个 query 对应的正样本位置。它与encodecompute_score一起构成训练闭环。

六、基类 AbsEmbedderModel:负样本策略、蒸馏与跨设备扩展

BiDecoderOnlyEmbedderModel自身只负责"表示",而训练期的损失组装逻辑在基类AbsEmbedderModel.forward中完成,理解建模必须连同基类一起看。forward(queries, passages, teacher_scores, no_in_batch_neg_flag)的执行顺序为:

  1. q_reps = self.encode(queries)p_reps = self.encode(passages),得到查询与段落表示;
  2. 训练模式下根据no_in_batch_neg_flagnegatives_cross_device选择损失函数:
    • _compute_no_in_batch_neg_loss:不使用任何 batch 内负样本,仅对每组 query 对应的group_size个 passage 做交叉熵;
    • _compute_in_batch_neg_loss(默认):batch 内所有 passage 互为负样本,目标为idxs * group_size
    • _compute_cross_device_neg_loss:通过_dist_gather_tensor将各进程的表示 all-gather 后计算全局得分,扩大负样本规模;
  3. 若提供teacher_scores(蒸馏),先softmax为软标签,再叠加蒸馏损失。

蒸馏损失由静态方法distill_loss实现,支持两种类型:

  • 'kl_div':学生对log_softmax得分与教师软标签的 KL 散度,即-mean(sum(log_softmax(student) * teacher_targets))
  • 'm3_kd_loss':BGE-M3 风格的加权交叉熵,按教师软标签权重对每个正样本组的交叉熵加权求和,并在组间用掩码屏蔽已计算的得分位置。

MRL 模式下(use_mrl=True),forward会对mrl_dims中每个维度分别调用损失函数并取平均,从而让每个子维度都具备可检索能力。此外,若negatives_cross_device=True但分布式环境未初始化,构造函数会抛出ValueError提醒。

七、配套参数与模型加载:arguments.py 与 load_model.py

训练脚本通过 FlagEmbedding/finetune/embedder/decoder_only/base/arguments.py 中的DecoderOnlyEmbedderModelArguments配置模型行为,除继承抽象基类的model_name_or_pathtokencache_dirtrust_remote_codeconfig_name等通用项外,核心参数如下:

参数默认值说明
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 的注意力与 FFN 投影层
modules_to_saveNone需要在 checkpoint 中完整保存的模块列表
use_flash_attnFalse是否使用 Flash Attention 2 加速
use_slow_tokenizerFalse是否使用慢速分词器
peft_model_path''PEFT 初始化 checkpoint 路径
from_peftNone从已有 PEFT 适配器继续训练
raw_peftNone加载并合并原始 PEFT 权重
additional_special_tokensNone额外特殊 token(如自定义 query/passage 前缀)
save_merged_lora_modelFalse训练后合并 LoRA 并保存完整模型
only_merge_lora_modelFalse仅执行合并不训练

对应地,FlagEmbedding/finetune/embedder/decoder_only/base/load_model.py 中的get_model负责组装模型,关键流程包括:

  • 通过AutoConfig+AutoModel加载骨干(use_flash_attn=True时指定attn_implementation="flash_attention_2",并设置config.use_cache=False适配训练);
  • 支持raw_peft预合并:先加载自定义embedding/emb.pth输入嵌入,再用PeftModel.from_pretrained加载并merge_and_unload()
  • 支持resize词表扩展(如为additional_special_tokens预留 token 位),并将新输入嵌入保存为output_dir/embedding/emb.pth
  • 未提供from_peftuse_lora=True时,用LoraConfig(task_type=TaskType.FEATURE_EXTRACTION, ...)包裹为 PEFT 模型。

八、保存、合并与训练器集成

模型的持久化由save方法完成:将state_dict克隆到 CPU 后调用self.model.save_pretrained(output_dir, state_dict=state_dict),避免直接保存引用带来后续训练污染。训练侧 FlagEmbedding/finetune/embedder/decoder_only/base/trainer.py 中的DecoderOnlyEmbedderTrainer._save在每次 checkpoint 时依次保存模型权重(调用self.model.save(output_dir))、分词器与training_args.bin

若配置了save_merged_lora_model,训练结束后可运行save_merged_model:重新加载骨干与训练产出的 LoRA(自动通过find_largest_checkpoint回退到最大的checkpoint-*),合并后连同 tokenizer 一起保存到output_dir/merged_model,得到可直接用于 FlagEmbedding/finetune/embedder/decoder_only/base 之外推理场景的完整模型。

九、从建模到实战:把类接入训练脚本

BiDecoderOnlyEmbedderModel本身是训练管线的一环,完整的微调入口由同目录下的 FlagEmbedding/finetune/embedder/decoder_only/base/main.py 与 runner.py 提供,对应的可直接运行的 shell 示例见 examples/finetune/embedder/decoder_only。典型调用形如:

python -m FlagEmbedding.finetune.embedder.decoder_only.base \ --model_name_or_path Qwen/Qwen2-0.5B \ --use_lora True \ --lora_rank 64 \ --lora_alpha 16 \ --sentence_pooling_method last_token \ --temperature 1.0 \ --train_data ./train.jsonl \ --output_dir ./output \ --save_merged_lora_model True

其中sentence_pooling_methodtemperatureuse_mrlmrl_dimsnegatives_cross_device等建模级参数会直接注入本文所述的模型类构造函数,控制最终嵌入的质量与训练显存/收敛特性。

十、相关文件索引

  • 建模文档(本文依据):docs/source/API/finetune/embedder/decoder_only/base/modeling.rst
  • 核心建模类:FlagEmbedding/finetune/embedder/decoder_only/base/modeling.py
  • 抽象基类(负样本/蒸馏/MRL 逻辑):FlagEmbedding/abc/finetune/embedder/AbsModeling.py
  • 参数定义:FlagEmbedding/finetune/embedder/decoder_only/base/arguments.py
  • 模型加载与合并:FlagEmbedding/finetune/embedder/decoder_only/base/load_model.py
  • 训练器:FlagEmbedding/finetune/embedder/decoder_only/base/trainer.py
  • 运行示例:examples/finetune/embedder/decoder_only

【免费下载链接】FlagEmbeddingRetrieval and Retrieval-augmented LLMs项目地址: https://gitcode.com/GitHub_Trending/fl/FlagEmbedding

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

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

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

立即咨询