sentence-transformers 训练参数完全指南:SentenceTransformerTrainingArguments 详解
2026/9/21 19:31:10 网站建设 项目流程
  • 人工智能
  • NLP
  • Embedding
  • 微调

【免费下载链接】sentence-transformers

State-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 ↑ BaseTrainingArguments(sentence_transformers.base.training_args) ↑ SentenceTransformerTrainingArguments(sentence_transformers.sentence_transformer.training_args)
  • transformers.TrainingArguments:来自 Hugging Face Transformers 库,提供了绝大多数通用训练参数(num_train_epochsper_device_train_batch_sizewarmup_stepsfp16/bf16eval_strategysave_strategylogging_stepsrun_name等)。
  • BaseTrainingArguments:定义在 sentence_transformers/base/training_args.py,在 Transformers 参数基础上补充了 sentence-transformers 特有的五个参数:promptsbatch_samplermulti_dataset_batch_samplerrouter_mappinglearning_rate_mapping
  • SentenceTransformerTrainingArguments:几乎完全复用BaseTrainingArguments的字段,作为句向量模型的专用入口,同时从 sentence_transformers/base/sampler.py 导出BatchSamplersMultiDatasetBatchSamplers两个枚举供用户直接使用(见 training_args.py 的__all__)。
from sentence_transformers import SentenceTransformerTrainingArguments from sentence_transformers.sentence_transformer.training_args import BatchSamplers, MultiDatasetBatchSamplers

从源码结构看,cross_encodersparse_encoder的训练参数类(如CrossEncoderTrainingArgumentsSparseEncoderTrainingArguments)同样继承自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"、"document")
learning_rate_mappingDict[str, float]{}按参数名正则表达式为不同模块设置不同学习率
num_train_epochsper_device_train_batch_sizeper_device_eval_batch_sizewarmup_stepsfp16/bf16eval_strategyeval_stepssave_strategysave_stepssave_total_limitlogging_stepsrun_namedataloader_num_workersdataloader_drop_last继承自transformers.TrainingArguments与 Transformers 一致通用训练、评估、保存与调度配置

三、必需参数:output_dir

output_dir是所有训练参数中唯一必须提供的参数(继承自 Transformers),用于指定模型检查点与训练产物的写入目录:

args = SentenceTransformerTrainingArguments( output_dir="checkpoints", num_train_epochs=1, per_device_train_batch_size=16, per_device_eval_batch_size=16, warmup_steps=0.1, fp16=True, # GPU 不支持 FP16 时请改为 False bf16=False, # GPU 支持 BF16 时可改为 True eval_strategy="steps", eval_steps=100, save_strategy="steps", save_steps=100, save_total_limit=2, logging_steps=100, run_name="my-sts-training", # 安装 wandb 时用于 W&B 实验名 )

上面的写法直接取自仓库官方示例 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 中给出了完整说明:

  1. str:单一提示词,应用于所有数据集的所有列。无论数据集是datasets.Dataset还是datasets.DatasetDict均适用。
  2. Dict[str, str]:列名到提示词的映射,例如{"query": "Query: ", "document": "Document: "},同样适用于DatasetDatasetDict
  3. Dict[str, str]:数据集名到提示词的映射,仅当数据集是DatasetDictDataset的字典时使用。
  4. 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_mappinglearning_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_hashes=True),用 xxhash 预计算哈希加速查重,需安装xxhash库,占用少量额外内存同上,尤其推荐用于图像/音频等媒体数据集
BatchSamplers.GROUP_BY_LABELGroupByLabelBatchSampler,每个 batch 至少包含 2 个不同标签、每个标签至少 2 个样本batch 内三元组挖掘类损失

与损失函数的搭配建议

源码 docstring 明确给出了推荐组合(sampler.py):

  • NO_DUPLICATES 系列推荐搭配使用 batch 内负样本(in-batch negatives)的损失:MultipleNegativesRankingLossCachedMultipleNegativesRankingLossMultipleNegativesSymmetricRankingLossCachedMultipleNegativesSymmetricRankingLossMegaBatchMarginLossGISTEmbedLossCachedGISTEmbedLoss(源码位于 sentence_transformers/sentence_transformer/losses/)。
  • GROUP_BY_LABEL推荐搭配 batch 内三元组挖掘损失:BatchAllTripletLossBatchHardSoftMarginTripletLossBatchHardTripletLossBatchSemiHardTripletLoss

官方示例中的典型用法(adaptive_layer_nli.py):

args = SentenceTransformerTrainingArguments( output_dir=output_dir, batch_sampler=BatchSamplers.NO_DUPLICATES, # MultipleNegativesRankingLoss 受益于 batch 内无重复样本 ... )

NO_DUPLICATES 的实现细节

NoDuplicatesBatchSampler(sampler.py)在__iter__中基于打乱后的索引构建单链表,逐样本检查其值集合是否与当前 batch 重叠,重叠则推迟到后续 batch,从而保证每个 batch 内样本值跨列唯一。当一轮完整遍历产出的 batch 数少于__len__承诺的批次数(drop_last=True且重复值很多时可能发生),会重新打乱并补充采样最多 2 轮,同时输出一次性警告。NO_DUPLICATES_HASHED变体则通过datasets.map预先计算每行各列的 xxhash64 值,将重复检查从“逐行读数据集、重新解码媒体”变为 O(1) 哈希比对,precompute_num_proc默认取min(8, cpu_count - 1)

GROUP_BY_LABEL 的实现细节

GroupByLabelBatchSampler(sampler.py)要求batch_size为大于等于 4 的偶数,且数据集中至少存在 2 个各含 2 个以上样本的标签,否则抛出ValueError。每个 batch 由多个标签轮流各出 2 个样本构成,保证 batch 内标签多样,是 batch 内三元组挖掘的前提。注意valid_label_columns指定候选标签列名,采样器取第一个在数据集中实际存在的列作为标签来源。

自定义 batch sampler

batch_sampler还接受自定义实现(sampler.py):

  • 子类化DefaultBatchSampler,把类本身(而非实例)传给batch_sampler参数;
  • 或传入一个接受datasetbatch_sizedrop_lastvalid_label_columnsgeneratorseed并返回DefaultBatchSampler实例的工厂函数。

在 training_args.py 的__post_init__中,若传入的是字符串,会被自动转换为对应的枚举值;to_dict(training_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,直到某个数据集耗尽;每个数据集被平等对待,但小数据集可能会用不完所有样本

底层MultiDatasetDefaultBatchSampler(sampler.py)接收一个ConcatDataset和一组子采样器,并将set_epoch传递给所有子采样器以保证每个 epoch 的可复现打乱。RoundRobinBatchSampler__len__min(len(sampler)) * 数量ProportionalBatchSampler__len__为各子采样器长度之和(sampler.py)。

多数据集 + 自定义采样器的用法(sampler.py):子类化MultiDatasetDefaultBatchSampler后把传给参数;或传入接受datasetConcatDataset)、batch_samplers(各子数据集采样器列表)、generatorseed的工厂函数。

典型用法(来自 sampler.py 的官方示例,配合CoSENTLoss训练跨领域 STS):

from 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": ["It's nice weather outside today.", "He drove to work."], "sentence_B": ["It's 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_dir="checkpoints", multi_dataset_batch_sampler=MultiDatasetBatchSamplers.PROPORTIONAL, ) trainer = SentenceTransformerTrainer( model=model, args=args, train_dataset=train_dataset, loss=loss, ) trainer.train()

多数据集时的 dataset_name 列

当数据集是DatasetDict,且满足以下任一条件时(trainer.py),Trainer 会自动为数据集添加dataset_name列:损失是一个字典(按数据集区分损失)、prompts是字典、或router_mapping是“数据集名 → 列 → 路由”的嵌套字典。该列是_resolve_prompts_resolve_router_mapping按数据集名解析配置的关键。

七、router_mapping:为 Router 模型指定路由

router_mapping用于将数据集列映射到 Router 模块的路由(route),例如 "query" 或 "document",从而让 Router 模型在训练时把不同输入列送到正确的子编码器。接受两种格式:

  1. Dict[str, str]:列名到路由的映射,如{"question": "query", "positive": "document", "negative": "document"}
  2. Dict[str, Dict[str, str]]:数据集名 → 列名 → 路由 的嵌套映射,用于多数据集训练/评估。

强制校验

如果模型包含Router模块但未提供router_mappingget_data_collator会直接抛出ValueError(trainer.py),并提示正确的映射示例;对应测试见 tests/base/modules/test_router.py。数据整理器在每批预处理时通过_resolve_router_mapping(data_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 的组合

promptsrouter_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

  1. re.search(parameter_pattern, n)loss_model.named_parameters()中找出所有匹配的参数;
  2. 将这些参数从既有优化器组中剔除(避免同一参数出现在多个组);
  3. 若正则没有匹配到任何参数,则抛出ValueError提示检查模式;
  4. 为匹配参数单独建组,设置lr=learning_rate,并按是否属于衰减参数(get_decay_parameter_names)决定是否应用weight_decay

因此learning_rate_mapping的关键限制是:每个正则模式必须至少匹配到一个参数,否则训练启动即报错。

官方示例

来自 router.py 的 docstring 示例——Router 模型中对稀疏静态嵌入使用更高学习率、其余部分保持低学习率:

args = SentenceTransformerTrainingArguments( ..., learning_rate=2e-5, learning_rate_mapping={ r"SparseStaticEmbedding\.*": 1e-3, }, )

命令行字符串兼容

router_mapping一样,learning_rate_mapping支持从命令行传入 JSON 字符串(training_args.py);若json.loads解析失败会抛出ValueError提示其必须是一个“正则 → 学习率”的字典。所有以 dict 类型出现的参数(含promptsrouter_mappingfsdp_configdeepspeed等)都被登记在_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_only=True:因为SentenceTransformerTrainercompute_loss只计算预测损失,这里强制开启以避免额外开销(training_args.py)。
  • ddp_broadcast_buffers=False:避免基于 BertModel 的模型在 DDP 训练时触发 “variable needed for gradient computation has been modified by an inplace operation” 错误(training_args.py)。
  • 分布式提示:非分布式下使用多卡会提示改用 DDP;DDP 模式下若未设置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_dir="checkpoints", # 通用训练参数(继承自 transformers.TrainingArguments) num_train_epochs=3, per_device_train_batch_size=32, per_device_eval_batch_size=32, warmup_steps=0.1, fp16=True, eval_strategy="steps", eval_steps=250, save_strategy="steps", save_steps=250, save_total_limit=2, logging_steps=100, # sentence-transformers 专属参数 prompts={ "web": {"question": "Query: ", "positive": "Document: "}, "medical": {"question": "Query: ", "positive": "Document: "}, }, batch_sampler=BatchSamplers.NO_DUPLICATES, # MultipleNegativesRankingLoss 推荐 multi_dataset_batch_sampler=MultiDatasetBatchSamplers.PROPORTIONAL, learning_rate=2e-5, learning_rate_mapping={r"transformer\.encoder\.layer\.11\.*": 5e-6}, # 最后一层用更低学习率 ) trainer = SentenceTransformerTrainer( model=model, args=args, train_dataset=train_dataset, loss=loss, ) trainer.train()

十一、参考与延伸阅读

  • 参数类定义与 docstring:sentence_transformers/sentence_transformer/training_args.py
  • 公共基类实现:sentence_transformers/base/training_args.py
  • 批采样器枚举与实现:sentence_transformers/base/sampler.py
  • Router 模块与路由映射:sentence_transformers/base/modules/router.py
  • Trainer 对参数的消费逻辑: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
  • 人工智能
  • NLP
  • Embedding
  • 微调

【免费下载链接】sentence-transformers

State-of-the-Art Embeddings, Retrieval, and Reranking

项目地址:https://gitcode.com/gh_mirrors/se/sentence-transformers
点击查看免费下载

相关推荐

上一篇:Bluebird Promise.any 详解:以 count=1 语义快速获取首个成功结果
下一篇:Shairport Sync中的服务健康监控:告警与通知机制

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询