拓冰建站拓冰建站
首页 / 资讯中心 / 正文

FlagEmbedding 微调训练核心:AbsEmbedderTrainer 抽象训练器的设计、损失计算与工程实践

FlagEmbedding 微调训练核心AbsEmbedderTrainer 抽象训练器的设计、损失计算与工程实践【免费下载链接】FlagEmbeddingRetrieval and Retrieval-augmented LLMs项目地址: https://gitcode.com/GitHub_Trending/fl/FlagEmbedding导读本文围绕 FlagEmbedding 项目abc抽象基类架构中 embedder 微调训练的核心组件AbsEmbedderTrainer展开。该抽象训练器定义了嵌入模型Embedder训练的统一接口——包括抽象保存逻辑_save与核心损失计算compute_loss是所有编码器/解码器嵌入模型微调 Trainer 的公共基类。读完本文你将掌握 FlagEmbedding 训练流水线中 Trainer 与 Model、Dataset、Collator 的协作关系理解compute_loss如何在 in-batch negatives、cross-device negatives、知识蒸馏KD与 MRL 等场景下工作并能在实际微调脚本中正确配置相关训练参数。AbsEmbedderTrainer 在 FlagEmbedding 架构中的位置FlagEmbedding 的微调代码采用抽象定义 具体实现两层结构FlagEmbedding/abc/finetune/embedder/目录定义了与具体模型架构无关的抽象接口包括AbsArguments参数、AbsDataset数据集与 Collator、AbsModeling模型、AbsTrainer训练器与AbsRunner运行器五件套FlagEmbedding/finetune/embedder/目录下则分为encoder_only编码器与decoder_only解码器两大分支每个分支的base基础与m3/icl等子目录提供具体实现。AbsEmbedderTrainer正是这一层抽象中的训练器基类定义于 AbsTrainer.py。从类声明可以看到它同时继承了两条血脉class AbsEmbedderTrainer(ABC, Trainer): Abstract class for the trainer of embedder.ABCPython 抽象基类通过abstractmethod强制子类实现关键行为transformers.Trainer直接复用 Hugging Face Transformers 的训练循环、断点续训、日志、评估与保存等成熟能力。这种设计意味着任何具体的嵌入模型训练器都天然继承Trainer的完整工程能力如train()、save_model()、resume_from_checkpoint、分布式支持、deepspeed/fsdp 支持只需针对嵌入检索任务的特点覆写少量方法即可。抽象保存接口_save模型保存的统一契约AbsEmbedderTrainer声明了唯一的抽象方法_saveabstractmethod def _save(self, output_dir: Optional[str] None, state_dictNone): pass该方法要求子类实现如何把当前嵌入模型持久化到磁盘。具体实现分别在DecoderOnlyEmbedderTrainer._save解码器分支EncoderOnlyEmbedderTrainer._save编码器分支。以编码器分支为例其保存逻辑包含四个关键步骤确定输出目录output_dir output_dir if output_dir is not None else self.args.output_dir未显式传入时回退到TrainingArguments.output_dir保存模型调用self.model.save(output_dir)并在模型不具备save接口时抛出NotImplementedError——这保证了保存逻辑与具体模型实现解耦保存 tokenizer仅在self.is_world_process_zero()分布式下的主进程时通过tokenizer.save_pretrained(output_dir)保存词表保存训练参数torch.save(self.args, os.path.join(output_dir, training_args.bin))把完整TrainingArguments序列化到training_args.bin便于后续复现实验或加载推理。从源码结构还可以看到分支实现中保留了为 sentence-transformers 保存 checkpoint的注释代码save_ckpt_for_sentence_transformers表明该位置原本设计有生态互操作能力但当前仓库版本中处于未启用状态。compute_loss嵌入模型训练的前向与损失计算入口AbsEmbedderTrainer.compute_loss是transformers.Trainer训练循环中每个 batch 都会调用的核心方法其实现位于 AbsTrainer.pydef compute_loss(self, model, inputs, return_outputsFalse, **kwargs): outputs model(**inputs) loss outputs.loss return (loss, outputs) if return_outputs else loss它的职责与约定参数model为待训练的AbsEmbedderModelinputs是 DataCollator 产出的张量字典return_outputs控制是否同时返回模型输出。行为将inputs直接透传给模型前向model(**inputs)从返回的EmbedderOutput中取loss字段默认所有模型把 loss 放在输出第一个元素子类可通过覆写实现自定义行为。返回值return_outputsFalse时仅返回torch.Tensor形式的 loss为True时返回(loss, outputs)元组其中outputs为EmbedderOutput。也就是说训练器本身不关心 loss 怎么算它只是把模型内部算好的 loss取出来交给优化器。真正的损失计算发生在模型侧即AbsEmbedderModel.forward中见 AbsModeling.py。EmbedderOutput训练器与模型之间的数据结构约定模型前向返回的是EmbedderOutput定义于 AbsModeling.py一个继承自transformers.file_utils.ModelOutput的 dataclassdataclass class EmbedderOutput(ModelOutput): q_reps: Optional[Tensor] None # 查询向量 (batch_size, dim) p_reps: Optional[Tensor] None # 段落向量 (batch_size * group_size, dim) loss: Optional[Tensor] None # 训练损失 scores: Optional[Tensor] None # 相似度分数训练模式下forward会返回带loss的EmbedderOutput这正是compute_loss取用outputs.loss的依据推理非训练模式下lossNone。模型内部的三条损失计算路径AbsEmbedderModel.forward会根据配置与数据形态选择不同的损失计算函数这是理解compute_loss背后语义的关键。整体逻辑如下若训练数据标注no_in_batch_neg_flag例如聚类、分类任务数据走_compute_no_in_batch_neg_loss只用每组内query group_size-1个负例不与其他样本构成 in-batch 负样本否则若开启negatives_cross_device走_compute_cross_device_neg_loss通过_dist_gather_tensor把所有 GPU 上的 query/passage 表示收集到一起计算分数实现跨设备共享负样本即share negatives across devices大幅扩大负样本规模默认走_compute_in_batch_neg_loss每个 query 与该 batch 内所有 passage共batch_size * group_size个计算分数用 batch 内其他样本作为负样本。在_compute_cross_device_neg_loss中_dist_gather_tensorAbsModeling.py通过dist.all_gather收集各进程张量并拼接配合process_rank/world_size切回本地 batch 计算局部损失。这也解释了AbsEmbedderModel.__init__中的约束开启negatives_cross_device时必须先初始化分布式环境否则直接抛出ValueError。训练超参数如何影响损失计算compute_loss的效果高度依赖训练参数。AbsEmbedderTrainingArgumentsAbsArguments.py在标准transformers.TrainingArguments基础上扩展了以下关键字段参数默认值作用negatives_cross_deviceFalse是否跨设备共享负样本开启后走 cross-device 负样本损失路径temperature0.02相似度分数缩放系数模型计算scores similarity / temperature编码器实现在 modeling.pyfix_position_embeddingFalse冻结位置编码参数Runner 中会按参数名position_embeddings置requires_gradFalsesentence_pooling_methodcls句向量池化方式可选cls/mean/last_tokennormalize_embeddingsTrue是否对输出向量做 L2 归一化sub_batch_sizeNone编码时的子批大小用于显存受限场景kd_loss_typekl_div蒸馏损失类型可选kl_div/m3_kd_lossuse_mrlFalse是否启用 MRLMatryoshka Representation Learning训练mrl_dims[]MRL 各层维度列表use_mrlTrue时必填模型侧会校验非空其中temperature直接影响对比学习的锐利程度温度越小softmax 分布越尖锐对难负样本的惩罚越强BGE 系列默认使用0.02。use_mrl开启后forward会对mrl_dims中的每个维度分别计算损失并取平均见 AbsModeling.py。知识蒸馏KD与 m3_kd_loss当训练数据包含pos_scores/neg_scores教师模型打出的分数时forward会把teacher_scores转成概率分布teacher_targets F.softmax(teacher_scores, dim-1)然后进入蒸馏分支。蒸馏损失由AbsEmbedderModel.distill_lossAbsModeling.py计算支持两种模式kl_div学生 log-softmax 分数与教师概率的 KL 散度通常还会叠加一层常规对比损失loss self.compute_loss(...)m3_kd_loss逐位置对教师概率加权后的交叉熵损失BGE-M3 专用内部通过torch.scatter屏蔽已计算的标签位置。数据侧对应的开关是AbsEmbedderDataArguments.knowledge_distillation默认False开启后 Collator 才从样本中读取pos_scores/neg_scores字段。Trainer 的装配从 Runner 到实际训练AbsEmbedderTrainer不会被直接实例化而是由各分支 Runner 装配。以解码器分支为例runner.py 中的load_trainer展示了标准装配方式trainer DecoderOnlyEmbedderTrainer( modelself.model, argsself.training_args, train_datasetself.train_dataset, data_collatorself.data_collator, processing_classself.tokenizer, ) if self.data_args.same_dataset_within_batch: trainer.add_callback(EmbedderTrainerCallbackForDataRefresh(self.train_dataset))要点训练器四个核心组件模型、训练参数、数据集、数据整理器Collator——其中train_dataset与data_collator由AbsEmbedderRunner.load_train_dataset/load_data_collatorAbsRunner.py按same_dataset_within_batch标志分别构造为普通或同数据集同批变体数据刷新回调开启same_dataset_within_batch时注册EmbedderTrainerCallbackForDataRefresh用于每个 epoch 刷新数据集保证同 batch 样本始终来自同一数据集训练时 batch size 会被强制置 1由数据集内部自行拼批输出目录保护AbsEmbedderRunner.__init__会在output_dir非空且do_train且未设置overwrite_output_dir时抛出ValueError防止误覆盖已有训练产物。实战在微调脚本中驱动 AbsEmbedderTrainer实际训练时用户不需要直接触碰 Trainer 代码而是通过 CLI 入口与参数驱动。仓库提供的示例脚本 base.sh 演示了编码器基础模型的完整微调命令其中与训练器损失行为直接相关的参数配置如下training_args\ --output_dir ./test_encoder_only_base_bge-large-en-v1.5 \ --overwrite_output_dir \ --learning_rate 1e-5 \ --fp16 \ --num_train_epochs 4 \ --per_device_train_batch_size 2 \ --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 kl_div \ cmdtorchrun --nproc_per_node 2 \ -m FlagEmbedding.finetune.embedder.encoder_only.base \ $model_args $data_args $training_args eval $cmd数据侧则通过--train_group_size 8、--query_max_len 512、--passage_max_len 512、--query_instruction_for_retrieval等AbsEmbedderDataArguments参数控制示例数据位于 examples/finetune/embedder/example_data/。脚本中的--negatives_cross_device与--temperature 0.02正是上文分析中决定损失计算路径的两个核心开关前者让两卡共享负样本后者控制对比分数尺度。命令执行后AbsEmbedderRunner.run()或解码器 Runner 的run()见 runner.py会依次完成创建输出目录 →trainer.train(resume_from_checkpoint...)→trainer.save_model()触发前述_save逻辑→ 需要时合并 LoRA 权重save_merged_model。如何扩展自定义训练器若要在 FlagEmbedding 框架内实现自定义训练逻辑推荐遵循如下模式继承AbsEmbedderTrainer实现抽象方法_save可参考EncoderOnlyEmbedderTrainer的保存流程按需覆写compute_loss例如增加辅助损失、梯度惩罚等注意保持从模型输出中取 loss或自行计算并返回(loss, outputs)的契约在 Runner 的load_trainer中装配自定义训练器并在run中编排训练与保存流程若自定义模型需要新的损失计算方式则在AbsEmbedderModel子类中覆写compute_loss/compute_score/encode三个抽象方法训练器无需改动。从源码结构看abc层抽象的价值正在于此模型结构编码器 vs 解码器与训练基础设施Trainer/Runner彻底解耦新增一种模型架构只需实现Modeling与Trainer两个文件训练流程、分布式与保存逻辑全部复用。总结AbsEmbedderTrainer是 FlagEmbedding 嵌入模型微调的调度中枢它通过compute_loss把损失计算委托给模型内部的对比/蒸馏逻辑通过抽象方法_save统一各架构的保存契约同时完整继承transformers.Trainer的训练工程能力。理解它就等于理解了 FlagEmbedding 微调流水线从参数、数据、模型到训练器的完整协作关系——无论你使用编码器还是解码器嵌入模型是常规微调、跨设备负样本训练还是 BGE-M3 式多阶段蒸馏最终都汇聚到这套统一的训练器接口之上。关键源码索引抽象训练器AbsTrainer.py抽象模型与损失计算AbsModeling.py训练参数定义AbsArguments.py编码器 Trainer 实现trainer.py解码器 Trainer 实现trainer.pyRunner 装配逻辑AbsRunner.py完整微调示例base.sh【免费下载链接】FlagEmbeddingRetrieval and Retrieval-augmented LLMs项目地址: https://gitcode.com/GitHub_Trending/fl/FlagEmbedding创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门