FlagEmbedding ABC 微调数据管线解析:AbsDataset 数据集与 Collator 深度指南
FlagEmbedding ABC 微调数据管线解析AbsDataset 数据集与 Collator 深度指南【免费下载链接】FlagEmbeddingRetrieval and Retrieval-augmented LLMs项目地址: https://gitcode.com/GitHub_Trending/fl/FlagEmbedding本文以 FlagEmbedding 仓库abcAbstract Base Classes抽象层中的AbsDataset模块为主线系统讲解 Embedder 微调fine-tuning场景下的数据集定义、批处理Batching与数据整理Collation机制从标准AbsEmbedderTrainDataset到支持同数据集分批的AbsEmbedderSameDatasetTrainDataset以及配套的AbsEmbedderCollator、AbsEmbedderSameDatasetCollator和训练回调。读完本文你将掌握 FlagEmbedding 中query/pos/neg训练数据的组织格式、train_group_size、知识蒸馏分数、no_in_batch_neg标记等关键配置的底层含义能够为自己的检索模型微调任务正确准备数据并选用合适的数据管线。一、AbsDataset 在微调流程中的位置FlagEmbedding 的abc目录是面向 Embedder 与 Reranker 微调/推理的抽象基类层其中 Embedder 微调由四个核心文件协同工作AbsArguments.py定义模型、数据、训练三类 dataclass 参数AbsModeling.py定义抽象模型与前向/损失计算AbsDataset.py本文主角负责训练数据的加载、采样与批处理AbsTrainer.py继承 Hugging FaceTrainer的抽象训练器。它们之间的装配关系由 AbsRunner.py 完成Runner 在初始化时调用load_train_dataset()与load_data_collator()根据data_args.same_dataset_within_batch选择标准或同数据集两套数据方案并将对应的 dataset 与 collator 交给 Trainer。因此理解AbsDataset是理解 FlagEmbedding 微调数据流向的入口。AbsDataset 模块对外暴露五个公开类对应 API 文档 AbsDataset.rst类名职责AbsEmbedderTrainDataset标准训练数据集混合所有数据源按 (query, passages, teacher_scores) 逐条返回AbsEmbedderCollator标准数据整理器将 query/passage 分词、截断、padding 并组装模型输入AbsEmbedderSameDatasetTrainDataset同数据集训练数据集保证一个 batch 内所有样本来自同一数据集AbsEmbedderSameDatasetCollator同数据集数据整理器配合上述数据集使用EmbedderTrainerCallbackForDataRefresh训练回调每个 epoch 结束时刷新重打乱数据二、AbsEmbedderTrainDataset标准训练数据集AbsEmbedderTrainDataset继承torch.utils.data.Dataset构造函数接收两个参数argsAbsEmbedderDataArguments与tokenizerPreTrainedTokenizer。其核心设计目标是从一个或多个 JSON/JSONL 文件中加载query/pos/neg结构的数据并在__getitem__中完成正负样本的在线采样。2.1 多数据源加载与过滤构造函数遍历args.train_data中给出的所有路径支持文件或目录两种形式目录会被递归遍历其中的.json/.jsonl文件逐个加载后调用datasets.concatenate_datasets拼接为统一的 Hugging Face Dataset。空数据集会被跳过源码if len(temp_dataset) 0: continue。2.2 _load_dataset加载与蒸馏分数检查_load_dataset是文档指定的核心方法之一实现要点如下通过datasets.load_dataset(json, data_filesfile_path, splittrain, cache_dirself.args.cache_path)加载数据若样本数超过args.max_example_num_per_dataset则随机采样max_example_num_per_dataset条用于控制单个数据集的上限依据knowledge_distillation参数决定pos_scores/neg_scores两列的去留不启用蒸馏时直接remove_columns删除这两列避免污染训练输入启用蒸馏时若数据中缺少这两列则抛出ValueError提示使用知识蒸馏必须提供 pos_scores 与 neg_scores。从源码结构看该检查在标准数据集与同数据集两个子类中行为一致只是标准数据集还额外做了列删除后的蒸馏校验。2.3 _shuffle_text文本打乱增强_shuffle_text实现了一种文本级数据增强当shuffle_ratio 0、文本长度超过 100 且随机概率命中时将文本按len(text)//3 1的块大小切分成约 3 段并随机重排以 连接后返回。shuffle_ratio默认0.0关闭可在 AbsArguments.py 中调整。该增强适用于对顺序不敏感的长文本如文档检索场景。2.4getitem正负样本在线采样每次取一条样本时执行以下流程query 指令格式化若设置了query_instruction_for_retrieval则以query_instruction_format默认{}{}拼接指令与查询若样本自带prompt字段则优先使用样本级prompt正样本采样pos是列表随机选择一个下标pos_idx并对该正样本执行_shuffle_text负样本采样目标组大小为train_group_size即每条 query 需要train_group_size - 1个负样本。若neg数量不足则先复制列表num ceil((train_group_size - 1) / len(neg))次再random.sample若足够则直接采样蒸馏分数对齐启用knowledge_distillation时同步收集与正负样本一一对应的pos_scores[pos_idx]与各neg_scores并校验全部为 int/float否则抛ValueErrorpassage 指令格式化若设置passage_instruction_for_retrieval对所有 passage 套用passage_instruction_format返回三元组(query, passages, teacher_scores)。可以看出train_group_size决定每个训练组的大小正 1 负 N其数值必须与负样本数量协同设置例如仓库示例脚本中统一使用--train_group_size 8。三、AbsEmbedderCollator标准数据整理器AbsEmbedderCollator继承transformers.DataCollatorWithPadding是一个dataclass自带三个字段query_max_len: int 32、passage_max_len: int 128、sub_batch_size: int -1。其__call__逻辑将 batch 中的(queries, passages, teacher_scores)三个字段拆出若teacher_scores[0] is None则整体置为None否则展平列表对 queries 与 passages 分别调用tokenizer(..., truncationTrue, max_length...)完成分词与截断——query 截断到query_max_lenpassage 截断到passage_max_lenpadding 阶段分两种模式sub_batch_size 0默认整体一次tokenizer.padsub_batch_size 0按该大小把分词结果切成多个子块分别 padding返回q_collated/d_collated列表。这是为显存受限时子批次编码预留的分块接口返回字典{queries: q_collated, passages: d_collated, teacher_scores: teacher_scores, no_in_batch_neg_flag: False}。该输出直接喂给AbsEmbedderModel.forward见 AbsModeling.pyq_reps encode(queries)、p_reps encode(passages)其中 passage 数量为batch_size * group_size与train_group_size严格对应。四、AbsEmbedderSameDatasetTrainDataset同数据集分批训练这是AbsDataset中实现最复杂的类用于满足同一 batch 内所有样本来自同一数据集的训练需求。该模式适用于对称语义任务STS、聚类以及no_in_batch_neg数据集——这些场景下若 batch 内混入其他数据集样本会破坏 in-batch 负例的语义甚至造成训练信号错误。4.1 构造函数数据集分类与 no_in_batch_neg 标记构造时遍历args.train_data对每个数据源解析出no_in_batch_neg标记——约定是在文件/目录名扩展名前追加no_in_batch_neg后缀例如classification-no_in_batch_neg/、clustering-no_in_batch_neg/仓库示例目录 example_data 即遵循此命名。该标记会通过 collator 一路传入模型最终决定走无 in-batch 负例的损失分支。同时按small_threshold/drop_threshold两个阈值对数据集分类处理样本数 small_threshold的称为小数据集同一目录下的小数据集会被合并合并后总样本数 drop_threshold则整体丢弃其余按正常数据集加入训练。4.2 _get_file_batch_size按数据特性调整 batch 大小该方法优先读取数据中的batch_size列取第一条记录的整数值否则若存在type列且类型含symmetric则将default_batch_size减半——从源码注释看这是让对称数据使用更小的 batch 大小默认返回default_batch_size。4.3 refresh_epoch每轮次的确定性打乱refresh_epoch在初始化时和每个 epoch 结束时被调用完成三件事用np.random.default_rng(seed)确定性生成器打乱数据集顺序在每个数据集内部打乱样本下标并按batch_size_idxs[dataset] * num_processes切分为批次最后一个不满的批次被丢弃打乱批次顺序后存入self.batch_datas并将self.step归零。其中数据集顺序、数据集内样本、批次顺序三重打乱均基于固定 seed保证分布式训练下可复现。4.4getitem按进程切分批次由于一个 batch 是先按num_processes份拼接再切分的__getitem__中会按self.process_index取出当前进程对应的子批调用_create_batch_data生成(queries, passages, teacher_scores, no_in_batch_neg_flag)四元组并递增self.step。4.5 _get_train_group_size任务类型感知的组大小根据 batch 内首条记录的type字段决定组大小only_1neg固定返回 2正 1 负 1symmetric_class返回min(len(neg[0]) 1, train_group_size)其他类型返回args.train_group_size无type但有train_group_size列且为正整数使用该值否则回退到args.train_group_size。同时该方法校验同一 batch 内type必须一致源码assert batch_raw_data[type][i] data_type避免混批。4.6 _create_batch_data完整样本构造逐条 query 执行query 指令格式化优先样本prompt→ 随机选正样本并_shuffle_text→ 负样本采样不足时复制采样→ 蒸馏分数收集仅当启用knowledge_distillation→ 指令拼接。这里有一个值得注意的细节当data_type为symmetric_sts或symmetric_clustering时passage 也使用 query 指令格式而非passage_instruction_format因为对称任务中正样本与 query 处于同一语义空间其他任务才使用 passage 指令。五、AbsEmbedderSameDatasetCollator同数据集整理器AbsEmbedderSameDatasetCollator与标准 collator 结构几乎一致差异在于由于 dataset 的__getitem__每次已返回一个完整 batch而非单条样本__call__直接取features[0]的四个元素。其 docstring 明确要求配套的训练参数training_args.per_device_train_batch_size 1 training_args.dataloader_num_workers 0 # avoid multi-processing这两项配置由 AbsRunner.py 在加载同数据集模式时自动强制设置用户无需也不应手动覆盖——因为一个 batch 一个数据集的逻辑已在 dataset 层完成DataLoader 层只需每次取一个元素。六、EmbedderTrainerCallbackForDataRefresh轮次刷新回调EmbedderTrainerCallbackForDataRefresh继承transformers.TrainerCallback构造时接收AbsEmbedderSameDatasetTrainDataset实例其on_epoch_end钩子对应 API 文档中的方法在每个 epoch 结束时调用self.train_dataset.refresh_epoch()实现每轮重新打乱数据。该回调由具体 Runner 在构建 Trainer 时注册是整个同数据集数据管线闭环的关键一环。七、训练数据格式规范与仓库示例综合源码FlagEmbedding Embedder 微调的标准 JSONL 格式为{query: 查询文本, pos: [正样本1, ...], neg: [负样本1, 负样本2, ...], prompt: 任务提示, type: normal}可选字段字段类型说明pos_scoresList[float]正样本教师分数启用knowledge_distillation时必需neg_scoresList[float]负样本教师分数同上promptstr样本级指令优先于query_instruction_for_retrievaltypestr任务类型normal/only_1neg/symmetric_class/symmetric_sts/symmetric_clusteringbatch_sizeint该数据集专用 batch 大小可选train_group_sizeint该数据集专用组大小可选仓库提供可直接参考的真实数据retrieval/msmarco.jsonl普通检索、带蒸馏分数与 prompt、sts/sts.jsonlsymmetric_sts类型、classification-no_in_batch_neg/与clustering-no_in_batch_neg/通过目录名后缀声明不使用 in-batch 负例。八、数据相关参数全表AbsEmbedderDataArguments以下参数均定义于 AbsArguments.py是配置数据管线时最常用的开关参数默认值作用train_data无必填一个或多个数据路径支持文件或目录query/pos/neg为必需字段目录不存在时抛FileNotFoundErrorcache_pathNoneHugging Face datasets 的缓存目录train_group_size8每条 query 的正负样本组大小含正样本query_max_len32query 截断长度passage_max_len128passage 截断长度pad_to_multiple_ofNonepadding 到该值的整数倍max_example_num_per_dataset100000000每个数据集最多加载的样本数超出随机采样query_instruction_for_retrievalNone默认 query 指令query_instruction_format{}{}query 指令拼接模板\n会被转义为换行passage_instruction_for_retrievalNone默认 passage 指令passage_instruction_format{}{}passage 指令拼接模板knowledge_distillationFalse是否启用知识蒸馏要求数据含pos_scores/neg_scoresshuffle_ratio0.0长文本100 字符打乱增强概率same_dataset_within_batchFalse是否启用同数据集分批模式small_threshold0小数据集判定阈值同目录小数据集合并且设更小 batchdrop_threshold0合并后小数据集样本数低于该值则丢弃九、训练脚本中的实际用法仓库示例脚本 decoder_only/base.sh 展示了标准模式的参数组合--knowledge_distillation False而 decoder_only/base_same_dataset.sh 则展示了同数据集 蒸馏的组合--train_data /path/to/train_data \ --train_group_size 8 \ --knowledge_distillation True \ --same_dataset_within_batch True \对应地encoder_only 下也有base.sh、base_same_dataset.sh、m3.sh、m3_same_dataset.sh等同类脚本覆盖编码器与解码器两类架构可作为配置参考。十、数据流向下游collator 输出如何驱动模型训练理解AbsDataset之后可以顺带看清它与模型侧的契约。AbsEmbedderModel.forward接收 collator 返回的四元组teacher_scores非空时被 reshape 为(batch_size, group_size)并做softmax得到teacher_targets随后由distill_loss支持kl_div与m3_kd_loss两种计算蒸馏损失no_in_batch_neg_flagTrue时走_compute_no_in_batch_neg_loss仅使用组内正负样本不引入 batch 内其他 query 的 passage 作为负例为False时根据negatives_cross_device决定走_compute_in_batch_neg_loss还是跨卡聚合的_compute_cross_device_neg_loss。因此标准模式下AbsEmbedderCollator输出固定no_in_batch_neg_flagFalse配合negatives_cross_device见 AbsArguments.py实现 in-batch 负例含跨卡训练而同数据集模式则通过文件名约定把该标记逐层传递到损失函数实现精细的负例控制。小结AbsDataset模块是 FlagEmbedding Embedder 微调数据侧的完整抽象标准数据集与同数据集两套方案分别服务混合检索数据与对称/无 in-batch 负例两类训练诉求collator 负责把文本组织为模型可直接消费的 tensor 字典epoch 回调保证同数据集模式每轮重排而AbsEmbedderDataArguments中的十余个参数则提供了从截断长度、组大小到蒸馏开关的细粒度控制。对想要基于 FlagEmbedding 微调自定义检索模型的开发者而言掌握本模块的数据格式约定与参数语义是搭建稳定、可复现训练流程的第一步。【免费下载链接】FlagEmbeddingRetrieval and Retrieval-augmented LLMs项目地址: https://gitcode.com/GitHub_Trending/fl/FlagEmbedding创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考