- 人工智能
- NLP
- Embedding
- 微调
【免费下载链接】sentence-transformers
State-of-the-Art Embeddings, Retrieval, and Reranking
导读
本文围绕 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_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_dir | str | 无(必填) | 模型检查点输出目录 |
prompts | str/Dict[str, str]/Dict[str, Dict[str, str]] | None | 为数据集各列(或各数据集)指定提示词模板 |
batch_sampler | BatchSamplers/str/ 采样器类 / 工厂函数 | BatchSamplers.BATCH_SAMPLER | 单数据集下的批采样策略 |
multi_dataset_batch_sampler | MultiDatasetBatchSamplers/str/ 采样器类 / 工厂函数 | MultiDatasetBatchSamplers.PROPORTIONAL | 多数据集下的批次调度策略 |
router_mapping | Dict[str, str]/Dict[str, Dict[str, str]] | {} | 将数据集列映射到 Router 路由(如 "query"、"document") |
learning_rate_mapping | Dict[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_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 中给出了完整说明:
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_DUPLICATES | NoDuplicatesBatchSampler,保证 batch 内样本值(跨列)唯一 | 依赖 batch 内负样本的损失函数 |
BatchSamplers.NO_DUPLICATES_HASHED | NoDuplicatesBatchSampler(precompute_hashes=True),用 xxhash 预计算哈希加速查重,需安装xxhash库,占用少量额外内存 | 同上,尤其推荐用于图像/音频等媒体数据集 |
BatchSamplers.GROUP_BY_LABEL | GroupByLabelBatchSampler,每个 batch 至少包含 2 个不同标签、每个标签至少 2 个样本 | batch 内三元组挖掘类损失 |
与损失函数的搭配建议
源码 docstring 明确给出了推荐组合(sampler.py):
- NO_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.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参数; - 或传入一个接受
dataset、batch_size、drop_last、valid_label_columns、generator、seed并返回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_ROBIN | RoundRobinBatchSampler | 各数据集轮流各取一个 batch,直到某个数据集耗尽;每个数据集被平等对待,但小数据集可能会用不完所有样本 |
底层MultiDatasetDefaultBatchSampler(sampler.py)接收一个ConcatDataset和一组子采样器,并将set_epoch传递给所有子采样器以保证每个 epoch 的可复现打乱。RoundRobinBatchSampler的__len__为min(len(sampler)) * 数量,ProportionalBatchSampler的__len__为各子采样器长度之和(sampler.py)。
多数据集 + 自定义采样器的用法(sampler.py):子类化MultiDatasetDefaultBatchSampler后把类传给参数;或传入接受dataset(ConcatDataset)、batch_samplers(各子数据集采样器列表)、generator、seed的工厂函数。
典型用法(来自 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 模型在训练时把不同输入列送到正确的子编码器。接受两种格式:
Dict[str, str]:列名到路由的映射,如{"question": "query", "positive": "document", "negative": "document"}。Dict[str, Dict[str, str]]:数据集名 → 列名 → 路由 的嵌套映射,用于多数据集训练/评估。
强制校验
如果模型包含Router模块但未提供router_mapping,get_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 的组合
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提示检查模式; - 为匹配参数单独建组,设置
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 类型出现的参数(含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_only=True:因为SentenceTransformerTrainer的compute_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
相关推荐
sentence-transformers 训练故障排查全指南:train-sentence-transformers Skill 症状索引式排障手册
sentence transformers 训练故障排查全指南:train sentence transformers Skill 症状索引式排障手册 训练嵌入
人工智能AI 技能/插件大模型AI 评测《sentence-transformers模型的参数设置详解》
《sentence transformers模型的参数设置详解》 引言 在自然语言处理(NLP)领域,模型参数设置的重要性不言而喻。参数的选择和调整直接影响模型
RetroWrite性能优化:理解零开销二进制插桩背后的设计原理与实现
RetroWrite性能优化:理解零开销二进制插桩背后的设计原理与实现 RetroWrite作为一款强大的二进制重写框架,通过静态插桩技术为COTS(商业现货)
人工智能NLPEmbedding微调
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考