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 文档,通过autoclass与automethod指令直接绑定 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列出的encode、compute_score、compute_loss、gradient_checkpointing_enable、enable_input_require_grads、save、_sentence_embedding、_compute_similarity这八个方法,正是该类对外暴露的核心接口,下文逐一剖析。
二、类的定位:继承自抽象嵌入模型的 decoder-only 实现
BiDecoderOnlyEmbedderModel继承自 FlagEmbedding/abc/finetune/embedder/AbsModeling.py 中的AbsEmbedderModel(抽象基类),后者本身继承ABC与torch.nn.Module,并强制要求子类实现四个抽象方法:encode、compute_loss、compute_score、save。继承关系可以概括为:
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 骨干 |
tokenizer | None | 分词器,用于编码输入文本 |
negatives_cross_device | False | 是否启用跨设备负样本(所有 GPU 上的 batch 互为负样本),启用后计算量随world_size线性增长 |
temperature | 1.0 | 温度系数,缩放相似度得分后再进入损失计算,控制 softmax 分布的锐利程度 |
sub_batch_size | -1 | 编码时的子批次大小;为正数时将 batch 切分为子批次逐块编码再拼接,用于显存受限场景;负数表示不切分 |
kd_loss_type | 'kl_div' | 知识蒸馏损失类型,支持'kl_div'与'm3_kd_loss' |
use_mrl | False | 是否启用 Matryoshka Representation Learning(MRL)多维表示学习 |
mrl_dims | [] | MRL 层的维度列表,例如[512, 256, 128];use_mrl=True时必须非空,否则基类会抛出ValueError |
sentence_pooling_method | 'last_token' | 句向量池化方式:'cls'/'mean'/'last_token' |
normalize_embeddings | False | 是否对最终嵌入向量做 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_method从last_hidden_state提取句向量,源码位于 FlagEmbedding/finetune/embedder/decoder_only/base/modeling.py:
cls:直接取序列首 token 的隐状态last_hidden_state[:, 0],与 encoder-only 架构的 [CLS] 池化一致。mean:对last_hidden_state按attention_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 对应的正样本位置。它与encode、compute_score一起构成训练闭环。
六、基类 AbsEmbedderModel:负样本策略、蒸馏与跨设备扩展
BiDecoderOnlyEmbedderModel自身只负责"表示",而训练期的损失组装逻辑在基类AbsEmbedderModel.forward中完成,理解建模必须连同基类一起看。forward(queries, passages, teacher_scores, no_in_batch_neg_flag)的执行顺序为:
q_reps = self.encode(queries)、p_reps = self.encode(passages),得到查询与段落表示;- 训练模式下根据
no_in_batch_neg_flag与negatives_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 后计算全局得分,扩大负样本规模;
- 若提供
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_path、token、cache_dir、trust_remote_code、config_name等通用项外,核心参数如下:
| 参数 | 默认值 | 说明 |
|---|---|---|
use_lora | True | 是否使用 LoRA 参数高效微调 |
lora_rank | 64 | LoRA 秩 |
lora_alpha | 16 | LoRA 缩放系数 |
lora_dropout | 0.1 | LoRA 模块 dropout |
target_modules | ['v_proj','q_proj','k_proj','gate_proj','down_proj','o_proj','up_proj'] | 应用 LoRA 的注意力与 FFN 投影层 |
modules_to_save | None | 需要在 checkpoint 中完整保存的模块列表 |
use_flash_attn | False | 是否使用 Flash Attention 2 加速 |
use_slow_tokenizer | False | 是否使用慢速分词器 |
peft_model_path | '' | PEFT 初始化 checkpoint 路径 |
from_peft | None | 从已有 PEFT 适配器继续训练 |
raw_peft | None | 加载并合并原始 PEFT 权重 |
additional_special_tokens | None | 额外特殊 token(如自定义 query/passage 前缀) |
save_merged_lora_model | False | 训练后合并 LoRA 并保存完整模型 |
only_merge_lora_model | False | 仅执行合并不训练 |
对应地,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_peft且use_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_method、temperature、use_mrl、mrl_dims、negatives_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),仅供参考