资讯动态

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

发布时间:2026/9/15 17:41:38 来源:尧图企业网站定制
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抽象层负责训练循环中的前向、负样本与蒸馏逻辑BiDecoderOnlyEmbedderModeldecoder-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_typekl_div知识蒸馏损失类型支持kl_div与m3_kd_lossuse_mrlFalse是否启用 Matryoshka Representation LearningMRL多维表示学习mrl_dims[]MRL 层的维度列表例如[512, 256, 128]use_mrlTrue时必须非空否则基类会抛出ValueErrorsentence_pooling_methodlast_token句向量池化方式cls/mean/last_tokennormalize_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。3MRL 分支与归一化若use_mrlTrue对每个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.pycls直接取序列首 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(dim1) - 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(reductionmean)在训练时由基类的损失函数调用目标为 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_lossBGE-M3 风格的加权交叉熵按教师软标签权重对每个正样本组的交叉熵加权求和并在组间用掩码屏蔽已计算的得分位置。MRL 模式下use_mrlTrueforward会对mrl_dims中每个维度分别调用损失函数并取平均从而让每个子维度都具备可检索能力。此外若negatives_cross_deviceTrue但分布式环境未初始化构造函数会抛出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_loraTrue是否使用 LoRA 参数高效微调lora_rank64LoRA 秩lora_alpha16LoRA 缩放系数lora_dropout0.1LoRA 模块 dropouttarget_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_pathPEFT 初始化 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负责组装模型关键流程包括通过AutoConfigAutoModel加载骨干use_flash_attnTrue时指定attn_implementationflash_attention_2并设置config.use_cacheFalse适配训练支持raw_peft预合并先加载自定义embedding/emb.pth输入嵌入再用PeftModel.from_pretrained加载并merge_and_unload()支持resize词表扩展如为additional_special_tokens预留 token 位并将新输入嵌入保存为output_dir/embedding/emb.pth未提供from_peft且use_loraTrue时用LoraConfig(task_typeTaskType.FEATURE_EXTRACTION, ...)包裹为 PEFT 模型。八、保存、合并与训练器集成模型的持久化由save方法完成将state_dict克隆到 CPU 后调用self.model.save_pretrained(output_dir, state_dictstate_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),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价