ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

FlagEmbedding ABC 微调框架 Embedder 抽象层全解析:从 AbsArguments 到 AbsTrainer 的模块化训练 API

FlagEmbedding ABC 微调框架 Embedder 抽象层全解析:从 AbsArguments 到 AbsTrainer 的模块化训练 API FlagEmbedding ABC 微调框架 Embedder 抽象层全解析从 AbsArguments 到 AbsTrainer 的模块化训练 API【免费下载链接】FlagEmbeddingRetrieval and Retrieval-augmented LLMs项目地址: https://gitcode.com/GitHub_Trending/fl/FlagEmbedding导读本文围绕 FlagEmbedding 开源仓库中 abc/finetune/embedder.rst 所定义的 Embedder 微调抽象层展开系统讲解FlagEmbedding.abc.finetune.embedder模块中五个核心抽象组件——AbsArguments、AbsDataset、AbsModeling、AbsTrainer、AbsRunner——的设计意图、参数语义与源码实现。读完本文你将掌握 FlagEmbedding 统一微调框架的分层思想理解训练数据如何被加载、采样、拼接为 batch模型如何编码并计算对比学习/蒸馏损失以及如何基于这套抽象基类扩展出自定义 Embedder 微调方案。一、模块定位为什么需要一个 ABC 微调层FlagEmbedding 的目录结构遵循抽象基类abc— 具体实现finetune/ inference / evaluation的分层设计。abc 包 之下并列了finetune、inference、evaluation三个子模块其中 abc/finetune/embedder 与 abc/finetune/reranker 一一对应分别覆盖检索模型embedder与重排模型reranker两种训练范式。从仓库结构看抽象层通过统一的接口约束让上层具体实现如 finetune/embedder/encoder_only/base、finetune/embedder/decoder_only/base、finetune/reranker 等共享相同的数据契约、损失计算逻辑与训练流程从而支持 BGE 系列、BGE-M3、LLM 基座decoder-only等多种模型的低成本接入。embedder.rst的 5 个子文档AbsArguments.rst、AbsDataset.rst、AbsModeling.rst、AbsTrainer.rst、AbsRunner.rst正是这五个抽象组件的 API 参考索引下文逐一展开。二、AbsArguments训练参数的三大契约AbsArguments.py 定义了三个 dataclass分别对应模型参数、数据参数与训练参数是微调脚本命令行入参的解析基准训练入口通过 HfArgumentParser 组合使用。2.1 AbsEmbedderModelArguments模型与分词器初始化参数定义于 AbsArguments.py#L8-L40参数默认值说明model_name_or_path必填用于初始化的模型 checkpoint本地路径或 HF Hub 模型名config_nameNone与模型名不同时的预训练 config 名称/路径tokenizer_nameNone与模型名不同时的分词器名称/路径cache_dirNone预训练模型下载缓存目录trust_remote_codeFalse是否信任远端自定义代码use_fast_tokenizerTrue是否使用 fast tokenizertoken环境变量HF_TOKEN访问受限模型的 HF token默认从环境变量读取2.2 AbsEmbedderDataArguments数据加载与组批参数定义于 AbsArguments.py#L43-L130核心参数如下参数默认值说明train_dataNone一个或多个训练数据路径nargs数据必须含query: str、pos: List[str]、neg: List[str]字段cache_pathNone数据集缓存目录train_group_size8每个训练组包含的文本数1 正例 若干负例query_max_len32query 侧最大 token 数超长截断passage_max_len128passage 侧最大 token 数超长截断pad_to_multiple_ofNone若设置padding 后序列长度对齐为该值的整数倍max_example_num_per_dataset100000000每个数据集最多采样的样本数超限随机下采样query_instruction_for_retrievalNone检索场景的 query 指令query_instruction_format{}{}query 指令拼接格式如Instruct: {}{}\nknowledge_distillationFalse开启后读取数据的pos_scores/neg_scores用于蒸馏passage_instruction_for_retrievalNonepassage 侧指令passage_instruction_format{}{}passage 指令拼接格式shuffle_ratio0.0文本打乱比例对长文本按 chunk 随机重排用于长文本增强same_dataset_within_batchFalse同一 batch 内样本是否来自同一数据集small_threshold0小数据集阈值同一目录下小于该值的文件合并为一个数据集drop_threshold0合并后的数据集样本数低于该值时直接丢弃值得注意的源码细节__post_init__AbsArguments.py#L120-L130会在初始化时把指令格式中的字面量\\n替换为真实换行符并逐一校验train_data路径存在否则抛出FileNotFoundError从源头拦截错误路径。2.3 AbsEmbedderTrainingArguments训练超参与损失策略继承自 HFTrainingArguments定义于 AbsArguments.py#L133-L143额外扩展了 Embedder 专属超参参数默认值可选值/说明negatives_cross_deviceFalse是否跨设备共享负例多卡时扩大负例规模temperature0.02相似度得分缩放温度fix_position_embeddingFalse是否冻结位置编码参数sentence_pooling_methodcls池化方式cls/mean/last_tokennormalize_embeddingsTrue是否对嵌入做 L2 归一化sub_batch_sizeNone编码时的子 batch 大小内存紧张时拆分kd_loss_typekl_div蒸馏损失类型kl_div/m3_kd_lossuse_mrlFalse是否启用 MRLMatryoshka Representation Learning训练mrl_dims[]MRL 各层输出维度列表开启 MRL 时必填三、AbsDataset训练数据管线抽象AbsDataset.py 提供两类数据集实现及其配套 collator外加一个 epoch 刷新回调。3.1 AbsEmbedderTrainDataset 与 AbsEmbedderCollator标准检索训练对AbsEmbedderTrainDatasetAbsDataset.py#L23-L151在初始化时遍历args.train_data中的每个路径若为文件.json/.jsonl则直接加载若为目录则逐个加载目录下的 json/jsonl 文件最终用datasets.concatenate_datasets合并为一个大数据集。_load_datasetAbsDataset.py#L54-L81的关键逻辑仅在 rank 0 打印加载日志dist.get_rank()未初始化时按 0 处理通过datasets.load_dataset(json, ...)加载支持cache_path缓存样本数超过max_example_num_per_dataset时随机采样截断未开启蒸馏时主动移除pos_scores/neg_scores列开启蒸馏但数据缺少这两列时抛ValueError提示保证训练前就暴露数据契约问题。_shuffle_textAbsDataset.py#L83-L100实现了shuffle_ratio的语义仅当比例大于 0、文本长度超过 100 且随机命中时将文本按三分之一长度切块后随机重排拼接用于构造长文本扰动样本。__getitem__AbsDataset.py#L105-L151完成单样本三元组构造query 侧若配置了指令优先使用数据自带的prompt字段否则用query_instruction_for_retrieval再经query_instruction_format拼接从pos列表随机选一个正例同样可经_shuffle_text扰动负例采样若neg数量不足train_group_size - 1则通过random.sample(neg_all_idx * num, train_group_size - 1)循环补足——这是实现组内固定规模的关键技巧蒸馏模式下同步采样对应的pos_scores/neg_scores并校验得分必须为数值passage 侧指令按passage_instruction_for_retrieval/passage_instruction_format统一应用。AbsEmbedderCollatorAbsDataset.py#L153-L242继承DataCollatorWithPadding将 query、passage 分别按query_max_len/passage_max_len截断并 padding。它支持sub_batch_size分块当sub_batch_size 0时把一批 query/passage 拆成多个子 batch 分别 pad返回queries/passages的列表结构否则返回单个拼好的张量。最终输出固定为{queries, passages, teacher_scores, no_in_batch_neg_flag}四元组其中no_in_batch_neg_flag默认False。3.2 AbsEmbedderSameDatasetTrainDataset同数据集组批策略AbsEmbedderSameDatasetTrainDatasetAbsDataset.py#L245-L510实现同一 batch 内的样本来自同一数据集same_dataset_within_batchTrue时启用用于多源异构数据混合训练时避免跨域负例干扰。其核心差异小数据集合并初始化时对同一目录下样本数小于small_threshold的文件用concatenate_datasets合并合并后样本数仍低于drop_threshold则丢弃no_in_batch_neg 标记通过给文件/目录名添加no_in_batch_neg后缀AbsDataset.py#L287、AbsDataset.py#L304声明该数据集不使用 batch 内负例每数据集独立 batch size_get_file_batch_sizeAbsDataset.py#L360-L377优先读取数据自带的batch_size列若无则看type列含symmetric的数据类型 batch 减半否则用默认 batch size每 epoch 重建 batchrefresh_epochAbsDataset.py#L379-L401用确定性随机数生成器np.random.default_rng(seed)打乱数据集顺序与组内样本按batch_size * num_processes切块并丢弃不足整 batch 的尾部保证多进程切分严格对齐组规模自适应_get_train_group_sizeAbsDataset.py#L415-L439按数据type动态决策——only_1neg固定为 21 query 1 negsymmetric_class取min(len(neg)1, train_group_size)也支持数据自带的train_group_size列否则回落为全局train_group_size对称任务处理_create_batch_dataAbsDataset.py#L441-L510对symmetric_sts/symmetric_clustering类型passage 也使用 query 侧指令格式拼接其余类型才使用 passage 指令。配套的AbsEmbedderSameDatasetCollatorAbsDataset.py#L513-L604在 docstring 中明确要求使用后需将training_args.per_device_train_batch_size 1且dataloader_num_workers 0避免多进程破坏预构造的 batch并透传no_in_batch_neg_flag。3.3 EmbedderTrainerCallbackForDataRefreshepoch 数据刷新钩子EmbedderTrainerCallbackForDataRefreshAbsDataset.py#L607-L625继承 HFTrainerCallback在on_epoch_end中调用数据集的refresh_epoch()从而在每个 epoch 结束时依据 seed 重新打乱并重建 batch——这正是同数据集训练模式每轮数据重排的机制来源。四、AbsModeling模型编码与损失计算抽象AbsModeling.py 定义了统一的模型输出结构与抽象模型基类。4.1 EmbedderOutput标准化的前向输出EmbedderOutputAbsModeling.py#L16-L24继承transformers.file_utils.ModelOutput包含四个可选张量字段q_repsquery 表示、p_repspassage 表示、loss、scores是各具体实现 forward 返回的统一载体。4.2 AbsEmbedderModel抽象模型基类AbsEmbedderModelAbsModeling.py#L27-L365同时继承ABC与nn.Module构造函数接收base_model、tokenizer及negatives_cross_device、temperature、sub_batch_size、kd_loss_type、use_mrl、mrl_dims等超参。两个显式校验值得注意开启negatives_cross_device但分布式未初始化时抛ValueError开启 MRL 但mrl_dims为空时同样报错。四个抽象方法是子类必须实现的核心契约encode(features)AbsModeling.py#L72-L79输入特征、输出嵌入compute_loss(scores, target)AbsModeling.py#L81-L89基于得分与目标计算损失compute_score(q_reps, p_reps)AbsModeling.py#L91-L99计算 query 与 passage 表示间的相似度得分矩阵save(output_dir)AbsModeling.py#L101-L108保存模型到指定目录。三种负例策略是检索对比学习的核心分支由forward依据no_in_batch_neg_flag与negatives_cross_device自动路由_compute_no_in_batch_neg_lossAbsModeling.py#L149-L169不使用 batch 内负例只对每个 query 与其组内 passage 计算局部得分(batch_size, group_size)_compute_in_batch_neg_lossAbsModeling.py#L171-L201标准 in-batch 负例得分矩阵为(batch_size, batch_size * group_size)目标为idxs * group_size即每行正例所在列索引这是 InfoNCE 类对比损失在稠密检索中的标准实现_compute_cross_device_neg_lossAbsModeling.py#L203-L241通过_dist_gather_tensorAbsModeling.py#L344-L365用dist.all_gather收集所有 rank 的 q/p 表示拼接为(world_size * batch_size, ...)得分矩阵扩大到全卡规模从而用跨卡负例显著提升对比学习负例数量。get_local_score/compute_local_scoreAbsModeling.py#L110-L147负责从全量得分矩阵中按group_size对角提取每个 query 的组内得分供蒸馏或非 in-batch 场景使用。知识蒸馏KDforward将teacher_scores转为张量并做softmax得到teacher_targets后由静态方法distill_lossAbsModeling.py#L303-L342计算kl_div学生得分log_softmax与教师目标逐元素相乘取负均值m3_kd_loss对每个组内位置构造带 mask 的交叉熵并按教师目标加权求和——这是 BGE-M3 多向量蒸馏使用的损失变体。对kl_div类型还会叠加常规对比损失compute_loss实现蒸馏 硬负例联合训练。MRLMatryoshka Representation Learning当use_mrlTrue时encode需返回多个维度的表示列表forward对mrl_dims中每个维度分别计算损失并取平均AbsModeling.py#L287-L293使单一模型同时支持不同维度的嵌入输出。五、AbsRunner 与 AbsTrainer训练流程编排AbsRunner.rst 定义了AbsEmbedderRunner的训练编排接口其方法清单即训练流水线的五个阶段load_tokenizer_and_model装配模型参数与分词器构建AbsEmbedderModel实例load_train_dataset依据same_dataset_within_batch选择AbsEmbedderTrainDataset或AbsEmbedderSameDatasetTrainDatasetload_data_collator配套加载对应的 collator含sub_batch_size与same_dataset模式下的 batch size 约束load_trainer组装AbsEmbedderTrainer并在同数据集模式下注册EmbedderTrainerCallbackForDataRefreshrun启动训练入口。AbsTrainer.rst 中的AbsEmbedderTrainer继承自 HFTrainer重写compute_loss以调用模型的forward将 collator 输出的queries/passages/teacher_scores/no_in_batch_neg_flag传入并返回loss。在仓库中这套抽象层的具体落地实现位于 finetune/embedder/encoder_only/baseBGE 系列与 finetune/embedder/decoder_only/baseLLM 基座等目录并配套提供了可直接运行的示例脚本见 examples/finetune/embedder 下的.sh文件与样例数据example_data目录中的 jsonl 文件。六、抽象层到具体实现的组合使用综合各抽象组件一次典型的 Embedder 微调可以归纳为如下数据流参数层AbsEmbedderModelArgumentsAbsEmbedderDataArgumentsAbsEmbedderTrainingArguments三个 dataclass 经 HfArgumentParser 从命令行解析__post_init__提前校验数据路径与指令格式数据层AbsEmbedderRunner.load_train_dataset按same_dataset_within_batch选择数据集类AbsEmbedderTrainDataset.__getitem__产出(query, passages, teacher_scores)三元组collator 完成截断、padding可选 sub-batch 分块并输出统一四元组模型层AbsEmbedderModel.forward编码 query/passage按no_in_batch_neg_flag→negatives_cross_device→ 默认 in-batch 的顺序选择损失函数融合蒸馏kl_div/m3_kd_loss与 MRL 多维度损失训练层AbsEmbedderTrainer.compute_loss驱动反向传播同数据集模式下由EmbedderTrainerCallbackForDataRefresh在每个 epoch 结束时刷新数据排布。具体模型子类只需实现encode、compute_loss、compute_score、save四个抽象方法即可无缝接入该训练框架——例如 encoder-only 实现中通过sentence_pooling_methodcls/mean/last_token控制池化通过fix_position_embedding控制是否冻结位置编码normalize_embeddings控制输出归一化这些超参最终都经由AbsEmbedderTrainingArguments贯穿到模型与损失计算的全链路。七、小结FlagEmbedding.abc.finetune.embedder抽象层是 FlagEmbedding 微调体系的标准化接口层AbsArguments以三个 dataclass 收敛全部入参AbsDataset以两类数据集与两个 collator 解决多源数据加载、同数据集组批、长文本扰动与知识蒸馏数据校验AbsModeling以四个抽象方法加三种负例策略、两种蒸馏损失与 MRL 支持覆盖了稠密检索训练的核心算法面AbsRunner/AbsTrainer则将以上组件编排为可复用的训练流水线。理解这一层是阅读或二次开发 BGE 系列、decoder-only LLM 嵌入模型微调代码的最佳起点——后续的具体实现如 encoder_only/base 与 decoder_only/base均在此骨架之上填充模型细节。【免费下载链接】FlagEmbeddingRetrieval and Retrieval-augmented LLMs项目地址: https://gitcode.com/GitHub_Trending/fl/FlagEmbedding创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表