ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

sentence-transformers 训练参数完全指南:SentenceTransformerTrainingArguments 详解

sentence-transformers 训练参数完全指南:SentenceTransformerTrainingArguments 详解 人工智能NLPEmbedding微调【免费下载链接】sentence-transformersState-of-the-Art Embeddings, Retrieval, and Reranking项目地址https://gitcode.com/gh_mirrors/se/sentence-transformers点击查看免费下载导读本文围绕 sentence-transformers 仓库中 docs/package_reference/sentence_transformer/training_args.md 所定义的SentenceTransformerTrainingArguments展开系统讲解 Sentence Transformer 训练时的全部核心参数从必须的output_dir到针对提示词prompts、批采样器batch sampler、Router 路由与分组学习率的进阶配置。读完本文你将掌握如何为 SentenceTransformerTrainer 正确构造训练参数对象并根据不同的损失函数、多数据集场景与模型架构挑选合适的参数组合。一、参数类继承关系与定位SentenceTransformerTrainingArguments是 sentence-transformers 中面向句向量模型SentenceTransformer专用的训练参数数据类dataclass定义于 sentence_transformers/sentence_transformer/training_args.py。它的继承链为transformers.TrainingArguments ↑ BaseTrainingArgumentssentence_transformers.base.training_args ↑ SentenceTransformerTrainingArgumentssentence_transformers.sentence_transformer.training_argstransformers.TrainingArguments来自 Hugging Face Transformers 库提供了绝大多数通用训练参数num_train_epochs、per_device_train_batch_size、warmup_steps、fp16/bf16、eval_strategy、save_strategy、logging_steps、run_name等。BaseTrainingArguments定义在 sentence_transformers/base/training_args.py在 Transformers 参数基础上补充了 sentence-transformers 特有的五个参数prompts、batch_sampler、multi_dataset_batch_sampler、router_mapping、learning_rate_mapping。SentenceTransformerTrainingArguments几乎完全复用BaseTrainingArguments的字段作为句向量模型的专用入口同时从 sentence_transformers/base/sampler.py 导出BatchSamplers与MultiDatasetBatchSamplers两个枚举供用户直接使用见 training_args.py 的__all__。from sentence_transformers import SentenceTransformerTrainingArguments from sentence_transformers.sentence_transformer.training_args import BatchSamplers, MultiDatasetBatchSamplers从源码结构看cross_encoder与sparse_encoder的训练参数类如CrossEncoderTrainingArguments、SparseEncoderTrainingArguments同样继承自BaseTrainingArguments因此本文介绍的五类 sentence-transformers 专属参数在其他编码器训练中也通用本文以句向量场景为主线讲解。二、参数总览表SentenceTransformerTrainingArguments可接受的参数分为三层下表先给出 sentence-transformers 专属参数其余通用参数继承自 Transformers参数类型默认值作用output_dirstr无必填模型检查点输出目录promptsstr/Dict[str, str]/Dict[str, Dict[str, str]]None为数据集各列或各数据集指定提示词模板batch_samplerBatchSamplers/str/ 采样器类 / 工厂函数BatchSamplers.BATCH_SAMPLER单数据集下的批采样策略multi_dataset_batch_samplerMultiDatasetBatchSamplers/str/ 采样器类 / 工厂函数MultiDatasetBatchSamplers.PROPORTIONAL多数据集下的批次调度策略router_mappingDict[str, str]/Dict[str, Dict[str, str]]{}将数据集列映射到 Router 路由如 query、documentlearning_rate_mappingDict[str, float]{}按参数名正则表达式为不同模块设置不同学习率num_train_epochs、per_device_train_batch_size、per_device_eval_batch_size、warmup_steps、fp16/bf16、eval_strategy、eval_steps、save_strategy、save_steps、save_total_limit、logging_steps、run_name、dataloader_num_workers、dataloader_drop_last等继承自transformers.TrainingArguments与 Transformers 一致通用训练、评估、保存与调度配置三、必需参数output_diroutput_dir是所有训练参数中唯一必须提供的参数继承自 Transformers用于指定模型检查点与训练产物的写入目录args SentenceTransformerTrainingArguments( output_dircheckpoints, num_train_epochs1, per_device_train_batch_size16, per_device_eval_batch_size16, warmup_steps0.1, fp16True, # GPU 不支持 FP16 时请改为 False bf16False, # GPU 支持 BF16 时可改为 True eval_strategysteps, eval_steps100, save_strategysteps, save_steps100, save_total_limit2, logging_steps100, run_namemy-sts-training, # 安装 wandb 时用于 WB 实验名 )上面的写法直接取自仓库官方示例 examples/sentence_transformer/training/avg_word_embeddings/training_stsbenchmark_avg_word_embeddings.py 与 examples/sentence_transformer/training/adaptive_layer/adaptive_layer_nli.py可作为日常训练的标准骨架。四、prompts为数据集列配置提示词prompts用于在训练、评估、测试阶段为数据集的不同列指定提示词模板适用于带指令的instruction-based句向量模型。它接受四种格式源码 docstring 在 training_args.py 中给出了完整说明str单一提示词应用于所有数据集的所有列。无论数据集是datasets.Dataset还是datasets.DatasetDict均适用。Dict[str, str]列名到提示词的映射例如{query: Query: , document: Document: }同样适用于Dataset与DatasetDict。Dict[str, str]数据集名到提示词的映射仅当数据集是DatasetDict或Dataset的字典时使用。Dict[str, Dict[str, str]]数据集名 → 列名 → 提示词的嵌套映射同样仅用于多数据集DatasetDict场景。底层解析机制在 sentence_transformers/base/data_collator.py 中_resolve_prompts负责按批次解析提示词若prompts是非空字典且批次带有dataset_name列则优先取prompts[dataset_name]随后_get_prompt_for_column再按列名取出该列对应的提示词若prompts本身就是字符串则所有列共用该提示词。在 sentence_transformers/base/trainer.py 的get_data_collator中args.prompts被直接传入数据整理器data collator从而在每个 batch 的预处理阶段生效。关于字符串形式的兼容处理__post_init__中training_args.py对prompts做了兼容处理若通过命令行传入字符串会先尝试用json.loads解析为字典若解析失败则回退为“对所有列生效的单一提示词字符串”。这与router_mapping、learning_rate_mapping解析失败直接抛错的行为不同。五、batch_sampler单数据集批采样策略batch_sampler控制训练样本如何被分组成 batch默认为BatchSamplers.BATCH_SAMPLER等价于 PyTorch 原生BatchSampler。其合法取值定义在 sentence_transformers/base/sampler.py 的BatchSamplers枚举中枚举值底层采样器适用场景BatchSamplers.BATCH_SAMPLER默认DefaultBatchSampler等价于 PyTorchBatchSampler常规训练BatchSamplers.NO_DUPLICATESNoDuplicatesBatchSampler保证 batch 内样本值跨列唯一依赖 batch 内负样本的损失函数BatchSamplers.NO_DUPLICATES_HASHEDNoDuplicatesBatchSampler(precompute_hashesTrue)用 xxhash 预计算哈希加速查重需安装xxhash库占用少量额外内存同上尤其推荐用于图像/音频等媒体数据集BatchSamplers.GROUP_BY_LABELGroupByLabelBatchSampler每个 batch 至少包含 2 个不同标签、每个标签至少 2 个样本batch 内三元组挖掘类损失与损失函数的搭配建议源码 docstring 明确给出了推荐组合sampler.pyNO_DUPLICATES 系列推荐搭配使用 batch 内负样本in-batch negatives的损失MultipleNegativesRankingLoss、CachedMultipleNegativesRankingLoss、MultipleNegativesSymmetricRankingLoss、CachedMultipleNegativesSymmetricRankingLoss、MegaBatchMarginLoss、GISTEmbedLoss、CachedGISTEmbedLoss源码位于 sentence_transformers/sentence_transformer/losses/。GROUP_BY_LABEL推荐搭配 batch 内三元组挖掘损失BatchAllTripletLoss、BatchHardSoftMarginTripletLoss、BatchHardTripletLoss、BatchSemiHardTripletLoss。官方示例中的典型用法adaptive_layer_nli.pyargs SentenceTransformerTrainingArguments( output_diroutput_dir, batch_samplerBatchSamplers.NO_DUPLICATES, # MultipleNegativesRankingLoss 受益于 batch 内无重复样本 ... )NO_DUPLICATES 的实现细节NoDuplicatesBatchSamplersampler.py在__iter__中基于打乱后的索引构建单链表逐样本检查其值集合是否与当前 batch 重叠重叠则推迟到后续 batch从而保证每个 batch 内样本值跨列唯一。当一轮完整遍历产出的 batch 数少于__len__承诺的批次数drop_lastTrue且重复值很多时可能发生会重新打乱并补充采样最多 2 轮同时输出一次性警告。NO_DUPLICATES_HASHED变体则通过datasets.map预先计算每行各列的 xxhash64 值将重复检查从“逐行读数据集、重新解码媒体”变为 O(1) 哈希比对precompute_num_proc默认取min(8, cpu_count - 1)。GROUP_BY_LABEL 的实现细节GroupByLabelBatchSamplersampler.py要求batch_size为大于等于 4 的偶数且数据集中至少存在 2 个各含 2 个以上样本的标签否则抛出ValueError。每个 batch 由多个标签轮流各出 2 个样本构成保证 batch 内标签多样是 batch 内三元组挖掘的前提。注意valid_label_columns指定候选标签列名采样器取第一个在数据集中实际存在的列作为标签来源。自定义 batch samplerbatch_sampler还接受自定义实现sampler.py子类化DefaultBatchSampler把类本身而非实例传给batch_sampler参数或传入一个接受dataset、batch_size、drop_last、valid_label_columns、generator、seed并返回DefaultBatchSampler实例的工厂函数。在 training_args.py 的__post_init__中若传入的是字符串会被自动转换为对应的枚举值to_dicttraining_args.py在序列化时会剔除不可 pickle 的可调用采样器。六、multi_dataset_batch_sampler多数据集调度策略当使用DatasetDict或多个Dataset组成的训练集时multi_dataset_batch_sampler决定从各数据集取 batch 的顺序默认为MultiDatasetBatchSamplers.PROPORTIONAL。合法取值定义在 sampler.py枚举值底层采样器行为MultiDatasetBatchSamplers.PROPORTIONAL默认ProportionalBatchSampler按各数据集大小成比例采样所有样本都会被用到大数据集被采样得更频繁MultiDatasetBatchSamplers.ROUND_ROBINRoundRobinBatchSampler各数据集轮流各取一个 batch直到某个数据集耗尽每个数据集被平等对待但小数据集可能会用不完所有样本底层MultiDatasetDefaultBatchSamplersampler.py接收一个ConcatDataset和一组子采样器并将set_epoch传递给所有子采样器以保证每个 epoch 的可复现打乱。RoundRobinBatchSampler的__len__为min(len(sampler)) * 数量ProportionalBatchSampler的__len__为各子采样器长度之和sampler.py。多数据集 自定义采样器的用法sampler.py子类化MultiDatasetDefaultBatchSampler后把类传给参数或传入接受datasetConcatDataset、batch_samplers各子数据集采样器列表、generator、seed的工厂函数。典型用法来自 sampler.py 的官方示例配合CoSENTLoss训练跨领域 STSfrom datasets import Dataset, DatasetDict from sentence_transformers import SentenceTransformer, SentenceTransformerTrainer, SentenceTransformerTrainingArguments from sentence_transformers.sentence_transformer.training_args import MultiDatasetBatchSamplers from sentence_transformers.sentence_transformer.losses import CoSENTLoss model SentenceTransformer(microsoft/mpnet-base) train_general Dataset.from_dict({ sentence_A: [Its nice weather outside today., He drove to work.], sentence_B: [Its so sunny., He took the car to the bank.], score: [0.9, 0.4], }) train_medical Dataset.from_dict({ sentence_A: [The patient has a fever., The doctor prescribed medication.], sentence_B: [The patient feels hot., The medication was given to the patient.], score: [0.8, 0.6], }) train_dataset DatasetDict({general: train_general, medical: train_medical}) loss CoSENTLoss(model) args SentenceTransformerTrainingArguments( output_dircheckpoints, multi_dataset_batch_samplerMultiDatasetBatchSamplers.PROPORTIONAL, ) trainer SentenceTransformerTrainer( modelmodel, argsargs, train_datasettrain_dataset, lossloss, ) trainer.train()多数据集时的 dataset_name 列当数据集是DatasetDict且满足以下任一条件时trainer.pyTrainer 会自动为数据集添加dataset_name列损失是一个字典按数据集区分损失、prompts是字典、或router_mapping是“数据集名 → 列 → 路由”的嵌套字典。该列是_resolve_prompts与_resolve_router_mapping按数据集名解析配置的关键。七、router_mapping为 Router 模型指定路由router_mapping用于将数据集列映射到 Router 模块的路由route例如 query 或 document从而让 Router 模型在训练时把不同输入列送到正确的子编码器。接受两种格式Dict[str, str]列名到路由的映射如{question: query, positive: document, negative: document}。Dict[str, Dict[str, str]]数据集名 → 列名 → 路由 的嵌套映射用于多数据集训练/评估。强制校验如果模型包含Router模块但未提供router_mappingget_data_collator会直接抛出ValueErrortrainer.py并提示正确的映射示例对应测试见 tests/base/modules/test_router.py。数据整理器在每批预处理时通过_resolve_router_mappingdata_collator.py按dataset_name解析嵌套映射再经_get_task_for_column为每个输入列打上任务标签task stamp。官方示例来自 router.py 的 docstring 示例——训练一个查询/文档不对称的双塔模型args SentenceTransformerTrainingArguments( ..., router_mapping{ question: query, positive: document, negative: document, }, )对应测试用例可参考 tests/base/test_trainer.py 与 tests/base/test_data_collator.py。与 prompts 的组合prompts与router_mapping可以同时生效get_data_collator会把两者一并传给数据整理器trainer.py使同一列既能套用特定提示词又能路由到正确的子编码器。多向量编码器MultiVectorEncoder场景中任务标签还会被MultiVectorMask等模块读取见 sentence_transformers/multi_vector_encoder/losses/multiple_negatives_ranking.py 的注释。八、learning_rate_mapping按模块设置分组学习率learning_rate_mapping通过“参数名正则表达式 → 学习率”的映射为模型不同部分设置不同学习率例如{SparseStaticEmbedding\\.*: 1e-3}表示对SparseStaticEmbedding模块的所有参数使用1e-3其余部分仍用全局learning_rate。这在冻结/微调混合场景如稀疏编码器、Router 多路由中非常实用。底层实现在 sentence_transformers/base/trainer.py 中Trainer 构造优化器参数组时会遍历args.learning_rate_mapping用re.search(parameter_pattern, n)在loss_model.named_parameters()中找出所有匹配的参数将这些参数从既有优化器组中剔除避免同一参数出现在多个组若正则没有匹配到任何参数则抛出ValueError提示检查模式为匹配参数单独建组设置lrlearning_rate并按是否属于衰减参数get_decay_parameter_names决定是否应用weight_decay。因此learning_rate_mapping的关键限制是每个正则模式必须至少匹配到一个参数否则训练启动即报错。官方示例来自 router.py 的 docstring 示例——Router 模型中对稀疏静态嵌入使用更高学习率、其余部分保持低学习率args SentenceTransformerTrainingArguments( ..., learning_rate2e-5, learning_rate_mapping{ rSparseStaticEmbedding\.*: 1e-3, }, )命令行字符串兼容与router_mapping一样learning_rate_mapping支持从命令行传入 JSON 字符串training_args.py若json.loads解析失败会抛出ValueError提示其必须是一个“正则 → 学习率”的字典。所有以 dict 类型出现的参数含prompts、router_mapping、fsdp_config、deepspeed等都被登记在_VALID_DICT_FIELDS列表中training_args.py以便 CLI 解析。九、post_init中的自动化行为构造参数对象时BaseTrainingArguments.__post_init__training_args.py会执行若干对用户透明的修正warmup 兼容层Transformers v5 移除了warmup_ratio仅保留warmup_steps可接受浮点比例旧版本则相反。__post_init__自动在两者之间转换并给出弃用警告training_args.py。字符串 → 枚举转换自动把字符串形式的batch_sampler/multi_dataset_batch_sampler转为对应枚举。prediction_loss_onlyTrue因为SentenceTransformerTrainer的compute_loss只计算预测损失这里强制开启以避免额外开销training_args.py。ddp_broadcast_buffersFalse避免基于 BertModel 的模型在 DDP 训练时触发 “variable needed for gradient computation has been modified by an inplace operation” 错误training_args.py。分布式提示非分布式下使用多卡会提示改用 DDPDDP 模式下若未设置dataloader_drop_last会自动置为True以避免最后一批不均匀导致挂起training_args.py。DataLoader worker 提示当dataloader_num_workers 0且使用spawn启动方式而未开启dataloader_persistent_workers时会警告每个 worker 都要重新导入 sentence-transformers 带来额外开销training_args.py。十、实战组合示例以下是一个同时使用多数据集、Router 路由、分组学习率与提示词的综合配置各字段的底层依据分别对应上文第五至八节from sentence_transformers import ( SentenceTransformer, SentenceTransformerTrainer, SentenceTransformerTrainingArguments, ) from sentence_transformers.sentence_transformer.training_args import ( BatchSamplers, MultiDatasetBatchSamplers, ) from sentence_transformers.sentence_transformer.losses import MultipleNegativesRankingLoss from datasets import Dataset, DatasetDict model SentenceTransformer(microsoft/mpnet-base) train_web Dataset.from_dict({ question: [What is the capital of France?, How do I bake bread?], positive: [Paris., Mix flour, water and yeast, then bake.], }) train_medical Dataset.from_dict({ question: [What causes a fever?, How is hypertension treated?], positive: [Infection or inflammation., Lifestyle changes and medication.], }) train_dataset DatasetDict({web: train_web, medical: train_medical}) loss MultipleNegativesRankingLoss(model) args SentenceTransformerTrainingArguments( # 必需参数 output_dircheckpoints, # 通用训练参数继承自 transformers.TrainingArguments num_train_epochs3, per_device_train_batch_size32, per_device_eval_batch_size32, warmup_steps0.1, fp16True, eval_strategysteps, eval_steps250, save_strategysteps, save_steps250, save_total_limit2, logging_steps100, # sentence-transformers 专属参数 prompts{ web: {question: Query: , positive: Document: }, medical: {question: Query: , positive: Document: }, }, batch_samplerBatchSamplers.NO_DUPLICATES, # MultipleNegativesRankingLoss 推荐 multi_dataset_batch_samplerMultiDatasetBatchSamplers.PROPORTIONAL, learning_rate2e-5, learning_rate_mapping{rtransformer\.encoder\.layer\.11\.*: 5e-6}, # 最后一层用更低学习率 ) trainer SentenceTransformerTrainer( modelmodel, argsargs, train_datasettrain_dataset, lossloss, ) trainer.train()十一、参考与延伸阅读参数类定义与 docstringsentence_transformers/sentence_transformer/training_args.py公共基类实现sentence_transformers/base/training_args.py批采样器枚举与实现sentence_transformers/base/sampler.pyRouter 模块与路由映射sentence_transformers/base/modules/router.pyTrainer 对参数的消费逻辑sentence_transformers/base/trainer.py数据整理器对 prompts / router_mapping 的解析sentence_transformers/base/data_collator.py官方训练示例examples/sentence_transformer/training/如 training_stsbenchmark_avg_word_embeddings.py、adaptive_layer_nli.py相关测试tests/base/test_trainer.py、tests/base/samplers/test_round_robin_batch_sampler.py、tests/base/modules/test_router.py训练总览docs/sentence_transformer/training_overview.md、docs/sentence_transformer/training/examples.rst赞分享人工智能NLPEmbedding微调【免费下载链接】sentence-transformersState-of-the-Art Embeddings, Retrieval, and Reranking项目地址https://gitcode.com/gh_mirrors/se/sentence-transformers点击查看免费下载相关推荐sentence-transformers 训练故障排查全指南train-sentence-transformers Skill 症状索引式排障手册sentence transformers 训练故障排查全指南train sentence transformers Skill 症状索引式排障手册 训练嵌入人工智能AI 技能/插件大模型AI 评测《sentence-transformers模型的参数设置详解》《sentence transformers模型的参数设置详解》 引言 在自然语言处理NLP领域模型参数设置的重要性不言而喻。参数的选择和调整直接影响模型RetroWrite性能优化理解零开销二进制插桩背后的设计原理与实现RetroWrite性能优化理解零开销二进制插桩背后的设计原理与实现 RetroWrite作为一款强大的二进制重写框架通过静态插桩技术为COTS商业现货人工智能NLPEmbedding微调上一篇Bluebird Promise.any 详解以 count1 语义快速获取首个成功结果下一篇Shairport Sync中的服务健康监控告警与通知机制创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表